| 629 | """ |
| 630 | |
| 631 | def __init__( |
| 632 | self, |
| 633 | dim: int, |
| 634 | dim_out: Optional[int] = None, |
| 635 | mult: int = 4, |
| 636 | dropout: float = 0.0, |
| 637 | activation_fn: str = "geglu", |
| 638 | final_dropout: bool = False, |
| 639 | inner_dim=None, |
| 640 | bias: bool = True, |
| 641 | ): |
| 642 | super().__init__() |
| 643 | if inner_dim is None: |
| 644 | inner_dim = int(dim * mult) |
| 645 | dim_out = dim_out if dim_out is not None else dim |
| 646 | linear_cls = LoRACompatibleLinear if not USE_PEFT_BACKEND else nn.Linear |
| 647 | |
| 648 | if activation_fn == "gelu": |
| 649 | act_fn = GELU(dim, inner_dim, bias=bias) |
| 650 | if activation_fn == "gelu-approximate": |
| 651 | act_fn = GELU(dim, inner_dim, approximate="tanh", bias=bias) |
| 652 | elif activation_fn == "geglu": |
| 653 | act_fn = GEGLU(dim, inner_dim, bias=bias) |
| 654 | elif activation_fn == "geglu-approximate": |
| 655 | act_fn = ApproximateGELU(dim, inner_dim, bias=bias) |
| 656 | |
| 657 | self.net = nn.ModuleList([]) |
| 658 | # project in |
| 659 | self.net.append(act_fn) |
| 660 | # project dropout |
| 661 | self.net.append(nn.Dropout(dropout)) |
| 662 | # project out |
| 663 | self.net.append(linear_cls(inner_dim, dim_out, bias=bias)) |
| 664 | # FF as used in Vision Transformer, MLP-Mixer, etc. have a final dropout |
| 665 | if final_dropout: |
| 666 | self.net.append(nn.Dropout(dropout)) |
| 667 | |
| 668 | def forward(self, hidden_states: torch.Tensor, scale: float = 1.0) -> torch.Tensor: |
| 669 | compatible_cls = (GEGLU,) if USE_PEFT_BACKEND else (GEGLU, LoRACompatibleLinear) |