Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion CONTRIBUTING.md
Original file line number Diff line number Diff line change
Expand Up @@ -361,7 +361,7 @@ If your feature needs a new training configuration (e.g., a new parallelism comb

### New training hook

Hooks extend the training loop without modifying `scripts/train.py`:
Hooks extend the training loop without modifying it:

```python
from kempnerforge.training.hooks import TrainingHook, StepContext
Expand Down
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,7 @@ flowchart TB
CLI[CLI overrides] --> JC

%% Layer 2: training loop
JC --> TL[scripts/train.py<br/>training loop]
JC --> TL[training/loop.py<br/>training loop]

%% Layer 3: subsystems (each a compact vertical stack)
TL --> Model
Expand Down
4 changes: 2 additions & 2 deletions configs/train/vlm_7b_freeze_schedule.toml
Original file line number Diff line number Diff line change
Expand Up @@ -6,13 +6,13 @@
# max_seq_len allocation (CA): max_text_len only. The residual stream is
# text-only; image features flow as K/V into separate CrossAttentionBlocks.
#
# Exercises the FreezeStage hook in scripts/train.py:
# Exercises the FreezeStage hook in training/loop.py:
# - step 0..9: vision encoder frozen, everything else trainable.
# - step 10: freeze adapter (typical "pretrain CA blocks first" recipe).
# - step 20: unfreeze adapter (typical "now align embeddings" recipe).
#
# checkpoint.interval matches one transition step (10) to exercise
# the async-save fence: the FreezeStage hook in scripts/train.py
# the async-save fence: the FreezeStage hook in training/loop.py
# calls flush_pending_save() before applying the transition, so the
# in-flight save's metadata.json lands with the pre-transition spec
# (adapter trainable) and only the next save records the post-
Expand Down
2 changes: 1 addition & 1 deletion docs/checkpointing/auto-resume.md
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,7 @@ training starts from step 0.

## Where it's called

`scripts/train.py` calls this **once**, right after creating the
`training/entry.py::restore_checkpoint` calls this **once**, right after creating the
`CheckpointManager`:

```python
Expand Down
4 changes: 2 additions & 2 deletions docs/checkpointing/dcp-model.md
Original file line number Diff line number Diff line change
Expand Up @@ -97,7 +97,7 @@ Each stage also gets a process group scoped to that stage's DP + TP
ranks:

```python
# scripts/train.py
# training/entry.py::build_checkpoint_manager
non_pp_dims = [d for d in device_mesh.mesh_dim_names if d != "pp"]
if len(non_pp_dims) == 1:
ckpt_pg = device_mesh[non_pp_dims[0]].get_group()
Expand Down Expand Up @@ -135,7 +135,7 @@ directory.
The returned future is an `AsyncCheckpointerFuture`; `wait()` calls
`.result()` which blocks until the background writer has flushed to
disk. Training calls `ckpt_mgr.wait()` once before shutdown
(`scripts/train.py` line ~788) to flush any pending save.
(`training/loop.py::run_training_loop`) to flush any pending save.

## Process groups

Expand Down
39 changes: 19 additions & 20 deletions docs/checkpointing/train-state.md
Original file line number Diff line number Diff line change
Expand Up @@ -83,45 +83,44 @@ See [Training § Schedulers](../training/schedulers.md).

## Dataloader state

The infrastructure supports dataloader state (via
`dataloader.state_dict()` / `load_state_dict()` on KempnerForge's
`StatefulDataLoader`), but the shipped training loop in
`scripts/train.py` **does not currently pass the dataloader** to
`ckpt_mgr.save()`:
`StatefulDataLoader` exposes `state_dict()` / `load_state_dict()`, and the
training loop passes `dataloader=` at every save site, so `build_train_state`
records the loader position under the `dataloader` key:

```python
# training/loop.py::run_training_loop
ckpt_mgr.save(
step=step,
tokens_seen=tokens_seen,
scheduler=scheduler,
extra=ckpt_extra, # note: no dataloader=...
dataloader=data.dataloader,
extra=ckpt_extra,
)
```

On resume, the dataloader restarts at the beginning of its epoch.
Combined with sampler resume logic (see
[Data § Stateful dataloader](../data/index.md)) and
deterministic RNG, this is usually fine for most pretraining
workflows — you may replay a few batches on resume but loss is not
affected long-term.

If exact-batch-level reproducibility matters, pass `dataloader=` into
`ckpt_mgr.save()` yourself; `build_train_state` will pick it up
automatically.
A checkpoint taken mid-epoch carries e.g.
`{'epoch': 7, 'batches_yielded': 14, 'sampler': {...}}`. On resume
`StatefulDataLoader.__iter__` re-applies the skip, so training continues at the
same sample boundary rather than restarting the epoch. Plain
`torch.utils.data.DataLoader` (the streaming and eval paths) has no
`state_dict`, so it is skipped and those loaders do restart.

## The `extra` dict

Anything that isn't step / tokens / RNG / scheduler / dataloader
can be threaded through `extra`:

```python
# scripts/train.py — around the checkpoint call
ckpt_extra = {"phase_idx": current_phase_idx} if active_phases else {}
# training/loop.py::checkpoint_extra
extra = {"phase_idx": phases.next_idx} if phases.phases else {}
if config.metrics.wandb_run_id:
ckpt_extra["wandb_run_id"] = config.metrics.wandb_run_id
extra["wandb_run_id"] = config.metrics.wandb_run_id
return extra

# training/loop.py::run_training_loop
ckpt_extra = checkpoint_extra(config, step, phases)
ckpt_mgr.save(step=step, tokens_seen=tokens_seen, scheduler=scheduler,
extra=ckpt_extra)
dataloader=data.dataloader, extra=ckpt_extra)
```

On load, `restore_train_state` strips out the "standard" keys
Expand Down
2 changes: 1 addition & 1 deletion docs/configuration/validation-rules.md
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ Validation runs in two passes:
Checks fields that can be validated in isolation.
2. **`JobConfig.validate(world_size)`** — called explicitly by the
launchers
([`scripts/train.py`](https://github.com/KempnerInstitute/KempnerForge/blob/main/scripts/train.py),
([`training/runtime.py`](https://github.com/KempnerInstitute/KempnerForge/blob/main/kempnerforge/training/runtime.py),
[`scripts/eval.py`](https://github.com/KempnerInstitute/KempnerForge/blob/main/scripts/eval.py))
once `world_size` is known. Checks rules that cross section
boundaries.
Expand Down
2 changes: 1 addition & 1 deletion docs/data/huggingface.md
Original file line number Diff line number Diff line change
Expand Up @@ -61,7 +61,7 @@ tokenizer_path = "meta-llama/Llama-3-8B"
Stream mode is an `IterableDataset`:

- No `__len__`, no sampler (the dataset is its own iterator).
- `scripts/train.py` wraps it in a plain `torch.utils.data.DataLoader`,
- `build_data_pipeline` wraps it in a plain `torch.utils.data.DataLoader`,
not `StatefulDataLoader`.
- Distributed sharding happens inside the iterator:
`if doc_idx % world_size != rank: continue`. Each rank takes every
Expand Down
2 changes: 1 addition & 1 deletion docs/data/index.md
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ partitions work across data-parallel ranks, the stateful dataloader
that makes checkpoint resume possible, and the mixing / annealing
hooks used for curriculum training.

Entry point: `scripts/train.py` picks one of four paths based on
Entry point: `training/data_pipeline.py::build_data_pipeline` picks one of four paths based on
config — pre-tokenized mmap, HuggingFace eager, HuggingFace streaming,
or multi-source mixture.

Expand Down
4 changes: 2 additions & 2 deletions docs/data/memory-mapped.md
Original file line number Diff line number Diff line change
Expand Up @@ -52,7 +52,7 @@ At training time, `__getitem__(idx)`:
2. Slices `seq_len` tokens starting at `local_idx * seq_len`.
3. Returns `{"input_ids": tokens[:-1], "labels": tokens[1:]}`.

Note: `scripts/train.py` passes `seq_len + 1` to the constructor so
Note: the data pipeline passes `seq_len + 1` to the constructor so
the sliced window contains one extra token, and the `[:-1]` / `[1:]`
split produces inputs and next-token labels of length `seq_len` each.

Expand All @@ -69,7 +69,7 @@ pack_sequences = false # see below
tokenizer_path = "" # only required when pack_sequences=true
```

`scripts/train.py` picks this path when `data.dataset_path` is set
`build_data_pipeline` picks this path when `data.dataset_path` is set
and `data.datasets` is empty.

## Sequence packing
Expand Down
6 changes: 3 additions & 3 deletions docs/data/mixing-and-annealing.md
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,7 @@ and `data.hf_dataset_name` — the single-source paths are ignored.

## What happens at construction

In `scripts/train.py`:
In `training/data_pipeline.py`:

1. Each `DatasetSource` builds a sub-dataset — `MemoryMappedDataset`
or `HuggingFaceDataset` — with the global `tokenizer_path` and
Expand Down Expand Up @@ -147,7 +147,7 @@ single-phase list at startup. Using both `data.phases` and

## How phases execute

Inside `scripts/train.py`:
Inside `training/loop.py`:

```python
# --- Phase activation in the training loop ---
Expand Down Expand Up @@ -175,7 +175,7 @@ Two effects:

### Resume into the right phase

On auto-resume, `scripts/train.py` replays the phase activations
On auto-resume, `build_phase_state` replays the phase activations
against the current `step` so the sampler and LR scale are correct
before the training loop starts:

Expand Down
2 changes: 1 addition & 1 deletion docs/data/sampler.md
Original file line number Diff line number Diff line change
Expand Up @@ -146,7 +146,7 @@ size, not transplanted from the save. `seed` is already in
Evaluation uses `DistributedSampler(eval_dataset, shuffle=False)`:

```python
# scripts/train.py
# training/data_pipeline.py::build_eval_dataloader
eval_sampler = DistributedSampler(
eval_dataset, num_replicas=dp_size, rank=dp_rank,
shuffle=False, seed=tc.seed,
Expand Down
24 changes: 10 additions & 14 deletions docs/data/stateful-dataloader.md
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ wraps `torch.utils.data.DataLoader` with three additions:
## Construction

```python
# scripts/train.py
# training/data_pipeline.py
dataloader = StatefulDataLoader(
dataset,
batch_size=tc.batch_size,
Expand Down Expand Up @@ -90,30 +90,26 @@ Effect: on the next `__iter__()`, the sampler yields its index list
minus the first `batches_yielded * batch_size` elements. Training
picks up at the exact same sample boundary.

## The wiring gap
## Save and restore

The infrastructure works, but `scripts/train.py` **does not currently
pass the dataloader to `ckpt_mgr.save()`**:
The training loop passes `dataloader=` at every save site, so the loader
position travels with the checkpoint:

```python
# scripts/train.py
# training/loop.py::run_training_loop
ckpt_mgr.save(
step=step,
tokens_seen=tokens_seen,
scheduler=scheduler,
dataloader=data.dataloader,
extra=ckpt_extra,
# no dataloader=...
)
```

Consequence: on resume, the dataloader restarts from batch 0 of the
current epoch. With deterministic seeding and a shuffled sampler, this
means a few batches get replayed but the loss trajectory is otherwise
indistinguishable from an uninterrupted run.

For exact-batch-level reproducibility, pass `dataloader=dataloader`
into `ckpt_mgr.save` yourself — `build_train_state` picks it up
automatically. See
An emergency checkpoint taken mid-epoch records e.g.
`{'epoch': 7, 'batches_yielded': 14, 'sampler': {...}}`, and `__iter__`
re-applies the skip on resume, so training picks up at the same sample
boundary rather than at batch 0. See
[Checkpointing § Train state](../checkpointing/train-state.md#dataloader-state).

## Worker RNG
Expand Down
2 changes: 1 addition & 1 deletion docs/distributed/device-mesh.md
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ shape and dimension order of that mesh are load-bearing.
## Dimension order

```python
# scripts/train.py
# training/runtime.py::setup_distributed
device_mesh = init_distributed(config.distributed, seed=config.train.seed)
```

Expand Down
2 changes: 1 addition & 1 deletion docs/distributed/fsdp2.md
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,7 @@ apply_fsdp2(model, device_mesh, reshard_after_forward=True)
| `False` | All-gather once, keep gathered across the step | Pipeline parallel: `1F1B` schedule sends many microbatches through each stage; resharding between them would be wasted work |
| `int` | Rate-limit concurrency: keep at most `N` blocks gathered at once | Rarely used; middle-ground for memory tuning |

[`scripts/train.py`](https://github.com/KempnerInstitute/KempnerForge/blob/main/scripts/train.py)
[`training/entry.py`](https://github.com/KempnerInstitute/KempnerForge/blob/main/kempnerforge/training/entry.py)
calls `apply_fsdp2(model, device_mesh, mp_policy=mp_policy)` on all
paths without overriding `reshard_after_forward`, so both dense and
PP runs default to `True`. The docstring on `pipeline_parallel.py`
Expand Down
6 changes: 3 additions & 3 deletions docs/distributed/pipeline-parallelism.md
Original file line number Diff line number Diff line change
Expand Up @@ -52,7 +52,7 @@ and resume with `pp=2`.
## Application order with PP

The non-PP path goes through `build_parallel_model`. PP has its own
path in `scripts/train.py`:
path in `training/entry.py::build_model`:

```python
# PP + TP
Expand Down Expand Up @@ -123,7 +123,7 @@ phase (only warmup and drain) and the bubble overhead is maximal;
## Training step under PP

The PP branch in
[`scripts/train.py`](https://github.com/KempnerInstitute/KempnerForge/blob/main/scripts/train.py)
[`training/loop.py`](https://github.com/KempnerInstitute/KempnerForge/blob/main/kempnerforge/training/loop.py)
looks very different from the non-PP step:

```python
Expand Down Expand Up @@ -160,7 +160,7 @@ That's not ideal: 1F1B sends many microbatches through each stage, and
resharding between them triggers a fresh all-gather every
microbatch. The docstring on `pipeline_parallel.py` recommends
`reshard_after_forward=False` for PP to amortize the all-gather over
`n_microbatches`, but the current `scripts/train.py` doesn't thread
`n_microbatches`, but the current `build_model` doesn't thread
the flag through. If you're running PP at scale and can't fit the
extra all-gathers, pass it manually. See [FSDP2](fsdp2.md) — the
**`reshard_after_forward`** section.
Expand Down
2 changes: 1 addition & 1 deletion docs/how-to/build-a-model.md
Original file line number Diff line number Diff line change
Expand Up @@ -225,7 +225,7 @@ build_parallel_model(model_config, device, mesh, ...)

Every box that isn't "hardcoded class" goes through a registry lookup.
Pipeline parallelism (not shown) is applied outside this flow — it
splits the model into stages in `scripts/train.py` *before*
splits the model into stages in `training/entry.py::build_model` *before*
`build_parallel_model` runs on each stage.

## See also
Expand Down
8 changes: 4 additions & 4 deletions docs/how-to/data-mixing-annealing.md
Original file line number Diff line number Diff line change
Expand Up @@ -56,7 +56,7 @@ Each `[[data.datasets]]` block is a
At least one of `path` / `hf_name` must be set per source, enforced by
`DataConfig.__post_init__`.

When any `[[data.datasets]]` is present, `scripts/train.py` builds a
When any `[[data.datasets]]` is present, `build_data_pipeline` builds a
[`MixtureDataset`](https://github.com/KempnerInstitute/KempnerForge/blob/main/kempnerforge/data/dataset.py)
over the sub-datasets and drives it with a
[`MixtureSampler`](https://github.com/KempnerInstitute/KempnerForge/blob/main/kempnerforge/data/sampler.py).
Expand All @@ -77,7 +77,7 @@ per-dataset metrics.
### Per-dataset metrics

When the mixture is active and the metrics interval fires,
`scripts/train.py` emits two series per dataset name:
The training loop emits two series per dataset name:

- `loss/{name}` — mean loss of samples from that dataset in the
accumulation window.
Expand Down Expand Up @@ -138,7 +138,7 @@ Constraints (validated in `DataConfig.__post_init__`):

On the first training step where `step >= phase.start_step`, the loop
in
[`scripts/train.py`](https://github.com/KempnerInstitute/KempnerForge/blob/main/scripts/train.py)
[`training/data_pipeline.py`](https://github.com/KempnerInstitute/KempnerForge/blob/main/kempnerforge/training/data_pipeline.py)
does two things:

1. Builds a new `weights` list by overriding the original declared
Expand Down Expand Up @@ -182,7 +182,7 @@ anneal_start_step = 40_000
anneal_weights = { web = 0.1, code = 0.9 }
```

`scripts/train.py` converts this into a one-element `TrainingPhase`
`build_phase_state` converts this into a one-element `TrainingPhase`
list internally:

```python
Expand Down
2 changes: 1 addition & 1 deletion docs/how-to/end-to-end-training-run.md
Original file line number Diff line number Diff line change
Expand Up @@ -208,4 +208,4 @@ Extensions from here:
- [Generation](../training/generation.md) — `generate()` internals
and KV-cache API.
- [Training loop](../training/training-loop.md) — what
`scripts/train.py` actually does at each step.
the training loop actually does at each step.
2 changes: 1 addition & 1 deletion docs/how-to/fp8-training.md
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@ One field in `[train]`:
mixed_precision = "fp8"
```

That's the whole surface. `scripts/train.py` picks it up via
That's the whole surface. `build_model` picks it up via
[`TrainConfig.is_fp8`](https://github.com/KempnerInstitute/KempnerForge/blob/main/kempnerforge/config/training.py),
`build_parallel_model` calls `apply_float8(model)` after TP/EP and
before AC/FSDP2, and FSDP2 picks bf16 master weights (FP8 is a
Expand Down
2 changes: 1 addition & 1 deletion docs/how-to/prepare-tokenized-data.md
Original file line number Diff line number Diff line change
Expand Up @@ -197,7 +197,7 @@ different mechanisms:
picks up at sample `3500 × batch_size × world_size` without replay.
- **HF streaming (Path B)**: wrapped in a plain `DataLoader`;
`StreamingHuggingFaceDataset` tracks its own position and exposes
`load_state_dict` / `state_dict`, which `scripts/train.py` drives
`load_state_dict` / `state_dict`, which the training loop drives
directly. Same guarantee — no replay, no skip — via a different
code path.

Expand Down
2 changes: 1 addition & 1 deletion docs/how-to/scaling-guide.md
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ exact sequence:
TP → EP → [FP8] → AC → FSDP (dp_shard)
```

Pipeline parallelism is applied separately in `scripts/train.py`
Pipeline parallelism is applied separately in `training/entry.py::build_model`
(split into stages *before* `build_parallel_model` runs on each
stage), and DP replicate / CP are mesh dimensions without a dedicated
apply step — they fall out of how the `DeviceMesh` is constructed.
Expand Down
2 changes: 1 addition & 1 deletion docs/metrics-and-profiling/index.md
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ loss smoothing, GPU memory, kernel traces. Two modules:

| Component | Type | Enabled by |
|-----------|------|-----------|
| `MetricsTracker` | always-on per-step metrics | created unconditionally in `scripts/train.py` |
| `MetricsTracker` | always-on per-step metrics | created unconditionally in `training/entry.py` |
| `WandBBackend` | cloud logging | `metrics.enable_wandb = true` |
| `TensorBoardBackend` | local event files | `metrics.enable_tensorboard = true` |
| `MLflowBackend` | Databricks / MLflow server | `metrics.enable_mlflow = true` |
Expand Down
Loading
Loading