Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
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
2 changes: 1 addition & 1 deletion tests/cli/base_rl_pipeline_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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})
Expand Down
19 changes: 13 additions & 6 deletions tests/experimental/orchestrator/algorithm_adapter_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand All @@ -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)

Expand All @@ -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)
Expand Down Expand Up @@ -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,
)
Expand Down Expand Up @@ -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)
Expand All @@ -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)
Expand Down
79 changes: 64 additions & 15 deletions tests/experimental/orchestrator/rl_program_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -361,7 +361,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(
Expand All @@ -379,7 +379,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(
Expand Down Expand Up @@ -871,8 +871,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 &
Expand Down Expand Up @@ -2139,7 +2139,48 @@ 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(
return_value=[
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)

Expand Down Expand Up @@ -3303,7 +3344,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
Expand All @@ -3312,16 +3353,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,
Expand All @@ -3334,14 +3378,15 @@ 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(
dataset=("p0",),
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,
Expand All @@ -3354,7 +3399,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)
Expand Down Expand Up @@ -4930,6 +4975,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,
Expand Down
3 changes: 2 additions & 1 deletion tests/experimental/worker/trainer_worker_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -141,7 +141,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
Expand Down
2 changes: 1 addition & 1 deletion tunix/cli/grpo_main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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})

Expand Down
13 changes: 6 additions & 7 deletions tunix/experimental/examples/deepswe_dist/run_deepswe_dist.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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...")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
7 changes: 4 additions & 3 deletions tunix/experimental/orchestrator/algorithm_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
11 changes: 9 additions & 2 deletions tunix/experimental/orchestrator/rl_program.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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)
Expand Down
Loading
Loading