Repository navigation
feat: Phase 4 — Stage 1 training loop, real checkpoint save/load - #4
Merged
Merged
Conversation
src/checkpoint.rs [NEW]
- save_checkpoint(vm, cfg, tokenizer, dir) → model.safetensors +
config.json + tokenizer.json via VarMap::save + existing helpers
- load_checkpoint(dir, device) → (ModelConfig, BpeTokenizer,
VarBuilder<'static>) via from_mmaped_safetensors (correct candle
pattern for disk-loaded inference)
- 3 milestone tests: three-files, round-trip logits, config-matches
src/train.rs [NEW]
- Trainer { model, vm: VarMap, optimizer: AdamW, cfg, model_cfg, step }
owns model + VarMap + AdamW together (correct lifetime ordering)
- step_batch(input_ids, target_ids) → f32 loss: forward [B,T,V],
flatten [B*T,V], cross_entropy, backward_step
- warmup_lr(): linear 0→lr over warmup_steps, then constant
- train_stage1(): encode JSONL corpus → batches → run loop,
log every eval_steps, return (Trainer, losses)
- 2 milestone tests: no-nan loss, loss-decreases in 60 steps
src/main.rs [MODIFIED]
- mod checkpoint; mod train; declarations
- TrainStage1Args: --config (nano|small), --steps, --data-dir,
--checkpoint-dir
- cmd_train_stage1: validates corpus, selects preset, trains BPE
tokenizer, calls train_stage1, prints loss summary, saves checkpoint
- Phase 3 marked ✅ done in long_about
docs:
- CHANGELOG.md: Phase 4 entry with arch decisions + test table
- ROADMAP.md: all [ ] → [x], milestone block filled, phase map ✅
- README.md: status → Phase 4 done, phase table updated
Milestone result:
cargo run --release -- train-stage1 --config nano --steps 200
first step loss: 41.78 → last step loss: 10.44
checkpoints/stage1/ (5.1 MB safetensors + 160 B config + 15 KB tokenizer)
32/32 tests pass · 0 clippy warnings
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
Phase 4 of the continual-learning-poc roadmap. Implements the full Stage 1 training pipeline: AdamW loop on the names-facts corpus, checkpoint save/load as real
model.safetensors+config.json+tokenizer.json, and thetrain-stage1CLI subcommand.What Changed
src/checkpoint.rs[NEW]Single-responsibility checkpoint module:
save_checkpoint(vm, cfg, tokenizer, dir)VarMap::save()→model.safetensors;cfg.save_json()→config.json;tokenizer.save()→tokenizer.jsonload_checkpoint(dir, device)(ModelConfig, BpeTokenizer, VarBuilder<'static>)viafrom_mmaped_safetensors— correct candle pattern for disk-load without a VarMapKey design decision:
load_checkpointreturnsVarBuilder<'static>(not a reconstructedPocModel) so the caller retains full control over construction. This matches how Phase 5 (eval) and Phase 6 (warm-start) will use it.src/train.rs[NEW]Training loop. Owns model + VarMap + optimizer as a unit:
candle_nn::loss::cross_entropyon flattened[B*T, V]logitsinput = ids[0..T-1],target = ids[1..T]lroverwarmup_steps, then constantoptimizer.backward_step(&loss)— single call for backward + weight updatesrc/main.rs[MODIFIED]Added
mod checkpoint; mod train;and thetrain-stage1subcommand:Tests — 32/32 pass
5 new tests added (27 carried over from Phase 3):
checkpoint::test_save_creates_three_filescheckpoint::test_checkpoint_round_trip_logitsto_bits())checkpoint::test_load_config_matchesModelConfig== original nano config field-for-fieldtrain::test_step_batch_no_nantrain::test_loss_decreasesMilestone Smoke Test
CI
cargo fmt --check✅cargo clippy --all-targets -- -D warnings✅ (0 warnings)cargo test --no-fail-fast✅ 32/32 passcargo build --release✅Docs Updated
CHANGELOG.md— Phase 4 entry with arch decisions, test table, milestone resultROADMAP.md— all[ ]→[x], milestone block filled, phase map marked ✅ DONEREADME.md— status section + phase table updated, Phase 4 ✅ DoneNext: Phase 5 — Evaluation Harness
eval.rs: probe accuracy + perplexity against any checkpoint;evalCLI subcommand.