nemo_rl.experience.payload#

Producer-side payload helpers for the async-RL TQ path.

Module Contents#

Functions#

_violation_counts

Count invalid tool calls / malformed thinking over flagged assistant turns.

_add_message_violation_masks

Attach token-aligned masks for generated assistant violations.

record_to_train_batch

Convert one prompt group’s record into a packed BatchedDataDict of N rows.

pack_payload

Pack a producer batch into (sample_ids, fields, tags) for put_samples.

Data#

API#

nemo_rl.experience.payload.VIOLATION_TAG_KEYS#

(‘num_invalid_tool_calls’, ‘num_malformed_thinking’, ‘num_assistant_messages’, ‘num_routed_experts_b…

nemo_rl.experience.payload._VIOLATION_COUNTS_KEY#

‘violation_counts’

nemo_rl.experience.payload._violation_counts(
message_log: nemo_rl.data.interfaces.LLMMessageLogType | nemo_rl.data.interfaces.VLMMessageLogType,
) → dict[str, int]#

Count invalid tool calls / malformed thinking over flagged assistant turns.

nemo_rl.experience.payload._add_message_violation_masks(
message_logs: list[nemo_rl.data.interfaces.LLMMessageLogType | nemo_rl.data.interfaces.VLMMessageLogType],
) → None#

Attach token-aligned masks for generated assistant violations.

This must run before the generic message normalizer fills missing generation_logprobs on prompt and environment messages, because field presence distinguishes generated assistant turns.

nemo_rl.experience.payload.record_to_train_batch(
record: nemo_rl.experience.interfaces.PromptGroupRecord,
*,
pad_value_dict: collections.abc.Mapping[str, int],
include_message_violation_fields: bool,
) → nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any]#

Convert one prompt group’s record into a packed BatchedDataDict of N rows.

Parameters:
  • record – Rollout’s PromptGroupRecord with N completions to flatten into rows.

  • pad_value_dict – Field-name → pad value used by batched_message_log_to_flat_message.

  • include_message_violation_fields – Whether to tensorize message violation flags for configured advantage penalties.

Returns:

BatchedDataDict with input IDs and lengths, generation log probabilities, token and prompt-level sample masks, raw mask_sample and truncated flags, prompt IDs for advantage computation, rewards, and violation counts. Optional fields include routed experts, message-violation masks, and any packed or per-token multimodal model inputs carried by the completions.

nemo_rl.experience.payload.pack_payload(
train_batch: collections.abc.Mapping[str, Any],
*,
weight_version: int,
group_id: str,
prompt_idx: int,
) → tuple[list[str], tensordict.TensorDict, list[dict[str, Any]]]#

Pack a producer batch into (sample_ids, fields, tags) for put_samples.

Parameters:
  • train_batch – Mapping with at least input_lengths plus the tensor/object fields to send.

  • weight_version – Trainer weight version stamped on every row’s tag.

  • group_id – Per-group identifier used as the sample_id prefix and stamped on every row’s tag; the caller owns uniqueness.

  • prompt_idx – Stable dataset prompt index stamped on every row’s tag.

Returns:

Sample IDs of the form {group_id}_g{i}, a jagged-packed TensorDict containing tensor fields and encoded multimodal wire fields, and per-row tags. Tags carry the weight version, prompt index, group id, violation counts, and <field>__row_shapes metadata required to reconstruct packed multimodal rows.