diff --git a/tests/utils/test_linear_cross_entropy.py b/tests/utils/test_linear_cross_entropy.py index 801eaff27c5..103aa51799b 100644 --- a/tests/utils/test_linear_cross_entropy.py +++ b/tests/utils/test_linear_cross_entropy.py @@ -395,6 +395,63 @@ def reference(hidden, weight, labels): torch.testing.assert_close(ref_ent, ker_ent, atol=1e-3, rtol=1e-3, msg=f"entropy mismatch: {desc}") +def test_lce_all_negative_logits_are_shift_invariant(): + """Regression test for max reductions incorrectly anchored at zero.""" + if not torch.cuda.is_available() or is_torch_npu_available(check_device=False): + return + + num_tokens = 32 + hidden_size = 128 + vocab_size = 1153 + temperature = 1.0 + + hidden_data = torch.zeros((1, num_tokens, hidden_size), dtype=torch.bfloat16, device="cuda") + hidden_data[..., 0] = 1 + hidden_data[..., 1] = 1 + weight_data = torch.zeros((vocab_size, hidden_size), dtype=torch.bfloat16, device="cuda") + weight_data[:, 0] = -334 + weight_data[:, 1] = torch.linspace(-1, 1, vocab_size, dtype=torch.bfloat16, device="cuda") + labels = (torch.arange(num_tokens, device="cuda") * 37 % vocab_size).unsqueeze(0).contiguous() + + reference_hidden = hidden_data.clone().requires_grad_() + reference_weight = weight_data.clone().requires_grad_() + reference_logprobs, reference_entropy = run_torch_entropy(reference_hidden, reference_weight, labels, temperature) + + kernel_hidden = hidden_data.clone().requires_grad_() + kernel_weight = weight_data.clone().requires_grad_() + kernel_logprobs, kernel_entropy = linear_cross_entropy(kernel_hidden, kernel_weight, labels, temperature) + + assert torch.isfinite(kernel_logprobs).all() + assert torch.isfinite(kernel_entropy).all() + torch.testing.assert_close(kernel_logprobs, reference_logprobs, atol=1e-3, rtol=1e-3) + torch.testing.assert_close(kernel_entropy, reference_entropy, atol=1e-3, rtol=1e-3) + + shifted_weight = weight_data.clone() + shifted_weight[:, 0] += 336 + shifted_logprobs, shifted_entropy = linear_cross_entropy( + hidden_data, shifted_weight.contiguous(), labels, temperature + ) + torch.testing.assert_close(kernel_logprobs, shifted_logprobs, atol=1e-3, rtol=1e-3) + torch.testing.assert_close(kernel_entropy, shifted_entropy, atol=1e-3, rtol=1e-3) + + grad_logprobs = torch.linspace(-0.5, 0.5, num_tokens, device="cuda") + grad_entropy = torch.linspace(0.5, -0.5, num_tokens, device="cuda") + reference_grads = torch.autograd.grad( + (reference_logprobs, reference_entropy), + (reference_hidden, reference_weight), + (grad_logprobs, grad_entropy), + ) + kernel_grads = torch.autograd.grad( + (kernel_logprobs, kernel_entropy), + (kernel_hidden, kernel_weight), + (grad_logprobs, grad_entropy), + ) + + assert all(torch.isfinite(grad).all() for grad in kernel_grads) + for kernel_grad, reference_grad in zip(kernel_grads, reference_grads, strict=True): + torch.testing.assert_close(kernel_grad, reference_grad, atol=2e-2, rtol=4e-2) + + if __name__ == "__main__": # torch.cuda.memory._record_memory_history() @@ -406,5 +463,6 @@ def reference(hidden, weight, labels): test.check_storage_all() test_lce_non_divisible_vocab_padding() + test_lce_all_negative_logits_are_shift_invariant() # torch.cuda.memory._dump_snapshot("test_linear_cross_entropy.pkl") diff --git a/tests/utils/test_special_linear_cross_entropy_tp.py b/tests/utils/test_special_linear_cross_entropy_tp.py index 9c1f868a93e..c2a9563c906 100644 --- a/tests/utils/test_special_linear_cross_entropy_tp.py +++ b/tests/utils/test_special_linear_cross_entropy_tp.py @@ -48,7 +48,7 @@ compute_entropy_from_logits = torch.compile(verl_F.entropy_from_logits, dynamic=True) -MAX_TEST_CASES = os.environ.get("MAX_TEST_CASES", 5) +MAX_TEST_CASES = int(os.environ.get("MAX_TEST_CASES", 5)) VERIFY_TORCH_SELF = os.environ.get("VERIFY_TORCH_SELF", False) LOW_MEMORY = os.environ.get("LOW_MEMORY", False) LOW_MEMORY_DIV_FACTOR = os.environ.get("LOW_MEMORY_DIV_FACTOR", 16) @@ -204,6 +204,70 @@ def cleanup(self): gc.collect() torch.cuda.synchronize() + def verify_all_negative_logits(self): + """Check TP max reductions with the large negative offset seen in ChemGraph.""" + self.cleanup() + num_tokens = 32 + hidden_size = 128 + local_vocab_size = 1153 + temperature = 1.0 + + hidden_data = torch.zeros((1, num_tokens, hidden_size), dtype=torch.bfloat16, device="cuda") + hidden_data[..., 0] = 1 + hidden_data[..., 1] = 1 + dist.broadcast(hidden_data, src=0, group=self.group) + + weight_data = torch.zeros((local_vocab_size, hidden_size), dtype=torch.bfloat16, device="cuda") + weight_data[:, 0] = -334 + weight_data[:, 1] = torch.linspace(-1, 1, local_vocab_size, dtype=torch.bfloat16, device="cuda") + weight_data[:, 1] += 0.25 * self.local_rank + labels = (torch.arange(num_tokens, device="cuda") * 137 % (local_vocab_size * self.world_size)).unsqueeze(0) + labels = labels.contiguous() + dist.broadcast(labels, src=0, group=self.group) + + reference_hidden = hidden_data.clone().requires_grad_() + reference_weight = weight_data.clone().requires_grad_() + reference_logprobs, reference_entropy = run_torch_entropy_tp( + reference_hidden, reference_weight, labels, temperature, self.group + ) + + kernel_hidden = hidden_data.clone().requires_grad_() + kernel_weight = weight_data.clone().requires_grad_() + kernel_logprobs, kernel_entropy = linear_cross_entropy( + kernel_hidden, kernel_weight, labels, temperature, "none", self.group + ) + + assert torch.isfinite(kernel_logprobs).all() + assert torch.isfinite(kernel_entropy).all() + torch.testing.assert_close(kernel_logprobs, reference_logprobs, atol=1e-3, rtol=1e-3) + torch.testing.assert_close(kernel_entropy, reference_entropy, atol=1e-3, rtol=1e-3) + + grad_logprobs = torch.linspace(-0.5, 0.5, num_tokens, device="cuda") + grad_entropy = torch.linspace(0.5, -0.5, num_tokens, device="cuda") + reference_grads = torch.autograd.grad( + (reference_logprobs, reference_entropy), + (reference_hidden, reference_weight), + (grad_logprobs, grad_entropy), + ) + kernel_grads = torch.autograd.grad( + (kernel_logprobs, kernel_entropy), + (kernel_hidden, kernel_weight), + (grad_logprobs, grad_entropy), + ) + + reference_hidden_grad, reference_weight_grad = reference_grads + kernel_hidden_grad, kernel_weight_grad = kernel_grads + dist.all_reduce(reference_hidden_grad, op=dist.ReduceOp.SUM, group=self.group) + dist.all_reduce(kernel_hidden_grad, op=dist.ReduceOp.SUM, group=self.group) + + assert torch.isfinite(kernel_hidden_grad).all() + assert torch.isfinite(kernel_weight_grad).all() + torch.testing.assert_close(kernel_hidden_grad, reference_hidden_grad, atol=2e-2, rtol=4e-2) + torch.testing.assert_close(kernel_weight_grad, reference_weight_grad, atol=2e-2, rtol=4e-2) + + if self.local_rank == 0: + print("[PASS]: TP all-negative-logits regression test.") + def generate_hyper(self): global LOW_MEMORY, LOW_MEMORY_DIV_FACTOR, MAX_TEST_CASES @@ -502,6 +566,7 @@ def check_kernel_storage(self): # set_backward_method(BackwardEnum._Split_Dlogits_N) test = TestLinearCrossEntropy_TensorParallel() + test.verify_all_negative_logits() for test_case_idx in range(MAX_TEST_CASES): print(f"[INFO] Running test case {test_case_idx}") test.initialize(test_case_idx) diff --git a/verl/utils/kernel/kernels.py b/verl/utils/kernel/kernels.py index 3eca28cb6fd..846b6c7c155 100644 --- a/verl/utils/kernel/kernels.py +++ b/verl/utils/kernel/kernels.py @@ -381,14 +381,18 @@ def efficient_entropy_triton_kernel_epilogue( pid_m = tl.program_id(axis=0) offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) - global_max = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32) + global_max = tl.full((BLOCK_SIZE_M,), -float("inf"), dtype=tl.float32) global_accu = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32) global_entropy_b = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32) for pid_n in range(0, tl.cdiv(num_splits, BLOCK_SIZE_N)): offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) max_ptrs = max_ptr + offs_m[:, None] * stride_max_m + offs_n[None, :] * stride_max_n - _max = tl.load(max_ptrs, mask=(offs_m[:, None] < num_tokens) & (offs_n[None, :] < num_splits), other=0.0) + _max = tl.load( + max_ptrs, + mask=(offs_m[:, None] < num_tokens) & (offs_n[None, :] < num_splits), + other=-float("inf"), + ) accu_ptrs = accu_ptr + offs_m[:, None] * stride_accu_m + offs_n[None, :] * stride_accu_n _accu = tl.load(accu_ptrs, mask=(offs_m[:, None] < num_tokens) & (offs_n[None, :] < num_splits), other=0.0) @@ -468,7 +472,7 @@ def efficient_entropy_triton_kernel_epilogue_tp( offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) - global_max = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32) + global_max = tl.full((BLOCK_SIZE_M,), -float("inf"), dtype=tl.float32) global_accu = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32) global_entropy_b = tl.zeros((BLOCK_SIZE_M,), dtype=tl.float32) for pid_n in range(0, tl.cdiv(num_splits, BLOCK_SIZE_N)): @@ -477,12 +481,12 @@ def efficient_entropy_triton_kernel_epilogue_tp( _reduced_max = tl.load( reduced_max_ptr + offs_m[:, None] * stride_reduced_max_m + offs_n[None, :] * stride_reduced_max_n, mask=(offs_m[:, None] < num_tokens) & (offs_n[None, :] < num_splits), - other=0.0, + other=-float("inf"), ) _original_max = tl.load( original_max_ptr + offs_m[:, None] * stride_original_max_m + offs_n[None, :] * stride_original_max_n, mask=(offs_m[:, None] < num_tokens) & (offs_n[None, :] < num_splits), - other=0.0, + other=-float("inf"), ) _accu = tl.load( accu_ptr + offs_m[:, None] * stride_accu_m + offs_n[None, :] * stride_accu_n,