[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:
Will Lin
2026-08-21 10:06:12 +00:00
co-authored by Claude Fable 5
parent b158388733
commit ac98869aa1
2 changed files with 52 additions and 1 deletions
@@ -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)