Repository navigation
feat: Phase 3 — decoder-only transformer forward pass - #3
Merged
Merged
Conversation
Implements the full model architecture as specified in ARCHITECTURE.md §5:
token_ids [B,T]
→ TokenPositionEmbedding (token embed + learned absolute position embed)
→ TransformerBlock × N (pre-norm RMSNorm + CausalSelfAttention + SwiGluFfn + residuals)
→ RMSNorm (final)
→ LM head (n_embd → vocab_size, weight-tied to token embedding)
→ logits [B, T, vocab_size]
New files
─────────
src/model/mod.rs — module declarations + PocModel re-export
src/model/embedding.rs — TokenPositionEmbedding: two Embedding tables summed
src/model/ffn.rs — SwiGluFfn: silu(x@Wgate)*(x@Wup)@wdown, no bias
src/model/attention.rs — CausalSelfAttention: MHA + causal mask (no GQA)
src/model/block.rs — TransformerBlock: pre-norm + residual both sublayers
src/model/transformer.rs — PocModel: full pipeline + PocModel::random_init helper
Modified files
──────────────
src/main.rs — add 'mod model;' declaration
Architecture decisions
──────────────────────
- Standard MHA (no GQA) — simpler weight layout for forgetting measurement
- Learned absolute position embeddings (not RoPE) — fewer confounds at seq≤128
- Weight tying via reshape+matmul: x[B,T,C]→[B*T,C] @ embed_weight^T → [B,T,V]
- Causal mask: additive -inf upper triangle, built per-forward from x.device()
- Candle broadcasting note: expand() required explicitly; + does not auto-broadcast
Tests (27/27 pass, 0 clippy warnings)
──────────────────────────────────────
✓ Milestone 1: forward shape [B,T,vocab_size] — nano + small configs
✓ Milestone 2: causal mask — logits pos 0-1 unchanged when pos 2+ tokens differ
✓ Milestone 3: no NaN/Inf on random input — nano + small configs
✓ All 11 prior tests (Phase 0–2) still pass
Closes Phase 3 milestone. Next: Phase 4 — Stage 1 training + checkpoint save.
git tag v0.3.0
CHANGELOG.md
- Added Phase 3 entry under [Unreleased]:
all 6 new model files, architecture decisions (no GQA/RoPE, weight
tying via reshape+matmul, causal mask expand), candle gotchas, and
full test results table (27/27 pass, 0 clippy warnings)
ROADMAP.md
- Phase 3 marked ✅ DONE in phase map quick-reference
- All task checkboxes [x] — embedding, attention, ffn, block, transformer
- Test checkboxes [x] — shape, causal mask, NaN/Inf
- Milestone block updated: 27/27 tests pass, commit done, tag pending
README.md
- Project Structure tree expanded with model/ subfiles and phase labels
- Status section rewritten: Phase 3 complete, PR #3 linked
- Phase progress table added (0–12) with Done/Next/Planned indicators
Also: cargo fmt applied to all model source files (fixes CI fmt check)
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
Implements Phase 3 — Model Architecture from
ROADMAP.md(lines 203–243).Full decoder-only transformer forward pass: token IDs in,
[batch, seq, vocab_size]logits out — verified on bothnanoandsmallconfigs with random weights.Architecture (ARCHITECTURE.md §5)
Files Added / Changed
src/model/mod.rsPocModelre-exportsrc/model/embedding.rsTokenPositionEmbedding— twoEmbeddingtables summed element-wisesrc/model/ffn.rsSwiGluFfn—silu(x@Wgate)*(x@Wup)@Wdown, no biassrc/model/attention.rsCausalSelfAttention— standard MHA + additive causal mask, no GQAsrc/model/block.rsTransformerBlock— pre-norm + residual for both sublayerssrc/model/transformer.rsPocModel— full pipeline +PocModel::random_inittest helpersrc/main.rsmod model;declarationKey Design Decisions
Weight tying
When
cfg.tie_embeddings == true(default for both presets), the LM head reuses the token embedding weight matrix via reshape + matmul:No separate
lm_head.weighttensor is allocated, halving the output-layer parameter count.Causal mask
Built inside
forwardfromx.device()so it always lives on the correct device. Uses additive-infmask (not multiplicative) applied before softmax. Explicitlyexpand()-ed to[B, n_head, T, T]— candle does not auto-broadcast via+.No GQA, no RoPE
Standard MHA with full Q/K/V per head. Learned absolute position embeddings. Both choices follow ARCHITECTURE.md §5.2 and §5.4 — they minimise confounds in the forgetting measurement that starts at Phase 7.
Tests — 27/27 pass
test_forward_shape_nano[2,16] → [2,16,2000]✅test_forward_shape_small[1,32] → [1,32,4000]✅test_causal_mask_end_to_endtest_no_nan_inf_nanotest_no_nan_inf_smalltest_weight_tying_enabled_by_defaulttie_embeddings=true✅test_causal_mask_isolation(attention unit)CI
The existing CI (
.github/workflows/ci.yml) runs on this PR:cargo checkon Rust 1.89.0cargo fmt --check→cargo check→cargo clippy -D warnings→cargo test→ release build → CLI smokerustsec/audit-checkon dependenciesAll three jobs are expected to pass — local verification already confirms build, clippy, and test results match CI requirements.
Next Phase
Phase 4 — Stage 1 training + checkpoint save
src/checkpoint.rs—save_checkpoint/load_checkpoint(real.safetensors)src/train.rs— AdamW training loop,train_stage1main.rs—train-stage1CLI subcommandMilestone:
cargo run -- train-stage1 --config nano --steps 2000converges and writescheckpoints/stage1/.