File size: 2,053 Bytes
a2ec932
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
"""
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))