nemo_rl.experience.payload#
Producer-side payload helpers for the async-RL TQ path.
Module Contents#
Functions#
Count invalid tool calls / malformed thinking over flagged assistant turns. |
|
Attach token-aligned masks for generated assistant violations. |
|
Convert one prompt group’s record into a packed BatchedDataDict of N rows. |
|
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( ) 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],
Attach token-aligned masks for generated assistant violations.
This must run before the generic message normalizer fills missing
generation_logprobson 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,
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_sampleandtruncatedflags, 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,
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_shapesmetadata required to reconstruct packed multimodal rows.