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
58 changes: 58 additions & 0 deletions tests/utils/test_linear_cross_entropy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand All @@ -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")
67 changes: 66 additions & 1 deletion tests/utils/test_special_linear_cross_entropy_tp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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)
Expand Down
14 changes: 9 additions & 5 deletions verl/utils/kernel/kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)):
Expand All @@ -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,
Expand Down