nemo_rl.utils.train_data_dump#
Stream untruncated training tensors without retaining a second step batch.
Module Contents#
Classes#
Write chunks to a partial file, publishing only a completed optimizer step. |
API#
- class nemo_rl.utils.train_data_dump.TrainDataDump(log_dir: str, *, shard_id: str = '0')#
Write chunks to a partial file, publishing only a completed optimizer step.
Sequence columns are trimmed to input_lengths (padding only). Scalar columns are written per row as given, so variable-length values such as prompt_ids must be passed as jagged tensors rather than padded. Masked rows are retained. Values use the legacy train_data JSONL singleton-batch shape. A failed step leaves part files, never a completed-looking JSONL file.
The advantage stage that produces these rows runs either in the controller or across a pool of Ray actors, so a step’s rows are written by one writer per shard and merged on publish. Each writer owns
train_data_step<N>.jsonl.part-<shard_id>and never touches another’s, which is what makes concurrent shards safe;finish_stepconcatenates them in shard order and assignsidxacross the merged file. Writers therefore omitidx– only the merge sees every row, so only the merge can number them.Sharded writes assume the pool shares a filesystem with the controller, which is already true of the log dir the trainer checkpoints into.
Initialization
- _part_path(step: int, shard_id: str) pathlib.Path#
- add_chunk(
- *,
- step: int,
- sample_ids: list[str],
- tags: list[dict[str, Any]] | None,
- input_lengths: torch.Tensor,
- sequences: dict[str, torch.Tensor],
- scalars: dict[str, torch.Tensor],
- finish_step(step: int, expected_rows: int | None = None) None#
Merge every shard’s part file and publish the step atomically.
Called on the controller, which may not itself have written anything: with a pool the rows come from the actors, so the parts on disk are the only record of what the step produced.
expected_rowsis what the shards reported writing. A pool that does not share this filesystem with the controller still runs and still reports, and its parts are simply not here – which would publish a short dump that looks complete. Checking the count is what turns that into a failure.