Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
30 commits
Select commit Hold shift + click to select a range
b7fa6e3
feat(nemo-gym): support multimodal rollouts with tokenizer_config plu…
rohitrango Aug 3, 2026
edb54a6
feat(nemo-gym): plumb multimodal rollouts through the async single-co…
rohitrango Jul 22, 2026
4cb8023
chore(examples): add Nemotron-Omni gym-v smoke configs (tangram + pol…
rohitrango Jul 29, 2026
37098a9
fix(nemo-gym): drop unused nemo_gym_row arg from _postprocess_nemo_gy…
rohitrango Jul 30, 2026
5b43d98
delete scratchspace configs
rohitrango Aug 4, 2026
faa185a
consolidate into one entrypoint
rohitrango Aug 4, 2026
98de6f8
reverted processor design to be consistent with vlm_grpo
rohitrango Aug 4, 2026
a50f772
chore: clean up multimodal processor plumbing
rohitrango Aug 5, 2026
f6300f8
change per-turn images to get results from tool-calls, etc (anything
rohitrango Aug 4, 2026
629a1d4
fix(nemo_gym): flush per-turn image bucket on any trainable item
rohitrango Aug 4, 2026
171506f
docs: add Google-style docstrings to image encoding helpers
rohitrango Aug 4, 2026
a104ae6
change non-default config option
rohitrango Aug 4, 2026
d3e02b4
build: preserve main dependency configuration
rohitrango Aug 5, 2026
428f5ab
(chore): add copyright notice to test
rohitrango Aug 5, 2026
33b70ec
(chore): undo vllm chat request change
rohitrango Aug 5, 2026
ef3f163
fix: address NeMo Gym multimodal image indexing issues
rohitrango Aug 5, 2026
1c5bfbb
lint fixes
rohitrango Aug 5, 2026
062e999
allow mixed (multimodal, text) batches from nemo-gym batch rollouts
rohitrango Aug 5, 2026
0edbe98
chore: apply ruff format to llm_message_utils
rohitrango Aug 6, 2026
a70eacc
feat(recipes): add Nemotron-Omni 30B circle-click 2n8g VLM-GRPO recipe
rohitrango Aug 6, 2026
68a538b
feat(nemo_gym): assert placeholder-style processor at actor init
rohitrango Aug 6, 2026
b9e1a27
chore(recipes): inherit from vlm_grpo_3B_megatron exemplar in circle-…
rohitrango Aug 6, 2026
4bb7e20
minimized config
rohitrango Aug 6, 2026
aa5ea0b
test(vlm): add circle-click gym driver script, disabled for now
yfw Aug 6, 2026
8ba3fe6
Merge branch 'main' into rohit/gymv-mm-integration-v2
aroshanghias-nvd Aug 6, 2026
f27dc6d
Merge branch 'main' into rohit/gymv-mm-integration-v2
rohitrango Aug 6, 2026
9195924
feat: deduplicate multimodal GRPO payloads
aroshanghias-nvd Aug 6, 2026
29de7b5
chore: trim multimodal dedup scope
aroshanghias-nvd Aug 6, 2026
d95bb1a
perf(grpo): pre-cast pixel payloads to bf16
aroshanghias-nvd Aug 6, 2026
36aa640
Merge main into multimodal dedup branch
aroshanghias-nvd Aug 7, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
defaults: ../../vlm_grpo_3B_megatron.yaml
grpo:
deduplicate_multimodal_data: true
num_prompts_per_step: 1
num_val_generations_per_prompt: 1
max_num_steps: 100
Expand Down
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
defaults: ../../vlm_grpo_3B.yaml
grpo:
deduplicate_multimodal_data: true
num_prompts_per_step: 32
val_at_start: true
checkpointing:
Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
defaults: ../../vlm_grpo_3B_megatron.yaml
grpo:
deduplicate_multimodal_data: true
loss_fn:
reference_policy_kl_penalty: 0.0
checkpointing:
Expand Down
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
defaults: ../../vlm_grpo_3B.yaml
grpo:
deduplicate_multimodal_data: true
num_prompts_per_step: 32
overlong_filtering: true
seq_logprob_error_threshold: 2
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@ grpo:
num_prompts_per_step: 512
overlong_filtering: true
zero_variance_prompt_filtering: false
deduplicate_multimodal_data: false
deduplicate_multimodal_data: true
loss_fn:
ratio_clip_max: 0.28
use_on_policy_kl_approximation: true
Expand Down
2 changes: 2 additions & 0 deletions examples/configs/vlm_grpo_3B.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@
defaults: "grpo_math_1B.yaml"

grpo:
deduplicate_multimodal_data: false
debug_payload_metrics: false
num_prompts_per_step: 8
reward_shaping:
overlong_buffer_length: 512
Expand Down
2 changes: 2 additions & 0 deletions examples/nemo_gym/run_grpo_nemo_gym.py
Original file line number Diff line number Diff line change
Expand Up @@ -315,6 +315,7 @@ def main() -> None:
max_trajectory_age_steps=config.grpo.async_grpo.max_trajectory_age_steps,
teacher_worker_groups=teacher_worker_groups,
alias_to_group_alias=alias_to_group_alias,
processor=processor,
)
else:
print("🚀 Running synchronous GRPO training")
Expand All @@ -333,6 +334,7 @@ def main() -> None:
checkpointer,
grpo_state,
master_config,
processor=processor,
)


Expand Down
1 change: 1 addition & 0 deletions examples/run_vlm_grpo.py
Original file line number Diff line number Diff line change
Expand Up @@ -146,6 +146,7 @@ def main() -> None:
checkpointer,
grpo_state,
master_config,
processor=processor,
)


Expand Down
14 changes: 14 additions & 0 deletions nemo_rl/algorithms/async_utils/interfaces.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,20 @@ def load_state_dict(
"""Restore state produced by ``state_dict``."""
...

def save_to_path(self, path: str) -> int:
"""Serialize state directly from the replay actor."""
...

def load_from_path(
self,
path: str,
num_prompts_per_step: int | None = None,
current_training_step: int | None = None,
max_age_steps: int | None = None,
) -> dict[str, int]:
"""Restore state directly in the replay actor."""
...

def get_trajectories_needed(
self,
target_step: int,
Expand Down
46 changes: 45 additions & 1 deletion nemo_rl/algorithms/async_utils/replay_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
# limitations under the License.

import asyncio
import gc
import statistics
import threading as _threading
import uuid
Expand All @@ -21,11 +22,16 @@
from typing import Any, Iterable, Optional

import ray
import torch

from nemo_rl.algorithms.async_utils.interfaces import ReplayBufferProtocol
from nemo_rl.data_plane import KVBatchMeta
from nemo_rl.data_plane.schema import ROUTED_EXPERTS_FIELD
from nemo_rl.experience.interfaces import PromptGroupRecord
from nemo_rl.experience.interfaces import (
NEMO_GYM_TASK_INDEX_KEY,
NEXT_NEMO_GYM_TASK_INDEX_KEY,
PromptGroupRecord,
)
from nemo_rl.experience.payload import pack_payload, record_to_train_batch
from nemo_rl.utils.r3_trace import trace_rollout_payload

Expand Down Expand Up @@ -340,6 +346,44 @@ def state_dict(self) -> dict[str, Any]:
"max_size": self.max_size,
}

def save_to_path(self, path: str) -> int:
"""Serialize inside the actor without materializing the buffer on the driver."""
state = self.state_dict()
torch.save(state, path)
num_trajectories = len(state["trajectories"])
del state
gc.collect()
return num_trajectories

def load_from_path(
self,
path: str,
num_prompts_per_step: int | None = None,
current_training_step: int | None = None,
max_age_steps: int | None = None,
) -> dict[str, int]:
"""Restore inside the actor and return only compact coordination metadata."""
state = torch.load(path, weights_only=False)
saved_task_indices = [
int(trajectory[NEMO_GYM_TASK_INDEX_KEY])
for trajectory in state.get("trajectories", [])
if trajectory.get(NEMO_GYM_TASK_INDEX_KEY) is not None
]
next_task_index = max(saved_task_indices, default=-1) + 1
num_trajectories = len(state["trajectories"])
self.load_state_dict(
state,
num_prompts_per_step=num_prompts_per_step,
current_training_step=current_training_step,
max_age_steps=max_age_steps,
)
del state
gc.collect()
return {
"num_trajectories": num_trajectories,
NEXT_NEMO_GYM_TASK_INDEX_KEY: next_task_index,
}

def load_state_dict(
self,
state: dict[str, Any],
Expand Down
55 changes: 53 additions & 2 deletions nemo_rl/algorithms/async_utils/trajectory_collector.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,10 +38,16 @@
)
from nemo_rl.experience.rollouts import (
RolloutGroupResult,
attach_initial_nemo_gym_image_payloads,
run_async_multi_turn_rollout_groups,
)
from nemo_rl.models.generation.interfaces import GenerationConfig, GenerationInterface
from nemo_rl.utils.logger import should_log_nemo_gym_full_result_tables
from nemo_rl.utils.multimodal_payload_metrics import (
collect_multimodal_payload_metrics,
drain_multimodal_payload_metrics,
print_multimodal_payload_metrics,
)
from nemo_rl.utils.timer import ThreadSafeTimer

TokenizerType = PreTrainedTokenizerBase
Expand All @@ -66,6 +72,7 @@ def __init__(
alias_to_group_alias: Optional[dict[str, str]] = None,
on_policy_distillation_cfg: Optional[dict[str, Any]] = None,
next_nemo_gym_task_index: int = 0,
processor: Any = None,
):
self.policy_generation = policy_generation
self.tokenizer = tokenizer
Expand All @@ -75,6 +82,7 @@ def __init__(
self.teacher_worker_groups = teacher_worker_groups or {}
self.alias_to_group_alias = alias_to_group_alias or {}
self.on_policy_distillation_cfg = on_policy_distillation_cfg or {}
self.processor = processor
self._has_distillation_teachers = bool(self.teacher_worker_groups)
self._teacher_seq_pad_multiple = teacher_seq_pad_multiple(
self.teacher_worker_groups,
Expand Down Expand Up @@ -428,7 +436,23 @@ def _process_batch(self, batch: BatchedDataDict[DatumSpec]) -> None:
rollout_batch = batch.slice(0, num_prompts_to_generate)
if use_nemo_gym:
self._stamp_nemo_gym_task_indices(rollout_batch)
repeated_batch = rollout_batch.repeat_interleave(num_generations)
if self.master_config.grpo.deduplicate_multimodal_data:
attach_initial_nemo_gym_image_payloads(
rollout_batch, self.processor
)
repeated_batch = rollout_batch.repeat_interleave(
num_generations,
share_immutable_media=(
self.master_config.grpo.deduplicate_multimodal_data
),
)
print_multimodal_payload_metrics(
collect_multimodal_payload_metrics(
repeated_batch,
"prompt_repeat_async",
enabled=self.master_config.grpo.debug_payload_metrics,
)
)

def _run_rollout_batch() -> None:
asyncio.run(
Expand Down Expand Up @@ -605,6 +629,15 @@ def get_efficiency_metrics(self) -> dict[str, float]:
self._efficiency_timer.get_timing_metrics(reduction_op="sum"),
)

async def drain_payload_metrics(self) -> dict[str, int | float]:
"""Close one drain-to-drain collector/Gym telemetry interval.

Rollout collection is concurrent with training, so the interval is not
claimed to own the sampled training batch. Call-normalized metrics make
intervals comparable even when their background transfer counts differ.
"""
return drain_multimodal_payload_metrics()

def get_rollouts_state(self) -> dict[str, int]:
"""Get collector-side rollout state for checkpointing."""
return {NEXT_NEMO_GYM_TASK_INDEX_KEY: self._next_nemo_gym_task_index}
Expand Down Expand Up @@ -777,6 +810,10 @@ async def _iter_rollout_groups(
mask_env_flagged_samples=should_mask_flagged_samples(
self.master_config.env
),
deduplicate_multimodal_data=(
self.master_config.grpo.deduplicate_multimodal_data
),
debug_payload_metrics=self.master_config.grpo.debug_payload_metrics,
):
task_index = rollout_result.task_index
if task_index is None:
Expand All @@ -801,6 +838,9 @@ async def _iter_rollout_groups(
num_generations=num_generations,
max_rollout_turns=self.master_config.grpo.max_rollout_turns,
greedy=False,
deduplicate_multimodal_data=(
self.master_config.grpo.deduplicate_multimodal_data
),
):
yield rollout_result

Expand Down Expand Up @@ -922,11 +962,22 @@ async def _enqueue_rollout_group(
}
if rollout_result.task_index is not None:
trajectory_group[NEMO_GYM_TASK_INDEX_KEY] = rollout_result.task_index

backoff_delay = 0.01
backoff_started_at: float | None = None
try:
while self.running:
# Every retry is a distinct Ray submission of the full payload.
print_multimodal_payload_metrics(
collect_multimodal_payload_metrics(
(
trajectory_group,
generation_weight_version,
target_weight_version,
),
"replay_push",
enabled=self.master_config.grpo.debug_payload_metrics,
)
)
status = await self.replay_buffer.add.remote(
trajectory_group,
generation_weight_version,
Expand Down
Loading
Loading