Skip to content

ELMo: biLM, ScalarMix, linear probes

ModuleL3.5 · side (optional, D31) · Python · Pass 4 · 3 to 4 h
You buildpython/tinyllm/rnn/elmo.py: BiLM (layers, forward), bilm_loss, ScalarMix, fit_linear_probe, probe_accuracy
Contractcourse/contracts/py/tinyllm/rnn/elmo.pyi
Testscourse/tests/L3.5/test_elmo.py (what they check: section 4), golden values from torch 2.14.1 in course/fixtures/L3.5/elmo_torch.npz (course/oracle/L3.5/elmo_torch.py), a learning bar in course/fixtures/ref-thresholds.tsv
NeedsL3.2 LSTM · L3.4 reverse_padded · L0.3 cross_entropy · L0.4 Embedding, Linear, ModuleList · L0.2 ops · L0.1 Tensor · M06.3 PCG32 · tests: M10.3 AdamW · reading: L3.6 (an RNN language model), math/07-probability-statistics (logistic regression) (or --ref-deps)
Used byno module: the side quest sq.elmo-probes (probing every zoo checkpoint) builds on it
Milestonenone (optional; not part of MS-L3)
Optional depthPeters et al., “Deep contextualized word representations” (NAACL 2018); Tenney et al., “What do you learn from context? Probing for sentence structure in contextualized word representations” (ICLR 2019); Hewitt and Liang, “Designing and Interpreting Probes with Control Tasks” (EMNLP 2019)
  • A biLM is two language models over the same sentence: the forward one predicts the next token, the backward one the previous token, and they share the embedding and the softmax (test_bilm_loss_targets, test_golden_torch).
  • Each direction must see only its own side: forward states never depend on later tokens, backward states never on earlier ones, and the backward stack reverses inside each length (test_forward_layers_never_see_the_future, test_backward_layers_never_see_the_past, test_padding_never_leaks).
  • ScalarMix is a softmax-weighted sum of the layers times γ\gamma, learned per task (test_hand_example_scalar_mix, test_scalar_mix_starts_uniform_and_trains).
  • A linear probe measures what a frozen layer encodes: the embedding layer cannot tell “saw” the noun from “saw” the verb, the first LSTM layer can (test_contextual_layers_disambiguate).
Terminal window
ol start L3.5 # stubs elmo.py
ol tests L3.5 # the course tests
ol check L3.5 # exit code is the verdict
ol diff L3.5 # after passing: your code against the reference

Your word vectors so far (L2.3, the embedding table of every model) give a word one vector, whatever the sentence: “saw” in “the saw cuts wood” and in “they saw the dog” is the same point. You have also just trained recurrent language models (L3.6) whose hidden states do depend on the sentence. ELMo (2018) made that the representation: run a forward and a backward LM, keep every layer’s states, and let each downstream task learn a mix of them. It was the step between word2vec and BERT (L6.2), and the linear probe that shows what it learned is the tool interpretability work still uses. The module is optional because nothing in the final system calls it (D31); sq.elmo-probes applies the probe to every checkpoint in the zoo.

SymbolMeaningType / shape
x1:Tx_{1:T}, ℓ\elltoken ids of a sentence, its real lengthint[B, T], int[B]
ddembedding and hidden widthint
LLnumber of LSTM layers per direction (n_layers)int
et=E[xt]e_t = E[x_t]token embeddingfloat32[d]
ftkf^k_t, btkb^k_tforward and backward layer-kk states at position tt, in source orderfloat32[d]
Rt0=[et;et]R^0_t = [e_t ; e_t], Rtk=[ftk;btk]R^k_t = [f^k_t ; b^k_t]the representation of layer kkfloat32[2d]
ss, γ\gammaScalarMix parametersfloat32[L + 1], float32[1]
XX, yyprobe features (one row per token) and labelsfloat[N, D], int[N]

A forward LM’s state ftf_t summarizes x1:tx_{1:t}; a backward LM’s state btb_t summarizes xt:ℓx_{t:\ell}. Concatenated, [ft;bt][f_t ; b_t] describes xtx_t in its whole sentence. They cannot be one bidirectional network trained as a language model: if position tt saw xt+1x_{t+1} it would just copy the answer. So the two directions are trained separately, each to predict in its own direction, and only their states are joined.

Each direction is a stack of LL single-direction LSTMs (L3.2). The forward stack reads ee left to right; the backward stack reads ee reversed inside each sentence’s length with L3.4’s reverse_padded, so it starts at the last real token, and its outputs are reversed back to source order. The forward LM predicts xt+1x_{t+1} from ftLf^L_t, the backward LM xt−1x_{t-1} from btLb^L_t, both through the same output layer:

L=12(mean⁡t≤ℓ−2CE(WftL,xt+1)+mean⁡t≥1CE(WbtL,xt−1)).\mathcal{L} = \tfrac12 \Big( \operatorname{mean}_{t \le \ell - 2} \mathrm{CE}(W f^L_t, x_{t+1}) + \operatorname{mean}_{t \ge 1} \mathrm{CE}(W b^L_t, x_{t-1}) \Big) .

Different layers carry different information: lower ones more about the word, higher ones more about the context. ELMo lets each task choose:

ELMot=γ∑k=0Lsoftmax⁡(s)k Rtk.\mathrm{ELMo}_t = \gamma \sum_{k=0}^{L} \operatorname{softmax}(s)_k \, R^k_t .

The softmax makes the weights positive and sum to 1; γ\gamma rescales the result to whatever the task’s next layer expects. At initialization s=0s = 0 (every layer weighted 1/(L+1)1/(L+1)) and γ=1\gamma = 1, the plain average. Both are trained with the task.

To ask “does layer kk encode X?”, freeze the layer, take its vectors for many tokens, and fit the simplest classifier: softmax regression. If a linear map from RkR^k predicts the label, the information is there in an easily usable form; a deep probe could learn the task itself and say little about the representation (Hewitt and Liang). fit_linear_probe standardizes the features (each column to mean 0 and standard deviation 1, so one learning rate suits every column), runs full-batch gradient descent from zero on the cross-entropy with a small L2 penalty, and folds the standardization back into WW and bb so the probe applies to raw features. From a zero start with full batches it is deterministic, so an accuracy is a property of the layer. (The catalog’s probe is M07.7’s IRLS logistic regression, taught in Pass 5; this module uses its own gradient-descent version.)

ScalarMix over two layers, s=(0,ln⁡3)s = (0, \ln 3), γ=2\gamma = 2, with R0=(1,2)R^0 = (1, 2) and R1=(5,6)R^1 = (5, 6).

  1. softmax⁡(s)=(e0,eln⁡3)/(1+3)=(1/4,3/4)\operatorname{softmax}(s) = (e^0, e^{\ln 3}) / (1 + 3) = (1/4, 3/4).
  2. The mix: 14(1,2)+34(5,6)=(0.25+3.75,0.5+4.5)=(4,5)\tfrac14 (1, 2) + \tfrac34 (5, 6) = (0.25 + 3.75, 0.5 + 4.5) = (4, 5).
  3. Times γ\gamma: (8,10)(8, 10).

With the raw scalars as weights instead of their softmax (0⋅R0+1.0986⋅R10 \cdot R^0 + 1.0986 \cdot R^1) the result would be (10.99,13.18)(10.99, 13.18). This is test_hand_example_scalar_mix.

class ScalarMix(Module):
def __init__(self, n_layers: int) -> None: ...
def weights(self) -> NDArray: ...; def forward(self, layers: Sequence[Tensor]) -> Tensor: ...
class BiLM(Module):
def __init__(self, vocab: int, d: int, n_layers: int = 2, rng=None) -> None: ...
def layers(self, ids, lengths=None) -> list[Tensor]: ... # [R^0 .. R^L], each [B, T, 2d]
def forward(self, ids, lengths=None) -> tuple[Tensor, Tensor]: ... # fwd and bwd logits [B, T, V]
def bilm_loss(model, ids, lengths=None) -> Tensor: ...
def fit_linear_probe(features, labels, n_classes, l2=1e-3, steps=200, lr=0.5) -> tuple[NDArray, NDArray]: ...
def probe_accuracy(W, b, features, labels) -> float: ...
TestKINDChecksWhy it matters downstream
test_hand_example_scalar_mixunitsection 3: (8,10)(8, 10), weights (1/4,3/4)(1/4, 3/4)you and the test agree on the mix
test_scalar_mix_starts_uniform_and_trainsgradcheckthe average at init; gradients of ss, γ\gamma, and every layer in float64a task can learn its mix
test_golden_torchgoldenlayers, both logits, the loss, a mix, and every gradient against torch’s packed nn.LSTMsyour biLM is ELMo’s
test_forward_layers_never_see_the_futurepropertychanging later tokens leaves ftkf^k_t bitwise equalthe forward LM cannot copy its answer
test_backward_layers_never_see_the_pastpropertychanging earlier tokens leaves btkb^k_t bitwise equal, with lengths 6, 4, 2the backward LM reads from each sentence’s own end
test_padding_never_leakspropertypadding content changes no real position; padded RkR^k are 0; a row alone equals it in a batchbatches of mixed lengths
test_bilm_loss_targetsunitthe loss recomputed with explicit targets; length-1 rows only raiseboth directions are scored on the right token
test_probe_separates_separable_dataunitaccuracy 1 on separable clusters far from the origin; folded weights work on raw featuresthe probe measures the features
test_probe_is_deterministic_and_validatesboundaryidentical probes from identical data; bad shapes, labels, one classprobe accuracies are reproducible
test_contextual_layers_disambiguatelearningafter 150 steps: loss at the reference bar; probe on R0R^0 at the majority rate, on R1R^1 at least 0.95the reason contextual vectors exist
PitfallSymptomCaught by
1. a forward layer that also reads the backward statesthe forward LM sees the future and its loss collapses to 0test_forward_layers_never_see_the_future (mutant s03)
2. reversing the padded tensor, or not reversing back, or running the backward LSTMs without lengthsthe backward LM reads padding first, or states land at the wrong positionstest_padding_never_leaks (mutant s04), test_backward_layers_never_see_the_past (mutants s05, s06)
3. scoring a direction on the current tokeneach LM is trained to predict what it already readtest_bilm_loss_targets (mutants s07, s08)
4. returning probe weights for standardized featurespredictions on raw features are wrongtest_probe_separates_separable_data (mutant s09)
raw scalars as mix weightsnegative or unnormalized layer weightstest_hand_example_scalar_mix (mutant s01)
γ\gamma never appliedthe task cannot rescale the mixtest_hand_example_scalar_mix (mutant s02)
gradient ascent in the probeaccuracy falls as it trainstest_probe_separates_separable_data (mutant s10)
DirectionModuleHow it uses this
BackL3.2LSTM, one per layer per direction
BackL3.4reverse_padded reverses each sentence inside its length
BackL0.3cross_entropy with ignore_index for both directions
BackL0.4Embedding, Linear, ModuleList
BackL0.2the op library
BackL0.1Tensor
BackM06.3PCG32 for initialization
BackM10.3AdamW in the learning test
Forwardsq.elmo-probesthe side quest probes every zoo checkpoint layer by layer for part of speech and position
ForwardL6.2BERT replaces two one-directional LMs with one masked LM that reads both sides at every layer

No module calls this code (it is a side module, D31). It is still worth doing: the probe is the cheapest way to see what a network has learned, and the biLM’s “each side only sees its own side” rule is exactly the causality property L5.2 tests for transformer masks.

Your pieceProduction equivalentWhat it addsWhere to look
BiLMAllenNLP’s ELMoa character CNN instead of a token embedding, projections and residual connections between layers, 4096-wide LSTMsallennlp/modules/elmo.py, elmo_lstm.py
ScalarMixAllenNLP ScalarMixoptional layer normalization of each layer before mixingallennlp/modules/scalar_mix.py
fit_linear_probeedge probing, control tasksspan-pair probes, selectivity against random control labelsTenney et al. 2019; Hewitt and Liang 2019
contextual vectorsBERT, sentence-transformersone bidirectional encoder trained by masked LM, pooled sentence embeddingsL6.2; ag.06’s embedding endpoint