Skip to content

Commit f9b5c37

Browse files
yuki-97RayenTianterrykong
authored
feat(sc): NeMo-Gym path (#3267)
Signed-off-by: Yuki Huang <yukih@nvidia.com> Signed-off-by: ruit <ruit@nvidia.com> Signed-off-by: Terry Kong <terryk@nvidia.com> Co-authored-by: ruit <ruit@nvidia.com> Co-authored-by: Terry Kong <terryk@nvidia.com>
1 parent 3ee5fa3 commit f9b5c37

9 files changed

Lines changed: 309 additions & 21 deletions

File tree

examples/nemo_gym/run_distillation_nemo_gym.py

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -32,9 +32,7 @@
3232
from nemo_rl.algorithms.utils import get_tokenizer
3333
from nemo_rl.data.utils import setup_response_data
3434
from nemo_rl.distributed.virtual_cluster import init_ray
35-
from nemo_rl.environments.nemo_gym import (
36-
setup_nemo_gym_config,
37-
)
35+
from nemo_rl.environments.nemo_gym import setup_nemo_gym_config
3836
from nemo_rl.models.generation import configure_generation_config
3937
from nemo_rl.utils.config import (
4038
load_config,

examples/run_grpo_single_controller.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,7 @@
3434
)
3535
from nemo_rl.algorithms.utils import get_tokenizer
3636
from nemo_rl.distributed.virtual_cluster import init_ray
37+
from nemo_rl.environments.nemo_gym import setup_nemo_gym_config
3738
from nemo_rl.models.generation import configure_generation_config
3839
from nemo_rl.utils.config import (
3940
load_config,
@@ -117,6 +118,10 @@ def main() -> None:
117118
trains_mtp=trains_mtp,
118119
)
119120

121+
# NeMo-Gym specific config setup.
122+
if bool(config.env.get("should_use_nemo_gym")):
123+
setup_nemo_gym_config(config, tokenizer)
124+
120125
actor_args = setup_single_controller(config, tokenizer)
121126

122127
print("🚀 Launching SingleControllerActor")

nemo_rl/algorithms/single_controller_utils/config.py

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -118,6 +118,26 @@ def validate_single_controller_config(master_config: MasterConfig) -> None:
118118
sampler_name=async_config.sampler.name,
119119
)
120120

121+
# A non-zero reference-policy KL penalty makes the loss read
122+
# ``reference_policy_logprobs``, but the SC train pump only computes them
123+
# when ``skip_reference_policy_logprobs_calculation`` is false (see
124+
# SingleControllerActor._reference_logprobs_required). Catch the
125+
# inconsistent pair at setup instead of a mid-training KeyError.
126+
reference_policy_kl_penalty = getattr(
127+
master_config.loss_fn, "reference_policy_kl_penalty", 0
128+
)
129+
if reference_policy_kl_penalty and master_config.grpo.get(
130+
"skip_reference_policy_logprobs_calculation"
131+
):
132+
raise ValueError(
133+
"loss_fn.reference_policy_kl_penalty="
134+
f"{reference_policy_kl_penalty} requires reference_policy_logprobs, "
135+
"but grpo.skip_reference_policy_logprobs_calculation=true skips "
136+
"computing them on the SingleController path. Set "
137+
"grpo.skip_reference_policy_logprobs_calculation=false, or set "
138+
"loss_fn.reference_policy_kl_penalty=0."
139+
)
140+
121141

122142
# ── Internal SingleController configs ────────────────────────────────────
123143

nemo_rl/algorithms/single_controller_utils/setup.py

Lines changed: 38 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -49,11 +49,16 @@
4949
from nemo_rl.data_plane import DataPlaneClient, build_data_plane_client
5050
from nemo_rl.distributed.virtual_cluster import RayVirtualCluster
5151
from nemo_rl.environments.interfaces import EnvironmentInterface
52+
from nemo_rl.environments.nemo_gym import spinup_nemo_gym_actor
5253
from nemo_rl.experience.rollout_manager import RolloutManager
5354
from nemo_rl.models.generation.sglang.config import SGLangConfig
5455
from nemo_rl.models.generation.sglang.sglang_generation import SGLangGeneration
5556
from nemo_rl.models.generation.vllm import VllmGeneration
5657
from nemo_rl.models.generation.vllm.config import VllmConfig
58+
from nemo_rl.models.megatron.router_replay import (
59+
configure_vllm_for_router_replay,
60+
router_replay_enabled,
61+
)
5762
from nemo_rl.models.policy.tq_policy import TQPolicy
5863
from nemo_rl.weight_sync import WeightSynchronizer, create_weight_synchronizer
5964

@@ -167,6 +172,7 @@ def _build_generation(
167172
vllm_config.setdefault("vllm_kwargs", {})["hf_overrides"] = (
168173
master_config.policy.get("hf_config_overrides", {})
169174
)
175+
configure_vllm_for_router_replay(master_config.policy)
170176
gen = VllmGeneration(cluster=inference_cluster, config=vllm_config)
171177
elif backend == "sglang":
172178
sglang_config = cast(SGLangConfig, generation_config)
@@ -318,16 +324,25 @@ def setup_single_controller(
318324
# Setup Dataset & Environments
319325
# ==========================
320326
# TODO: add validate dataset wiring.
321-
if _should_use_nemo_gym(cast(GrpoMasterConfig, master_config)):
327+
use_nemo_gym = _should_use_nemo_gym(cast(GrpoMasterConfig, master_config))
328+
if use_nemo_gym and generation_config["backend"] != "vllm":
322329
raise NotImplementedError(
323-
"NeMo-Gym integration for SingleController is not supported yet; "
324-
"it will land in https://github.com/NVIDIA-NeMo/RL/pull/3267."
330+
"SC NeMo-Gym integration currently supports the vllm backend "
331+
f"only; got {generation_config['backend']!r}"
325332
)
326-
response_data = setup_response_data(
327-
tokenizer, data_config, env_configs=master_config.env
328-
)
329-
assert len(response_data) == 4
330-
dataset, _val_dataset, env_handles, _val_env_handles = response_data
333+
if use_nemo_gym:
334+
# NeMo-Gym creates the env actor outside setup_response_data; we wire
335+
# it in after generation is up (it needs the OpenAI server URLs).
336+
response_data = setup_response_data(tokenizer, data_config, env_configs=None)
337+
assert len(response_data) == 2
338+
dataset, _val_dataset = response_data
339+
env_handles: dict[str, EnvironmentInterface] = {}
340+
else:
341+
response_data = setup_response_data(
342+
tokenizer, data_config, env_configs=master_config.env
343+
)
344+
assert len(response_data) == 4
345+
dataset, _val_dataset, env_handles, _val_env_handles = response_data
331346
dataloader = StatefulDataLoader(
332347
dataset,
333348
batch_size=grpo_config["num_prompts_per_step"],
@@ -363,6 +378,20 @@ def setup_single_controller(
363378
generation = gen_future.result()
364379
policy = policy_future.result()
365380

381+
# ==========================
382+
# NeMo-Gym actor (after generation is up so OpenAI URLs are available)
383+
# ==========================
384+
if use_nemo_gym:
385+
# TODO(#2625): Mirror GRPO's deferred vLLM load so NeMo-Gym spinup
386+
# overlaps model loading instead of running serially afterward.
387+
enable_router_replay = router_replay_enabled(master_config.policy)
388+
env_handles["nemo_gym"] = spinup_nemo_gym_actor(
389+
env_configs=master_config.env,
390+
base_urls=generation.dp_openai_server_base_urls,
391+
model_name=generation_config["model_name"],
392+
enable_router_replay=enable_router_replay,
393+
)
394+
366395
# ==========================
367396
# Setup Data Plane Client & Weight Sync
368397
# ==========================
@@ -403,6 +432,7 @@ def setup_single_controller(
403432
max_rollout_turns=grpo_config.get("max_rollout_turns"),
404433
policy_generation=generation,
405434
generation_config=generation_config,
435+
use_nemo_gym=use_nemo_gym,
406436
tq_buffer=tq_buffer,
407437
)
408438

nemo_rl/environments/nemo_gym.py

Lines changed: 80 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,12 +18,14 @@
1818
from collections import Counter
1919
from collections.abc import AsyncGenerator
2020
from pathlib import Path
21-
from typing import Any, Dict, List, NotRequired, TypedDict
21+
from typing import Any, Dict, List, NotRequired, Optional, TypedDict
2222

2323
import ray
2424
import torch
25+
from ray.util.scheduling_strategies import NodeAffinitySchedulingStrategy
2526
from transformers import PreTrainedTokenizerBase
2627

28+
from nemo_rl.distributed.ray_actor_environment_registry import get_actor_python_env
2729
from nemo_rl.distributed.virtual_cluster import (
2830
DEFAULT_GYM_PORT_RANGE_HIGH,
2931
DEFAULT_GYM_PORT_RANGE_LOW,
@@ -32,6 +34,7 @@
3234
)
3335
from nemo_rl.environments.interfaces import EnvironmentInterface
3436
from nemo_rl.utils.timer import Timer
37+
from nemo_rl.utils.venvs import create_local_venv_on_each_node
3538

3639
# Kept local (not imported from models.generation) so the gym actor stays free of
3740
# generation-module imports. Must cover every name resolve_routed_experts_dtype
@@ -600,3 +603,79 @@ def setup_nemo_gym_config(config, tokenizer) -> None:
600603
# Stop strings or token ids are not supported
601604
generation_config["stop_strings"] = None
602605
generation_config["stop_token_ids"] = None
606+
607+
608+
def spinup_nemo_gym_actor(
609+
env_configs: dict[str, Any],
610+
base_urls: list[Optional[str]],
611+
model_name: str,
612+
enable_router_replay: bool,
613+
) -> Any:
614+
"""Spin up the NeMo-Gym actor against the given generation server URLs.
615+
616+
When ``env_configs["nemo_gym"]["num_gpu_nodes"] > 0``, the actor is
617+
scheduled with soft NodeAffinity to the current Ray node so its colocated
618+
GPU resources land where the caller expects.
619+
620+
Args:
621+
env_configs: The ``master_config.env`` mapping; ``env_configs["nemo_gym"]``
622+
supplies the Gym global config plus NeMo-RL detection knobs
623+
(``invalid_tool_call_patterns``, ``thinking_tags``, ``num_gpu_nodes``).
624+
base_urls: Per-DP-rank OpenAI-compatible server base URLs from the
625+
generation backend.
626+
model_name: Served model name the Gym rollouts should target.
627+
enable_router_replay: Sets ``require_routed_experts`` on the
628+
``NemoGymConfig``.
629+
630+
Returns:
631+
The spun-up ``NemoGym`` Ray actor handle (``_spinup`` already awaited).
632+
"""
633+
nemo_gym_dict = dict(env_configs["nemo_gym"])
634+
635+
# NeMo-RL-side detection knobs are top-level NemoGymConfig fields
636+
# (where the detector reads them), not part of Gym's global config.
637+
invalid_tool_call_patterns = nemo_gym_dict.pop("invalid_tool_call_patterns", None)
638+
thinking_tags = nemo_gym_dict.pop("thinking_tags", None)
639+
640+
# Pass prebuilt cache + venv dirs through the global config so the gym reuses
641+
# image-baked venvs instead of rebuilding them.
642+
uv_cache_dir = get_nemo_gym_uv_cache_dir()
643+
if uv_cache_dir is not None:
644+
nemo_gym_dict.setdefault("uv_cache_dir", uv_cache_dir)
645+
uv_venv_dir = get_nemo_gym_venv_dir()
646+
if uv_venv_dir is not None:
647+
nemo_gym_dict.setdefault("uv_venv_dir", uv_venv_dir)
648+
649+
nemo_gym_cfg = NemoGymConfig(
650+
model_name=model_name,
651+
base_urls=base_urls,
652+
invalid_tool_call_patterns=invalid_tool_call_patterns,
653+
thinking_tags=thinking_tags,
654+
require_routed_experts=enable_router_replay,
655+
initial_global_config_dict=nemo_gym_dict,
656+
)
657+
658+
nemo_gym_py_exec = get_actor_python_env("nemo_rl.environments.nemo_gym.NemoGym")
659+
if nemo_gym_py_exec.startswith("uv"):
660+
nemo_gym_py_exec = create_local_venv_on_each_node(
661+
nemo_gym_py_exec, "nemo_rl.environments.nemo_gym.NemoGym"
662+
)
663+
664+
nemo_gym_opts: dict[str, Any] = {}
665+
if nemo_gym_dict.get("num_gpu_nodes", 0):
666+
nemo_gym_opts["scheduling_strategy"] = NodeAffinitySchedulingStrategy(
667+
node_id=ray.get_runtime_context().get_node_id(),
668+
soft=True,
669+
)
670+
nemo_gym_opts["runtime_env"] = {
671+
"py_executable": nemo_gym_py_exec,
672+
"env_vars": {
673+
**os.environ,
674+
"VIRTUAL_ENV": nemo_gym_py_exec,
675+
"UV_PROJECT_ENVIRONMENT": nemo_gym_py_exec,
676+
},
677+
}
678+
679+
actor = NemoGym.options(**nemo_gym_opts).remote(nemo_gym_cfg)
680+
ray.get(actor._spinup.remote())
681+
return actor

tests/functional/L1_Functional_Tests_SingleController.sh

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -35,6 +35,7 @@ run_test() {
3535
}
3636

3737
run_test fast uv run --no-sync bash ./tests/functional/grpo_dp_single_controller.sh
38+
run_test fast uv run --no-sync bash ./tests/functional/grpo_async_gym_single_controller.sh
3839

3940
cd ${PROJECT_ROOT}/tests
4041
if compgen -G ".coverage*" > /dev/null; then
Lines changed: 114 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,114 @@
1+
#!/bin/bash
2+
# SingleController + NeMo-Gym e2e smoke. Mirrors grpo_async_gym.sh but
3+
# routes everything through the SC path (TransferQueue data plane +
4+
# SingleControllerActor) instead of async_grpo_train.
5+
6+
SCRIPT_DIR=$( cd -- "$( dirname -- "${BASH_SOURCE[0]}" )" &> /dev/null && pwd)
7+
PROJECT_ROOT=$(realpath $SCRIPT_DIR/../..)
8+
# Mark the current repo as safe, since wandb fetches metadata about the repo
9+
git config --global --add safe.directory $PROJECT_ROOT
10+
11+
set -eou pipefail
12+
13+
EXP_NAME=$(basename $0 .sh)
14+
EXP_DIR=$SCRIPT_DIR/$EXP_NAME
15+
LOG_DIR=$EXP_DIR/logs
16+
JSON_METRICS=$EXP_DIR/metrics.json
17+
RUN_LOG=$EXP_DIR/run.log
18+
CHECKPOINT_DIR=$EXP_DIR/checkpoints
19+
DATA_DIR=$EXP_DIR/data
20+
export PYTHONPATH=${PROJECT_ROOT}:${PYTHONPATH:-}
21+
22+
rm -rf $EXP_DIR $LOG_DIR
23+
mkdir -p $EXP_DIR $LOG_DIR $CHECKPOINT_DIR $DATA_DIR
24+
25+
# clean up checkpoint directory on exit
26+
trap "rm -rf $CHECKPOINT_DIR" EXIT
27+
28+
cd $PROJECT_ROOT
29+
30+
# Follow nemo-gym instructions here to get this data:
31+
# https://docs.nvidia.com/nemo/gym/0.1.0/tutorials/nemo-rl-grpo/setup.html#training-nemo-rl-grpo-setup
32+
cd 3rdparty/Gym-workspace/Gym
33+
34+
# We need HF_TOKEN to download the data from huggingface
35+
if [[ ! -f env.yaml ]]; then
36+
if [[ -z "${HF_TOKEN:-}" ]]; then
37+
echo "[ERROR] HF_TOKEN is not set"
38+
exit 1
39+
fi
40+
echo "hf_token: $HF_TOKEN" >> env.yaml
41+
fi
42+
43+
uv run ng_prepare_data "+config_paths=[resources_servers/workplace_assistant/configs/workplace_assistant.yaml]" \
44+
+output_dirpath=data/workplace_assistant \
45+
+mode=train_preparation \
46+
+should_download=true \
47+
+data_source=huggingface
48+
cd -
49+
50+
# This trimming of the workplace assistant dataset is necessary b/c with all the tools the first prompt is >4000 tokens
51+
# which will cause vllm to return nothing on the first prompt and crash RL. Since we want to keep this test short to
52+
# smoke test, we trim all but the first tool
53+
TRAIN_PATH=$DATA_DIR/workplace_assistant_train.jsonl
54+
VALIDATION_PATH=$DATA_DIR/workplace_assistant_validation.jsonl
55+
jq -c '.responses_create_params.tools |= (.[0:1])' 3rdparty/Gym-workspace/Gym/data/workplace_assistant/train.jsonl > $TRAIN_PATH
56+
jq -c '.responses_create_params.tools |= (.[0:1])' 3rdparty/Gym-workspace/Gym/data/workplace_assistant/validation.jsonl > $VALIDATION_PATH
57+
58+
uv run coverage run -a --data-file=$PROJECT_ROOT/tests/.coverage --source=$PROJECT_ROOT/nemo_rl \
59+
$PROJECT_ROOT/examples/run_grpo_single_controller.py \
60+
--config $PROJECT_ROOT/examples/nemo_gym/grpo_qwen3_30ba3b_instruct.yaml \
61+
policy.model_name=Qwen/Qwen3-0.6B \
62+
policy.dtensor_cfg.enabled=false \
63+
policy.megatron_cfg.enabled=true \
64+
policy.megatron_cfg.tensor_model_parallel_size=1 \
65+
policy.megatron_cfg.pipeline_model_parallel_size=1 \
66+
policy.megatron_cfg.expert_model_parallel_size=1 \
67+
policy.megatron_cfg.context_parallel_size=1 \
68+
policy.megatron_cfg.sequence_parallel=false \
69+
policy.generation.vllm_cfg.tensor_parallel_size=1 \
70+
policy.generation.vllm_cfg.async_engine=true \
71+
policy.max_total_sequence_length=512 \
72+
policy.generation.colocated.enabled=false \
73+
policy.generation.colocated.resources.num_nodes=1 \
74+
policy.generation.colocated.resources.gpus_per_node=1 \
75+
grpo.num_prompts_per_step=4 \
76+
grpo.num_generations_per_prompt=2 \
77+
grpo.max_num_steps=10 \
78+
grpo.val_period=-1 \
79+
grpo.val_at_start=false \
80+
policy.train_global_batch_size=8 \
81+
policy.train_micro_batch_size=1 \
82+
cluster.gpus_per_node=2 \
83+
loss_fn.reference_policy_kl_penalty=0.01 \
84+
grpo.skip_reference_policy_logprobs_calculation=false \
85+
loss_fn.use_importance_sampling_correction=true \
86+
logger.tensorboard_enabled=true \
87+
logger.log_dir=$LOG_DIR \
88+
logger.wandb_enabled=false \
89+
logger.monitor_gpus=true \
90+
checkpointing.enabled=false \
91+
data.train.data_path=$TRAIN_PATH \
92+
data.validation.data_path=$VALIDATION_PATH \
93+
++data_plane.enabled=true \
94+
++data_plane.impl=transfer_queue \
95+
++data_plane.backend=simple \
96+
++data_plane.storage_capacity=1000000 \
97+
++data_plane.num_storage_units=2 \
98+
++data_plane.claim_meta_poll_interval_s=0.5 \
99+
++data_plane.global_segment_size=549755813888 \
100+
++data_plane.local_buffer_size=68719476736 \
101+
++async_rl.sampler.name=in_order \
102+
++async_rl.sampler.max_lookahead_versions=0 \
103+
++async_rl.min_groups_for_streaming_train=4 \
104+
++async_rl.max_inflight_prompts=4 \
105+
++async_rl.max_buffered_rollouts=4 \
106+
$@ \
107+
2>&1 | tee $RUN_LOG
108+
109+
uv run tests/json_dump_tb_logs.py $LOG_DIR --output_path $JSON_METRICS
110+
111+
# Observed to be between 0.8-1.3
112+
uv run tests/check_metrics.py $JSON_METRICS \
113+
'median(data["train/gen_kl_error"]) < 1.3' \
114+
'max(data["train/reward"]) > 0'

tests/unit/single_controller/test_run_grpo_single_controller.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@ def main_context(monkeypatch: pytest.MonkeyPatch) -> SimpleNamespace:
3131
"draft": {"enabled": False},
3232
"megatron_cfg": {"mtp_num_layers": 2},
3333
},
34+
env={},
3435
data_plane={"enabled": True},
3536
logger={"log_dir": "/tmp/logs"},
3637
checkpointing={"enabled": False},

0 commit comments

Comments
 (0)