nemo_rl.data_plane.codec#

Wire <-> trainer codec — jagged-on-the-wire bridge.

  • Writer side: variable-length fields are encoded as torch.nested.nested_tensor with layout=torch.jagged before put_samples. Padding tax is paid only when a consumer needs a rectangular tensor.

  • Reader side: :func:materialize accepts the wire TensorDict and, when layout='padded', calls

func:

torch.nested.to_padded_tensor on any nested leaves using the per-field padding value supplied in pad_value_dict. Trainer code consumes the padded BatchedDataDict unchanged.

  • Worker write-backs that produce response-shaped outputs use

func:

response_from_nested to extract the response slice from a (prompt+response) nested tensor.

  • Non-tensor object fields ride as NonTensorStack / NonTensorData leaves (TQ-native passthrough). :func:materialize decodes them back to np.ndarray(dtype=object) for the trainer.

Module Contents#

Functions#

record_codec_s

Record one pad/unpad measurement, in seconds.

timed_codec

Time a pad/unpad block, recording on every exit path.

drain_codec_ms

Milliseconds spent packing and unpacking since the last drain.

to_nested_by_length

Strip right-padding off a rectangular tensor using per-row lengths.

stack_or_nest

Stack equal-shape rows; reconstruct as jagged nested when ragged.

unwrap_wire_stripped_payload

Recover the payload of a possibly wire-stripped NonTensorData.

pack_jagged_fields

Pack a column dict into the wire layout expected by put_samples.

pack_per_token_field

Force-jaggedize a known per-token field, tolerating SP padding.

response_from_nested

Extract the response slice from a (prompt+response) nested tensor.

materialize

Convert a wire TensorDict to a BatchedDataDict.

Data#

API#

nemo_rl.data_plane.codec._CODEC_TIMER#

‘ThreadSafeTimer(…)’

nemo_rl.data_plane.codec.record_codec_s(phase: str, elapsed_s: float) → None#

Record one pad/unpad measurement, in seconds.

Prefer :func:timed_codec; this is for callers that already measured.

Parameters:
  • phase – "pack" or "unpack".

  • elapsed_s – Seconds spent, as returned by time.perf_counter() deltas.

nemo_rl.data_plane.codec.timed_codec(phase: str) → collections.abc.Iterator[None]#

Time a pad/unpad block, recording on every exit path.

Records in finally because the blocks it wraps return from more than one place – an early return once slipped past a hand-written bracket and silently dropped every no-op unpack from the metric.

Not :meth:Timer.time: Timer.start raises if the label is already running, and both phases run concurrently (the single-controller loop dispatches through asyncio.to_thread).

nemo_rl.data_plane.codec.drain_codec_ms() → dict[str, float]#

Milliseconds spent packing and unpacking since the last drain.

Timer.drain pops and sums under one lock: reduce then reset would drop any sample recorded between them, and both phases run concurrently.

Not every packing process has a reader – the rollout actor calls pack_jagged_fields but is not on the policy worker group, so nothing drains it. Its samples accumulate unread, which is why the caller that does drain should do so every step.

Returns:

{"pack": ms, "unpack": ms}, omitting a phase that did not run.

nemo_rl.data_plane.codec.to_nested_by_length(
padded: torch.Tensor,
lengths: torch.Tensor,
) → torch.Tensor#

Strip right-padding off a rectangular tensor using per-row lengths.

Used by the producer side: convert

Func:

batched_message_log_to_flat_message output (already padded) into the wire format before put_samples.

Parameters:
  • padded – Rectangular tensor of shape (N, S, ...).

  • lengths – Per-row valid lengths, shape (N,). CUDA tensors are moved to CPU once to avoid per-row syncs.

Returns:

A torch.jagged nested tensor whose i-th row is padded[i, :lengths[i], ...].

nemo_rl.data_plane.codec.stack_or_nest(tensors: list[torch.Tensor]) → torch.Tensor#

Stack equal-shape rows; reconstruct as jagged nested when ragged.

Parameters:

tensors – Per-row tensors; assumed to share leading dims modulo an optional ragged seq dim. Empty list returns torch.empty(0).

Returns:

A regular tensor when all rows share shape; otherwise a torch.jagged nested tensor.

nemo_rl.data_plane.codec.unwrap_wire_stripped_payload(item: Any) → Any#

Recover the payload of a possibly wire-stripped NonTensorData.

TQ’s MsgpackEncoder._encode_tensordict serializes any TensorDictBase via dict(obj.items()) — only the tensor backing dict. NonTensorData stores its payload in _non_tensordict["data"], so it round-trips through ZMQ as an empty TensorDict({}, batch_size=[]). We map only that exact signature to None; any other TensorDictBase (with tensor fields, non-scalar batch, or a salvageable _non_tensordict payload) passes through unchanged so we never drop real data.

nemo_rl.data_plane.codec.pack_jagged_fields(
fields: dict[str, torch.Tensor | np.ndarray],
*,
lengths: torch.Tensor | None,
token_aligned_fields: set[str] | frozenset[str] | None = None,
) → tensordict.TensorDict#

Pack a column dict into the wire layout expected by put_samples.

Zero-copy where possible: explicitly named per-token tensors become torch.jagged views via :func:pack_per_token_field; all other tensors pass through rectangular; np.ndarray(dtype=object) is forwarded as-is. This is a layout transform, not serialization — the on-wire bytes are produced later by the TQ backend’s msgpack encoder. Centralizing the transform here makes it the single source of truth for both :func:kv_first_write and :func:write_columns.

Parameters:
  • fields – Column name → tensor or object array. Other value types raise TypeError.

  • lengths – Per-row valid lengths used by :func:pack_per_token_field. None disables jagged conversion entirely.

  • token_aligned_fields –

    Field names known to be per-token. These use

    func:

    pack_per_token_field, which tolerates extra padded columns and slices each row to lengths.

Returns:

TensorDict with batch_size=[N] (N from lengths if given, else 0) ready for put_samples.

nemo_rl.data_plane.codec.pack_per_token_field(
val: torch.Tensor,
lengths: torch.Tensor,
) → torch.Tensor#

Force-jaggedize a known per-token field, tolerating SP padding.

This function is invoked at write sites where the caller already knows the field is per-token (e.g. prev_logprobs, reference_policy_logprobs). mcore SP rounds the forward output’s seq dim up to a multiple of TP, so the value can be 1+ tokens wider than max(lengths); :func:to_nested_by_length slices each row to its own length and drops the trailing SP padding cleanly.

Parameters:
  • val – Per-token tensor. Falls back to rectangular when it cannot be jaggedized (wrong batch dim, < 2D, or seq dim shorter than max(lengths)).

  • lengths – Per-row valid lengths, shape (N,).

Returns:

A torch.jagged nested tensor when the shape allows; otherwise val passed through as a rectangular tensor.

nemo_rl.data_plane.codec.response_from_nested(
full: torch.Tensor,
response_mask: torch.Tensor,
) → torch.Tensor#

Extract the response slice from a (prompt+response) nested tensor.

Used on the worker side for logprob / ref-logprob write-back where only the response-token slice is interesting downstream. The “left-shift by one token” convention is applied (so logprobs at output position i correspond to the prediction of input token i+1).

Parameters:
  • full – Jagged nested tensor of shape (N, prompt_len + response_len).

  • response_mask – Jagged nested tensor of shape (N, response_len); its offsets().diff() gives the per-row response length.

Returns:

Jagged nested tensor of shape (N, response_len) containing the left-shifted response slice.

nemo_rl.data_plane.codec.materialize(
td: tensordict.TensorDict,
layout: nemo_rl.data_plane.schema.Layout = 'padded',
pad_value_dict: dict[str, int | float] | None = None,
pad_to_seqlen: int = 0,
tags: list[dict[str, Any]] | None = None,
) → BatchedDataDict[Any]#

Convert a wire TensorDict to a BatchedDataDict.

Trainer/worker code expects rectangular tensors — this is the bridge from the on-wire nested format.

The lazy BatchedDataDict / multimodal_utils imports keep import nemo_rl.data_plane cheap for unit tests that don’t actually call this function — both transitively pull PIL, requests and a few hundred transformers submodules.

Parameters:
  • td – Wire TensorDict to materialize.

  • layout –

    "padded" (default) pads nested-tensor leaves via

    func:

    torch.nested.to_padded_tensor using pad_value_dict[k] (or 0 if unspecified); rectangular leaves pass through. "jagged" passes nested leaves through — use only when the caller knows how to consume them.

  • pad_value_dict – Per-field pad value used when layout='padded'.

  • pad_to_seqlen – When > 0, right-pad the seq dim up to this absolute length after to_padded_tensor. Worker-side _fetch passes its forward-pass target here (rounded up to sequence_length_round for Megatron’s microbatch iterator); driver-side read_columns leaves it 0 and consumes the natural-padded shape. Default 0 disables.

Returns:

BatchedDataDict with rectangular tensors for padded layout, nested tensors for jagged layout, and np.ndarray(dtype=object) for NonTensorStack leaves (TQ-native non-tensor passthrough).