ereniko commited on
Commit
a726540
·
verified ·
1 Parent(s): 3e6f95c

Delete feedforward.py

Browse files
Files changed (1) hide show
  1. feedforward.py +0 -27
feedforward.py DELETED
@@ -1,27 +0,0 @@
1
- import torch
2
- import torch.nn as nn
3
- import torch.nn.functional as F
4
-
5
-
6
- class SwiGLU(nn.Module):
7
- """SwiGLU feed-forward block (Section 4.6), as used in Llama/PaLM.
8
-
9
- Standard formulation: down_proj(silu(gate_proj(x)) * up_proj(x))
10
- The inner dim is scaled down from the naive 4x so that SwiGLU's extra
11
- gate_proj matrix doesn't blow the parameter budget relative to a plain MLP
12
- of the same nominal "4x" size -- this matches how Llama-style models size it.
13
- """
14
-
15
- def __init__(self, hidden_dim: int, mult: float = 4.0):
16
- super().__init__()
17
- # standard correction: 4 * hidden * (2/3) keeps param count comparable
18
- # to a plain (non-gated) 4x MLP, rounded to a clean multiple of 8.
19
- inner_dim = int(hidden_dim * mult * 2 / 3)
20
- inner_dim = ((inner_dim + 7) // 8) * 8
21
-
22
- self.gate_proj = nn.Linear(hidden_dim, inner_dim, bias=False)
23
- self.up_proj = nn.Linear(hidden_dim, inner_dim, bias=False)
24
- self.down_proj = nn.Linear(inner_dim, hidden_dim, bias=False)
25
-
26
- def forward(self, x: torch.Tensor) -> torch.Tensor:
27
- return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))