nemo_rl.data_plane.worker_mixin#
TransferQueue awareness for policy workers, isolated from the base class.
Mix into a worker class to add per-rank TQ-mediated entrypoints
(:meth:train_presharded, :meth:get_logprobs_presharded,
- meth:
get_reference_policy_logprobs_presharded, and the frozen-teacher variant) without touchingBasePolicyWorker. Subclasses that don’t need TQ keep their bare inheritance and stay zero-cost.
Subclasses must implement :meth:_get_replica_group (returns the
NCCL group of TP×CP×PP siblings within this DP rank, or None for
TP=CP=PP=1) and inherit train / get_logprobs /
get_reference_policy_logprobs from the worker base.
Module Contents#
Classes#
Adds TransferQueue per-rank fetch/write-back to a policy worker. |
Functions#
Broadcast a BatchedDataDict from |
|
Materialize a fetched TensorDict with the reader that matches the writer. |
Data#
API#
- nemo_rl.data_plane.worker_mixin.FetchPolicy#
None
- nemo_rl.data_plane.worker_mixin._broadcast_batched_data_dict(
- data: Optional[nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any]],
- *,
- is_leader: bool,
- src: int,
- group: Any,
Broadcast a BatchedDataDict from
srcto all ranks ingroup.Two-phase to avoid pickling tensor payloads on the hot path: a small descriptor (per-key dtype/shape) ships via
broadcast_object_listfirst, then each tensor’s data ships viabroadcaston its current device. The leader suppliesdata; non-leaders passNoneand get an empty BatchedDataDict filled in-place.
- nemo_rl.data_plane.worker_mixin._materialize_fetched(
- td: Any,
- *,
- local_batch: bool,
- layout: nemo_rl.data_plane.schema.Layout,
- pad_value_dict: dict[str, int | float] | None,
- pad_to_seqlen: int,
- tags: list[dict[str, Any]] | None = None,
Materialize a fetched TensorDict with the reader that matches the writer.
materializeandmaterialize_localare not interchangeable. The local adapter stores each non-tensor column as oneNonTensorDatawithbatch_size=(N,), andmaterializereads aNonTensorDataas a single row, so calling it on a local batch collapses N rows into 1 without raising. Both readers are picked from the samelocal_batchflag, so this only guards against a future edit that changes one of the two branches; it costs the TQ path oneisinstanceper column.
- class nemo_rl.data_plane.worker_mixin.TQWorkerMixin#
Adds TransferQueue per-rank fetch/write-back to a policy worker.
The driver-side
TQPolicyfans out per-rankKVBatchMeta; each worker callsself._fetch(meta, ...)to pull its slice from TQ and runs the existing per-rank method body.- _dp_client: Optional[nemo_rl.data_plane.interfaces.DataPlaneClient]#
None
- _route_fallback_counts: collections.Counter[str]#
‘Counter(…)’
- setup_data_plane( ) None#
Create this worker process’s configured data-plane client.
Called once by the driver after worker construction. Idempotent.
- mooncake_checkpoint(
- body: dict[str, Any],
Run an owner-local checkpoint command; return metadata, never payloads.
- _require_dp_client() nemo_rl.data_plane.interfaces.DataPlaneClient#
- _get_replica_group() Optional[Any]#
NCCL group of TP×CP×PP siblings within this DP rank.
Nonemeans “no siblings” (TP=CP=PP=1). Subclasses must override using their parallelism state (DTensordevice_mesh, Megatronparallel_state). ReturningNonemakes- Meth:
_fetchuse independent fetch; returning a group makes it use leader-fetch + NCCL broadcast.
- abstractmethod _routed_experts_dimensions() tuple[int, int]#
Return model-owned
(num_moe_layers, top_k)route dimensions.
- _pad_value_dict() dict[str, Any]#
Per-field pad value used by :func:
materializeto detile the jagged wire format.Token-id fields use the tokenizer pad id.
- _forward_pad_seqlen(meta: nemo_rl.data_plane.KVBatchMeta) int#
Cross-DP forward pad target, minted by :meth:
TQPolicy._stamp_pad_seqlen.
- get_data_plane_snapshot() dict[str, Any] | None#
This rank’s data-plane counters, for cluster-wide aggregation.
Returns
Nonewhen observability is off or no client exists, so the driver can filter rather than special-case. The payload is counters only (about 1 kB), not tensors.Closes this rank’s step window (
step_wall_ms,step_max_ms) as it reads, since the driver calls this once per step. Neither a sum the cluster reduces with a max nor a max itself can be differenced out of a cumulative counter, so without the reset the cluster’s per-step figures would latch at the worst call ever seen.
- _fetch(
- meta: nemo_rl.data_plane.KVBatchMeta,
- *,
- layout: nemo_rl.data_plane.schema.Layout = 'padded',
- fetch_policy: nemo_rl.data_plane.worker_mixin.FetchPolicy = 'auto',
- preprocess: Optional[Any] = None,
- dp_aligned_seq_len: bool = True,
Fetch this rank’s slice from TQ and return a BatchedDataDict.
- Parameters:
meta –
Per-rank
KVBatchMetafrom :func:shard_meta_for_dp. Forward-pass pad target is read frommeta.extra_info[GLOBAL_FORWARD_PAD_SEQLEN]minted by- meth:
TQPolicy._stamp_pad_seqlen.
layout – Materialization layout (
"padded"or"jagged").fetch_policy –
"auto"uses leader-fetch + NCCL broadcast when- meth:
_get_replica_groupreturns a group, else independent fetch (cheapest for TP=CP=PP=1)."independent"forces every sibling to fetch."leader_broadcast"forces the broadcast path and asserts a replica group exists.
preprocess – Optional
(worker, td) -> tdapplied between materialize and return.dp_aligned_seq_len – When True (default), right-pad the seq dim for the forward pass. Disabled in tests that want to observe per-rank local-pad behavior.
- Returns:
BatchedDataDictof this rank’s slice.
- _fetch_route_fragments(
- *,
- keys: list[str],
- partition_id: str,
Fetch a unique key set in one request and preserve request identity.
- _route_fragments_by_row(
- plans: list[Any],
Use one normal-path batch read; isolate error retries per rollout.
- _maybe_assemble_routed_experts(
- meta: nemo_rl.data_plane.KVBatchMeta,
- data: nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any],
Materialize deferred routes at the policy worker consumption boundary.
- _apply_packing_prep( ) nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any]#
Re-derive
micro_batch_indices/micro_batch_lengthson the local slice.Uses
shard_by_batch_size(shards=1, ...). The legacy DP path computes those as a side effect of the DP-shard call; the TQ presharded path receives a per-rank slice without them set, so we recompute here usingself.cfg.
- _attach_or_repack_pack_metadata(
- data: nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any],
- meta: nemo_rl.data_plane.KVBatchMeta,
Trust driver-supplied packing metadata or re-derive locally.
When the driver pre-balanced packing across DP ranks it ships
micro_batch_indices/micro_batch_lengths(and optionallyelem_counts_per_gb) inmeta.extra_info. Locally re-packing produces variable bin counts across DP groups and desyncs Megatron’s per-microbatch collectives — trust the driver when it provided the metadata.
- abstractmethod _local_coords() dict[str, int]#
This worker’s (axis -> local-rank) mapping.
Subclasses MUST override: DTensor reads
device_mesh, Megatron readsparallel_state. There’s no honest default — a missing impl would silently make every rank a writeback leader and re-create the-601 ILLEGAL_CLIENTduplicate-write bug.
- _is_replica_leader() bool#
True iff this rank should perform per-DP-rank-unique side-effects.
Examples include TQ write-back. Shares the same predicate the driver uses to gate dispatch (:meth:
NamedSharding.is_axis_zero) — fed by per-worker :meth:_local_coordsinstead ofNamedSharding.get_worker_coords; same answer either way.
- _is_stage_local_writer() bool#
True iff this rank is the TP/CP-zero rank of its own pipeline stage.
Unlike :meth:
_is_replica_leaderthis does not pin the pipeline stage, so it selects one rank per stage rather than one per DP rank. Callers must therefore only write outputs that exist on a single stage; see- Meth:
_write_back_stage_local.
- _write_back_stage_local(
- meta: nemo_rl.data_plane.KVBatchMeta,
- fields: dict[str, torch.Tensor],
Write fields produced on exactly one pipeline stage.
The ordinary :meth:
_write_backwrites from the replica leader, which sits on stage 0. Outputs that only the last stage holds – notably the full-vocabulary MOPD teacher payload – would then have to be broadcast backwards just to be written, which for a per-token payload means moving gigabytes across the pipeline group for nothing. This writes from the stage that already owns the data instead.Single-writer safety comes from the caller: it must pass fields that are absent (
None) on every other stage, so exactly one stage reaches this and :meth:_is_stage_local_writerpicks one rank within it.- Parameters:
meta – Per-rank
KVBatchMetafor this slice.fields – Map of field name to tensor to write back.
- _write_back(
- meta: nemo_rl.data_plane.KVBatchMeta,
- fields: dict[str, torch.Tensor],
Leader-only
put_samples(meta.sample_ids, fields=...).Per-token fields are jagged-packed via :func:
pack_per_token_fieldso they land with the same row lengths as the initial put; without this a worker write-back (rectangular[N, S]) would mismatch the jaggedinput_idson the next read.- Parameters:
meta – Per-rank
KVBatchMetafor this slice.fields – Map of field name to tensor to write back.
- _write_back_result_field(
- meta: nemo_rl.data_plane.KVBatchMeta,
- result: Any,
- *,
- result_key: str,
- tq_field: str,
Single chokepoint for
*_preshardedwrite-backs.resultis checked via theMappingABC becauseBatchedDataDictis aUserDict(notdict).- Parameters:
meta – Per-rank
KVBatchMetafor this slice.result – Worker output containing
result_key.result_key – Key into
resultfor the tensor to write back.tq_field – Field name on the TQ side.
- train_presharded(
- meta: nemo_rl.data_plane.KVBatchMeta,
- loss_fn: Any,
- eval_mode: bool = False,
- gbs: Optional[int] = None,
- mbs: Optional[int] = None,
Per-rank training entrypoint. Fetch → packing prep → delegate.
- get_logprobs_presharded(
- meta: nemo_rl.data_plane.KVBatchMeta,
- micro_batch_size: Optional[int] = None,
Per-rank logprob entrypoint. Fetch → packing prep → run → write back.
Returns
None— the per-token tensor is committed to TQ via- Meth:
_write_back_result_fieldunderprev_logprobs; when the worker narrowstoken_mask, that column is rewritten too. Callers fetch it through :meth:TQPolicy.read_from_dataplane— skipping the Ray plasma roundtrip on the (B, S) tensor.del resultdrops the local reference before returning so the worker doesn’t carry the tensor into the next dispatch.
- get_reference_policy_logprobs_presharded(
- meta: nemo_rl.data_plane.KVBatchMeta,
- micro_batch_size: Optional[int] = None,
Per-rank reference-policy logprob entrypoint.
See :meth:
get_logprobs_preshardedfor the contract. Tensor lives in TQ underreference_policy_logprobs.
- get_teacher_logprobs_presharded(
- meta: nemo_rl.data_plane.KVBatchMeta,
- micro_batch_size: Optional[int] = None,
- opd_full_payload: Optional[str] = None,
- opd_full_payload_dtype: Optional[str] = None,
- opd_full_payload_field: Optional[str] = None,
- opd_full_teacher_index: Optional[int] = None,
- opd_full_teacher_index_field: Optional[str] = None,
Per-rank frozen-teacher logprob entrypoint for SingleController MOPD.
- Parameters:
meta – Per-rank
KVBatchMetafor this DP shard.micro_batch_size – Overrides the configured logprob batch size.
opd_full_payload – When set (
"hidden_states"or"logits"), also emit the full-vocabulary teacher payload from the same forward.opd_full_payload_dtype – Torch dtype name for that payload.
opd_full_payload_field – Data-plane column the payload is written to.
opd_full_teacher_index – This teacher group’s stable index (see
create_teacher_worker_groups), tagged onto every row this call writes so the student can select the matching LM head.opd_full_teacher_index_field – Data-plane column the index is written to;
Nonewhen the run doesn’t need per-sample teacher routing (logits payload, or opd_full off).
- Raises:
ValueError – If a payload is requested without a target column, or if a teacher-index column is requested without an index.
RuntimeError – If batching metadata was not planned driver-side.
- get_values_presharded(
- meta: nemo_rl.data_plane.KVBatchMeta,
- micro_batch_size: Optional[int] = None,
Per-rank value-forward entrypoint. Fetch → packing prep → run → write back.
Same contract as get_logprobs_presharded, and only the value workers mix it in: only the PPO critic implements get_values.
- begin_train_step_presharded(
- loss_fn: Any,
- gbs: Optional[int] = None,
- mbs: Optional[int] = None,
Open a logical train step. No fetch — pure lifecycle.
The backend stores
loss_fn/gbs/mbs, clears gradients, and initialises accumulators forlocal_valid_seqs/local_valid_toksand any per-microbatch metrics. Only one step can be open at a time — the backend raises on a secondbegin— so no step identifier is needed. Optimizer state is untouched here.
- train_microbatch_presharded(
- meta: nemo_rl.data_plane.KVBatchMeta,
Per-rank microbatch entrypoint. Fetch → packing prep → forward+backward.
Gradients accumulate into
.gradacross calls; nooptimizer.stephere. Returns nothing — per-microbatch metrics accumulate in the backend’s open-step state and surface once viafinish_train_step_presharded.
- finish_train_step_presharded() dict[str, Any]#
Close a logical train step. No fetch — pure lifecycle.
Backend all-reduces accumulated
local_valid_seqs/toks, rescales gradients to the final global normalization, runs grad clip, steps the optimizer + scheduler, then zeros gradients. Returns the aggregated step result (loss,grad_norm,all_mb_metrics, …).Tags the result with
is_replica_leaderso the driver-side aggregator can dedupe TP/CP/non-last-PP-stage twins that hold identical copies of this DP shard’s metrics. Without it the driver’srun_all_workers_single_datareturns one dict per GPU and the metric list ends up TP×CP×PP times too long, which inflates every per-token aggregate (gen_kl_error, probs_ratio, etc.) by that same factor.Also pops this step’s deferred-route fallback counts (only the replica leader ever populates them, same as the metrics above) so they report a per-step rate instead of accumulating silently for the worker’s lifetime.
- abort_train_step_presharded() None#
Discard partial train-step state without stepping the optimizer.
Used when SC decides the logical batch will not complete (e.g. weight-sync triggered mid-step). Backend drops accumulators and zeros gradients.