Skip to content

Fix SM120 BF16 GEMM shared-memory reuse ordering - #453

Open
Sunt-ing wants to merge 2 commits into
deepseek-ai:nv_devfrom
Sunt-ing:fix/sm120-bf16-stage-reuse
Open

Sunt-ing wants to merge 2 commits into
deepseek-ai:nv_devfrom
Sunt-ing:fix/sm120-bf16-stage-reuse

Conversation

@Sunt-ing

@Sunt-ing Sunt-ing commented Sep 21, 2026 •

Copy link
Copy Markdown

Summary

Fixes intermittent silent data corruption in SM120 BF16 GEMM caused by a generic-to-async proxy WAR (Write-After-Read) hazard.

Problem

Consumer warps read smem tiles via the generic proxy and signal the empty barrier without an async proxy fence. Producer TMA loads (async proxy) can subsequently overwrite the smem stage before generic reads complete across all consumer threads.

Fix

Add cutlass::arch::fence_view_async_shared() before the empty-barrier arrival. Placed outside the lane-0 check so all consumer threads issue the fence for their own smem reads.

Targets nv_dev (based on 572557e7ae9ad5331b81a1c250f141fba2c57962).

Reproduction and validation

Environment: RTX 5090 (SM120), Driver 590.48.01, PyTorch 2.13.0+cu130.

Added test_bf16_repeatability in tests/test_bf16.py ([4096, 1024] × [1024, 896], seed 233, 6,000 iterations against initial output).

To run the standalone regression:

cd tests
python -c "from test_bf16 import test_bf16_repeatability; test_bf16_repeatability()"

Note: The corruption is an intermittent race condition; a single baseline run passing does not guarantee correctness.

Repeatability (A/B/A Test)

  • Original: 2 / 6,000 mismatches
  • Patched: 0 / 6,000 mismatches
  • Reverted: 2 / 6,000 mismatches

Performance Overhead

CUDA graph timing on [4096, 1024] × [1024, 896] (median per-call time across 9 batches of 200 calls):

Version Median Latency
Original 50.67 μs
Patched 52.01 μs
Reverted 50.68 μs

Adds ~2.6% latency overhead on this shape from the fence.

Supplemental End-to-End Inference Validation

Internal TP2 inference test (weights/prompts omitted; public verification relies on the standalone GEMM test above).

  • Setup: 32,640-token prompt, 89 decode steps -> re-prefill -> 16 decode steps vs. continuous generation. Unified prefill/decode attention was held constant; CUDA graphs were disabled.
  • Original / Reverted: 6 / 12 runs with divergent logprobs.
  • Patched: 0 / 12 runs diverged (token IDs and logprobs match strictly).

Related work

Comment thread deep_gemm/include/deep_gemm/impls/sm120_bf16_gemm.cuh
Comment thread tests/test_bf16.py
Comment thread tests/test_bf16.py
Comment thread tests/test_bf16.py
Comment thread tests/test_bf16.py
@ds-review-bot

Copy link
Copy Markdown
Collaborator

🤖 ds-review-bot Code Review

v6

未发现本次变更新引入的可确认缺陷。fence 位于共享内存读取之后、empty barrier 到达之前,且由所有消费者线程执行;新增重复性测试也会调用目标内核。静态检查通过,但当前环境缺少 PyTorch 和 CUDA 工具链,未执行 GPU 验证。

v5

本 MR 在 SM120 BF16 GEMM 的消费者释放路径中,于 empty barrier arrive 之前加入 cutlass::arch::fence_view_async_shared()(fence.proxy.async.shared::cta),修复 generic proxy(ldmatrix 读 smem)与 async proxy(TMA 写 smem)之间的 WAR 竞争。修复位置正确:fence 放在 lane_idx == 0 判断之外,保证每个消费者线程都对自身的 smem 读做排序后再由 lane 0 发出 arrive(fence.proxy.async 仅对执行线程自身访存生效,若只由 lane 0 执行则其余 lane 的读仍会乱序);cutlass/arch/barrier.h 已在文件头包含,符号可正常解析。新增 test_bf16_repeatability 回归测试(seed 233,[4096,1024]×[1024,896],6000 次 torch.equal 对比)设计合理,能在无修复时以约 2/6000 概率复现问题,且对非 SM120 架构直接返回、不影响其他架构 CI。A/B/A 复现数据、端到端 TP2 验证与 ~2.6% 的性能开销说明充分,改动最小且聚焦。建议合并。遗留问题:同类 lane-0 arrive 无 fence 的模式仍存在于其他 SM120 内核,建议在后续 MR(#447)中统一处理,详见评论。

v4

The change is correct and well-targeted. It fixes a real generic-to-async proxy WAR hazard in the SM120 BF16 GEMM consumer path: after the consumer warps finish reading the smem A/B tiles, they signal stage reuse via empty_barriers[stage]->arrive() without ordering those generic-proxy reads against the producer's subsequent async-proxy TMA write, so a fast producer can overwrite a stage while consumer reads are still in flight. Adding cutlass::arch::fence_view_async_shared() (fence.proxy.async.shared::cta) immediately before the release is the right primitive, and placing it outside the lane_idx == 0 guard is correct because the fence is a per-thread ordering primitive and every lane that issued smem loads must drain them before lane 0's release publishes the stage as reusable. <cutlass/arch/barrier.h> is already included, so no new include is needed. The added repeatability test is a reasonable stress test for the race (though probabilistic). The main concerns are: (1) the identical unfenced release pattern remains in the sibling SM120 GEMM consumers (FP8/FP4 1D1D, BMK/BNK MN, TF32 prenorm), so this class of silent corruption is only partially closed; and (2) the new soak test may pass on the buggy build a non-trivial fraction of the time and synchronizes on every iteration.

Files reviewed: 2
Issues found: 🔴 3 critical | 🟡 1 warning | 🔵 4 suggestion
Inline comments posted: 5
General comments (无法定位到 diff): 3


📍 未定位到 diff 的评论

🔴 critical deep_gemm/include/deep_gemm/impls/sm120_fp8_fp4_gemm_1d1d.cuh:L666: Same unfenced release pattern remains here (also at lines 814 and 1143). The producer issues TMA loads via tma::copy + arrive_and_expect_tx (lines 321-331), the consumer reads the tiles through the generic proxy (load_a_fragment/load_b_fragment/ldmatrix) and then releases with if (lane_idx == 0) empty_barriers[stage]->arrive(); with no async proxy fence. This is the same generic->async WAR hazard fixed for BF16 and can produce the same intermittent silent corruption. Please apply cutlass::arch::fence_view_async_shared() before each of these arrivals (or factor the release into a helper that always fences). 🤖 v4

🔴 critical deep_gemm/include/deep_gemm/impls/sm120_bmk_bnk_mn.cuh:L202: This consumer has the identical structure to the BF16 kernel: the producer TMA-copies A/B into stages (lines 126-136), the consumer ldmatrix-reads them (lines 178-198), then releases the stage with a bare if (lane_idx == 0) empty_barriers[stage_idx]->arrive();. No proxy fence is present, so the producer can overwrite the stage before all generic reads complete. It should receive the same fix; otherwise the corruption class this MR addresses is only partially closed. 🤖 v4

🔴 critical deep_gemm/include/deep_gemm/impls/sm120_tf32_hc_prenorm_gemm.cuh:L231: Same pattern as the other SM120 consumers: TMA producer (lines 129-136) plus generic smem reads followed by a lane-0 empty-barrier arrival with no async proxy fence. Please audit and apply fence_view_async_shared() here as well. Given the references to prior similar fixes (#389, #447), centralizing the release so this fence cannot be forgotten again would prevent recurrence. 🤖 v4

@Sunt-ing

Copy link
Copy Markdown
Author

@LyricZhao @zheanxu Could you please take a look at this fix? We found this while investigating reproducibility in LLM evaluation with SGLang on RTX 5090s. With a 32K prompt, continuing decode and re-prefilling the same token prefix sometimes produced different logprobs. Tracing the difference through the model led to a BF16 projection where repeated GEMM calls with identical inputs occasionally returned different values. We then reproduced it with the public random-matrix test in this PR.

@lucifer1004

Copy link
Copy Markdown
Collaborator

Have you checked if this is covered by #447?

@Sunt-ing

Copy link
Copy Markdown
Author

I checked #447, but I couldn’t find the corresponding fix in the SM120 BF16 kernel. At its current head (1f78b85), the consumer still calls empty_barriers[stage]->arrive() without a preceding async proxy fence.

lucifer1004 added a commit to lucifer1004/DeepGEMM-sm120 that referenced this pull request Sep 21, 2026
Port of deepseek-ai/DeepGEMM#453 (Sunt-ing) plus an audit of the same
hazard across all vendored SM120 kernels: 14 sites in 8 headers.

The SM120 kernels read TMA-written smem through ldmatrix/ld_shared --
generic-proxy accesses -- so the empty-barrier arrive that hands the
stage back to the TMA producer must be ordered after those reads with
cutlass::arch::fence_view_async_shared(). Without the fence the next
TMA into the stage can overlap the consumer's reads (generic/async
proxy WAR hazard). Upstream's SM100 kernels carry the same fence for
the analogous UMMA case ("Release KV smem only after UMMA commits
TMEM", sm100_mqa_logits.cuh); the SM120 ports inherited the barrier
protocol but not the fences.

Sites: bf16_gemm (the #453 report), bmk_bnk_mn, tf32_hc_prenorm_gemm,
fp8_fp4_gemm_1d1d (3 arrives), fp8/fp4 mqa_logits (empty_kv + empty_q
each), fp8/fp4 paged_mqa_logits (empty_q stage release + empty_kv
each). sm120_fp8_fp4_sparse_mqa_logits is exempt: its producer is
legacy cp.async (generic proxy), which needs no proxy fence.

Evidence:
- #453 author: 2/6000 failures on RTX 5090 without the fence, 0/6000
  with it (A/B/A, m=4096 n=896 k=1024 bf16).
- Local RTX PRO 6000: 30k-iteration soak does not reproduce either way;
  the fence is required by the memory model, not by local observation.
- Perf A/B (bf16 GEMM, fp8 dense/paged MQA, 3 runs): all deltas within
  run-to-run noise, no regression.

Bump vendored tag to v0.1.5.

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>
lucifer1004 added a commit to lucifer1004/DeepGEMM that referenced this pull request Sep 21, 2026
Vendor lucifer1004/DeepGEMM-sm120 v0.1.5 (b31a688): 14
fence_view_async_shared() sites across 8 SM120 kernel headers, ordering
generic-proxy smem reads (ldmatrix/ld_shared) before the empty-barrier
arrive that hands each stage back to the TMA producer. Ports
deepseek-ai#453 (bf16 GEMM) and fixes the same hazard found by
audit in bmk_bnk_mn, tf32_hc_prenorm, fp8_fp4_gemm_1d1d, and the
fp8/fp4 dense+paged MQA logits kernels. Sparse MQA is exempt (legacy
cp.async producer).

Perf A/B on RTX PRO 6000 (bf16 GEMM, fp8 dense/paged MQA, 3 runs): all
deltas within run-to-run noise. Fork suite: 235 passed (GPU4, fresh JIT
cache).

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>
lucifer1004 added a commit to lucifer1004/DeepGEMM that referenced this pull request Sep 21, 2026
Vendor lucifer1004/DeepGEMM-sm120 v0.1.5 (b31a688): 14
fence_view_async_shared() sites across 8 SM120 kernel headers, ordering
generic-proxy smem reads (ldmatrix/ld_shared) before the empty-barrier
arrive that hands each stage back to the TMA producer. Ports
deepseek-ai#453 (bf16 GEMM) and fixes the same hazard found by
audit in bmk_bnk_mn, tf32_hc_prenorm, fp8_fp4_gemm_1d1d, and the
fp8/fp4 dense+paged MQA logits kernels. Sparse MQA is exempt (legacy
cp.async producer).

Perf A/B on RTX PRO 6000: no regression beyond run-to-run noise.

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>
@lucifer1004

Copy link
Copy Markdown
Collaborator

@Sunt-ing Thanks for the careful A/B/A reproduction — confirmed the hazard class. Your fix has now been ported into #447 at 381e2d7, via the canonical SM120 device-layer repo (lucifer1004/DeepGEMM-sm120 v0.1.5, b31a688), which both SM120 forks vendor byte-for-byte; the same fix is also in the vllm fork's SM120 layer (vllm-project#10).

While merging it we audited all SM120 kernels for the same pattern — TMA-produced smem consumed via ldmatrix/ld_shared (generic proxy), with the empty-barrier arrive handing the stage back to the async proxy — and added the missing fence_view_async_shared() at 14 sites across 8 headers: bf16_gemm (your report), bmk_bnk_mn, tf32_hc_prenorm_gemm, fp8_fp4_gemm_1d1d (×3), fp8/fp4 dense MQA logits (empty_kv + empty_q each), and fp8/fp4 paged MQA logits (empty_q stage release + empty_kv each). The sparse MQA kernel is exempt (its producer is legacy cp.async, generic proxy). Upstream's SM100 kernels carry the same fence for the analogous UMMA case ("Release KV smem only after UMMA commits TMEM"), which corroborates the pattern.

For transparency on evidence: on our RTX PRO 6000 a 30k-iteration soak of the affected shapes did not reproduce the corruption either way, so the fix rests on the memory model plus your 5090 reproduction (2/6000 → 0/6000). Perf A/B (bf16 GEMM, fp8 dense/paged MQA, 3 runs each) shows all deltas within run-to-run noise — consistent with your +2.6% report.

@Sunt-ing

Copy link
Copy Markdown
Author

Thanks for porting the fix and checking the other kernels!

zyongye pushed a commit to vllm-project/DeepGEMM that referenced this pull request Sep 22, 2026
* feat(sm120): vendor DeepGEMM-sm120 v0.1.3 device layer

Replace the in-tree SM120 device headers with byte-identical vendored
copies from lucifer1004/DeepGEMM-sm120 v0.1.3 (merged superset of both
lineages). Marker-free vendoring: no manifests, tooling, or header
comments enter this fork; provenance and drift control live on the
canonical repo's side (per-tag manifest, check_vendor.py --fork,
scheduled downstream watcher).

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

* feat(sm120): adapt host launchers to vendored kernel signatures

The vendored DeepGEMM-sm120 v0.1.0 device layer widens several kernel
ABIs; adapt the SM120 host glue with neutral mappings:

- sm120_bf16_gemm: pass the epilogue operator as a runtime argument
  (EpilogueArgs, same marshalling as SM100) and split the old
  int64_t stride_cd_m/stride_cd_batch into u32 stride_d_m/stride_c_m/
  stride_d_batch; stride_c_m=0 keeps C sharing D's row stride.
- sm120_fp8_fp4_gemm_1d1d: pass shape_cd_m (= shape_m), the runtime
  epilogue argument, and stride_c_m (= 0) after stride_cd_batch.
- sm120_split_k_reduce: the kWithAccumulation template flag is gone;
  instantiate the 3-arg template and map accumulation onto D in place
  to the new runtime C operand (gmem_c = D with D's own strides,
  with_alpha = false).
- Masked m-grouped GEMM (bf16 + fp8/fp4): disable the TMA-store
  epilogue (swizzle_cd_mode = 0) when m % BLOCK_M != 0. The vendored
  kernels dropped the old masked-boundary scalar fallback, so a
  boundary tile's full-tile TMA store would spill into the next
  group's rows; the scalar store path carries the masked row bounds.
  This mirrors the proven nv_dev-lineage launchers.

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

* fix(sm120): guard k-grouped GEMM TMA store against unaligned M

The vendored v0.1.0 GEMM kernels have no boundary-tile scalar fallback
in the TMA-store epilogue, and the k-grouped D descriptor is one flat
2D map over all groups, so TMA can only clamp at the outermost M dim:
with m % BLOCK_M != 0, a group's partial tail M tile TMA-store spills
into the next group's slab. Mirror the masked-launcher guard (and the
nv_dev-lineage reference glue) in both k-grouped launchers: force
swizzle_cd_mode = 0 (group-bounded scalar store path) for unaligned m.

Repro evidence (num_sms=2, groups=2, m=48, n=128, ks=[8192, 128]; the
asymmetric K makes group 0's tile store last, so the spill overwrites
group 1's head rows deterministically):
- before: bf16/bf16-out 5/5 iterations contaminated (d[1][0:16] = 0.0,
  expected 768); fp8/fp32-out 5/5 contaminated (d[1][0:16] = 0.0,
  expected 771)
- bf16/fp32-out was already clean (0/5): the bf16 kernel gates its
  TMA-store epilogue on sizeof(cd_dtype_t) <= 2, so fp32 outputs take
  the scalar path; the fp8 kernel has no such dtype gate
- after: all three configurations 0/5 contaminated

Adds test_sm120_kgroup_unaligned_m_tail_tile_isolation to
tests/test_sm120_bf16.py (fp32 and bf16 outputs) and
tests/test_sm120_fp8_fp4.py (NT and TN layouts), asserting per-group
correctness including slab head rows plus flat-storage guard regions.
Full SM120 suite: 48 passed (46 baseline nodes + 2 new). Memcheck on
the two affected files: 20 passed, 0 sanitizer errors.

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

* feat(sm120): vendor DeepGEMM-sm120 v0.1.4 device layer

FP8 paged MQA logits gains PAGE_KV=32 (BLOCK_KV derived as
min(PAGE_KV, 64), mirroring the FP4 sibling). Device-only update: this
fork's host launcher still restricts paged FP8 to page 64, so page32
stays inert until host glue opts in (#14).

Validated on sm_120a: test_sm120_mqa.py + test_sm120_fp8_fp4.py 23/23
passed from a fresh JIT cache (only the pre-existing test_filter
collection quirk remains, deepseek-ai#446).

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

* fix(sm120): vendor DeepGEMM-sm120 v0.1.5 generic-proxy WAR fences

Vendor lucifer1004/DeepGEMM-sm120 v0.1.5 (b31a688): 14
fence_view_async_shared() sites across 8 SM120 kernel headers, ordering
generic-proxy smem reads (ldmatrix/ld_shared) before the empty-barrier
arrive that hands each stage back to the TMA producer. Ports
deepseek-ai#453 (bf16 GEMM) and fixes the same hazard found by
audit in bmk_bnk_mn, tf32_hc_prenorm, fp8_fp4_gemm_1d1d, and the
fp8/fp4 dense+paged MQA logits kernels. Sparse MQA is exempt (legacy
cp.async producer).

Perf A/B on RTX PRO 6000: no regression beyond run-to-run noise.

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

---------

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>
RayWang96 pushed a commit that referenced this pull request Sep 23, 2026
#447)

* Public Release 26/09 (#432)

* Public release 26/09

* Update News

* feat(sm120): add sparse MQA support on upstream main APIs

Port contiguous and paged sparse kernels to DeepJIT while preserving SM100 dispatch, stream-local metadata, entry-balanced scheduling and register guards.

Restore histogram synchronization and capture-safe UE8M0 validation, with independent SM120 numerical and CUDA graph regression fixtures.

* feat(sm120): restore native GEMM paths on main APIs

Restore the original BF16/FP8/FP4 kernels, scheduler, heuristics, and split-K paths from 139f504 through DeepJIT adapters. Support alpha, independent C/D strides, GPU-only KPSUM, and bounded padding cleanup.

Separately repair inherited mixed-tail, producer handoff, grouped row ownership, and TMA output-width defects. Preserve the original TMA/MMA pipelines and dispatch choices except for the documented safety eligibility guard.

This is a bounded restoration checkpoint, not full migration or universal performance parity. BF16 odd-N, FP4 K-grouped, and remaining public integrations are unfinished.

* feat(sm120): preserve nv_dev features alongside main API SM120 kernels

* fix(kernels): address migration review and restore grouped TMA pipeline

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

* fix(sm120): harden GEMM execution and complete HeadSplits output mapping

- Wait for in-flight TMA descriptor reads before overwriting the
  per-CTA K-grouped BF16 descriptor, and add full-SM randomized
  alias/separate-C regressions.
- Require packed INT32 SFB on the SM120 skip-head path instead of
  silently reading FP32 scales as UE8M0.
- Apply HeadSplits index mapping with logical-column bounds in scalar
  stores and split-K reduction, replacing the former capability guards
  while keeping alpha/C/stride contracts.
- Store BF16 batched einsum and HC prenorm directly to supported
  strided outputs, keeping the existing fallback for unsupported
  layouts.

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

* test(sm120): collect SM120 tests into dedicated files

Restore tests/test_{attention,bf16,fp8_fp4,hyperconnection}.py to their
upstream byte state and consolidate all SM120-specific coverage into
tests/test_sm120_*.py plus the shared tests/sm120_exercise.py helper.
This eliminates per-release merge conflicts with upstream's wholesale
rewrites of the shared test files.

Switch gating from collection-time string skipif (which initialized
CUDA during pytest collection on non-SM120 machines and broke the
sanitizer runner's direct-call path) to upstream's call-time
test_filter convention, and give each standalone file its own
__main__ runner matching upstream script style. Moved test bodies are
AST-identical; test counts and results are unchanged (741 passed both
before and after on SM120, memcheck/racecheck subsets clean).

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

* feat(sm120): vendor device layer from DeepGEMM-sm120 v0.1.3

Adopt the standalone DeepGEMM-sm120 repository as the single source of
truth for the SM120 device headers: add the CUDA>=13 compile-time guard
(silent block_scale drop on pre-13 ptxas), relocate
tensor_map_replace_global_dim_in_smem into common/sm120_utils.cuh so the
kernels no longer depend on ptx/tma.cuh providing it, and inline the FP4
smem pack factor in layout/sparse_mqa_logits.cuh.

Marker-free vendoring: the files carry no canonical-source comments or
other foreign markers; provenance and drift control live on the canonical
repo's side (per-tag manifest, check_vendor.py --fork, scheduled
downstream watcher).

Verified: 741 SM120 tests pass on sm_120a from a fresh JIT cache.
Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

* fix(sm120): gate ptxas -O2 for sparse MQA logits on CUDA 13.3+

CUDA 13.3's ptxas crashes (SIGSEGV; DeepJIT compile subprocess exit 139)
when compiling the non-warp-specialized MXFP4 sparse MQA logits variants
at -O3 with any --register-usage-level (lucifer1004/DeepGEMM-sm120#2);
CUDA 13.2, the validated toolchain, is unaffected. Add a toolchain-gated
per-kernel override via DeepJIT CompilerOptions::extra_nvcc_flags: when
nvcc reports >= 13.3, the non-warp-specialized instantiations compile
with ptxas -O2, which avoids the crash and reproduces the 13.2 register
allocation (REG:101 for the reference instantiation). The warp-specialized
variant is not affected by the crash and keeps the default flags (its
compiled reg==64 contract also verified under -O2).

Validated on sm_120a: sparse MQA suite 9/9 under CUDA 13.3.73 with the
gate active (previously 3 failing tests), and 9/9 under CUDA 13.2.86
with the gate inactive (unchanged default path).

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

* feat(sm120): FP8 paged MQA logits PAGE_KV=32 support

Vendor DeepGEMM-sm120 v0.1.4, which derives BLOCK_KV = min(PAGE_KV, 64)
in the FP8 paged kernel (mirroring the FP4 sibling): a 32-row page gets
a 32-row compute tile and 4 KV groups at SPLIT_KV=128, instead of a
64-row tile straddling two non-contiguous physical pages. DSv4.1 indexer
caches mix 64- and 32-state pages (vllm-project#14).

Host glue: relax the FP8 launcher and fused-cache API gates to admit
block_kv=32 for arch 12, derive tile_kv = min(block_kv, 64) like the FP4
launcher, and drop to two KV stages for page32 + 64 heads + paired
queries, where three stages exceed the 99 KiB SMEM budget by 4 bytes.

Tests: the paged-MQA contract matrix gains (fp8, page32), and a focused
case covers the two-stage fallback (page32, 64 heads, paired/varlen).
Validated on sm_120a from a fresh JIT cache: paged MQA suite 20/20
passed, incl. graph replay and legacy-API cross-checks.

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

* perf(sm90): clamp paged MQA scheduler range once at construction

The stale/OOB-metadata protection added during the migration lived in the
per-task hot path (compound has_next_task/exist_q_atom_idx checks per
fetch), regressing paged MQA logits on SM90: ~1% on upstream shapes and
13-16% at next_n=4. Move the protection to the constructor: clamp
end_q_atom_idx to the in-batch sentinel (zeroing end_kv_idx at the
sentinel) and collapse out-of-range starts, then make the per-task end
test a single ordered 64-bit compare on (q_atom, kv) packed keys. The
ordered compare also keeps the hang fix: an advance overshooting the end
is caught by >=, never by exact equality. context_lens/indices are only
dereferenced when current_q_atom_idx is provably in batch (end is clamped
at or below the sentinel, and the sentinel range always has end_kv == 0).

Patch from local review. Validated on H200: test_paged_mqa_logits 24/24
cases pass; next_n=4 (H=64, D=128, page 64) 140.6us vs 154.6us pre-patch
at L=8192 and 1093.3us vs 1233.1us at L=65536 (-9%/-11%), upstream-shape
sweep at parity or better. Cross-compiled for sm_90a. Reviewer's full
105-case SM90 suite passes on their side.

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

* perf(layout): pow2 shift fast path in k-grouped SF pack row accounting

pack_fp32_into_ue8m0's per-group loop ran three runtime integer divisions
per group per thread on the psum path (align of previous end, align of
group K, ceil_div to SF rows). gran_k and k_alignment are powers of two
at every call site (host-asserted), so detect that once per block and use
shift forms instead; the generic division path is kept for any future
non-pow2 alignment. The earlier suggestion to drop the align entirely on
SM90/100 was unsafe: psum tail K sizes are not k_alignment-multiples, and
without the align the row count comes up short (B300 fp8 psum accumulate
regression, diff 0.024 vs 0.008 tolerance).

Validated: pack output bitwise identical to the division form on a
128-group psum shape (B300); kernel time 38.7 -> 22.1 us/call (-43%) on
the same shape. SM100 k-grouped fp8/bf16 contiguous suites pass on B300
(fresh JIT cache), SM120 k-grouped regressions pass on sm_120a, and the
TU cross-compiles for sm_103a/sm_120a.

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

* docs(layout): note why the pow2 shift forms are exact on SM90/100

psum end offsets are k_alignment-multiples and gran_k divides k_alignment
there, so aligned_group_k % gran_k == 0 and the rounding shift is exact
division; the ceil semantics matter only on SM120. Comment-only change.

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

* perf(layout): move gran_k/k_alignment into the JIT template for SF pack

Follow-up to 3aca951: the runtime pow2 branch regressed non-pow2
k_alignment (reported by RayWang96: k_alignment=384, a third of their
k-grouped cases, 91 -> 124 us vs main on B300). Instead of branching at
runtime, pass gran_k/k_alignment as JIT template parameters: the per-group
align/ceil_div fold to shifts (pow2) or multiply-high sequences (non-pow2)
at compile time, covering both alignment classes with no runtime branch.
Runtime kernel args drop accordingly.

Validated on B300: pack output bitwise identical to the pre-change form at
k_alignment 128 and 384; pack kernel 38.8 -> 14.0 us/call at 128 (-64%)
and 65.2 -> 32.3 us/call at 384 (-50%); k-grouped fp8/bf16 contiguous
suites pass from a fresh JIT cache. SM120 k-grouped regressions pass on
sm_120a; the TU cross-compiles for sm_103a/sm_120a.

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

* fix(sm120): vendor v0.1.5 generic-proxy WAR fences from DeepGEMM-sm120

Vendor lucifer1004/DeepGEMM-sm120 v0.1.5 (b31a688): 14
fence_view_async_shared() sites across 8 SM120 kernel headers, ordering
generic-proxy smem reads (ldmatrix/ld_shared) before the empty-barrier
arrive that hands each stage back to the TMA producer. Ports
#453 (bf16 GEMM) and fixes the same hazard found by
audit in bmk_bnk_mn, tf32_hc_prenorm, fp8_fp4_gemm_1d1d, and the
fp8/fp4 dense+paged MQA logits kernels. Sparse MQA is exempt (legacy
cp.async producer).

Perf A/B on RTX PRO 6000 (bf16 GEMM, fp8 dense/paged MQA, 3 runs): all
deltas within run-to-run noise. Fork suite: 235 passed (GPU4, fresh JIT
cache).

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>

---------

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>
Co-authored-by: Zhean Xu <94977922+zheanxu@users.noreply.github.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants