nemo_rl.models.policy.tq_policy#

TQ-mediated Policy: meta-driven 1-hop counterpart to Policy.

Exposes train_from_meta / get_logprobs_from_meta / get_reference_policy_logprobs_from_meta — same return shapes as Policy.{train, get_logprobs, get_reference_policy_logprobs} but accepting a KVBatchMeta instead of a BatchedDataDict. The meta names per-sample TQ keys minted once at rollout (:class:nemo_rl.experience.sync_rollout_actor.SyncRolloutActor); each dispatch slices the key list per DP rank via

func:

nemo_rl.data_plane.preshard.shard_meta_for_dp (no re-fan-out, no key minting). Workers fetch their slice from TQ via self._fetch(meta) and write deltas back via self._write_back_result_field(...). See nemo_rl/data_plane/README.md for the full design.

Module Contents#

Classes#

TQPolicy

TQ-mediated counterpart to :class:Policy.

Functions#

Data#

API#

nemo_rl.models.policy.tq_policy._aggregate_train_results(
results: list[dict[str, Any]],
) → dict[str, Any]#
nemo_rl.models.policy.tq_policy.logger#

‘getLogger(…)’

class nemo_rl.models.policy.tq_policy.TQPolicy(
*args: Any,
dp_cfg: nemo_rl.data_plane.interfaces.DataPlaneRuntimeConfig,
checkpointing: bool = False,
tq_partition_id: str = 'train',
**kwargs: Any,
)#

Bases: nemo_rl.data_plane.driver_mixin.TQDriverMixin, nemo_rl.models.policy.lm_policy.Policy

TQ-mediated counterpart to :class:Policy.

Constructor accepts an additional dp_cfg (the master_config["data_plane"] dict). Bootstraps the controller on the driver and forwards setup_data_plane(dp_cfg) to every worker so they can attach as clients (bootstrap=False).

checkpointing is an internal bootstrap mode derived from the existing checkpoint settings and resume path, not another user-facing switch. For Mooncake it enables hard-pinned memory, disables offload, and keeps the driver out of the storage topology; workers inherit the controller’s mode.

The partition lifecycle (register_partition / clear_samples) is the trainer’s responsibility — this class assumes the partition named by tq_partition_id (default "train") is open with a schema covering DP_TRAIN_FIELDS (the bulk schema written by the rollout actor at first put + driver-/worker-written deltas).

Initialization

load_data_plane_checkpoint(
checkpoint_dir: str | pathlib.Path,
) → dict[str, Any]#

Restore TQ through the clean bootstrap client during SC setup.

shutdown() → bool#

Close the TQ client before shutting down the worker group.

prepare_step(
num_samples: int,
group_size: Optional[int] = None,
) → None#

Register the per-step TQ partition.

Sync trainers call this at the start of each step. The static partition id "train" is cleared and reused across steps. The schema is the union of all consumer fields — producers write only the subset they have, consumers fetch via select_fields.

Parameters:
  • num_samples – Expected total samples this step.

  • group_size – GRPO group size for balanced sampling; None disables grouping.

prepare_val_partition(
num_samples: int,
*,
partition_id: str = 'val',
) → None#

Register a per-batch val partition (single consumer, no GRPO grouping).

Sync val trainers call this at the start of each val batch. Distinct from :meth:prepare_step because val has its own partition id and a single consumer task.

discard_samples(sample_ids: list[str], partition_id: str) → None#

Drop a set of uids from TQ.

Used both for step-end teardown (via :meth:finish_step) and mid-step filtering (e.g. dynamic sampling).

finish_step(meta: nemo_rl.data_plane.KVBatchMeta) → None#

Drop this step’s bulk from TQ. Mirror of :meth:prepare_step.

collect_data_plane_snapshots() → list[dict[str, Any]]#

This driver’s data-plane counters plus every worker rank’s.

The driver sees roughly a sixth of a step’s traffic — the rollout actor writes the batch and the workers read it back per DP rank, both in other processes with their own counters. Aggregating is what turns these series from one process’s slice into the cluster figure.

Best effort by design: a rank that cannot answer is dropped rather than failing the step, because a metrics fan-out must never be able to take training down. Measured at ~2.4 ms and ~1 kB per process.

get_data_plane_step_metrics(
step_time_s: float,
) → tuple[dict[str, float], str] | None#

This step’s data-plane cost and the scope it covers, or None.

None when observability is off, so the caller filters rather than repeating the check. The scope is the cluster’s – the driver’s counters plus every worker rank’s – and falls back to the driver’s alone when the fan-out reached only one process. Reported one way or the other, never both, so there is a single answer to “what did the data plane cost” rather than two that disagree by roughly the DP degree.

Only the cluster baseline lives here; the driver’s stays on the client that owns those counters. The driver reading is taken every step, even when the cluster view supersedes it, so that a step which falls back after N cluster steps differences against last step rather than reporting N steps’ accumulated history as one.

_with_route_fields(
meta: nemo_rl.data_plane.KVBatchMeta,
base_fields: tuple[str, ...],
*,
task_name: str,
want_routes: bool,
) → nemo_rl.data_plane.KVBatchMeta#

Resolve direct versus deferred route storage for one worker request.

Delegates to :meth:TQDriverMixin._isolated_meta so the narrowed meta also gets its per-dispatch forward-pad target minted.

_logprob_dispatch(
meta: nemo_rl.data_plane.KVBatchMeta,
*,
task_name: str,
worker_method: str,
timer_prefix: str,
timer: Optional[nemo_rl.utils.timer.Timer],
common_kwargs: dict[str, Any],
include_router_replay: bool = False,
) → None#

Shared body of get_logprobs_from_meta / get_reference_policy_logprobs_from_meta.

Logprob workers fetch LP_SEED_FIELDS plus the multimodal columns _isolated_meta unions in, so prev/ref logprobs see the same model inputs as the training forward, which is narrowed through the same helper. Narrowing the meta’s field list still keeps rollout-only payload (message-log bulk, content) in TQ. The same shape is used for both prev_lp and ref_lp. Workers compute the per-token tensor and commit it to TQ via the leader-rank _write_back_result_field; the Ray return is always None, so this dispatcher just waits for completion.

get_logprobs_from_meta(
meta: nemo_rl.data_plane.KVBatchMeta,
micro_batch_size: Optional[int] = None,
timer: Optional[nemo_rl.utils.timer.Timer] = None,
) → None#
get_reference_policy_logprobs_from_meta(
meta: nemo_rl.data_plane.KVBatchMeta,
micro_batch_size: Optional[int] = None,
timer: Optional[nemo_rl.utils.timer.Timer] = None,
) → None#
train_from_meta(
meta: nemo_rl.data_plane.KVBatchMeta,
loss_fn: nemo_rl.algorithms.loss.interfaces.LossFunction,
eval_mode: bool = False,
gbs: Optional[int] = None,
mbs: Optional[int] = None,
timer: Optional[nemo_rl.utils.timer.Timer] = None,
train_fields: tuple[str, ...] = DP_TRAIN_FIELDS,
) → dict[str, Any]#

1-hop counterpart to :meth:train.

meta names per-sample keys; columns written by the rollout actor + worker logprob deltas + driver-side advantage delta have all landed under the same keys at this point. Workers fetch the union via train_presharded → self._fetch(meta). No partition drain here — sync 1-hop’s trainer calls clear_samples once at end of step.

Parameters:
  • meta – Full-step KVBatchMeta (consumed by all DP ranks).

  • gbs – Global batch size; defaults to cfg["train_global_batch_size"].

  • mbs – Micro batch size; defaults to cfg["train_micro_batch_size"].

  • timer – Optional timer for nested policy_training/* measurements.

  • train_fields – TQ columns workers fetch this step; defaults to the full DP_TRAIN_FIELDS schema. Caller may narrow it to drop columns it skipped writing (e.g. prev_logprobs when force_on_policy_ratio=True).

Returns:

Aggregated training-step output dict.

begin_train_step(
loss_fn: nemo_rl.algorithms.loss.interfaces.LossFunction,
gbs: Optional[int] = None,
mbs: Optional[int] = None,
) → None#

Open a logical train step on every worker.

train_microbatches_from_meta(
meta: nemo_rl.data_plane.KVBatchMeta,
timer: Optional[nemo_rl.utils.timer.Timer] = None,
train_fields: tuple[str, ...] = DP_TRAIN_FIELDS,
) → None#

Dispatch one meta slice (DP-sharded) into an open train step.

Named plural because one call fans out to every DP rank and the backend then iterates its own internal (pipeline/packed) microbatches — with a 2x packing ratio a group of G generations is G/2 backend microbatches inside this single call, not G/2 calls.

Mirrors the sharding logic of :meth:train_from_meta but without a logical-batch sizing constraint: this routes meta to DP ranks and runs forward+backward; gradients accumulate in .grad. Returns nothing — per-microbatch metrics accumulate in the workers’ open-step state and surface once via

Meth:

finish_train_step.

Parameters:
  • meta – Data-plane metadata for the samples in this chunk.

  • timer – Optional timer for nested policy-training measurements.

  • train_fields – Columns produced for this step and fetched by workers.

train_placed_microbatches(
dp_metas: list[nemo_rl.data_plane.KVBatchMeta],
timer: Optional[nemo_rl.utils.timer.Timer] = None,
) → None#

Dispatch one producer-assigned metadata batch per logical DP rank.

The input order is the logical DP-rank order. Producer field lists remain unchanged because an SFT loader can provide a narrower schema than the rollout training path.

_stamp_placed_pad_seqlen(
dp_metas: list[nemo_rl.data_plane.KVBatchMeta],
) → list[nemo_rl.data_plane.KVBatchMeta]#

Mint one fresh forward padding target across all placed DP batches.

Returns new metadata rather than mutating the caller’s, and ignores any inherited target: reusing one would let it ratchet upward across steps and pad every later step to a historical maximum. This mirrors TQDriverMixin._isolated_meta, which pops the key for the same reason.

_dispatch_train_microbatches(
dp_metas: list[nemo_rl.data_plane.KVBatchMeta],
*,
timer: Optional[nemo_rl.utils.timer.Timer],
) → None#

Send prepared per-DP metadata into an open train step.

finish_train_step() → dict[str, Any]#

Close an open train step: all_reduce, rescale, optimizer.step.

Aggregates per-rank step results into the same shape as

Meth:

train_from_meta so callers don’t have to special-case the split path.

abort_train_step() → None#

Drop partial step state on every worker. No optimizer.step.