Skip to content
Merged
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
100 changes: 55 additions & 45 deletions op_tests/test_mla_v4_nm.py
Original file line number Diff line number Diff line change
Expand Up @@ -703,6 +703,22 @@ def _build_bf16_inputs(
}


def _gated_allclose(a, b, msg, rtol=3e-2, atol=3e-2, tol_err_ratio=0.02):
"""checkAllclose that fails the test on a `failed!` result.

checkAllclose only raises on a catastrophic delta; otherwise it logs
`failed!` and returns the mismatch fraction, so an ungated call lets a bad
accuracy result pass silently in both pytest and the script-mode sweep.
"""
err = checkAllclose(
a, b, rtol=rtol, atol=atol, tol_err_ratio=tol_err_ratio, msg=msg
)
assert err <= tol_err_ratio, (
f"{msg}: {err:.1%} of elements outside rtol={rtol} atol={atol} "
f"(limit {tol_err_ratio:.0%})"
)


def _run_one_point(
batch=2,
kv_seq_lens=64,
Expand Down Expand Up @@ -944,20 +960,13 @@ def _run_one_point(
num_rotate_args=1,
)

# Resolve the asm output to compare against. Three cases, all reading the
# buffer the wrapper actually populated (the 2b call above):
# out_16_nosplit=1 -> kernel writes packed-BF16 into the logits region;
# the wrapper unpacks it into output_buf (see
# mla_decode_fwd_v4_nm). Read output_buf directly.
# single-pass (fp32) -> kernel writes one FP32 partial to logits[:, 0],
# no stage2; cast it to BF16.
# multi-pass -> stage2 merge wrote merged BF16 to output_buf.
if out_16_nosplit != 0:
out_asm = output_buf # wrapper unpacked packed-BF16 here
elif num_kv_splits == 1:
out_asm = logits_buf[:, 0].to(dtypes.bf16) # [total_q, num_heads, dv]
else:
out_asm = output_buf # already [total_q, num_heads, dv] BF16
# output_buf is the authoritative result for every split count: single-pass
# writes packed BF16 straight into it, multi-pass gets it from the stage2
# merge. logits_buf must not be read here -- the dispatcher derives
# out_16_nosplit from num_kv_splits (ignoring the caller's value), so for a
# single split the raw kernel call above fills logits_buf with packed BF16,
# not FP32 partials.
out_asm = output_buf # [total_q, num_heads, dv] BF16

# ---- accuracy ----
# Two comparisons, run for BOTH single- and multi-split (split-kv is a perf
Expand All @@ -969,23 +978,17 @@ def _run_one_point(
f"\n[v4 nm accuracy] batch={batch} kv_seq_lens={kv_seq_lens} "
f"q_seq_logical={q_seq_logical} num_kv_splits={num_kv_splits} seed={seed}"
)
# Per-element check at checkAllclose's default 1% tolerance (rtol=atol=1e-2).
# checkAllclose prints pass/warning/failed with the offending-element ratio +
# max delta (it does not raise).
checkAllclose(
# Both are gated: golden vs fp8_ref catches a broken quant pipeline (e.g.
# all-zero e8m0 scales) that fp8_ref vs asm alone would miss, since the
# kernel and fp8_ref consume the same quantized bytes.
_gated_allclose(
out_golden.float(),
out_fp8_ref.float(),
rtol=3e-2,
atol=3e-2,
tol_err_ratio=0.02,
msg="mla_v4_nm [golden_bf16 vs fp8_ref]",
)
checkAllclose(
_gated_allclose(
out_fp8_ref.float(),
out_asm.float(),
rtol=3e-2,
atol=3e-2,
tol_err_ratio=0.02,
msg="mla_v4_nm [fp8_dequant_ref vs asm]",
)

Expand Down Expand Up @@ -1111,12 +1114,9 @@ def _run_varlen_point(kv_lens, gqa_ratio=128, seed=0, attn_sink=True):
f"\n[v4 nm varlen] gqa={gqa} kv_lens={kv_lens} total_kv={total_kv} "
f"resolved_splits={resolved}"
)
checkAllclose(
_gated_allclose(
out_ref.float(),
out_asm,
rtol=3e-2,
atol=3e-2,
tol_err_ratio=0.02,
msg=f"mla_v4_nm varlen [fp8_dequant_ref vs asm] kv_lens={kv_lens}",
)

Expand Down Expand Up @@ -1495,12 +1495,9 @@ def _run_pad_poison_point(fill, poison_q, poison_kv, gqa_ratio, kv_seq_lens, bat
f"kernel is reading past the {_QUANT_NUM_SCALE_BYTES}-byte scale field "
f"at offset {_QUANT_D_NOPE} into the padding at {pad_off}."
)
checkAllclose(
_gated_allclose(
out_ref.float(),
out_asm,
rtol=3e-2,
atol=3e-2,
tol_err_ratio=0.02,
msg=f"mla_v4_nm {who} pad=0x{fill:02X} [fp8_dequant_ref vs asm]",
)

Expand Down Expand Up @@ -2556,25 +2553,32 @@ def test_v4_nm_kv_tail_not_tile_multiple():
args = parser.parse_args()

perf_rows = []
failures = []
for batch, kv_seq_lens, q_seq_logical in itertools.product(
args.batch, args.kv_seq_lens, args.q_seq_logical
):
print(
f"\n========== batch={batch} kv_seq_lens={kv_seq_lens} "
f"q_seq_logical={q_seq_logical} =========="
)
us_asm, us_ref = _run_one_point(
batch=batch,
kv_seq_lens=kv_seq_lens,
q_seq_logical=q_seq_logical,
seed=args.seed,
num_iters=args.iters,
num_warmup=args.warmup,
num_kv_splits=args.split_kv,
gqa_ratio=args.gqa_ratio,
attn_sink=args.attn_sink,
out_16_nosplit=args.out_16_nosplit,
)
try:
us_asm, us_ref = _run_one_point(
batch=batch,
kv_seq_lens=kv_seq_lens,
q_seq_logical=q_seq_logical,
seed=args.seed,
num_iters=args.iters,
num_warmup=args.warmup,
num_kv_splits=args.split_kv,
gqa_ratio=args.gqa_ratio,
attn_sink=args.attn_sink,
out_16_nosplit=args.out_16_nosplit,
)
except AssertionError as e:
# Keep sweeping so one bad shape doesn't hide the rest; the exit
# code below is what fails CI (aiter_test.sh runs this as a script).
failures.append(((batch, kv_seq_lens, q_seq_logical), str(e)))
continue
perf_rows.append((batch, kv_seq_lens, q_seq_logical, us_asm, us_ref))

print("\n[v4 nm perf summary] (us; speedup = fp8_ref / asm_kernel)")
Expand All @@ -2584,3 +2588,9 @@ def test_v4_nm_kv_tail_not_tile_multiple():
)
for b, k, q, ua, ur in perf_rows:
print(f" {b:>6d} {k:>8d} {q:>6d} {ua:>10.2f} {ur:>12.2f} {ur / ua:>8.1f}x")

if failures:
print(f"\n[v4 nm accuracy] {len(failures)} shape(s) FAILED:")
for (b, k, q), reason in failures:
print(f" batch={b} kv_seq_lens={k} q_seq_logical={q}: {reason}")
sys.exit(1)
Loading