Download model/feed_forward.py from Vivid86/MiniTransformer-91M: direct link, hf CLI and curl.
- Browser
- Download file 2.05 kB
-
https://huggingface.co/Vivid86/MiniTransformer-91M/resolve/main/model/feed_forward.py
- Command line
-
hf download hf://Vivid86/MiniTransformer-91M/model/feed_forward.py
-
curl -L -o feed_forward.py https://huggingface.co/Vivid86/MiniTransformer-91M/resolve/main/model/feed_forward.py
2.05 kB
| """ | |
| SwiGLU Feed-Forward Network | |
| ββββββββββββββββββββββββββββ | |
| Replaces the original GELU/ReLU FFN with SwiGLU β the activation function | |
| used in Llama, Mistral, PaLM, and most modern open-weight LLMs. | |
| WHY SWIGLU OVER GELU/RELU: | |
| Standard FFN: FFN(x) = max(0, xW1 + b1) W2 + b2 (ReLU) | |
| SwiGLU FFN: FFN(x) = (SiLU(xW1) β xW3) W2 (SwiGLU) | |
| The key difference is the *gating* term (xW3). Instead of a fixed | |
| activation, SwiGLU uses a learned gate that modulates how much of the | |
| activated signal passes through. This gives the FFN extra expressiveness | |
| at the same parameter count. | |
| SiLU (Sigmoid Linear Unit) = x Β· sigmoid(x) | |
| Smooth, non-monotone; empirically outperforms GELU in this gated form. | |
| PARAMETER COUNT: | |
| Standard 4Γ FFN: 2 matrices (d_model β d_ff β d_model) | |
| SwiGLU FFN: 3 matrices (w1, w3: d_model β d_ff; w2: d_ff β d_model) | |
| To keep FLOPs equal to a 4Γ FFN, set d_ff β 8/3 Γ d_model (see ModelConfig). | |
| BIAS=FALSE: | |
| All projections omit biases following Llama/Mistral practice β biases add | |
| negligible capacity at scale and waste memory bandwidth. | |
| """ | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| class SwiGLUFeedForward(nn.Module): | |
| """ | |
| SwiGLU Feed-Forward Network. | |
| Computes: FFN(x) = W2(SiLU(W1 x) β W3 x) | |
| Args: | |
| d_model: Input/output dimension. | |
| d_ff: Hidden (intermediate) dimension. | |
| Use ModelConfig.d_ff (auto-sized to 8/3 Γ d_model β nearest 64). | |
| """ | |
| def __init__(self, d_model: int, d_ff: int): | |
| super().__init__() | |
| self.w1 = nn.Linear(d_model, d_ff, bias=False) # gate projection | |
| self.w3 = nn.Linear(d_model, d_ff, bias=False) # value projection | |
| self.w2 = nn.Linear(d_ff, d_model, bias=False) # output projection | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| # SiLU(gate) β value, then project back down | |
| return self.w2(F.silu(self.w1(x)) * self.w3(x)) | |