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 touching BasePolicyWorker. 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#

TQWorkerMixin

Adds TransferQueue per-rank fetch/write-back to a policy worker.

Functions#

_broadcast_batched_data_dict

Broadcast a BatchedDataDict from src to all ranks in group.

_materialize_fetched

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,
) → nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any]#

Broadcast a BatchedDataDict from src to all ranks in group.

Two-phase to avoid pickling tensor payloads on the hot path: a small descriptor (per-key dtype/shape) ships via broadcast_object_list first, then each tensor’s data ships via broadcast on its current device. The leader supplies data; non-leaders pass None and 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,
) → nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any]#

Materialize a fetched TensorDict with the reader that matches the writer.

materialize and materialize_local are not interchangeable. The local adapter stores each non-tensor column as one NonTensorData with batch_size=(N,), and materialize reads a NonTensorData as a single row, so calling it on a local batch collapses N rows into 1 without raising. Both readers are picked from the same local_batch flag, so this only guards against a future edit that changes one of the two branches; it costs the TQ path one isinstance per column.

class nemo_rl.data_plane.worker_mixin.TQWorkerMixin#

Adds TransferQueue per-rank fetch/write-back to a policy worker.

The driver-side TQPolicy fans out per-rank KVBatchMeta; each worker calls self._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(
cfg: nemo_rl.data_plane.interfaces.DataPlaneRuntimeConfig,
) → 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],
) → dict[str, Any] | None#

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.

None means “no siblings” (TP=CP=PP=1). Subclasses must override using their parallelism state (DTensor device_mesh, Megatron parallel_state). Returning None makes

Meth:

_fetch use 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:materialize to 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 None when 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,
) → nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any]#

Fetch this rank’s slice from TQ and return a BatchedDataDict.

Parameters:
  • meta –

    Per-rank KVBatchMeta from :func:shard_meta_for_dp. Forward-pass pad target is read from meta.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_group returns 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) -> td applied 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:

BatchedDataDict of this rank’s slice.

_fetch_route_fragments(
*,
keys: list[str],
partition_id: str,
) → dict[str, nemo_rl.experience.route_assembly.RouteFragment]#

Fetch a unique key set in one request and preserve request identity.

_route_fragments_by_row(
plans: list[Any],
) → tuple[list[dict[str, nemo_rl.experience.route_assembly.RouteFragment]], int, float]#

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],
) → nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any]#

Materialize deferred routes at the policy worker consumption boundary.

_apply_packing_prep(
data: nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any],
) → nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any]#

Re-derive micro_batch_indices / micro_batch_lengths on 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 using self.cfg.

_attach_or_repack_pack_metadata(
data: nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any],
meta: nemo_rl.data_plane.KVBatchMeta,
) → nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any]#

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 optionally elem_counts_per_gb) in meta.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 reads parallel_state. There’s no honest default — a missing impl would silently make every rank a writeback leader and re-create the -601 ILLEGAL_CLIENT duplicate-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_coords instead of NamedSharding.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_leader this 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],
) → None#

Write fields produced on exactly one pipeline stage.

The ordinary :meth:_write_back writes 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_writer picks one rank within it.

Parameters:
  • meta – Per-rank KVBatchMeta for 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],
) → None#

Leader-only put_samples(meta.sample_ids, fields=...).

Per-token fields are jagged-packed via :func:pack_per_token_field so they land with the same row lengths as the initial put; without this a worker write-back (rectangular [N, S]) would mismatch the jagged input_ids on the next read.

Parameters:
  • meta – Per-rank KVBatchMeta for 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,
) → None#

Single chokepoint for *_presharded write-backs.

result is checked via the Mapping ABC because BatchedDataDict is a UserDict (not dict).

Parameters:
  • meta – Per-rank KVBatchMeta for this slice.

  • result – Worker output containing result_key.

  • result_key – Key into result for 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,
) → dict[str, Any]#

Per-rank training entrypoint. Fetch → packing prep → delegate.

get_logprobs_presharded(
meta: nemo_rl.data_plane.KVBatchMeta,
micro_batch_size: Optional[int] = None,
) → 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_field under prev_logprobs; when the worker narrows token_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 result drops 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,
) → None#

Per-rank reference-policy logprob entrypoint.

See :meth:get_logprobs_presharded for the contract. Tensor lives in TQ under reference_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,
) → None#

Per-rank frozen-teacher logprob entrypoint for SingleController MOPD.

Parameters:
  • meta – Per-rank KVBatchMeta for 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; None when 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,
) → 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,
) → None#

Open a logical train step. No fetch — pure lifecycle.

The backend stores loss_fn / gbs / mbs, clears gradients, and initialises accumulators for local_valid_seqs / local_valid_toks and any per-microbatch metrics. Only one step can be open at a time — the backend raises on a second begin — so no step identifier is needed. Optimizer state is untouched here.

train_microbatch_presharded(
meta: nemo_rl.data_plane.KVBatchMeta,
) → None#

Per-rank microbatch entrypoint. Fetch → packing prep → forward+backward.

Gradients accumulate into .grad across calls; no optimizer.step here. Returns nothing — per-microbatch metrics accumulate in the backend’s open-step state and surface once via finish_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_leader so 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’s run_all_workers_single_data returns 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.