File size: 4,735 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
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
"""
TransformerBlock
─────────────────
A single decoder layer combining:
  - RMSNorm     (normalisation)
  - Pre-LayerNorm ordering  (stability)
  - CausalSelfAttention     (with RoPE + GQA + SDPA)
  - SwiGLU FeedForward      (gated activation)

─────────────────────────────────────────
WHY PRE-LAYERNORM (PRE-LN) vs POST-LN:
─────────────────────────────────────────

Original "Attention is All You Need" used Post-LN:
    x = LayerNorm(x + Attention(x))

The norm is applied AFTER the residual addition. This means the residual
signal must pass through the normaliser at every layer. With many layers,
gradients at the early layers become vanishingly small (the norm acts as a
bottleneck), causing instability or requiring extremely small learning rates.

Pre-LN applies the norm BEFORE the sublayer:
    x = x + Attention(LayerNorm(x))

Now the residual path is clean β€” the gradient flows straight through the
addition without being rescaled. This makes training far more stable at any
depth and is the standard in every modern LLM (Llama, Mistral, GPT-NeoX,
Falcon, Gemma).

─────────────────────────────────────────
WHY RMSNORM vs LAYERNORM:
─────────────────────────────────────────

LayerNorm computes:
    mean = mean(x)
    std  = std(x)
    out  = (x - mean) / std * gamma + beta

RMSNorm simplifies to:
    rms = sqrt(mean(xΒ²))
    out = x / rms * gamma

RMSNorm omits:
  1. Mean subtraction (re-centering) β€” empirically unnecessary for LLMs
  2. Beta parameter β€” saves a small amount of memory

Result: ~15% faster than LayerNorm with equivalent training dynamics.
Used in: Llama, Mistral, Falcon, Gemma, Qwen.
"""

import torch
import torch.nn as nn

from model.attention import CausalSelfAttention
from model.feed_forward import SwiGLUFeedForward


class RMSNorm(nn.Module):
    """
    Root Mean Square Layer Normalisation.

    Args:
        dim: Feature dimension to normalise over.
        eps: Small constant for numerical stability (prevents division by zero
             when the RMS is extremely small).
    """

    def __init__(self, dim: int, eps: float = 1e-5):
        super().__init__()
        self.eps    = eps
        self.weight = nn.Parameter(torch.ones(dim))   # learnable scale (Ξ³)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        # Compute RMS along the last dimension, normalise, then re-scale
        rms = x.pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
        return x * rms * self.weight


class TransformerBlock(nn.Module):
    """
    One decoder layer (Pre-LN + RMSNorm + CausalSelfAttention + SwiGLU FFN).

    Args:
        config: ModelConfig instance. All hyperparameters are read from here.

    Forward signature:
        x         β€” input tensor [B, T, d_model]
        cos, sin  β€” RoPE tables from RotaryEmbedding.forward(T)
        past_kv   β€” optional KV-cache tuple for inference

    Returns:
        (x, present_kv)
        x          β€” output tensor [B, T, d_model]
        present_kv β€” updated (k, v) tuple to pass back during inference
    """

    def __init__(self, config):
        super().__init__()
        self.attn_norm = RMSNorm(config.d_model, eps=config.norm_eps)
        self.ff_norm   = RMSNorm(config.d_model, eps=config.norm_eps)
        self.attn = CausalSelfAttention(
            d_model    = config.d_model,
            n_heads    = config.n_heads,
            n_kv_heads = config.n_kv_heads,
            dropout    = config.dropout,
        )
        self.ff = SwiGLUFeedForward(config.d_model, config.d_ff)

    def forward(
        self,
        x:       torch.Tensor,
        cos:     torch.Tensor,
        sin:     torch.Tensor,
        past_kv: tuple | None = None,
    ) -> tuple[torch.Tensor, tuple]:
        # ── Attention sublayer (Pre-LN) ───────────────────────────────────────
        # Normalise BEFORE attention so residual stream stays unscaled.
        attn_out, present_kv = self.attn(self.attn_norm(x), cos, sin, past_kv)
        x = x + attn_out                         # residual connection

        # ── Feed-Forward sublayer (Pre-LN) ────────────────────────────────────
        x = x + self.ff(self.ff_norm(x))         # residual connection

        return x, present_kv