07-03 Batching and Train Step
Why this matters
Batching makes optimization efficient and keeps tensor contracts consistent across updates.
Intuition first (no jargon)
Sample many windows, compute one loss, and apply one deterministic parameter update step.
Paper grounding
- The paper trains with Adam using
beta1 = 0.9,beta2 = 0.98, andepsilon = 1e-9. - It uses a warmup-based learning-rate schedule before inverse-square-root decay.
Worked example
- Inputs:
idslength20,batchSize = 2,blockSize = 4 - Shapes:
xbatch (inputs) is[2][4]ybatch (targets) is[2][4]- logits from MiniGPT are
[2][4][V] - loss is scalar
[]
- One computed step:
- sample start indices, build
x, and shift by one token fory.
- sample start indices, build
Code walkthrough
jsexport function getBatch(ids, batchSize, blockSize, rng) {} export function trainStep(batch, params, cfg) {}
getBatchsamples deterministic windows when an RNG is provided.trainStepruns forward pass, computes loss, and applies one parameter update.
Your task
Implement batch sampling and one optimization step.
- Sample random start positions for contexts.
- Build input and target tensors.
- Compute loss and apply parameter update.
Common mistakes
- Forgot one-token shift between
xandy - Direct
Math.randomusage in batching - Non-deterministic RNG path in tests
Precision note
This lesson uses simple SGD-style updates to keep mechanics visible. Paper training uses Adam with warmup scheduling.
Hints
- Keep RNG injectable for deterministic tests.
- Reuse
makeExampleslogic conceptually. - Log
step,loss, and learning rate.
Check your thinking
- Why random batches help generalization?
- Why is deterministic batching still useful during tests?
- Why monitor loss each step?
Stretch (optional)
Add gradient clipping for stability.
Likely test focus
- Correct batch shapes.
- Deterministic batches with seeded RNG.
- Train step lowers loss trend on toy corpus.
What should improve
You can now run reproducible batched training updates.
Bridge to next lesson
Next: add validation and checkpointing.