Skip to content

Commit cba2389

Browse files
committed
test(sc): cover nemo-gym env handle wiring and non-vllm backend guard
Signed-off-by: Yuki Huang <yukih@nvidia.com>
1 parent d435d32 commit cba2389

1 file changed

Lines changed: 42 additions & 0 deletions

File tree

tests/unit/single_controller/test_single_controller_setup.py

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -306,3 +306,45 @@ def test_megatron_train_iters_not_set_when_disabled(self, patched_factories):
306306
setup_single_controller(mc, MagicMock(pad_token_id=0))
307307

308308
assert "train_iters" not in mc.policy.get("megatron_cfg", {})
309+
310+
def test_nemo_gym_wires_env_handle(self, patched_factories):
311+
"""When _should_use_nemo_gym is True the nemo-gym actor is spun up and stored."""
312+
mc = _make_master_config(colocated=True, backend="vllm")
313+
mc.policy["generation"]["model_name"] = "test-model"
314+
mc.policy["generation"]["stop_strings"] = None
315+
mc.policy["generation"]["stop_token_ids"] = None
316+
mc.policy["generation"]["top_k"] = None
317+
patched_factories["setup_response_data"].return_value = (
318+
list(range(8)),
319+
None,
320+
)
321+
fake_gym_actor = MagicMock(name="nemo_gym_actor")
322+
323+
with (
324+
patch.object(sc_setup_mod, "_should_use_nemo_gym", return_value=True),
325+
patch.object(
326+
sc_setup_mod, "spinup_nemo_gym_actor", return_value=fake_gym_actor
327+
) as mock_spinup,
328+
patch.object(sc_setup_mod, "router_replay_enabled", return_value=False),
329+
):
330+
actor_args = setup_single_controller(mc, MagicMock(pad_token_id=0))
331+
332+
mock_spinup.assert_called_once()
333+
assert actor_args.env_handles["nemo_gym"] is fake_gym_actor
334+
335+
@pytest.mark.parametrize("backend", ["sglang", "megatron"])
336+
def test_nemo_gym_rejects_non_vllm_backend(self, patched_factories, backend):
337+
"""SC nemo-gym wiring only supports vLLM; every other backend must raise."""
338+
mc = _make_master_config(colocated=True, backend=backend)
339+
patched_factories["setup_response_data"].return_value = (
340+
list(range(8)),
341+
None,
342+
)
343+
344+
with (
345+
patch.object(sc_setup_mod, "_should_use_nemo_gym", return_value=True),
346+
patch.object(sc_setup_mod, "spinup_nemo_gym_actor") as mock_spinup,
347+
pytest.raises(NotImplementedError, match="vllm"),
348+
):
349+
setup_single_controller(mc, MagicMock(pad_token_id=0))
350+
mock_spinup.assert_not_called()

0 commit comments

Comments
 (0)