YAML Metadata Warning:empty or missing yaml metadata in repo card

Check out the documentation for more information.

ANLP Assignment 2 — Mixture of Experts, Optimizers, and Decoding Strategies

Checkpoints and full experimental record for three controlled studies on a small decoder-only transformer. Everything — attention, RoPE, RMSNorm, the MoE router, all five optimizers and all four decoding strategies — is implemented from basic PyTorch operations. nn.Transformer, nn.MultiheadAttention, the torch.optim update rules and model.generate() are not used anywhere.

Placeholders to fill before publishing: author name and roll number, code.iiit.ac.in repository URL, Weights & Biases run URL, and the exact checkpoint filenames listed under Repository contents.

Study Question Best result
Part 1 — Mixture of Experts What does sparsity cost, and what does it buy, under matched parameter budgets? C5 (4 experts, top-2, wide): 25.42 BLEU vs 24.96 dense at equal active parameters
Part 2 — Optimizers How much of a new optimizer's advantage survives fair per-optimizer tuning? MARS: 38.01 test perplexity; AdamW highest BLEU (2.27)
Part 3 — Decoding strategies What does each decoder actually optimise? top_k_10: 0.193 unigram F1; top_p_0.95: 0.650 distinct-2

Repository contents

Nine checkpoints in total, roughly 1.0 GB together. Part 3 uses a frozen public checkpoint (EleutherAI/pythia-160m) and produces none.

Part Checkpoints Source path in the training repo
Part 1 5 (one per FFN variant C1–C5) runs/part1/<config>/model.pt
Part 2 4 (one per optimizer) runs/part2/<optimizer>/model.pt

Note / to verify: Part 2 trains and reports five optimizers (AdamW, MARS, Lion, Muon, Sophia-G) but the checkpoint inventory records four. Confirm which optimizer's checkpoint is absent and either upload it or state the omission here.


Shared architecture and training configuration

Held fixed across every run of Parts 1 and 2, so that each ablation varies one thing at a time.

Component Choice
Architecture Pre-norm decoder-only transformer, 6 layers, d_model 512, 8 heads
Attention Causal multi-head self-attention with RoPE, written from scratch
Normalisation RMSNorm, plus a final norm before the head
Embeddings Tied input/output, 16k byte-level BPE
Precision bf16 autocast, fp32 master weights and loss
Schedule Linear warmup then cosine decay to 10% of peak
Initialisation N(0, 0.02); residual projections scaled by 1/sqrt(2L)
Seed 1337 — data order, initialisation and sampling are reproducible
Hardware Single NVIDIA RTX 6000 Ada, shared with unrelated jobs

Because the GPU was shared, wall-clock totals from the training runs reflect what else the device was doing. Per-step optimizer cost is therefore measured separately under controlled conditions (see Part 2).


Part 1 — Mixture of Experts

Task and data

Translation is posed as plain language modelling over a single sequence:

<bos> <vi> {source} <en> {target} <eos>

Cross-entropy is masked to the English target tokens and the closing <eos>; the source is read but never predicted. Because routing happens on the forward pass regardless of whether a token contributes to the loss, source tokens are still dispatched to experts — which is what makes the per-language routing analysis possible at all.

Each of the ~446k rows of belumind/en-vi-ja-curated-500k-triplets yields two training examples, a Vietnamese–English pair and a Japanese–English pair. Examples are packed back-to-back into 256-token blocks rather than padded individually: a typical pair is around 60 tokens, so per-example padding would have spent roughly three quarters of every batch on padding.

The tokenizer is a 16k byte-level BPE trained jointly on all three languages. Byte-level is a requirement rather than a preference here — the dataset is not romanised, so a vocabulary built on an ASCII assumption would fail on Vietnamese diacritics and on Japanese kana and kanji entirely.

Parameter-budget methodology

A single 2-layer MLP of hidden width d_ff costs 2 · d_model · d_ff weights. Writing E for the number of experts and k for how many are active per token, an MoE layer built from the same primitive costs E times that in total and k times it per token. The two budget rules therefore reduce to two choices of per-expert width: dividing d_ff by E holds the total fixed (C2–C4), and dividing by k holds the active count fixed instead (C5), which necessarily doubles the total.

Variant d_ff/expert Total params Active params Active % Tokens
C1 dense MLP 2048 27,073,024 27,073,024 100.0% 40.0M
C2 4 experts, top-1 512 27,085,312 17,648,128 65.2% 40.0M
C3 4 experts, top-2 512 27,085,312 20,793,856 76.8% 40.0M
C4 1 shared + 3 routed, top-1 512 27,082,240 20,790,784 76.8% 40.0M
C5 4 experts, top-2, wide 1024 39,668,224 27,085,312 68.3% 40.0M

C1–C4 agree on total parameters to within 12k (the routers, which C1 does not have); C5's active count matches C1's total to within 12k. All five consume an identical 40.0M training tokens in the same order, so the only thing that varies is the feed-forward block.

Routing

A linear router scores the experts per token; the top-k gates are renormalised to sum to one so that the block's output scale does not depend on k. Without that, top-2 would feed roughly twice the activation magnitude into the residual stream as top-1, and the comparison would be measuring a scale change as much as a routing change.

Two auxiliary terms are added to the loss: a Switch-style load-balancing term E · Σ_i f_i P_i and a router z-loss that penalises large logits. The balancing term normalises f_i by k as well as by the token count, so it bottoms out at 1.0 for any k; without that normalisation top-2 would carry twice the balancing pressure of top-1 at the same coefficient, confounding two changes in what is meant to be a one-variable ablation.

Results

Perplexity is per English target token; BLEU is corpus BLEU (sacreBLEU) over 800 greedy generations, 400 per source language.

Variant Test loss Perplexity BLEU BLEU vi→en BLEU ja→en Routing entropy Min expert share Train min
C1 dense MLP 1.5354 4.643 24.96 35.73 14.68 — — 5.8
C2 4 experts, top-1 1.7064 5.509 20.81 32.09 10.84 1.000 0.239 10.9
C3 4 experts, top-2 1.5429 4.678 23.81 35.90 13.14 1.000 0.225 6.0
C4 1 shared + 3 routed, top-1 1.5698 4.805 22.94 34.72 12.66 1.000 0.317 8.3
C5 4 experts, top-2, wide 1.5181 4.563 25.42 36.85 14.84 0.999 0.221 10.4

Analysis — what the parameter budget buys

Sparsity is not free, and top-1 pays the most for it. C2 routes each token to a single expert of width 512 and reaches perplexity 5.51 against the dense baseline's 4.64, losing 4.2 BLEU. It is the only variant whose loss curve separates visibly from the others, and the gap is established early and never closes. The cause is capacity per token rather than capacity overall: C2 holds the same 27.1M total parameters as C1, but any individual token is processed by a 512-wide MLP instead of a 2048-wide one, and a quarter of the feed-forward width is simply not enough for this task.

Top-2 recovers almost all of it. C3 uses the same experts as C2 and changes only how many fire per token, moving from 65.2% to 76.8% active parameters. That single change is worth 0.83 perplexity and 3.0 BLEU, landing C3 within 0.03 perplexity and 1.2 BLEU of the dense model while doing roughly three quarters of its feed-forward work. Combining two 512-wide experts is evidently a decent substitute for one 2048-wide one; combining one is not.

The shared expert did not help here. C4 activates exactly as many parameters as C3 (20.79M) but spends one of its two active slots on an expert that every token sees, and comes out 0.13 perplexity and 0.9 BLEU behind C3. The intended benefit of a shared expert is that it absorbs the generic transformation every token needs, freeing the routed experts to specialise; the cost is that the router chooses among only three experts instead of four, and only one routed expert's worth of specialised capacity reaches any given token. At this scale the cost dominates — half of C4's active feed-forward budget is spent on a pathway that is by construction identical for every token, which is the one thing a dense MLP already does well.

Holding active parameters fixed instead makes MoE a win. C5's active count (27.09M) matches the dense baseline's total to within 12k, so the two do essentially the same arithmetic per token, but C5 carries 39.7M parameters against C1's 27.1M. It is the best model in the table on every metric, and it wins on both language directions separately. This is the actual argument for MoE, and the ablation isolates it cleanly: extra parameters that are only conditionally used still improve the model, at no extra per-token compute. The margin is real but modest — around 0.5 BLEU — which is unsurprising at this budget, since the literature's large MoE gains come from expert counts in the tens or hundreds and token budgets far past the 40M used here.

Analysis — division of labour between experts

Load balancing worked, arguably too well. Normalised routing entropy (1.0 = every expert gets an equal share) is 1.000, 1.000, 1.000 and 0.999 for C2–C5. No expert collapsed and none was starved, which is the failure mode the auxiliary loss exists to prevent. The flip side is that the balancing pressure is itself a constraint on specialisation: the loss explicitly penalises the router for sending one language disproportionately to one expert, so whatever specialisation survives does so against that pressure.

Comparing raw minimum shares across variants would mislead, since C4 routes among three experts and the others among four. Measured against each variant's own balanced share, the least-used expert receives 0.954 of it in C2, 0.899 in C3, 0.951 in C4 and 0.885 in C5. Every variant is within 12% of perfect balance, and C5 — the best-performing one — is the furthest from it, a small hint that a little imbalance is the price of useful specialisation rather than a defect.

Specialisation is nonetheless visible, and it is by language. Conditioning on the input language (each column gives the share of that language's tokens going to each expert, so columns sum to one):

  • C2 is the clearest case, since a top-1 router has no way to hedge: expert 1 takes 33% of Japanese tokens but only 20% of English, while expert 2 takes 32% of English and just 12% of Japanese — less than half the balanced share. Those two experts have divided Japanese from English between them.
  • C3 shows the same axis more weakly (expert 0 at 31% of Japanese against 19% of English), which is what one would expect once each token gets two experts and can spread its bets.
  • C5 shows the tidiest structure of the four: expert 0 leans Vietnamese (30%), expert 1 leans Japanese (29%), expert 2 leans English (29%), and expert 3 sits near neutral on all three.
  • C4's routed experts show the same pattern on the ja/vi axis (expert 2 at 38% of Japanese against 28% of Vietnamese) with English left nearly uniform — plausibly because the always-on shared expert is already absorbing the English-generation work.

The dominant split throughout is script and language, not syntax or topic. That is the natural thing for a router reading token embeddings to key on: Vietnamese, Japanese and English occupy almost disjoint regions of a byte-level BPE vocabulary, so language identity is close to linearly decodable from the embedding alone and is by far the cheapest signal available to a linear router.

Translation quality by direction

Every variant translates Vietnamese far better than Japanese — 35.7 against 14.7 BLEU for the dense baseline, and the ~21-point gap holds across all five. Three things contribute: Vietnamese shares the Latin script and much subword structure with English, so the shared BPE transfers directly; Vietnamese word order is closer to English than Japanese's verb-final, heavily case-marked order, so less long-range reordering is needed; and Japanese text carries no whitespace, so the tokenizer has to infer segmentation that is given for free in the other two languages. The ranking of the five variants is identical on both directions, so the FFN ablation and the language difficulty are independent effects.


Part 2 — Optimizers

Five optimizers implemented from scratch, one from each of the first four categories of Table 1 of Fantastic Pretraining Optimizers and Where to Find Them, plus the bonus Hessian-based entry. torch.optim.Optimizer is used strictly as a base class, for its parameter-group and state bookkeeping; no torch.optim update rule is called anywhere.

The model is the Part 1 dense baseline retrained from scratch on browndw/human-ai-parallel-corpus for next-token prediction. That corpus contains several parallel versions of every document — the human original and one continuation per generator model — so the split is made on the base document id, with all versions of a document moving together. Splitting rows at random would have put near-paraphrases of the same text into both train and test and quietly flattered every validation number.

The update rules

Category Implementation Optimizer state
1. AdamW AdamW 2×
2. Variance-reduced AdamW variant MARS 3×
3. Memory-efficient Lion 1×
4. Matrix-based Muon (Newton–Schulz) 1× on 2D, 2× elsewhere
5. Hessian-based (bonus) Sophia-G 2×

AdamW — the reference point. Two exponential moving averages per parameter, of the gradient and of its square; the update divides one by the root of the other, making the step approximately scale-free per coordinate. The "W" is decoupled decay: the shrinkage is applied to the weights directly rather than folded into the gradient, so it is not itself rescaled by sqrt(v) the way an L2 penalty would be.

MARS — variance-reduced. AdamW builds its moments from raw minibatch gradients and so inherits the full minibatch noise. MARS inserts a STORM-style correction first, extrapolating along the difference between consecutive gradients: c_t = g_t + γ(β₁/(1-β₁))(g_t - g_{t-1}). Because consecutive gradients differ only by the minibatch draw plus one small parameter step, their difference cancels part of the shared sampling noise. The corrected gradient is clipped to unit norm — it is a difference of two noisy quantities and can blow up early — then fed to an otherwise ordinary AdamW. Cost: a third state tensor for the previous gradient.

Lion — memory-efficient. One state tensor instead of two, and the sign of an interpolated momentum. The detail worth noticing is that it uses two different betas: the direction comes from a fresher blend (β₁ = 0.9) than the one that is stored (β₂ = 0.99), so the step reacts to the current gradient while the memory stays long. Because sign makes every coordinate's step exactly ±lr, the update norm is fixed and much larger than AdamW's at the same learning rate — which is why Lion needs its own sweep grid rather than inheriting AdamW's, and why its selected rate lands an order of magnitude lower.

Muon — matrix-based. AdamW normalises each weight coordinate independently, which throws away the fact that a weight matrix is a linear map. Muon normalises the momentum matrix: it replaces M with the nearest semi-orthogonal matrix UVᵀ, so every singular direction is stepped equally and no direction dominates merely because it happened to carry a large singular value. An SVD per step would be hopeless, so the orthogonalisation uses a quintic Newton–Schulz iteration whose coefficients (3.4445, −4.7750, 2.0315) push singular values towards 1 aggressively in five steps. They do not converge to exactly 1, but Muon only needs the direction. Muon applies to 2D hidden matrices; embeddings, the tied head and the 1D gains fall back to AdamW, as in the reference implementation.

Sophia-G — Hessian-based (bonus). Sophia divides by an estimate of the Hessian diagonal rather than the gradient second moment — curvature is what actually says how far one can step before the loss turns back up. Two things make it affordable. The Gauss–Newton–Bartlett estimator: for a cross-entropy loss, sampling labels equally from the model's own predictive distribution and squaring the resulting gradient gives an unbiased estimate of the Gauss–Newton approximation to that diagonal, needing one extra backward pass and no second-order autograd. And clipping: the update is clip(m/max(h, ε), ρ), which bounds the step wherever the curvature estimate is tiny, noisy or negative. The clip is what makes a stale estimate safe, which in turn lets the estimator be refreshed only every k = 10 steps, holding the amortised overhead near 10%.

Correctness verification

All five are checked by a verification suite (src/part2/verify_optimizers.py):

  • This AdamW matches torch.optim.AdamW to 3.0e-08 after 25 steps from identical initialisations on identical batches — i.e. to float32 rounding. That pins down the shared scaffolding (parameter groups, decoupled decay, bias correction) as known-good for all five.
  • Lion's step is exactly ±lr per coordinate and uses the second beta for its stored momentum.
  • MARS clips its corrected gradient to unit norm.
  • Muon's Newton–Schulz iteration compresses the singular spectrum toward 1.
  • Sophia's clip stops saturating once the Gauss–Newton–Bartlett batch scaling is applied.

Hyperparameter sweep (bonus)

The five update rules disagree about the peak learning rate by two orders of magnitude, so handing them all AdamW's value would have measured tuning effort rather than the update rules. Each was swept over its own grid on a short slice of the regime — 4M tokens, about 8% of the full run.

Optimizer Grid searched Selected Val loss @4M At grid edge?
AdamW 1e-4, 3e-4, 6e-4, 1e-3, 3e-3 1.0e-03 5.5231 no
MARS 1e-4, 3e-4, 6e-4, 1e-3, 3e-3, 6e-3, 1e-2 3.0e-03 5.5003 no
Lion 3e-5, 1e-4, 3e-4, 1e-3 3.0e-04 5.8462 no
Muon 5e-3, 1e-2, 2e-2, 5e-2, 1e-1, 2e-1 5.0e-02 5.1898 no
Sophia-G 1e-5, 3e-5, 1e-4, 3e-4, 6e-4, 1e-3, 3e-3 1.0e-04 6.0415 no

The selected rates span a factor of ~60, which is the point: Lion's sign update has a fixed step norm and needs a much smaller rate, while Muon's orthogonalised update tolerates a much larger one. Three of the five grids had to be widened after a first pass selected a value sitting on a grid edge — an edge-valued optimum means the search found the boundary of the grid, not the minimum.

Training regime

Each optimizer then trains on one full pass over the corpus — 41.8M tokens, the "1× the dataset" point — with validation loss and test BLEU measured every 0.1×. For this 27.1M-parameter model the Chinchilla-optimal budget would be 541M tokens, so the whole corpus amounts to 0.077× Chinchilla. Every model here is therefore heavily under-trained relative to its size, and the rankings below are early-training dynamics, not converged results.

The BLEU panel is visibly noisier than the loss panel and the curves cross repeatedly, so it should not be read as a ranking. BLEU here scores a greedy 48-token continuation of a held-out document against the single continuation that actually followed — a task with many acceptable answers and one reference — and the absolute values sit near 2, where a handful of n-gram matches moves the score materially. Validation loss is measured over far more tokens and is the reliable signal at this scale.

Measuring cost honestly

The training runs shared a GPU with unrelated jobs, so their wall-clock totals say more about what else the device was doing than about the update rules: Muon finished in 3.5 minutes and AdamW in 7.1, which would imply Muon is twice as fast as AdamW when it in fact does strictly more work per step. Rather than report that, every optimizer was re-timed back-to-back on the same model, batch and device (src/part2/benchmark.py), separating the optimizer step from the forward/backward pass they all share. Those are the step ms numbers below.

Results

Optimizer Category LR Val loss Test ppl Test BLEU Step ms State
AdamW 1. AdamW 1.0e-03 3.7953 41.192 2.27 3.57 2.0×
MARS 2. Variance-reduced AdamW variant 3.0e-03 3.7181 38.007 2.18 6.12 3.0×
Lion 3. Memory-efficient 3.0e-04 3.9307 47.010 2.07 3.25 1.0×
Muon 4. Matrix-based 5.0e-02 3.7464 39.292 1.98 12.41 1.3×
Sophia-G 5. Hessian-based (bonus) 1.0e-04 3.9262 46.967 1.39 2.76 2.0×

Analysis

Ranking. By test perplexity the order is MARS (38.01), Muon (39.29), AdamW (41.19), Sophia-G (46.97), Lion (47.01). The spread between best and worst is 9.00 perplexity.

The learning rate matters more than the optimizer does. Across all five optimizers at their own best rates, final validation loss spans 0.21 nats. Within AdamW alone, varying nothing but the peak learning rate spans 0.85 nats — 4.0 times as much. The two figures are measured at different budgets (4M tokens for the sweep, the full corpus for the runs) so the ratio is indicative rather than exact, but the direction is not in doubt: choosing the learning rate badly costs far more than choosing the optimizer badly. This is the practical form of the argument in Wen et al. — an optimizer comparison conducted at a single shared learning rate mostly measures which optimizer that rate happened to suit, and a new method can look impressive simply by having been tuned more carefully than the baseline it is compared against.

Muon. The matrix-based update came out better than AdamW (39.29 vs 41.19 perplexity) while tolerating a learning rate 50× larger. That is the expected signature: orthogonalising the momentum fixes the scale of the update's singular values, so the learning rate no longer has to be set conservatively enough for the largest one. It also keeps only one state tensor for its 2D parameters, so it is cheaper in memory than AdamW as well — a rare combination, and consistent with the strong showing matrix-based methods get in the paper. Its cost is the Newton–Schulz iteration: five extra matmuls per 2D parameter per step, visible as 12.41 ms against AdamW's 3.57 ms.

Lion. Half the optimizer memory of AdamW (1.0× model size against 2.0×) for 5.82 perplexity worse. Its selected learning rate is 3× smaller than AdamW's, exactly as the sign update predicts: every coordinate moves by precisely the learning rate, so the update norm is fixed rather than adaptive and the rate has to absorb that. In a memory-bound setting the trade is attractive; here, where memory was not the constraint, it is simply a slightly weaker optimizer.

MARS. Variance reduction gave 3.19 perplexity better than AdamW at 3.0× model size in state — the highest memory cost of the five, since it must keep the previous gradient alongside both moments. The correction is most valuable when gradient noise dominates, which favours small batches; at the 24×512 batch used here the gradients are already fairly clean, so there is less noise for it to remove.

Sophia-G. The curvature-based method was the weakest here (46.97 perplexity) and needed the smallest learning rate of the five (1e-04). Part of that is structural: the clipped ratio caps every coordinate's step at ±lr, and the more often that cap binds the more the update resembles Lion's sign step, so the learning rate has to absorb a partly fixed update norm. The rest is regime. Sophia's advantage comes from dividing by curvature, but the Gauss–Newton estimate is least informative early in training, far from any minimum — which is exactly and only where this experiment lives, at 0.077× Chinchilla — while it still pays for an extra forward/backward every ten steps. Sophia is reported to pay off over long runs; one tenth of a Chinchilla budget is not a setting in which that could show up, and this result should be read as "not yet useful here" rather than "worse".

Caveat on the ranking. Only the peak learning rate was swept. Betas, weight decay, warmup and (for Sophia) ρ were left at published defaults, which favours whichever optimizer's defaults happen to suit a 27M-parameter model at 0.1× Chinchilla. The differences reported above are also single-seed. They should be read as directional, not as a benchmark.


Part 3 — Decoding Strategies

Greedy, top-k, top-p and beam search (widths 1, 2, 4) implemented on top of the frozen EleutherAI/pythia-160m checkpoint and evaluated on hamishivi/ROCStories. model.generate() is not used; each strategy drives model.forward() directly and maintains its own KV cache, which is what keeps the cost linear in the number of generated tokens instead of quadratic.

Implementation details that matter for correctness

Prompts in a batch have different lengths and are left-padded, so that every real final token sits at index −1 where the decoder expects it. But left padding makes absolute positions wrong, so position_ids is derived from the cumulative sum of the attention mask rather than from arange. In beam search, a hypothesis that has emitted EOS is retired by masking its row of log-probabilities to −∞, so it can never be extended again yet still competes for the final selection.

Sanity checks. Beam search at width 1 produces token-identical output to greedy decoding on every test prompt, as it must. Widths 2 and 4 return hypotheses whose length-normalised log-probability is greater than or equal to greedy's on every prompt — beam search is finding genuinely better-scoring sequences, not merely different ones. Sampling is reproducible under a fixed seed, and the top-p filter keeps 1 token on a peaked distribution and 2 on a uniform one at p = 0.5, as expected.

Metric definitions

Free-form generation has no single correct answer, so accuracy, precision, recall and F1 have to be defined rather than assumed. As implemented in src/part3/evaluate.py:

Metric Definition
Reference perplexity Perplexity of the gold continuation under teacher forcing. A property of the model, not the decoder, so it is identical for every strategy (30.52) and acts as a control.
Generation perplexity Perplexity of the model on its own output. Varies by strategy and measures how safely each one plays.
Precision / recall / F1 Unigram overlap with the reference by multiset intersection, micro-averaged over the corpus. Precision is matched/|generated|, recall is matched/|reference|.
Accuracy Position-wise token agreement with the reference over the overlapping prefix; exact-match rate reported alongside.
distinct-1/2, repetition Diversity: the share of generated unigrams/bigrams that are unique, and the share of bigrams that repeat within a generation.

Two deserve a caveat up front. Exact match is essentially always zero and is reported only to make the point that it is the wrong instrument for this task. And the overlap metrics systematically favour whichever strategy produces the most generic, highest-probability text, because generic text reuses the reference's common words; they cannot see whether the output is repetitive or dull, which is precisely why the diversity columns are there.

Results

All strategies on 1000 ROCStories test samples, 48 new tokens each. Reference perplexity (the control) is 30.52 for every row.

Strategy P R F1 Acc Gen ppl distinct-1 distinct-2 Repetition Len Sec
greedy 0.131 0.149 0.140 0.012 2.50 0.028 0.090 0.549 48.0 25.4
top_k_10 0.181 0.206 0.193 0.011 8.10 0.068 0.349 0.086 47.9 26.0
top_k_50 0.172 0.192 0.182 0.010 14.30 0.096 0.485 0.044 47.1 26.2
top_p_0.90 0.155 0.174 0.164 0.009 25.65 0.151 0.604 0.033 47.3 29.2
top_p_0.95 0.153 0.173 0.162 0.009 35.18 0.168 0.650 0.026 47.6 30.2
beam_1 0.131 0.149 0.140 0.012 2.50 0.028 0.090 0.549 48.0 31.2
beam_2 0.118 0.135 0.126 0.012 2.21 0.023 0.068 0.581 48.0 34.0
beam_4 0.113 0.129 0.121 0.011 2.03 0.022 0.060 0.611 48.0 45.3

Analysis of the generations

Greedy decoding degenerates, and the metrics show it twice over. beam_4 has a bigram repetition rate of 0.611 — more than half of all bigrams it emits are repeats of one it has already emitted in the same continuation — and a distinct-2 of 0.060. Reading the outputs makes the mechanism plain: given the prompt "David had achieved his lifelong goal.", beam search at width 4 returns "I'm not going to lie to you." five times in a row and nothing else. Once a high-probability phrase is in the context, repeating it becomes more probable still, and a decoder that only ever maximises probability has no mechanism to escape the loop it has walked into.

The consequence is that sampling beats greedy on the overlap metrics too — which is not the usual expectation. top_k_10 has the highest unigram F1 (0.193) against greedy's 0.140, a relative improvement of 38%. Precision is where greedy loses most (0.131 against 0.181): a generation that spends half its tokens repeating one clause contains few distinct words, so most of what it emits cannot match anything in the reference. The common claim that probability-maximising decoders win precision-style metrics while sampling wins diversity does not hold here, because degeneration costs greedy both at once.

Quality against truncation width is an inverted U, not a trade-off curve. Ordering the sampling runs by how much probability mass they keep — top-k 10, top-k 50, top-p 0.90, top-p 0.95 — F1 goes 0.193, 0.182, 0.164, 0.162, rising from greedy's 0.140 to a peak at top_k_10 and falling away again. Both extremes are bad for different reasons: too little truncation and the sampler reaches into the unreliable tail, too much and it collapses back towards the repetition that afflicts greedy. The best setting is a middle one, and it has to be found empirically.

Perplexity of a model's own output is the most revealing column in the table. The human reference continuations score 30.52 under this model, yet greedy decoding produces text of perplexity 2.50 — its own output is about 12× more predictable than genuine human writing. That single comparison is the clearest statement of why maximising probability is the wrong objective for open-ended generation: real text simply does not sit at the mode of the distribution. top_p_0.95 lands closest to the human figure (35.18 against 30.52), which is exactly what nucleus sampling was designed to achieve.

Top-k and top-p differ in how they choose their width. Top-k keeps a fixed number of candidates however peaked the distribution is; top-p keeps a fixed probability mass and so adapts, narrowing to one token where the model is confident and widening where it is not. That adaptivity shows up here as top-p reaching higher diversity (distinct-2 0.650 at p = 0.95 against 0.485 at k = 50) and landing nearer the human perplexity, at some cost in overlap.

That cost is not only statistical, and the sample generations show what it looks like. Top-p at 0.95 produces fluent, varied, non-repetitive English — and, on the same prompt, invents "Pandas", "Dorrit" and "the Kári complex", none of which has anything to do with David and his lifelong goal. The failure mode has moved rather than disappeared: greedy fails by saying the same true-ish thing forever, sampling fails by confidently wandering off topic. A 160M-parameter model has weak enough long-range coherence that both failure modes are easy to trigger, and no decoding strategy can repair a model that does not track the prompt well in the first place — decoding redistributes the errors, it does not remove them.

Positional accuracy is uninformative and exact match is degenerate. Accuracy sits near 0.012 for every strategy and the exact-match rate is 0.000 at best. Both are reported because the deliverable asks for them, but neither can distinguish the strategies: requiring the model to reproduce a specific human-written sentence token-for-token is the wrong question to ask of open-ended story continuation, where many different continuations are equally valid. The overlap and diversity columns are what carry the signal.

Reproduction

uv sync

uv run main.py params                    # parameter-budget table for the 5 variants
uv run main.py part1 --all               # train all five FFN variants
uv run main.py part2 --sweep             # learning-rate sweep (bonus)
uv run main.py part2 --all --use_sweep   # train with each optimizer
uv run main.py part3 --n_samples 1000    # decoding-strategy benchmark
uv run main.py plots                     # regenerate every figure into assets/

uv run python -m src.part2.verify_optimizers   # optimizer correctness suite
uv run python -m src.part2.benchmark           # controlled per-step cost

Individual parts also run directly, e.g. uv run python -m src.part1.train --config c3_moe_4e2a --token_budget 40000000.

Checkpoints and metrics land in runs/ (git-ignored); figures and the value tables behind them land in assets/.

Environment. The lockfile pins torch==2.6.0+cu124. The machine these runs were done on has a 550.x driver, which supports CUDA 12.4 but not the CUDA 13 wheels that a bare torch>=2.6 resolves to on PyPI, so that pin and the explicit [[tool.uv.index]] block in pyproject.toml are load-bearing.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support