Skip to content

Commit 2cf7301

Browse files
committed
fix(sm120): add labels-contract checker for m-grouped contiguous GEMM
Same fix as the nv_dev-lineage sibling: group boundaries in labels mode must be multiples of the runtime mk alignment, or a BLOCK_M tile straddles two groups and silently computes the straddled rows with the wrong group's B (reported as 'middle empty group corrupts later groups'; the empty group merely made the misalignment likely). The tile heuristics already guarantee BLOCK_M divides the runtime alignment, so contract-respecting layouts are always safe; the opt-in checker (DG_CHECK_CONTIGUOUS_LABELS=1) turns violations into a loud host error. Regression tests pin middle-empty patterns at al in {64, 128} and the rejection path. Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>
1 parent 0aecffa commit 2cf7301

2 files changed

Lines changed: 108 additions & 0 deletions

File tree

‎csrc/apis/gemm.hpp‎

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
#pragma once
22

3+
#include <format>
4+
35
#include "../utils/compatibility.hpp"
46

57
#include "../jit_kernels/impls/sm90_fp8_gemm_1d1d.hpp"
@@ -182,6 +184,37 @@ static void fp8_fp4_gemm_tt(const std::pair<torch::Tensor, torch::Tensor>& a,
182184
d, c, recipe, recipe_a, recipe_b, compiled_dims, disable_ue8m0_cast, alpha);
183185
}
184186

187+
// SM120 m-grouped contiguous (labels-mode) layout contract: each group's first row must
188+
// start at a multiple of `get_mk_alignment_for_contiguous_layout()`. The kernel selects B
189+
// (and SFB) per BLOCK_M tile from the label of the tile's first row, and the tile
190+
// heuristics guarantee BLOCK_M divides the runtime alignment, so contract-respecting
191+
// labels never put a group boundary inside a tile. Labels built at a finer granularity
192+
// than the runtime alignment violate this and silently compute the straddled rows with
193+
// the wrong group's B. Set DG_CHECK_CONTIGUOUS_LABELS=1 to turn violations into a loud
194+
// error (does one GPU->CPU copy of the labels; meant for debugging/integration).
195+
static void sm120_check_contiguous_labels_contract(const torch::Tensor& grouped_layout,
196+
const int& num_groups) {
197+
if (not deep_jit::get_env<bool>("DG_CHECK_CONTIGUOUS_LABELS", false))
198+
return;
199+
const int alignment = heuristics_runtime->get_mk_alignment_for_contiguous_layout();
200+
const auto labels = grouped_layout.cpu();
201+
const auto* data = labels.data_ptr<int>();
202+
const auto m = static_cast<int64_t>(labels.size(0));
203+
for (int64_t i = 0, prev_label = -1; i < m; ++ i) {
204+
const int label = data[i];
205+
DG_HOST_ASSERT(label >= -1 and label < num_groups
206+
and "m-grouped contiguous label out of range");
207+
if (label < 0 or label == prev_label)
208+
continue;
209+
if (i % alignment != 0)
210+
DG_HOST_UNREACHABLE(std::format(
211+
"m-grouped contiguous group {} starts at row {}, not a multiple of the "
212+
"runtime mk alignment {}; rebuild the labels with matching alignment or "
213+
"call set_mk_alignment_for_contiguous_layout()", label, i, alignment));
214+
prev_label = label;
215+
}
216+
}
217+
185218
static void m_grouped_fp8_fp4_gemm_nt_contiguous(const std::pair<torch::Tensor, torch::Tensor>& a,
186219
const std::pair<torch::Tensor, torch::Tensor>& b,
187220
const torch::Tensor& d,
@@ -246,6 +279,8 @@ static void m_grouped_fp8_fp4_gemm_nt_contiguous(const std::pair<torch::Tensor,
246279
num_groups, m, n, k, gran_k_a, gran_k_b, major_a, major_b,
247280
compiled_dims, use_psum_layout, ensure_zero_padding, expected_m_for_psum_layout);
248281
} else if (arch_major == 12 and sfa.scalar_type() == torch::kInt and sfb.scalar_type() == torch::kInt) {
282+
if (not use_psum_layout)
283+
sm120_check_contiguous_labels_contract(grouped_layout, num_groups);
249284
const auto b_data = sm120::to_k_major(b.first, major_b, n);
250285
const bool is_mixed_fp4 = (a.first.scalar_type() != b_data.scalar_type()) and
251286
(a.first.scalar_type() == kPackedFP4 or b_data.scalar_type() == kPackedFP4);
@@ -630,6 +665,8 @@ static void m_grouped_bf16_gemm_nt_contiguous(const torch::Tensor& a, const torc
630665
num_groups, m, n, k, major_a, major_b, compiled_dims,
631666
use_psum_layout, ensure_zero_padding, expected_m_for_psum_layout);
632667
} else if (arch_major == 12) {
668+
if (not use_psum_layout)
669+
sm120_check_contiguous_labels_contract(grouped_layout, num_groups);
633670
sm120_m_grouped_bf16_gemm_contiguous(a, b, d, grouped_layout,
634671
num_groups, m, n, k, major_a, major_b, compiled_dims,
635672
use_psum_layout, expected_m_for_psum_layout);

‎tests/test_sm120_fp8_fp4.py‎

Lines changed: 71 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -405,3 +405,74 @@ def test_sm120_kgroup_descriptor_reuse_at_default_sms() -> None:
405405
test_sm120_asymmetric_scale_recipe_swap()
406406
test_sm120_scale_dtype_validation()
407407
test_sm120_masked_physical_capacity()
408+
409+
410+
@test_filter(lambda: get_arch_major() == 12)
411+
def test_sm120_contiguous_labels_middle_empty_group() -> None:
412+
"""Labels-mode contiguous GEMM with an empty group in the middle: groups after the
413+
gap must still read their own B (regression for a host/kernel contract where group
414+
boundaries must align with the runtime mk alignment)."""
415+
old_alignment = deep_gemm.get_mk_alignment_for_contiguous_layout()
416+
try:
417+
for alignment in (64, 128):
418+
deep_gemm.set_mk_alignment_for_contiguous_layout(alignment)
419+
for lengths in ([100, 0, 130, 65], [128, 130, 65], [65, 0, 0, 66, 31]):
420+
n, k = 128, 256
421+
intervals, end = [], 0
422+
for length in lengths:
423+
start = (end + alignment - 1) // alignment * alignment
424+
end = start + length
425+
intervals.append((start, end))
426+
m = (end + alignment - 1) // alignment * alignment
427+
a = (torch.full((m, k), 2.0, device='cuda').to(torch.float8_e4m3fn),
428+
torch.ones(m, k // 128, device='cuda'))
429+
b = (torch.cat([torch.full((1, n, k), float(g + 1), device='cuda')
430+
for g in range(len(lengths))]).to(torch.float8_e4m3fn),
431+
torch.ones(len(lengths), n, k // 128, device='cuda'))
432+
labels = torch.full((m,), -1, dtype=torch.int32)
433+
for g, (start, end) in enumerate(intervals):
434+
labels[start:end] = g
435+
d = torch.full((m, n), 3.0, dtype=torch.bfloat16, device='cuda')
436+
deep_gemm.m_grouped_fp8_fp4_gemm_nt_contiguous(
437+
a, b, d, labels.cuda(), recipe=(1, 1, 128))
438+
for g, (start, end) in enumerate(intervals):
439+
if start == end:
440+
continue
441+
want = float(k * 2 * (g + 1))
442+
assert torch.all(d[start:end] == want), \
443+
(alignment, lengths, g, d[start:end].unique())
444+
finally:
445+
deep_gemm.set_mk_alignment_for_contiguous_layout(old_alignment)
446+
447+
448+
@test_filter(lambda: get_arch_major() == 12)
449+
def test_sm120_contiguous_labels_contract_rejection() -> None:
450+
"""Labels finer than the runtime mk alignment must fail loudly when the opt-in
451+
contract checker is enabled (DG_CHECK_CONTIGUOUS_LABELS=1), not corrupt silently."""
452+
import os
453+
old_env = os.environ.get('DG_CHECK_CONTIGUOUS_LABELS')
454+
os.environ['DG_CHECK_CONTIGUOUS_LABELS'] = '1'
455+
old_alignment = deep_gemm.get_mk_alignment_for_contiguous_layout()
456+
deep_gemm.set_mk_alignment_for_contiguous_layout(128)
457+
try:
458+
m, n, k = 256, 128, 256
459+
a = (torch.zeros(m, k, device='cuda').to(torch.float8_e4m3fn),
460+
torch.ones(m, k // 128, device='cuda'))
461+
b = (torch.zeros(2, n, k, device='cuda').to(torch.float8_e4m3fn),
462+
torch.ones(2, n, k // 128, device='cuda'))
463+
labels = torch.full((m,), -1, dtype=torch.int32)
464+
labels[0:64] = 0
465+
labels[64:128] = 1 # group 1 starts at 64: finer than the runtime alignment 128
466+
d = torch.zeros(m, n, dtype=torch.bfloat16, device='cuda')
467+
try:
468+
deep_gemm.m_grouped_fp8_fp4_gemm_nt_contiguous(a, b, d, labels.cuda(),
469+
recipe=(1, 1, 128))
470+
raise AssertionError('expected the contract checker to reject')
471+
except RuntimeError as e:
472+
assert 'not a multiple of the runtime mk alignment' in str(e), e
473+
finally:
474+
if old_env is None:
475+
os.environ.pop('DG_CHECK_CONTIGUOUS_LABELS', None)
476+
else:
477+
os.environ['DG_CHECK_CONTIGUOUS_LABELS'] = old_env
478+
deep_gemm.set_mk_alignment_for_contiguous_layout(old_alignment)

0 commit comments

Comments
 (0)