diff --git a/tests/cli/base_rl_pipeline_test.py b/tests/cli/base_rl_pipeline_test.py index 810b7975b..485c06e1e 100644 --- a/tests/cli/base_rl_pipeline_test.py +++ b/tests/cli/base_rl_pipeline_test.py @@ -75,7 +75,7 @@ def _create_agentic_dummy_config(self): ) # Strip helper keys that are not GRPOConfig fields - valid = {f.name for f in dataclasses.fields(DummyConfig)} + valid = {f.name for f in dataclasses.fields(DummyConfig) if f.init} cfg.pop("max_turns", None) return DummyConfig(**{k: v for k, v in cfg.items() if k in valid}) diff --git a/tests/experimental/orchestrator/algorithm_adapter_test.py b/tests/experimental/orchestrator/algorithm_adapter_test.py index 0fb92a51d..7d0c2f226 100644 --- a/tests/experimental/orchestrator/algorithm_adapter_test.py +++ b/tests/experimental/orchestrator/algorithm_adapter_test.py @@ -237,11 +237,11 @@ def test_grpo_build_gen_model_input_fn(self): num_generations=4, epsilon=0.25, beta=0.05, - temperature=0.8, loss_agg_mode="token-mean", kl_loss_mode="kld", kl_clamp_value=1.5, ) + algo_config.temperature = 0.8 adapter = algorithm_adapter.GRPOAdapter(algo_config=algo_config) gen_fn = adapter.build_gen_model_input_fn(pad_id=10, eos_id=20) self.assertTrue(callable(gen_fn)) @@ -263,13 +263,20 @@ def test_grpo_build_gen_model_input_fn(self): self.assertEqual(cfg.kl_loss_mode, "kld") self.assertEqual(cfg.kl_clamp_value, 1.5) + def test_grpo_config_rejects_temperature_init_kwarg(self): + with self.assertRaises(TypeError): + algorithm_config.GRPOConfig( # pyrefly: ignore[unexpected-keyword] + num_generations=2, + temperature=0.8, + ) + def test_grpo_build_gen_model_input_fn_fails_without_temperature(self): """Verifies build_gen_model_input_fn raises ValueError if temperature is unset.""" adapter = algorithm_adapter.GRPOAdapter( algo_config=algorithm_config.GRPOConfig(num_generations=2) ) with self.assertRaisesRegex( - ValueError, "Trainer temperature must be explicitly set" + ValueError, "Trainer temperature is unset on algo_config" ): adapter.build_gen_model_input_fn(pad_id=0, eos_id=1) @@ -278,12 +285,12 @@ def test_grpo_custom_algo_config(self): num_generations=4, epsilon=0.2, epsilon_high=0.3, - temperature=1.0, loss_algo="gspo-token", policy_loss_fn="grpo", advantage_estimator="drgrpo", kl_loss_mode="mse_kl", ) + config.temperature = 1.0 adapter = algorithm_adapter.GRPOAdapter(algo_config=config) self.assertEqual(adapter.algo_config.kl_loss_mode, "mse_kl") self.assertEqual(adapter.algo_config.epsilon_high, 0.3) @@ -519,11 +526,11 @@ def test_grpo_wraps_canonical_config(self): num_generations=4, beta=0.03, epsilon=0.15, - temperature=1.0, loss_agg_mode="token-mean", kl_loss_mode="low_var_kl", kl_clamp_value=10.0, ) + canonical_config.temperature = 1.0 adapter = algorithm_adapter.GRPOAdapter( algo_config=canonical_config, ) @@ -558,8 +565,8 @@ def test_grpo_asymmetric_and_dual_clipping(self): epsilon=0.2, epsilon_high=0.28, epsilon_c=0.1, - temperature=1.0, ) + config.temperature = 1.0 adapter = algorithm_adapter.GRPOAdapter(algo_config=config) self.assertEqual(adapter.algo_config.epsilon, 0.2) self.assertEqual(adapter.algo_config.epsilon_high, 0.28) @@ -575,9 +582,9 @@ def test_grpo_algo_config_is_single_source_of_truth(self): """Verifies algo_config carries temperature/use_rollout_logps unmodified.""" config = algorithm_config.GRPOConfig( num_generations=2, - temperature=0.8, use_rollout_logps=False, ) + config.temperature = 0.8 adapter = algorithm_adapter.GRPOAdapter(algo_config=config) # The adapter wraps the caller's config as-is, without copying or mutating. self.assertIs(adapter.algo_config, config) diff --git a/tests/experimental/orchestrator/rl_program_test.py b/tests/experimental/orchestrator/rl_program_test.py index 97c118c14..abd3a9f2d 100644 --- a/tests/experimental/orchestrator/rl_program_test.py +++ b/tests/experimental/orchestrator/rl_program_test.py @@ -22,6 +22,7 @@ import weakref from absl.testing import absltest +from flax import nnx import metrax.logging as metrax_logging import numpy as np from tunix.experimental.common import datatypes @@ -32,9 +33,13 @@ from tunix.experimental.orchestrator import rl_program from tunix.experimental.trajectory import in_memory_store from tunix.experimental.trajectory import trajectory as trajectory_lib +from tunix.experimental.worker import inference_worker as exp_inference_worker from tunix.experimental.worker import remote_execution +from tunix.rl import algorithm_config +from tunix.rl.inference import inference_worker as rl_inference_worker from tunix.sft import metrics_logger as metrics_logger_lib from tunix.sft import utils as sft_utils +from tunix.tests import test_common def _padding_stats( @@ -362,7 +367,7 @@ def test_unexpected_mini_batch_size_argument_raises_type_error(self): def test_program_use_rollout_logps_matching(self): self.mock_algo.algo_config = types.SimpleNamespace( - temperature=0.7, + temperature=None, use_rollout_logps=True, ) program = rl_program.StandardRLProgram( @@ -380,7 +385,7 @@ def test_program_use_rollout_logps_missing_in_generation_args_inherits_from_algo self, ): self.mock_algo.algo_config = types.SimpleNamespace( - temperature=0.7, + temperature=None, use_rollout_logps=True, ) program = rl_program.StandardRLProgram( @@ -872,8 +877,8 @@ async def _run(): sync_weights=False, ) - async def _fake_per_token_logps(role, items): - del role + async def _fake_per_token_logps(role, items, **kwargs): + del role, kwargs mb_id = int(items.prompt_ids[0, 0]) if mb_id == 2: # mb0 train_step must already be in flight when mb1 is packed & @@ -2140,12 +2145,182 @@ async def _run(): await program.run_async(self.mock_engine) self.mock_engine.per_token_logps.assert_called_once_with( - datatypes.Role.REFERENCE, items=mock_payload + datatypes.Role.REFERENCE, items=mock_payload, temperature=None ) self.assertEqual(program.step, 1) asyncio.run(_run()) + def test_reference_kl_logprobs_forwards_temperature_in_train_stage(self): + async def _run(): + self.mock_algo.requires_reference_kl = True + mock_payload = datatypes.RLTrainerPayload( + prompt_ids=np.array([[1, 2]], dtype=np.int32), + prompt_mask=np.ones((1, 2), dtype=np.float32), + completion_ids=np.array([[3, 4]], dtype=np.int32), + completion_mask=np.ones((1, 2), dtype=np.float32), + advantages=np.ones((1, 2), dtype=np.float32), + ref_per_token_logps=None, + old_per_token_logps=None, + ) + self.assembler.feed = mock.MagicMock( + side_effect=lambda _: iter([ + batch_assembly.AssembledBatch( + payload=mock_payload, + is_final_batch=True, + padding_stats=_padding_stats(), + trajectory_ids=(), + ) + ]) + ) + self.mock_engine.per_token_logps = mock.AsyncMock( + return_value=np.array([[-0.1, -0.2]], dtype=np.float32) + ) + + _set_mock_poll_batches(self.mock_engine, _make_trajectory_group()) + program = self._create_program( + dataset=["prompt_0"], + generation_args=datatypes.GenerationArgs(temperature=0.6), + ) + + await program.run_async(self.mock_engine) + + self.mock_engine.per_token_logps.assert_called_once_with( + datatypes.Role.REFERENCE, items=mock_payload, temperature=0.6 + ) + self.assertEqual(program.step, 1) + + asyncio.run(_run()) + + def test_reference_kl_logprobs_temperature_matches_actor_kl_term(self): + async def _run(): + config = test_common.ModelConfig(vocab_size=32, num_layers=2) + actor_model = test_common.ToyTransformer(config=config, rngs=nnx.Rngs(42)) + ref_model = test_common.ToyTransformer(config=config, rngs=nnx.Rngs(42)) + + ref_worker = exp_inference_worker.InferenceWorker( + rl_inference_worker.InferenceWorker({"reference": ref_model}), + worker_id="ref_0", + pad_id=0, + eos_id=2, + max_prompt_length=4, + max_response_length=4, + temperature=1.0, + ) + ref_worker.initialize() + ref_worker.start() + + algo_config = algorithm_config.GRPOConfig( + num_generations=2, + num_iterations=1, + beta=0.04, + kl_loss_mode="low_var_kl", + use_rollout_logps=False, + ) + algo = algorithm_adapter.GRPOAdapter( + algo_config=algo_config, + mini_batch_size=1, + train_micro_batch_size=2, + max_packed_len=8, + max_response_length=4, + ) + + async def _score_ref(role, items, temperature=None): + self.assertEqual(role, datatypes.Role.REFERENCE) + return ref_worker.per_token_logps(items, temperature=temperature) + + trained_batches: list[datatypes.RLTrainerPayload] = [] + + async def _capture_train_step(batch, **kwargs): + del kwargs + trained_batches.append(batch) + return "step_done" + + self.mock_engine.per_token_logps = mock.AsyncMock(side_effect=_score_ref) + self.mock_engine.train_step = mock.AsyncMock( + side_effect=_capture_train_step + ) + group = [ + datatypes.TrajectoryItem( + prompt_id="prompt_0", + group_index=0, + start_step=0, + traj={ + "trajectory_reward": 1.0, + "status": datatypes.TrajectoryStatus.SUCCEEDED, + "prompt_tokens": np.array([3, 4], dtype=np.int32), + "conversation_tokens": np.array([5, 6, 7], dtype=np.int32), + "conversation_masks": np.ones(3, dtype=np.float32), + }, + prompt_tokens=np.array([3, 4], dtype=np.int32), + completion_tokens=np.array([5, 6, 7], dtype=np.int32), + action_mask=np.ones(3, dtype=np.float32), + policy_version=0, + ), + datatypes.TrajectoryItem( + prompt_id="prompt_0", + group_index=1, + start_step=0, + traj={ + "trajectory_reward": 0.0, + "status": datatypes.TrajectoryStatus.SUCCEEDED, + "prompt_tokens": np.array([3, 4], dtype=np.int32), + "conversation_tokens": np.array([8, 9, 10], dtype=np.int32), + "conversation_masks": np.ones(3, dtype=np.float32), + }, + prompt_tokens=np.array([3, 4], dtype=np.int32), + completion_tokens=np.array([8, 9, 10], dtype=np.int32), + action_mask=np.ones(3, dtype=np.float32), + policy_version=0, + ), + ] + _set_mock_poll_batches(self.mock_engine, group) + + program = rl_program.StandardRLProgram( + algo=algo, + dataset=["prompt_0"], + max_steps=1, + generation_args=datatypes.GenerationArgs(temperature=0.7), + batch_size=1, + batch_config=batch_assembly.BatchConfig( + pad_id=0, + max_prompt_length=4, + max_response_length=4, + ), + sync_weights=False, + ) + await program.run_async(self.mock_engine) + program.close() + + self.assertLen(trained_batches, 1) + loss_fn = algo.loss_fn() + gen_input_fn = algo.build_gen_model_input_fn(pad_id=0, eos_id=2) + + # Post-fix: train_stage forwards temperature=0.7 to Role.REFERENCE, so + # identical actor & reference weights yield zero KL divergence. + loss_out_matched = loss_fn( + actor_model, **gen_input_fn(trained_batches[0]) + ) + kl_matched = float(loss_out_matched.aux_metrics["kl"].compute()) + kl_loss_matched = float(loss_out_matched.aux_metrics["kl_loss"].compute()) + self.assertAlmostEqual(kl_matched, 0.0, places=5) + self.assertAlmostEqual(kl_loss_matched, 0.0, places=5) + + # Pre-fix regression check: omitting temperature on Role.REFERENCE falls + # back to ref_worker._temperature (1.0), producing a spurious positive KL + # even when actor and reference weights are identical. + unmatched_ref_logps = ref_worker.per_token_logps( + trained_batches[0], temperature=None + ) + unmatched_batch = batch_assembly.with_ref_per_token_logps( + trained_batches[0], unmatched_ref_logps + ) + loss_out_unmatched = loss_fn(actor_model, **gen_input_fn(unmatched_batch)) + kl_unmatched = float(loss_out_unmatched.aux_metrics["kl"].compute()) + self.assertGreater(kl_unmatched, 1e-2) + + asyncio.run(_run()) + def test_reference_kl_raises_type_error_for_invalid_microbatch(self): async def _run(): self.mock_algo.requires_reference_kl = True @@ -3304,7 +3479,7 @@ async def _run(): asyncio.run(_run()) - def test_program_temperature_matching_sets_algo_config(self): + def test_program_raises_if_algo_config_temperature_is_set(self): mock_algo = mock.MagicMock(spec=algorithm_adapter.AlgorithmAdapter) mock_algo.num_generations = 2 mock_algo.mini_batch_size = 1 @@ -3313,16 +3488,19 @@ def test_program_temperature_matching_sets_algo_config(self): mock_algo.max_response_length = 1024 mock_algo.requires_reference_kl = False mock_algo.algo_config = mock.MagicMock( - temperature=0.8, use_rollout_logps=None + temperature=0.8, use_rollout_logps=False ) gen_args = datatypes.GenerationArgs(temperature=0.8) - rl_program.StandardRLProgram( - dataset=("p0",), - algo=mock_algo, - generation_args=gen_args, - ) - self.assertEqual(mock_algo.algo_config.temperature, 0.8) + with self.assertRaisesRegex( + ValueError, + "Do not set temperature on AlgorithmConfig", + ): + rl_program.StandardRLProgram( + dataset=("p0",), + algo=mock_algo, + generation_args=gen_args, + ) def test_program_temperature_missing_in_generation_args_leaves_none( self, @@ -3335,7 +3513,7 @@ def test_program_temperature_missing_in_generation_args_leaves_none( mock_algo.max_response_length = 1024 mock_algo.requires_reference_kl = False mock_algo.algo_config = mock.MagicMock( - temperature=0.8, use_rollout_logps=None + temperature=None, use_rollout_logps=False ) program = rl_program.StandardRLProgram( @@ -3343,6 +3521,7 @@ def test_program_temperature_missing_in_generation_args_leaves_none( algo=mock_algo, ) self.assertIsNone(program.generation_args.temperature) + self.assertIsNone(mock_algo.algo_config.temperature) def test_program_temperature_missing_in_algo_config_propagates_from_generation_args( self, @@ -3355,7 +3534,7 @@ def test_program_temperature_missing_in_algo_config_propagates_from_generation_a mock_algo.max_response_length = 1024 mock_algo.requires_reference_kl = False mock_algo.algo_config = mock.MagicMock( - temperature=None, use_rollout_logps=None + temperature=None, use_rollout_logps=False ) gen_args = datatypes.GenerationArgs(temperature=0.8) @@ -4931,6 +5110,10 @@ def setUp(self): self.mock_algo.mini_batch_size = 1 self.mock_algo.max_packed_len = 16 self.mock_algo.max_response_length = 1024 + self.mock_algo.algo_config = types.SimpleNamespace( + temperature=None, + use_rollout_logps=True, + ) self.assembler = batch_assembly.SequencePackedBatchAssembler( batch_size=1, num_generations=2, diff --git a/tests/experimental/worker/trainer_worker_test.py b/tests/experimental/worker/trainer_worker_test.py index 61029e57a..29adc206f 100644 --- a/tests/experimental/worker/trainer_worker_test.py +++ b/tests/experimental/worker/trainer_worker_test.py @@ -207,7 +207,8 @@ def test_with_loss_fn_chunked_grpo_loss_matches_unchunked(self): completion_mask=np.ones((batch, completion_len), dtype=np.int32), advantages=np.array([1.0, -1.0], dtype=np.float32), ) - algo_config = algorithm_config.GRPOConfig(beta=0.0, temperature=1.0) + algo_config = algorithm_config.GRPOConfig(beta=0.0) + algo_config.temperature = 1.0 worker = trainer_worker.TrainerWorker( trainer_factory=lambda: self.fake_trainer, logps_chunk_size=3 diff --git a/tunix/cli/grpo_main.py b/tunix/cli/grpo_main.py index c835eeba7..3227cb846 100644 --- a/tunix/cli/grpo_main.py +++ b/tunix/cli/grpo_main.py @@ -93,7 +93,7 @@ def _create_agentic_grpo_config(self): ) # Strip helper keys that are not GRPOConfig fields - valid = {f.name for f in dataclasses.fields(GRPOConfig)} + valid = {f.name for f in dataclasses.fields(GRPOConfig) if f.init} cfg.pop("max_turns", None) return GRPOConfig(**{k: v for k, v in cfg.items() if k in valid}) diff --git a/tunix/experimental/examples/deepswe_dist/run_deepswe_dist.py b/tunix/experimental/examples/deepswe_dist/run_deepswe_dist.py index fc1e45c3a..c5a9d9516 100644 --- a/tunix/experimental/examples/deepswe_dist/run_deepswe_dist.py +++ b/tunix/experimental/examples/deepswe_dist/run_deepswe_dist.py @@ -226,7 +226,6 @@ def _build_algo(args: argparse.Namespace) -> algorithm_adapter.GRPOAdapter: num_generations=args.num_generations, epsilon=args.epsilon, beta=args.beta, - temperature=args.temperature, use_rollout_logps=args.use_rollout_logps, ) return algorithm_adapter.GRPOAdapter( @@ -364,12 +363,6 @@ def main(argv: list[str], context: ProcessContext | None = None) -> None: trainer_handles = cluster.worker_handles(datatypes.Role.ACTOR) if len(trainer_handles) != 1: raise ValueError(f"Expected 1 trainer worker, got {len(trainer_handles)}.") - _configure_trainer_loss( - trainer_handles[0], - algo=algo, - pad_id=pad_id, - eos_id=eos_id, - ) metrics_logging_options = metrics_logger_lib.MetricsLoggerOptions( log_dir=args.log_dir, @@ -432,6 +425,12 @@ def main(argv: list[str], context: ProcessContext | None = None) -> None: result, ), ) + _configure_trainer_loss( + trainer_handles[0], + algo=algo, + pad_id=pad_id, + eos_id=eos_id, + ) try: logging.info("Bringing up remote workers through ClusterOrchestrator...") diff --git a/tunix/experimental/examples/frozenlake_dist/run_frozenlake_dist.py b/tunix/experimental/examples/frozenlake_dist/run_frozenlake_dist.py index a9921dec0..189429e57 100644 --- a/tunix/experimental/examples/frozenlake_dist/run_frozenlake_dist.py +++ b/tunix/experimental/examples/frozenlake_dist/run_frozenlake_dist.py @@ -227,7 +227,6 @@ def _build_algo(args: argparse.Namespace) -> algorithm_adapter.GRPOAdapter: epsilon=args.epsilon, epsilon_high=args.epsilon_high, beta=args.beta, - temperature=args.temperature, loss_algo=args.loss_algo, policy_loss_fn="grpo", advantage_estimator=args.advantage_estimator, @@ -340,9 +339,6 @@ def main(argv: list[str], context: ProcessContext | None = None) -> None: trainer_handles = cluster.worker_handles(datatypes.Role.ACTOR) if len(trainer_handles) != 1: raise ValueError(f"Expected 1 trainer worker, got {len(trainer_handles)}.") - _configure_trainer_loss( - trainer_handles[0], algo=algo, pad_id=pad_id, eos_id=eos_id - ) metrics_options = metrics_logger_lib.MetricsLoggerOptions( log_dir=args.log_dir, @@ -395,6 +391,9 @@ def main(argv: list[str], context: ProcessContext | None = None) -> None: "<<< FrozenLake step %d finished | %s", step, result ), ) + _configure_trainer_loss( + trainer_handles[0], algo=algo, pad_id=pad_id, eos_id=eos_id + ) try: cluster.bring_up_workers(dummy_data=None) diff --git a/tunix/experimental/examples/math_gsm8k_dist/run_gsm8k_dist_grpo.py b/tunix/experimental/examples/math_gsm8k_dist/run_gsm8k_dist_grpo.py index 06fd378d3..94c6446fe 100644 --- a/tunix/experimental/examples/math_gsm8k_dist/run_gsm8k_dist_grpo.py +++ b/tunix/experimental/examples/math_gsm8k_dist/run_gsm8k_dist_grpo.py @@ -252,7 +252,6 @@ def _build_algo(args: argparse.Namespace) -> algorithm_adapter.GRPOAdapter: num_generations=args.num_generations, epsilon=args.epsilon, beta=args.beta, - temperature=args.temperature, use_rollout_logps=args.use_rollout_logps, ) return algorithm_adapter.GRPOAdapter( diff --git a/tunix/experimental/orchestrator/algorithm_adapter.py b/tunix/experimental/orchestrator/algorithm_adapter.py index 6100c6e0a..07e11290f 100644 --- a/tunix/experimental/orchestrator/algorithm_adapter.py +++ b/tunix/experimental/orchestrator/algorithm_adapter.py @@ -336,9 +336,10 @@ def build_gen_model_input_fn( or self.algo_config.temperature is None ): raise ValueError( - "Trainer temperature must be explicitly set on algo_config to match" - " rollout generation temperature. Running with an unset temperature" - " biases policy gradient importance ratios." + "Trainer temperature is unset on algo_config. Configure temperature" + " via generation_args on StandardRLProgram (which propagates it to" + " algo_config.temperature) before building the trainer model input" + " fn." ) return functools.partial( _algo_model_input, diff --git a/tunix/experimental/orchestrator/rl_program.py b/tunix/experimental/orchestrator/rl_program.py index d38507bbc..59e884c89 100644 --- a/tunix/experimental/orchestrator/rl_program.py +++ b/tunix/experimental/orchestrator/rl_program.py @@ -304,6 +304,11 @@ def __init__( self.max_response_length = algo_max_response_length self.generation_args = generation_args or datatypes.GenerationArgs() + if self.algo.algo_config.temperature is not None: + raise ValueError( + "Do not set temperature on AlgorithmConfig; configure it via" + " generation_args on StandardRLProgram." + ) gen_temp = self.generation_args.temperature if gen_temp is not None: self.algo.algo_config.temperature = gen_temp @@ -1093,7 +1098,7 @@ async def _apply_sampler_trainer_agreement( IS/RS weights and trainer logps when configured. """ assert self.engine is not None - gen_temp = getattr(self.generation_args, "temperature", None) + gen_temp = self.generation_args.temperature logps_req = datatypes.LogprobsRequest( prompt_tokens=batch.prompt_ids, completion_tokens=batch.completion_ids, @@ -1303,7 +1308,9 @@ async def _maybe_save_checkpoint() -> None: f"{type(batch).__name__}." ) ref_logps = await self.engine.per_token_logps( - datatypes.Role.REFERENCE, items=batch + datatypes.Role.REFERENCE, + items=batch, + temperature=self.generation_args.temperature, ) batch = batch_assembly.with_ref_per_token_logps(batch, ref_logps) algo_config = getattr(self.algo, "algo_config", None) diff --git a/tunix/experimental/orchestrator/simple_orchestrator_nb.py b/tunix/experimental/orchestrator/simple_orchestrator_nb.py index fc4665fbe..5961a1702 100644 --- a/tunix/experimental/orchestrator/simple_orchestrator_nb.py +++ b/tunix/experimental/orchestrator/simple_orchestrator_nb.py @@ -151,7 +151,6 @@ def main(): algo_config = algorithm_config.GRPOConfig( num_generations=2, - temperature=1.0, ) algo = algorithm_adapter.GRPOAdapter( algo_config=algo_config, @@ -173,6 +172,7 @@ def main(): dataset=train_dataset, reward_fns=[lambda x: 1.0], assembler=assembler, + generation_args=datatypes.GenerationArgs(temperature=1.0), max_steps=2, ) orch.run(program=program) diff --git a/tunix/rl/algorithm_config.py b/tunix/rl/algorithm_config.py index 242deb577..3c0385b7d 100644 --- a/tunix/rl/algorithm_config.py +++ b/tunix/rl/algorithm_config.py @@ -61,7 +61,7 @@ class AlgorithmConfig: # probabilities and entropy in the loss function. # NB: This should not be configured manually, instead it will be set by the RL # engine based on the rollout config. - temperature: float | None = None + temperature: float | None = dataclasses.field(default=None, init=False) # Whether to use rollout-side log probabilities as old-policy log # probabilities. If False, recompute old-policy log probabilities on the # trainer actor. @@ -79,6 +79,7 @@ class AlgorithmConfig: seq_logprob_error_threshold: float | None = None def __post_init__(self): + self.temperature = None valid_algo_variants = [ "grpo", "drgrpo",