From f28ff8a34cf59a8ffa783c365b8b3785a9e1b7ef Mon Sep 17 00:00:00 2001 From: EazyReal <8047065+EazyReal@users.noreply.github.com> Date: Fri, 3 Jul 2026 01:38:23 +0000 Subject: [PATCH] fix: guard zero rollout temperature logprob scaling Greedy rollout can use rollout_temperature=0, but both training-side logprob paths divided logits whenever the temperature differed from 1. This produced infinite logits and NaN gradients. Skip temperature scaling when the value is zero while preserving the existing behavior for positive temperatures. Add CPU regression coverage for both call sites and register it in the PR test matrix. --- .github/workflows/pr-test.yml | 4 + .github/workflows/pr-test.yml.j2 | 1 + slime/backends/megatron_utils/loss.py | 4 +- tests/test_rollout_temperature_zero_guard.py | 85 ++++++++++++++++++++ 4 files changed, 92 insertions(+), 2 deletions(-) create mode 100644 tests/test_rollout_temperature_zero_guard.py diff --git a/.github/workflows/pr-test.yml b/.github/workflows/pr-test.yml index 761496b949..d8937fc83b 100644 --- a/.github/workflows/pr-test.yml +++ b/.github/workflows/pr-test.yml @@ -593,6 +593,10 @@ jobs: "num_gpus": 0, "test_file": "test_value_temperature.py" }, + { + "num_gpus": 0, + "test_file": "test_rollout_temperature_zero_guard.py" + }, { "num_gpus": 0, "test_file": "test_cispo_loss.py" diff --git a/.github/workflows/pr-test.yml.j2 b/.github/workflows/pr-test.yml.j2 index 415586cb84..114beab16f 100644 --- a/.github/workflows/pr-test.yml.j2 +++ b/.github/workflows/pr-test.yml.j2 @@ -73,6 +73,7 @@ {'test_file': 'test_loss_cp_invariance.py', 'num_gpus': 0}, {'test_file': 'test_logprob_response_spans.py', 'num_gpus': 0}, {'test_file': 'test_value_temperature.py', 'num_gpus': 0}, + {'test_file': 'test_rollout_temperature_zero_guard.py', 'num_gpus': 0}, {'test_file': 'test_cispo_loss.py', 'num_gpus': 0}, {'test_file': 'test_ppo_logprob_entropy.py', 'num_gpus': 0}, {'test_file': 'test_rm_f1.py', 'num_gpus': 0}, diff --git a/slime/backends/megatron_utils/loss.py b/slime/backends/megatron_utils/loss.py index adfa5d67e6..1d324f9911 100644 --- a/slime/backends/megatron_utils/loss.py +++ b/slime/backends/megatron_utils/loss.py @@ -87,7 +87,7 @@ def get_responses( assert logits.size(0) == 1, f"{logits.shape}" logits = logits.squeeze(0) - if apply_temperature and args.rollout_temperature != 1.0: + if apply_temperature and args.rollout_temperature > 0 and args.rollout_temperature != 1.0: logits = logits.div(args.rollout_temperature) cp_size = mpu.get_context_parallel_world_size() @@ -496,7 +496,7 @@ def get_log_probs_and_entropy( # Apply rollout temperature scaling to logits to match rollout-time log-probs. rollout_temperature = getattr(args, "rollout_temperature", 1.0) - if rollout_temperature != 1.0: + if rollout_temperature > 0 and rollout_temperature != 1.0: logits = logits / rollout_temperature logits = logits.contiguous() T = logits.size(0) diff --git a/tests/test_rollout_temperature_zero_guard.py b/tests/test_rollout_temperature_zero_guard.py new file mode 100644 index 0000000000..2252d18980 --- /dev/null +++ b/tests/test_rollout_temperature_zero_guard.py @@ -0,0 +1,85 @@ +import sys +import types +from argparse import Namespace + +import pytest +import torch + + +NUM_GPUS = 0 + + +def _install_mpu_stub(monkeypatch): + """Install a single-rank megatron.core.mpu stub and return the loss module.""" + sys.modules.pop("slime.backends.megatron_utils.loss", None) + sys.modules.pop("slime.backends.megatron_utils.cp_utils", None) + + mpu_stub = types.SimpleNamespace( + get_context_parallel_world_size=lambda: 1, + get_context_parallel_rank=lambda: 0, + get_tensor_model_parallel_group=lambda: None, + ) + megatron_mod = types.ModuleType("megatron") + core_mod = types.ModuleType("megatron.core") + core_mod.mpu = mpu_stub + monkeypatch.setitem(sys.modules, "megatron", megatron_mod) + monkeypatch.setitem(sys.modules, "megatron.core", core_mod) + + import slime.backends.megatron_utils.loss as loss + + return loss + + +def test_get_responses_zero_temperature_stays_finite(monkeypatch): + loss = _install_mpu_stub(monkeypatch) + + args = Namespace(rollout_temperature=0.0, allgather_cp=False, true_on_policy_mode=False) + logits = torch.tensor([[[0.1, 0.2], [0.3, 0.4], [0.5, 0.6], [0.7, 0.8]]], dtype=torch.float32) + tokens = [torch.tensor([10, 11, 12, 13], dtype=torch.long)] + + chunks = [ + chunk + for chunk, _ in loss.get_responses( + logits, + args=args, + unconcat_tokens=tokens, + total_lengths=[4], + response_lengths=[2], + ) + ] + + assert len(chunks) == 1 + out = chunks[0] + torch.testing.assert_close(out, logits.squeeze(0)[1:3]) + + +def test_get_log_probs_zero_temperature_stays_finite(monkeypatch): + loss = _install_mpu_stub(monkeypatch) + + captured = {} + + def fake_calculate(scaled_logits, tokens, tp_group, **kwargs): + captured["logits"] = scaled_logits + T = scaled_logits.size(0) + return scaled_logits.new_zeros((T, 1)), None + + monkeypatch.setattr(loss, "calculate_log_probs_and_entropy", fake_calculate) + + args = Namespace(rollout_temperature=0.0, allgather_cp=False, log_probs_chunk_size=-1) + logits = torch.tensor([[[0.1, 0.2], [0.3, 0.4], [0.5, 0.6], [0.7, 0.8]]], dtype=torch.float32) + tokens = [torch.tensor([10, 11, 12, 13], dtype=torch.long)] + + loss.get_log_probs_and_entropy( + logits, + args=args, + unconcat_tokens=tokens, + total_lengths=[4], + response_lengths=[2], + ) + + scaled = captured["logits"] + torch.testing.assert_close(scaled, logits.squeeze(0)) + + +if __name__ == "__main__": + raise SystemExit(pytest.main([__file__]))