Skip to content

feat: Phase 4 — Stage 1 training loop, real checkpoint save/load - #4

Merged
aarambh-darshan merged 1 commit into
mainfrom
feat/phase-4-training-checkpoint
Jul 30, 2026
Merged

aarambh-darshan merged 1 commit into
mainfrom
feat/phase-4-training-checkpoint

Conversation

@aarambh-darshan

Copy link
Copy Markdown
Member

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 the train-stage1 CLI subcommand.


What Changed

src/checkpoint.rs [NEW]

Single-responsibility checkpoint module:

Function What it does
save_checkpoint(vm, cfg, tokenizer, dir) VarMap::save() → model.safetensors; cfg.save_json() → config.json; tokenizer.save() → tokenizer.json
load_checkpoint(dir, device) Returns (ModelConfig, BpeTokenizer, VarBuilder<'static>) via from_mmaped_safetensors — correct candle pattern for disk-load without a VarMap

Key design decision: load_checkpoint returns VarBuilder<'static> (not a reconstructed PocModel) 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:

pub struct Trainer {
    pub model: PocModel,
    pub vm: VarMap,          // must outlive model + optimizer
    optimizer: AdamW,
    pub cfg: TrainConfig,
    pub model_cfg: ModelConfig,
    pub step: usize,
}
Item Details
Loss candle_nn::loss::cross_entropy on flattened [B*T, V] logits
Objective Causal LM: input = ids[0..T-1], target = ids[1..T]
LR schedule Linear warm-up from 0 → lr over warmup_steps, then constant
Optimizer step optimizer.backward_step(&loss) — single call for backward + weight update
Batch sampling Seeded StdRng (seed=42), sample-with-replacement from encoded corpus

src/main.rs [MODIFIED]

Added mod checkpoint; mod train; and the train-stage1 subcommand:

continual-learning-poc train-stage1 [OPTIONS]

  --config <PRESET>         nano | small  [default: nano]
  --steps <N>               Override max_steps  [default: 2000]
  --data-dir <DIR>          Source of stage1_names_facts.jsonl  [default: data]
  --checkpoint-dir <DIR>    Destination  [default: checkpoints/stage1]

Tests — 32/32 pass

5 new tests added (27 carried over from Phase 3):

Test Assertion
checkpoint::test_save_creates_three_files Dir contains exactly model.safetensors + config.json + tokenizer.json
checkpoint::test_checkpoint_round_trip_logits Logits before save == logits after load (bit-identical f32, checked with to_bits())
checkpoint::test_load_config_matches Loaded ModelConfig == original nano config field-for-field
train::test_step_batch_no_nan Single gradient step returns finite, non-negative loss
train::test_loss_decreases Smoothed loss steps 30–60 < steps 0–30 on a 10-sentence corpus

Milestone Smoke Test

$ cargo run --release -- generate-data
train corpus: 640 examples  →  data/stage1_names_facts.jsonl

$ cargo run --release -- train-stage1 --config nano --steps 200
step     1 / 200  loss=41.7793
step   200 / 200  loss=10.4383

Loss summary:
  first step : 41.7793
  last step  : 10.4383
  smooth-10% : 13.8132

       model.safetensors   5259280 bytes  (5.1 MB)
             config.json       160 bytes
          tokenizer.json     15026 bytes

CI

  • cargo fmt --check ✅
  • cargo clippy --all-targets -- -D warnings ✅ (0 warnings)
  • cargo test --no-fail-fast ✅ 32/32 pass
  • cargo build --release ✅

Docs Updated

  • CHANGELOG.md — Phase 4 entry with arch decisions, test table, milestone result
  • ROADMAP.md — all [ ] → [x], milestone block filled, phase map marked ✅ DONE
  • README.md — status section + phase table updated, Phase 4 ✅ Done

Next: Phase 5 — Evaluation Harness

eval.rs: probe accuracy + perplexity against any checkpoint; eval CLI subcommand.

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
@aarambh-darshan
aarambh-darshan merged commit 02e6ee3 into main Jul 30, 2026
3 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant