Train loop and input pipeline hooks for training a data-parallel replica on a submesh. - #5506
Open
copybara-service[bot] wants to merge 1 commit into
Open
copybara-service[bot] wants to merge 1 commit into
copybara-service[bot] wants to merge 1 commit into
Conversation
…ica on a submesh.
Threaded DiLoCo runs several independent learners, one per slice, in one process.
Each learner needs to build its train state on its own submesh and read a disjoint
share of the data. With the defaults, every change here is a no-op.
- `train_utils.setup_train_loop(..., mesh=None)`: optionally train on a prebuilt
mesh instead of the one derived from the config (also used as the fallback mesh
when restoring LoRA weights).
- `input_pipeline_utils.get_dataloading_shard(config, process_indices,
shard_offset=0, shards_per_host=1)`: one place that computes
`(shard_index, shard_count)` for this process. Grain (train, the
`expansion_factor_real_data` restore list, eval), HF (train/eval) and TFDS
(train/eval) now call it instead of computing the index inline. It further
splits each host's shard among `num_data_replicas_per_process` replicas.
Contract: a replica's `global_batch_size_to_load` is its own batch, while the
shard set covers all replicas. Everything sized for the shard set uses
`get_all_replicas_batch_size` (batch times replicas):
- the grain ElasticIterator's global batch, so each replica gets its own batch;
- the mmap_npy train and eval sample budgets, so no replica runs out early;
- mmap_npy blends: a shuffle window of one all-replica step runs before the
host/replica stride, so every replica reads a mix of all components and the
union of the replicas' batches is the single-replica step batch.
- New `DatasetGeneral` fields `num_data_replicas_per_process` (default 1) and
`data_replica_index` (default 0), validated so the index is below the count.
Replicas above 1 are rejected with `colocated_python_data_input`, `c4_mlperf`
and `olmo_grain`. Configs without these fields (duck-typed test configs) read
the defaults.
- `HyperParameters.replace(**updates)`: returns a copy with fields overridden
verbatim (pydantic `model_copy`), for deriving per-learner configs. Unknown keys
raise and the data-replica fields are re-validated; other validation and
derived fields (including output paths derived from `run_name`) are not
recomputed, and the docstring lists them.
Testing:
- `tests/unit/input_pipeline_utils_test.py`: `GetDataloadingShardTest` (one
replica matches the old host position, replicas split a host shard, several
shards per host, configs without the replica fields, index range checks).
- `tests/unit/pyconfig_test.py`: `replace` returns an updated copy, rejects
unknown fields and re-checks the replica fields; field boundaries and
incompatible pipelines are rejected.
- `tests/unit/grain_data_processing_test.py`: two replicas through a real grain
ElasticIterator get disjoint batches of the per-replica size whose union per
step is the single-replica batch.
- `tests/unit/mmap_data_processing_test.py`: every replica gets data for every
step through a real 2-component blend (train), the eval budget covers all
replicas, and each replica's batch is a mix whose per-step union equals the
single-replica batch. The existing `GrainMmapNpyEvalConfigTest` passes.
PiperOrigin-RevId: 988578527
copybara-service
Bot
requested review from
A9isha,
NuojCheng,
RissyRan,
SurbhiJainUSC,
abhinavclemson,
aireenmei,
bvandermoon,
darisoy,
dipannita08,
gagika,
gobbleturk,
hengtaoguo,
huytransformer,
igorts-git,
jiangjy1982,
khatwanimohit,
richjames0,
shralex,
shuningjin,
vipannalla and
xibinliu
as code owners
October 2, 2026 00:25
Codecov Report❌ Patch coverage is 📢 Thoughts on this report? Let us know! |
This branch has not been deployed
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.
Train loop and input pipeline hooks for training a data-parallel replica on a submesh.
Threaded DiLoCo runs several independent learners, one per slice, in one process.
Each learner needs to build its train state on its own submesh and read a disjoint
share of the data. With the defaults, every change here is a no-op.
train_utils.setup_train_loop(..., mesh=None): optionally train on a prebuiltmesh instead of the one derived from the config (also used as the fallback mesh
when restoring LoRA weights).
input_pipeline_utils.get_dataloading_shard(config, process_indices, shard_offset=0, shards_per_host=1): one place that computes(shard_index, shard_count)for this process. Grain (train, theexpansion_factor_real_datarestore list, eval), HF (train/eval) and TFDS(train/eval) now call it instead of computing the index inline. It further
splits each host's shard among
num_data_replicas_per_processreplicas.Contract: a replica's
global_batch_size_to_loadis its own batch, while theshard set covers all replicas. Everything sized for the shard set uses
get_all_replicas_batch_size(batch times replicas):host/replica stride, so every replica reads a mix of all components and the
union of the replicas' batches is the single-replica step batch.
DatasetGeneralfieldsnum_data_replicas_per_process(default 1) anddata_replica_index(default 0), validated so the index is below the count.Replicas above 1 are rejected with
colocated_python_data_input,c4_mlperfand
olmo_grain. Configs without these fields (duck-typed test configs) readthe defaults.
HyperParameters.replace(**updates): returns a copy with fields overriddenverbatim (pydantic
model_copy), for deriving per-learner configs. Unknown keysraise and the data-replica fields are re-validated; other validation and
derived fields (including output paths derived from
run_name) are notrecomputed, and the docstring lists them.
Testing:
tests/unit/input_pipeline_utils_test.py:GetDataloadingShardTest(onereplica matches the old host position, replicas split a host shard, several
shards per host, configs without the replica fields, index range checks).
tests/unit/pyconfig_test.py:replacereturns an updated copy, rejectsunknown fields and re-checks the replica fields; field boundaries and
incompatible pipelines are rejected.
tests/unit/grain_data_processing_test.py: two replicas through a real grainElasticIterator get disjoint batches of the per-replica size whose union per
step is the single-replica batch.
tests/unit/mmap_data_processing_test.py: every replica gets data for everystep through a real 2-component blend (train), the eval budget covers all
replicas, and each replica's batch is a mix whose per-step union equals the
single-replica batch. The existing
GrainMmapNpyEvalConfigTestpasses.