[bugfix] MiniMax-H3 fusions: int64 row offsets in the fused qknorm+RoPE kernel
_qknorm_partial_rope_kernel left tl.program_id(0) in int32, so row * head_dim wrapped once the flattened input reached 2**31 elements and the kernel read/wrote out of bounds (CUDA illegal memory access). The PR's other two kernels (modulation.py, swiglu.py) already cast tl.program_id(0).to(tl.int64); this one now matches, and seq_index / table_offset inherit int64 from row. Confirmed on GB200: (1, 8_500_000, 2, 128) bf16 (2.176e9 elements) crashed before the cast and matches eager after it; the just-under-2**31 control shape matched all along. For H3 (56 heads x 128 head_dim) the boundary is batch*seq >= 299_593 tokens per rank, reachable at SP=1. Adds a GPU regression test at the over-2**31 shape that compares the head and tail rows against eager. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
b158388733
commit
ac98869aa1
@@ -35,7 +35,12 @@ if HAVE_TRITON:
|
||||
eps,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
row = tl.program_id(0)
|
||||
# int64, like the sibling kernels: with int32 program ids,
|
||||
# ``row * head_dim`` wraps once the flattened input reaches 2**31
|
||||
# elements (H3's 56 heads x 128 head_dim crosses that at
|
||||
# batch*seq >= 299_593 tokens per rank) and the loads/stores below
|
||||
# become out-of-bounds. ``seq_index`` inherits int64 from ``row``.
|
||||
row = tl.program_id(0).to(tl.int64)
|
||||
seq_index = (row // num_heads) % seq_len
|
||||
cols = tl.arange(0, BLOCK_SIZE)
|
||||
head_mask = cols < head_dim
|
||||
@@ -128,6 +133,10 @@ def fused_qknorm_rope(
|
||||
RMSNorm reduction and RoPE arithmetic stay in FP32 registers until the
|
||||
final store. Triton's reduction order and the absence of eager's BF16
|
||||
intermediate materializations can produce small, expected rounding drift.
|
||||
|
||||
Row offsets are computed in int64, so inputs beyond 2**31 total elements
|
||||
(about 300k tokens per rank at H3's 56 heads x 128 head_dim) address
|
||||
correctly.
|
||||
"""
|
||||
batch, seq_len, num_heads, head_dim, rotary_dim = _validate_inputs(x, weight, cos, sin, eps)
|
||||
if not weight.is_contiguous():
|
||||
|
||||
@@ -148,3 +148,45 @@ def test_fused_qknorm_rope_matches_eager_bf16_cuda(
|
||||
assert actual.shape == x.shape
|
||||
assert actual.dtype == x.dtype
|
||||
assert actual.is_contiguous()
|
||||
|
||||
|
||||
def test_fused_qknorm_rope_matches_eager_beyond_int32_element_count() -> None:
|
||||
"""Regression: kernel row offsets must be int64.
|
||||
|
||||
With int32 offsets, ``row * head_dim`` wraps once the flattened input
|
||||
crosses 2**31 elements and the kernel reads/writes out of bounds (CUDA
|
||||
illegal memory access). ``(1, 8_500_000, 2, 128)`` is 2.176e9 elements,
|
||||
just past the boundary; for H3's 56 heads x 128 head_dim the equivalent
|
||||
is ``batch*seq >= 299_593`` tokens per rank, reachable at SP=1.
|
||||
|
||||
GPU assumption: needs ~16 GiB free CUDA memory (input + output at
|
||||
bf16 plus the fp32 rotary-table construction); skips below 20 GiB.
|
||||
"""
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA is required for the Triton fusion")
|
||||
if not HAVE_TRITON:
|
||||
pytest.skip("Triton is required for the fusion")
|
||||
free_bytes, _ = torch.cuda.mem_get_info()
|
||||
if free_bytes < 20 * 1024**3:
|
||||
pytest.skip("needs ~20 GiB free GPU memory for a >2**31-element input")
|
||||
|
||||
heads, head_dim, rotary_dim = 2, 128, 96
|
||||
seq_len = 8_500_000
|
||||
assert seq_len * heads * head_dim > 2**31
|
||||
|
||||
torch.manual_seed(9)
|
||||
device = torch.device("cuda")
|
||||
x = torch.randn(1, seq_len, heads, head_dim, dtype=torch.bfloat16, device=device)
|
||||
weight = (1.0 + 0.05 * torch.randn(head_dim, dtype=torch.bfloat16, device=device)).contiguous()
|
||||
cos, sin = _rotary_tables(seq_len, rotary_dim, dtype=x.dtype, device=device)
|
||||
|
||||
with torch.inference_mode():
|
||||
fused = fused_qknorm_rope(x, weight, cos, sin, 1e-6)
|
||||
torch.cuda.synchronize()
|
||||
# Compare only head/tail slices against eager: a full-tensor eager
|
||||
# reference would double peak memory for no extra coverage, and the
|
||||
# tail rows are exactly the ones an int32 wrap corrupts first.
|
||||
expected_head = _eager_qknorm_rope(x[:, :8], weight, cos[:8], sin[:8], 1e-6)
|
||||
expected_tail = _eager_qknorm_rope(x[:, -8:], weight, cos[-8:], sin[-8:], 1e-6)
|
||||
torch.testing.assert_close(fused[:, :8], expected_head, atol=2e-2, rtol=2e-2)
|
||||
torch.testing.assert_close(fused[:, -8:], expected_tail, atol=2e-2, rtol=2e-2)
|
||||
|
||||
Reference in New Issue
Block a user