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 viaself._fetch(meta)and write deltas back viaself._write_back_result_field(...). Seenemo_rl/data_plane/README.mdfor the full design.
Module Contents#
Classes#
TQ-mediated counterpart to :class: |
Functions#
Data#
API#
- nemo_rl.models.policy.tq_policy._aggregate_train_results(
- results: list[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.PolicyTQ-mediated counterpart to :class:
Policy.Constructor accepts an additional
dp_cfg(themaster_config["data_plane"]dict). Bootstraps the controller on the driver and forwardssetup_data_plane(dp_cfg)to every worker so they can attach as clients (bootstrap=False).checkpointingis 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 bytq_partition_id(default"train") is open with a schema coveringDP_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,
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,
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 viaselect_fields.- Parameters:
num_samples – Expected total samples this step.
group_size – GRPO group size for balanced sampling;
Nonedisables grouping.
- prepare_val_partition(
- num_samples: int,
- *,
- partition_id: str = 'val',
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_stepbecause 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,
This step’s data-plane cost and the scope it covers, or
None.Nonewhen 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,
Resolve direct versus deferred route storage for one worker request.
Delegates to :meth:
TQDriverMixin._isolated_metaso 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,
Shared body of get_logprobs_from_meta / get_reference_policy_logprobs_from_meta.
Logprob workers fetch
LP_SEED_FIELDSplus the multimodal columns_isolated_metaunions 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,
- 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,
- 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,
1-hop counterpart to :meth:
train.metanames 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 viatrain_presharded→self._fetch(meta). No partition drain here — sync 1-hop’s trainer callsclear_samplesonce 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_FIELDSschema. Caller may narrow it to drop columns it skipped writing (e.g.prev_logprobswhenforce_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,
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,
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_metabut without a logical-batch sizing constraint: this routesmetato 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,
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],
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],
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_metaso callers don’t have to special-case the split path.
- abort_train_step() None#
Drop partial step state on every worker. No optimizer.step.