Skip to content

Module system and basic layers

ModuleL0.4 · build · Python · Pass 2 · 4 to 5 h, plus your graded tests (rung R3)
You buildpython/tinyllm/nn/module.py: Module (registration, named_parameters, state_dict, load_state_dict, train/eval, zero_grad); python/tinyllm/nn/layers.py: Linear, Embedding, LayerNorm, Dropout, ReLU, Tanh, GELU, Sequential, ModuleList; and your own tests in python/tests/l0-4-module/, written first
Contractcourse/contracts/py/tinyllm/nn/module.pyi · course/contracts/py/tinyllm/nn/layers.pyi
Testscourse/tests/L0.4/ (what they check: section 4), golden values from torch 2.14 in course/fixtures/L0.4/layers_torch.npz; your tests are graded by mutation, threshold 0.70 plus one required fault, with a red-then-green journal
NeedsL0.1 Tensor (parameters) · L0.2 the ops every forward is written in · M07.3 normal_init · M06.3 PCG32 init and dropout streams · reading: craft.03 red then green (or --ref-deps)
Used byL0.5 BigramLogits is a Module, the loop trains Linear stacks · L0.6 a checkpoint is a module’s state_dict · later: L11.1, L2.2, L3.2, L3.3, L3.6, L4.1, L4.2, L4.3, L5.3, L5.4, L5.5, L6.1, L6.2, L6.3, L6.5, L6.6, L7.1, L7.2, L7.5, L7.6, L7.8, L7.9, L8.5
MilestoneMS-L0 (the digits MLP is two Linear layers)
Optional depththe PyTorch torch.nn.Module documentation and source (torch/nn/modules/module.py); Ba, Kiros, Hinton, “Layer Normalization” (2016)
  • A module finds its parameters and children by attribute assignment, in assignment order, and the dotted names this produces are the safetensors keys of every checkpoint in the course, identical to PyTorch’s (test_state_dict_names_match_torch).
  • load_state_dict matches by name, checks every key and shape before copying anything, and writes into the existing arrays, so an optimizer built earlier keeps training the loaded weights (test_load_state_dict_strict, test_load_is_in_place).
  • Linear stores one row per output, W∈Rout×inW \in \mathbb{R}^{\text{out} \times \text{in}}, and computes xW⊤+bxW^\top + b, PyTorch’s layout, which is what lets HF weights load by name (test_hand_example_linear, test_matches_torch).
  • LayerNorm uses the population variance with ϵ\epsilon inside the square root (test_layernorm_eps_inside_the_square_root).
  • eval() and train() reach every descendant, and a weight tied under two names is one parameter (test_train_eval_reaches_every_child, test_parameters_order_and_tying).
Terminal window
ol start L0.4 # stubs module.py and layers.py; prints your test path and rung (R3)
ol tests L0.4 # the course tests
# write ONE test in python/tests/l0-4-module/, then:
ol tdd red L0.4 # must FAIL against your current code: records the red
# make it pass, then:
ol tdd green L0.4 # must PASS with the same test files: records the green
# repeat for each test; then:
ol check L0.4 # course tests, the red-then-green journal, the mutation grade
ol mutate L0.4 # the full grade, cached by your test files' hash

With L0.1 to L0.3 you can train anything, but a model is still a loose set of Tensors you have to pass to the optimizer and to the checkpoint writer by hand, in the right order, every time. The digits MLP has four parameter arrays; a Llama layer (L7.9) has nine, a 30-layer model 270. Worse, the checkpoint format (formats/safetensors.md) names tensors by strings such as model.layers.3.self_attn.q_proj.weight, and the Rust engine (L10.1) and HF’s loaders find weights by those names. This module gives every model one way to list its parameters, name them exactly as PyTorch does, save and load them, and switch dropout off for evaluation.

SymbolMeaningType / shape
xxinput batch, rows of dind_{\text{in}} featuresfloat32[..., d_in]
WWLinear weight, one row per outputfloat32[d_out, d_in]
bbLinear biasfloat32[d_out]
y=xW⊤+by = xW^\top + bLinear outputfloat32[..., d_out]
yˉ\bar{y}upstream gradientshape of yy
μ,σ2\mu, \sigma^2mean and population variance of one row (over the last axis, dd entries)scalars per row
ϵ\epsilonLayerNorm’s guard, 10−510^{-5}float
γ,β\gamma, \betaLayerNorm’s affine weight and biasfloat32[d]
N(0,s2)\mathcal{N}(0, s^2)normal distribution with standard deviation ss (M07.3)

Registration by assignment. Module.__init__ creates two ordered dictionaries, parameters and children, before anything else. __setattr__ then sorts every assignment: a Tensor with requires_grad=True is a parameter, a Module is a child, anything else (a number, a constant Tensor such as a causal mask, None) is plain state and unregisters the name. That is why a subclass must call super().__init__() first: without the dictionaries there is nowhere to register, and the error should say so. Python dictionaries keep insertion order, so registration order is assignment order.

Names are the contract. named_parameters() yields the module’s own parameters first, then each child’s with the child’s name and a dot as prefix, depth first. A two-layer Sequential(Linear(2, 3), ReLU(), Linear(3, 1)) yields 0.weight, 0.bias, 2.weight, 2.bias: exactly torch.nn.Sequential’s state_dict keys. A Tensor registered under two names (tied input and output embeddings in L7.9) is listed once, under its first name, by identity; otherwise the optimizer would step it twice.

state_dict and load_state_dict. state_dict() is {name: copy of the array} in named_parameters order; copies, so a saved state does not change when training continues. load_state_dict(sd, strict=True) first compares keys: with strict, any missing or unexpected key is a KeyError that names all of them, raised before anything is copied (a half-loaded model is worse than none). A shape mismatch is a ValueError, also before copying. Then each array is written into the existing parameter array (p.data[...] = a), converted to the parameter’s dtype. The optimizer (M10.2, M10.3) holds references to those Tensors and arrays; replacing them would leave it updating arrays the model no longer uses. strict=False loads the names that match, which is how you load a pretrained backbone under a new head.

train and eval. training is a flag on every module; Dropout reads it. train(mode) sets it on the module and every descendant and returns the module; eval() is train(False). Setting it only on the top module leaves every nested dropout active during evaluation.

Linear in PyTorch’s layout. WW has shape [d_out, d_in], so y=xW⊤+by = xW^\top + b, and backward (from L0.1’s matmul VJP) gives Wˉ=yˉ⊤x\bar{W} = \bar{y}^\top x, bˉ=∑rowsyˉ\bar{b} = \sum_{\text{rows}} \bar{y}, xˉ=yˉW\bar{x} = \bar{y}W. Storing WW as [d_in, d_out] computes the same function from its own initialization, and then fails the day you load torch weights by name. Initialization follows M07.3: W∼N(0,1/din)W \sim \mathcal{N}(0, 1/d_{\text{in}}) (standard deviation 1/din1/\sqrt{d_{\text{in}}}), so unit-variance inputs give unit-variance outputs, and b=0b = 0. Embedding(n, d) is a table [n, d] with N(0,1)\mathcal{N}(0, 1) rows and F.embedding as its forward.

Seeded initialization. Each layer takes an optional rng, a PCG32 (M06.3), and draws its weights from it once, at construction, in parameter registration order. With no rng, layers use the spec’s sub-streams of seed 0: PCG32(0).substream("init") for weights and substream("dropout") for dropout masks, so adding a dropout layer does not change any initial weight.

LayerNorm. For each row (the last axis): x^=(x−μ)/σ2+ϵ\hat{x} = (x - \mu)/\sqrt{\sigma^2 + \epsilon}, then y=γ⊙x^+βy = \gamma \odot \hat{x} + \beta, with σ2\sigma^2 the population variance (F.var with correction 0, as torch). ϵ\epsilon belongs inside the square root: it guards rows whose variance is tiny, and outside the root it does almost nothing there. The forward is written with F.mean, F.var, and **, so its backward comes from L0.2 and needs no code here.

Containers. Sequential(*mods) registers its children as "0", "1", … and applies them in order. ModuleList(mods) only holds them (no forward); append registers the next index and returns the list.

Linear(2, 3) with W=[123456]W = \begin{bmatrix}1&2\\3&4\\5&6\end{bmatrix}, b=[0.5,−0.5,1]b = [0.5, -0.5, 1], and one input row x=[[1,1]]x = [[1, 1]].

Forward. xW⊤xW^\top takes the dot product of xx with each row of WW: [1+2,3+4,5+6]=[3,7,11][1 + 2, 3 + 4, 5 + 6] = [3, 7, 11]. Add bb: y=[[3.5,6.5,12]]y = [[3.5, 6.5, 12]].

Backward with yˉ=[[1,1,1]]\bar{y} = [[1, 1, 1]]:

  1. Wˉ=yˉ⊤x=[111][1,1]\bar{W} = \bar{y}^\top x = \begin{bmatrix}1\\1\\1\end{bmatrix}[1, 1], a 3×23 \times 2 matrix of ones: each weight WijW_{ij} multiplied xj=1x_j = 1 once.
  2. bˉ=[1,1,1]\bar{b} = [1, 1, 1]: the bias is added once per output.
  3. xˉ=yˉW=[1+3+5,2+4+6]=[[9,12]]\bar{x} = \bar{y}W = [1 + 3 + 5, 2 + 4 + 6] = [[9, 12]]: each input feeds all three outputs.

LayerNorm on a nearly constant row. x=[0,0.002]x = [0, 0.002]: μ=0.001\mu = 0.001, σ2=10−6\sigma^2 = 10^{-6}. With ϵ=10−5\epsilon = 10^{-5} inside the root, x^=±0.001/1.1×10−5=±0.30151\hat{x} = \pm 0.001/\sqrt{1.1 \times 10^{-5}} = \pm 0.30151, torch’s value. With ϵ\epsilon added after the root, ±0.001/(0.001+10−5)=±0.990\pm 0.001/(0.001 + 10^{-5}) = \pm 0.990.

These are test_hand_example_linear and test_layernorm_eps_inside_the_square_root.

python/tinyllm/nn/module.py
class Module:
training: bool
def __init__(self) -> None; def __setattr__(self, name, value) -> None
def forward(self, *args, **kwargs); def __call__(self, *args, **kwargs)
def named_parameters(self, prefix="") -> Iterator[tuple[str, Tensor]]; def parameters(self)
def named_modules(self, prefix="") -> Iterator[tuple[str, "Module"]]
def train(self, mode=True) -> "Module"; def eval(self) -> "Module"; def zero_grad(self) -> None
def state_dict(self) -> dict[str, NDArray]; def load_state_dict(self, sd, strict=True) -> None
# python/tinyllm/nn/layers.py
Linear(in_f, out_f, bias=True, rng=None); Embedding(n, d, rng=None); LayerNorm(d, eps=1e-5)
Dropout(p, rng=None); ReLU(); Tanh(); GELU(approximate="none"); Sequential(*mods); ModuleList(mods=())
TestKINDChecksWhy it matters downstream
test_hand_example_linearunitsection 3: y=[[3.5,6.5,12]]y = [[3.5, 6.5, 12]], Wˉ\bar{W} ones, bˉ\bar{b} ones, xˉ=[[9,12]]\bar{x} = [[9, 12]]you and the test agree on the layout
test_matches_torchgoldenevery layer with torch’s weights copied in by name gives torch’s outputs and gradientsHF weights load by name (L7.9)
test_state_dict_names_match_torchgoldenthe exact keys and order torch produces, through nested containersthe safetensors key contract
test_state_dict_roundtrippropertysave, rebuild with another seed, load: same function; saved arrays are copiescheckpoints (L0.6)
test_load_is_in_placeunitthe Tensors and arrays an optimizer holds are the ones loaded intoresuming a run (L0.6)
test_load_state_dict_strictboundarystrict names every missing and unexpected key and copies nothing; wrong shapes fail; non-strict loads the resta wrong checkpoint fails loudly
test_parameters_order_and_tyingunitown parameters before children’s; a tied Tensor onceweight tying (L7.9)
test_plain_state_is_not_a_parameterboundaryconstants and numbers are not parameters; None unregistersmasks and buffers in attention (L5)
test_forgetting_super_init_is_explainedboundarythe error names super().__init__()the most common first bug
test_train_eval_reaches_every_childuniteval() reaches nested dropout; both return the modelevaluation is deterministic (L0.5)
test_zero_gradunitevery parameter’s .grad is clearedgradients accumulate (L0.1)
test_layernorm_normalizespropertyrows come out with mean 0 and variance 1every transformer block (L5)
test_layernorm_eps_inside_the_square_rootboundarysection 3’s ±0.30151\pm 0.30151torch parity on small-variance rows
test_init_statistics_and_seedingstatisticalN(0,1/din)\mathcal{N}(0, 1/d_{\text{in}}) weights, zero bias, seed-reproduciblestable signal through depth (M07.3)
test_sequential_and_modulelistunitorder, indexing, appendstacks of blocks (L5, L7)

Rung R3 gives you the interface (above) and one test; you write the rest before the code that makes each pass. The given test:

python/tests/l0-4-module/test_module.py
import numpy as np
from tinyllm.autograd.tensor import Tensor
from tinyllm.nn.layers import Linear
def test_hand_example_linear():
"""W = [[1, 2], [3, 4], [5, 6]], b = [0.5, -0.5, 1], x = [[1, 1]]: y = [[3.5, 6.5, 12]], dx = [[9, 12]]."""
lin = Linear(2, 3)
lin.load_state_dict({"weight": np.array([[1, 2], [3, 4], [5, 6]]), "bias": np.array([0.5, -0.5, 1.0])})
x = Tensor([[1.0, 1.0]], requires_grad=True)
y = lin(x)
assert np.allclose(y.data, [[3.5, 6.5, 12.0]])
y.backward(np.ones((1, 3)))
assert np.allclose(x.grad, [[9.0, 12.0]])

Then, one at a time, red then green: state_dict names in registration order with no key for a missing bias; loading by name (a reordered dict gives the same parameters); loading in place; strict mode rejecting missing and unexpected keys without changing anything; a tied parameter listed once; eval() turning dropout off in nested modules; LayerNorm against its formula. Import only contract names (tinyllm.nn.module, tinyllm.nn.layers, tinyllm.autograd.tensor). ol check L0.4 requires, for every test file, a ol tdd red record before its last ol tdd green, a mutation score of at least 0.70, and the required fault killed (it is the one section 5’s last pitfall describes).

PitfallSymptomCaught by
1. assigning a parameter before super().__init__()an AttributeError about _params, far from the causetest_forgetting_super_init_is_explained (mutant m06)
2. train()/eval() setting only the top module’s flagdropout stays on in nested blocks during evaluationtest_train_eval_reaches_every_child (mutant s14)
3. load_state_dict replacing the arraysa resumed run’s optimizer updates arrays the model no longer uses: the loss stops movingtest_load_is_in_place (mutant s09)
4. loading by position, checking keys after copying, or skipping the shape checka reordered checkpoint loads into the wrong layers; a wrong one half-loads; a broadcastable array loads silentlytest_state_dict_roundtrip (mutant s08), test_load_state_dict_strict (mutants s10, s11)
5. LayerNorm with the sample variance, or ϵ\epsilon outside the rootsmall disagreements with torch everywhere, large ones on near-constant rowstest_layernorm_normalizes (mutant s03), test_layernorm_eps_inside_the_square_root (mutant s04)
6. Linear weight as [in, out], or a forward without the transposetrains fine from scratch, then torch and HF weights load transposed or failtest_hand_example_linear, test_matches_torch (mutants s01, s02)

| Forward | L11.1 | Registered call site uses this module. | | Forward | L2.2 | Registered call site uses this module. | | Forward | L3.2 | Registered call site uses this module. | | Forward | L3.3 | Registered call site uses this module. | | Forward | L3.6 | Registered call site uses this module. | | Forward | L4.1 | Registered call site uses this module. | | Forward | L4.2 | Registered call site uses this module. | | Forward | L4.3 | Registered call site uses this module. | | Forward | L5.3 | Registered call site uses this module. | | Forward | L5.4 | Registered call site uses this module. | | Forward | L5.5 | Registered call site uses this module. | | Forward | L6.1 | Registered call site uses this module. | | Forward | L6.2 | Registered call site uses this module. | | Forward | L6.3 | Registered call site uses this module. | | Forward | L6.5 | Registered call site uses this module. | | Forward | L6.6 | Registered call site uses this module. | | Forward | L7.1 | Registered call site uses this module. | | Forward | L7.2 | Registered call site uses this module. | | Forward | L7.5 | Registered call site uses this module. | | Forward | L7.6 | Registered call site uses this module. | | Forward | L7.8 | Registered call site uses this module. | | Forward | L7.9 | Registered call site uses this module. | | Forward | L8.5 | Registered call site uses this module. |

DirectionModuleHow it uses this
BackL0.1parameters are Tensors with requires_grad=True
BackL0.2Linear is F.matmul and F.transpose; LayerNorm is F.mean and F.var; Embedding is F.embedding
BackM07.3normal_init draws the initial weights
BackM06.3PCG32(0).substream("init") and substream("dropout") are the defaults
ForwardL0.5BigramLogits subclasses Module; the loop trains Sequential MLPs
ForwardL0.6save_checkpoint writes state_dict() and resume calls load_state_dict

If you skip this module, ol check L0.5 stops with L0.5 needs L0.4: build it, or rerun with --ref-deps.

Your pieceProduction equivalentWhat it addsWhere to look
Module registrationtorch.nn.Modulebuffers (saved, not trained), forward and backward hooks, _load_from_state_dict per module for versioned formatstorch/nn/modules/module.py
state_dict namesHF transformers key mappingrenames between checkpoint versions, sharded model-0000x-of-0000y.safetensors with an index filemodeling_utils.py, _load_state_dict_into_model
Linear, LayerNormfused kernels in Apex and LigerLayerNorm and RMSNorm forward and backward in one pass, fused with the residual addliger_kernel/ops/layer_norm.py
explicit rng for initJAX and Flax PRNGKey splittingevery random draw takes a key, so initialization is a pure function of the seedflax/linen/module.py (make_rng)