diff --git a/nemo_rl/algorithms/single_controller.py b/nemo_rl/algorithms/single_controller.py index 9effa5168ca..9dae17720d5 100644 --- a/nemo_rl/algorithms/single_controller.py +++ b/nemo_rl/algorithms/single_controller.py @@ -124,6 +124,7 @@ def __init__( # Built here, not on the driver: Logger backends (wandb/tb/...) hold # _thread.lock that Ray can't cloudpickle into the actor. self._logger = Logger(master_config.logger) # type: ignore + self._logger.log_hyperparams(master_config.model_dump()) self._timer = Timer() # Pin clusters so RayVirtualCluster.__del__ doesn't remove the PGs. diff --git a/tests/unit/single_controller/test_rollout_pump.py b/tests/unit/single_controller/test_rollout_pump.py index cfeda7b1f57..262115e2b04 100644 --- a/tests/unit/single_controller/test_rollout_pump.py +++ b/tests/unit/single_controller/test_rollout_pump.py @@ -31,6 +31,7 @@ WindowedSampler, WindowedSamplerConfig, ) +from nemo_rl.algorithms.loss import ClippedPGLossConfig from nemo_rl.algorithms.single_controller import SingleControllerActor from nemo_rl.algorithms.single_controller_utils.config import ( AsyncRLConfig, @@ -294,7 +295,7 @@ def test_rollout_pump_writes_expected_tq_data( "max_num_steps": 1, "max_num_epochs": 1, }, - loss_fn=SimpleNamespace(force_on_policy_ratio=False), + loss_fn=ClippedPGLossConfig(force_on_policy_ratio=False), async_rl=AsyncRLConfig( sampler=WindowedSamplerConfig(max_staleness_versions=1), min_groups_for_streaming_train=1, diff --git a/tests/unit/single_controller/test_single_controller.py b/tests/unit/single_controller/test_single_controller.py index 25121a84f9e..279c5301151 100644 --- a/tests/unit/single_controller/test_single_controller.py +++ b/tests/unit/single_controller/test_single_controller.py @@ -22,6 +22,7 @@ import torch import nemo_rl.algorithms.single_controller as single_controller +from nemo_rl.algorithms.loss import ClippedPGLossConfig from nemo_rl.algorithms.single_controller import SingleControllerActor from nemo_rl.algorithms.single_controller_utils.config import ( AdvantageConfig, @@ -77,18 +78,19 @@ def test_rejects_multiple_optimizer_steps_per_rl_step(monkeypatch) -> None: ) -def test_logs_concrete_weight_synchronizer( +def test_logs_hyperparameters_and_concrete_weight_synchronizer( monkeypatch, capsys: pytest.CaptureFixture[str], ) -> None: - monkeypatch.setattr(single_controller, "Logger", lambda _: object()) + logger = MagicMock() + monkeypatch.setattr(single_controller, "Logger", lambda _: logger) master_config = MasterConfig.model_construct( policy={"train_global_batch_size": 8}, grpo={ "num_prompts_per_step": 2, "num_generations_per_prompt": 4, }, - loss_fn=SimpleNamespace(force_on_policy_ratio=False), + loss_fn=ClippedPGLossConfig(force_on_policy_ratio=False), async_rl=AsyncRLConfig( min_groups_for_streaming_train=1, max_buffered_rollouts=4, @@ -116,6 +118,7 @@ def test_logs_concrete_weight_synchronizer( actor_args=actor_args, ) + logger.log_hyperparams.assert_called_once_with(master_config.model_dump()) output = capsys.readouterr().out assert "weight_sync=FakeWeightSynchronizer" in output assert "transport=stub" not in output diff --git a/tests/unit/single_controller/test_train_pump.py b/tests/unit/single_controller/test_train_pump.py index fe2b59ebbd0..54d9be6a385 100644 --- a/tests/unit/single_controller/test_train_pump.py +++ b/tests/unit/single_controller/test_train_pump.py @@ -28,6 +28,7 @@ from nemo_rl.algorithms.async_utils.replay_buffer import TQReplayBuffer from nemo_rl.algorithms.async_utils.staleness_sampler import WindowedSamplerConfig +from nemo_rl.algorithms.loss import ClippedPGLossConfig from nemo_rl.algorithms.single_controller import SingleControllerActor from nemo_rl.algorithms.single_controller_utils.config import ( AsyncRLConfig, @@ -305,7 +306,7 @@ def test_train_pump_drives_mcore_training_step( "max_num_steps": train_steps, "max_num_epochs": None, }, - loss_fn=SimpleNamespace(force_on_policy_ratio=False), + loss_fn=ClippedPGLossConfig(force_on_policy_ratio=False), async_rl=AsyncRLConfig( sampler=WindowedSamplerConfig(max_staleness_versions=1), min_groups_for_streaming_train=num_prompts,