"""Layer 11: the output head. ln_f + lm_head over the vocab. Supports weight tying with layer 1's wte. The 12-layer stack: layer 0 + layer 1 + 9x layer 2 + this. """ from dataclasses import dataclass import torch.nn as nn @dataclass class Layer11Config: vocab_size: int = 109 d_model: int = 512 class Layer11(nn.Module): def __init__(self, cfg: Layer11Config): super().__init__() self.cfg = cfg self.ln_f = nn.LayerNorm(cfg.d_model) self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False) def tie_to(self, layer1): """Share layer 1's wte: head.weight IS wte.weight (0 extra params).""" self.lm_head.weight = layer1.wte.weight def forward(self, x): return self.lm_head(self.ln_f(x))