Skip to content

Improve SM100 MX MQA logits clean-output performance - #436

Open
OutWhite wants to merge 1 commit into
deepseek-ai:mainfrom
OutWhite:fix/nqlin/sm100-indexer-clean-perf
Open

OutWhite wants to merge 1 commit into
deepseek-ai:mainfrom
OutWhite:fix/nqlin/sm100-indexer-clean-perf

Conversation

@OutWhite

@OutWhite OutWhite commented Sep 10, 2026 •

Copy link
Copy Markdown

On GB200 with NVCC 13.3.33, placing MXFP8/MXFP4 MQA-logits cleanup after the CTA's main computation improves clean_logits=True performance by about 1.6–2.3x for the H64/D128 cases below, with identical complete outputs. This PR proposes that scheduling change for the SM100 MX paths.

Keep the existing inline cleaner block for FP8. For MXFP8/MXFP4, run the same cleaning loop at the end of the core, after the CTA's producer/consumer work and TMEM release, using a separate named barrier (ID 1) from the math-only barrier (ID 0). Both paths support clean_logits=True; FP8 stays concurrent because deferring it caused a 5–8% slowdown in some exploratory cases. The two inline blocks preserve the desired code placement without introducing a helper or lambda. The cleanup ranges, in-coverage masks, metadata layout, and public API are unchanged.

Performance

Baseline: 66081d4c9c7d7c44f13fea402e5b622aa0f409c2. Same GB200 GPU (152 SMs), NVCC 13.3.33, default DeepGEMM JIT flags targeting sm_100f, driver 580.167.08, PyTorch 2.13.0+cu130; PDL disabled.

H=64, D=128, FP32 weights/output, uncompressed logits. Q/K/weights are all zero, MX scales are one, and each row uses [0, KV-Q+row). These are CUDA Graph GPU timings of the complete clean_logits=True call, including all cleaning. Metadata is precomputed outside timing. Five calls per graph with distinct retained output buffers; 40 samples, rotating/reversing variant order, immediate warmup before each timed replay; median per-call latency. The timing runs use normal GPU clocks. Separate Nsight captures below use controlled base clocks.

Q / KV tokens Format Metadata Before (us) After (us) Speedup
2048 / 32768 MXFP8 no 645.9 412.0 1.57x
2048 / 32768 MXFP8 yes 817.7 406.3 2.01x
2048 / 32768 MXFP4 no 806.4 407.9 1.98x
2048 / 32768 MXFP4 yes 780.2 403.0 1.94x
8192 / 65536 MXFP8 no 4994.6 3022.2 1.65x
8192 / 65536 MXFP8 yes 7093.2 3083.6 2.30x
8192 / 65536 MXFP4 no 5990.3 2968.9 2.02x
8192 / 65536 MXFP4 yes 6008.1 3054.5 1.97x

The clean_logits=False control stays within 0.2% between builds for these eight points. A further 18-shape screen (both metadata modes) includes H=8/32/64, Q=KV=8192, CP windows, and BF16 weights/output: MX latencies improve in every screened case; FP8 is within 0.2% of baseline. The improvement also appears with random inputs. The table uses an all-zero workload to reduce input-dependent power effects.

The deferred schedule is an empirical tuning choice. When changing SM100 hardware or NVCC, remeasure both schedules with the complete-output benchmark below. Prefer concurrent cleanup if it matches or outperforms the deferred schedule within measurement variability on the target workloads; the full-CTA barrier is required by the deferred schedule, not by the MQA operation itself.

Timing data, including p10/p90 (microseconds)
Q,KV,format,metadata,before_median_us,before_p10_us,before_p90_us,after_median_us,after_p10_us,after_p90_us
2048,32768,mxfp8,0,645.859,636.832,675.693,411.974,411.77,412.186
2048,32768,mxfp8,1,817.683,812.563,825.261,406.333,406.227,406.65
2048,32768,mxfp4,0,806.432,806.042,806.842,407.888,407.469,408.301
2048,32768,mxfp4,1,780.218,779.808,780.218,402.966,402.566,403.59
8192,65536,mxfp8,0,4994.589,4950.765,5020.403,3022.163,3021.344,3022.982
8192,65536,mxfp8,1,7093.178,7083.347,7106.08,3083.603,3082.784,3084.013
8192,65536,mxfp4,0,5990.317,5989.088,5991.52,2968.915,2966.861,2970.547
8192,65536,mxfp4,1,6008.093,6007.13,6008.755,3054.506,3052.454,3056.557

Investigation notes

A few local ablations helped distinguish cleanup bandwidth costs from sensitivity to code placement. These use H64/D128, Q=2048, KV=32768:

  • Compile out only the cleaner warp's work, keeping the main-loop invalid-column mask: MXFP8 drops from 641.3 to 397.2 us; MXFP4 from 806.0 to 416.0 us. This is a diagnostic variant, not a complete-output implementation.
  • Keep the exact same compiled kernel and disable the cleaner via a runtime gate: MXFP8 is 633.1 us with cleanup versus 632.8 us without; MXFP4 is 805.1 versus 805.5 us. The gate toggles the first row's end by one element; split counts and metadata remain unchanged (metadata was checked bitwise). This suggests that the stores themselves do not explain the timing difference.
  • Removing only the main-loop output mask gives similar no-metadata timings to baseline: 635.1 us for MXFP8 and 800.3 us for MXFP4.
  • Changing the ordering of the specialized-warp branches also changes performance, including for clean_logits=False. The cleaner-loop unrolling and inline-PTX store variants showed similar timings to baseline.
  • Keeping the cleaner in its original source location and adding a full-CTA named barrier before it, with the other warps joining at the end of the core, does not recover performance: MXFP8/MXFP4 take 645.6/813.4 us without metadata, versus 412.4/408.3 us with the inline tail placement. Thus the runtime ordering alone is insufficient for this tested build.

With the final patch, a controlled-clock Nsight Compute 2026.2 comparison for MXFP8, Q=2048, KV=32768, no metadata shows:

Metric Before After
SM clock (GHz) 1.132 1.128
Registers / thread 168.000 168.000
Dynamic shared memory (decimal KB / CTA) 222.208 222.208
DRAM reads (decimal MB) 22.218 22.249
DRAM writes (decimal MB) 211.368 209.558
Long-scoreboard stalls (per_issue_active ratio) 3.519 1.529
Tensor-pipe active (% of elapsed cycles) 34.927 53.430

Exploratory PC samples place much of the additional wait around the math-side TMEM-result barrier check. Together, these results suggest that cleaner placement affects main-loop execution; memory traffic and resource allocation are largely unchanged, and neither binary has register spills. The precise low-level dependency remains to be investigated. The proposed change avoids this sensitivity on the tested configuration while retaining the existing output semantics.

Validation

  • Current revision: all 48 complete outputs are bitwise identical to the previously validated PR version, with full mask checks and comparison against a quantized PyTorch reference. Covers FP8/MXFP8/MXFP4, FP32/BF16 weights and output, both metadata modes, random inputs, full/causal and irregular windows, nonzero starts, empty rows, non-aligned lengths (Q=129/KV=4100), and all-zero inputs.
  • Current revision: existing upstream test_paged_mqa_logits passes all 4 sampled configurations, exercising the shared core with paged schedulers.
  • For all 12 H64/D128 specializations (three formats, both metadata modes, cleanup on/off), the current inline-block revision produces identical SASS instructions and scheduling control words to the previously validated PR revision. Complete-call timings at Q=2048/KV=32768 agree within 0.1%.
  • The initial implementation also passed 32 sampled upstream test_mqa_logits configurations and Compute Sanitizer synccheck/memcheck on 32 compact MX output configurations. Those broader checks were not rerun for this source cleanup.

No permanent test-matrix or benchmark-file additions. Existing upstream tests can be run from a built checkout with:

CUDA_VISIBLE_DEVICES=0 DG_MQA_NUM_CASES=32 python - <<'PY'
import os, random, sys, torch, deep_gemm
sys.path.insert(0, 'tests')
import test_attention
torch.manual_seed(0)
random.seed(0)
deep_gemm.set_pdl(False)
test_attention.test_mqa_logits()
os.environ['DG_MQA_NUM_CASES'] = '4'
test_attention.test_paged_mqa_logits()
PY
Standalone performance reproducer (writes every timing sample to JSON)

Save as pr-reproduce.py and run it once with each built revision, using the same GPU and NVCC. Both builds and all JIT compilations must finish before timed sampling; the script warms every specialization before capture. It verifies the zero output and complete negative-infinity mask before timing.

"""Run with the baseline and PR DeepGEMM builds on the same GPU/toolchain.

CUDA_VISIBLE_DEVICES=0 python pr-reproduce.py --output before.json
CUDA_VISIBLE_DEVICES=0 python pr-reproduce.py --output after.json
"""
import argparse
import gc
import json
import statistics
from functools import partial

import torch
import deep_gemm as dg


def measure(functions, samples=40, calls=5):
    graphs = {}
    for name, fn in functions.items():
        for _ in range(3):
            fn()
        torch.cuda.synchronize()
        graph = torch.cuda.CUDAGraph()
        with torch.cuda.graph(graph):
            outputs = [fn() for _ in range(calls)]
        graphs[name] = (graph, outputs)
        for _ in range(5):
            graph.replay()
    torch.cuda.synchronize()
    begin, end = [torch.cuda.Event(enable_timing=True) for _ in range(2)]
    times = {name: [] for name in functions}
    names = list(functions)
    for sample in range(samples):
        order = names[sample % len(names):] + names[:sample % len(names)]
        if sample % 2:
            order.reverse()
        for name in order:
            graphs[name][0].replay()
            begin.record()
            graphs[name][0].replay()
            end.record()
            end.synchronize()
            times[name].append(begin.elapsed_time(end) * 1000 / calls)
    return {
        name: {
            'median_us': statistics.median(values),
            'p10_us': sorted(values)[int(.1 * (samples - 1))],
            'p90_us': sorted(values)[int(.9 * (samples - 1))],
            'samples_us': values,
        }
        for name, values in times.items()
    }


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument('--output', required=True)
    args = parser.parse_args()
    dg.set_pdl(False)
    result = []
    for m, n in ((2048, 32768), (8192, 65536)):
        for fmt in ('mxfp8', 'mxfp4'):
            fp4 = fmt == 'mxfp4'
            q = torch.zeros((m, 64, 64 if fp4 else 128), device='cuda',
                            dtype=torch.int8 if fp4 else torch.float8_e4m3fn)
            k = torch.zeros((n, q.shape[-1]), device='cuda', dtype=q.dtype)
            qs = torch.full((m, 64), 0x7f7f7f7f, device='cuda', dtype=torch.int32)
            ks = torch.full((n,), 0x7f7f7f7f, device='cuda', dtype=torch.int32)
            w = torch.zeros((m, 64), device='cuda', dtype=torch.float32)
            starts = torch.zeros(m, device='cuda', dtype=torch.int32)
            ends = torch.arange(m, device='cuda', dtype=torch.int32) + n - m
            metadata = dg.get_mqa_logits_metadata(starts, ends, n, 64)
            functions = {
                f'clean{int(clean)}_meta{int(use_meta)}': partial(
                    dg.fp8_fp4_mqa_logits, (q, qs), (k, ks), w, starts, ends,
                    clean_logits=clean, logits_dtype=torch.float32,
                    schedule_meta=metadata if use_meta else None,
                )
                for clean in (False, True) for use_meta in (False, True)
            }
            for use_meta in (False, True):
                out = functions[f'clean1_meta{int(use_meta)}']()
                cols = torch.arange(n, device='cuda')[None, :]
                for row in range(0, m, 128):
                    valid = cols < ends[row:row + 128, None]
                    chunk = out[row:row + 128]
                    assert torch.equal(chunk.isneginf(), ~valid)
                    assert (chunk.masked_fill(~valid, 0) == 0).all()
                del out
            timings = measure(functions)
            result.append(dict(shape=[m, n, 64, 128], format=fmt, timings=timings))
            with open(args.output, 'w') as f:
                json.dump(result, f, indent=2)
            print(m, n, fmt, {k: round(v['median_us'], 3) for k, v in timings.items()}, flush=True)
            del functions, q, k, qs, ks, w, starts, ends, metadata, chunk, valid, cols
            gc.collect()
            torch.cuda.empty_cache()


if __name__ == '__main__':
    main()

Comment thread deep_gemm/include/deep_gemm/impls/sm100_mqa_logits.cuh Outdated
Comment thread deep_gemm/include/deep_gemm/impls/sm100_mqa_logits.cuh
Comment thread deep_gemm/include/deep_gemm/impls/sm100_mqa_logits.cuh
Comment thread deep_gemm/include/deep_gemm/impls/sm100_mqa_logits.cuh Outdated
Comment thread deep_gemm/include/deep_gemm/impls/sm100_mqa_logits.cuh Outdated
Comment thread deep_gemm/include/deep_gemm/impls/sm100_mqa_logits.cuh
@ds-review-bot

Copy link
Copy Markdown
Collaborator

🤖 ds-review-bot Code Review

v6

未发现本次变更新引入的明确缺陷。MX 清理前的全 CTA 同步与数学线程 barrier 相互独立,清理范围及 FP8、paged 路径语义保持不变。当前环境缺少 NVCC 和 GPU,未复跑 CUDA 测试或性能验证。

v5

The MR moves MXFP8/MXFP4 MQA-logits cleanup in sm100_mqa_logits_core_impl from the concurrent specialized warp to after the CTA's producer/consumer work and TMEM release, reusing the existing cleaner via a kCleanLogits-guarded local lambda and synchronizing on a separate named barrier (user ID 1) that is distinct from the math-only barrier (user ID 0). FP8 keeps concurrent cleanup. Verified: (1) CUTLASS NamedBarrier(num_threads, id) offsets user IDs by ReservedNamedBarrierCount, so user IDs 0/1 do not collide with each other or with the hardware barrier used by __syncthreads(); (2) all five warp roles fall through to the common tail with no early returns, and producer/MMA warps end their branches with __syncwarp(), so every lane reaches barrier 1; the barrier thread count 128 + kNumMathThreads is a multiple of 32 as bar.sync requires; (3) the paged entry passes kCleanLogits=false, so the lambda body (which calls the contiguous-only make_cleaner()) is discarded and never instantiated for SM100PagedMQALogitsScheduler; (4) the lambda's by-reference captures (lane_idx, sm_idx, logits, logits_stride, make_scheduler) are live for the whole kernel; (5) the cleaner still runs on the warp that executed warpgroup_reg_dealloc&lt;56&gt;, so register budgets are unchanged; (6) cleanup ranges, the main-loop -inf mask, LogitsCleaner::fill_row semantics, and the public API are identical to baseline. The change is minimal, correct, and well-commented. Could not compile or run on hardware from this environment; performance claims rely on the MR's reported measurements and tests. Approve; only two minor comment/documentation suggestions.

v4

This MR relocates MXFP8/MXFP4 clean_logits cleanup out of the specialized cleaner warp's concurrent branch and runs it once after all producer/consumer work and TMEM release, synchronized on a dedicated user named barrier (ID 1). The existing grid-stride cleaner is factored into a local clean_logits lambda that both the concurrent FP8 path and the deferred MX path reuse, and the lambda's if constexpr (kCleanLogits) guard prevents the contiguous-only make_cleaner() from being instantiated for the paged scheduler. The cleanup ranges, in-coverage masks, metadata layout and public API are unchanged, and the MR reports bitwise-identical outputs together with 1.6-2.3x MX speedups on GB200/NVCC 13.3.33. The implementation looks correct: the deferred writes are disjoint from the math writes, barrier 1 is a separate hardware barrier from the math-only barrier 0 so the two full/partial CTA synchronizations cannot deadlock, and every warp role reaches the new barrier. My comments are about the toolchain-specific nature of the tuning, maintainability of the scheduler coupling, and the loss of a permanent regression guard rather than correctness.

Files reviewed: 1
Issues found: 🟡 1 warning | 🔵 5 suggestion
Inline comments posted: 6

@OutWhite
OutWhite force-pushed the fix/nqlin/sm100-indexer-clean-perf branch from f2f2a0d to 7710a5b Compare September 16, 2026 12:00
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.

2 participants