Skip to content

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
mainfrom
test_988578527
Open

copybara-service[bot] wants to merge 1 commit into
mainfrom
test_988578527

Conversation

@copybara-service

Copy link
Copy Markdown
Contributor

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 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.

…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

This branch has not been deployed

No deployments
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