Skip to content

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

Closed
OutWhite wants to merge 2 commits into
deepseek-ai:mainfrom
OutWhite:codex/sm100-mx-mqa-clean-performance
Closed

OutWhite wants to merge 2 commits into
deepseek-ai:mainfrom
OutWhite:codex/sm100-mx-mqa-clean-performance

Conversation

@OutWhite

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.

Move the MX cleaner after the CTA's producer/consumer work and TMEM release. Reuse the existing cleaner through a local lambda, and synchronize with a separate named barrier (ID 1) so it cannot overlap the math-only barrier (ID 0). The grid-stride cleanup ranges, in-coverage masks, metadata layout, and public API are unchanged. FP8 retains concurrent cleanup; moving FP8 cleanup as well caused a 5–8% slowdown in some exploratory cases. The lambda's kCleanLogits guard also avoids instantiating contiguous-only cleaner methods for paged schedulers.

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.

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.

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

  • Existing upstream test_mqa_logits: 32 sampled configurations pass (including its 20-repeat bitwise checks, scheduled-path equality, full mask checks, and reference comparisons).
  • Existing upstream test_paged_mqa_logits: 4 sampled configurations pass, exercising the shared core with paged schedulers.
  • One-off comparison: all 48 complete outputs are bitwise identical to baseline. 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; also checked against a quantized PyTorch reference.
  • Compute Sanitizer on 32 compact MX output configurations (Q=17/KV=260, both metadata modes, FP32/BF16, varied masks): synccheck PASS; memcheck PASS.

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
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
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

未发现本次变更新引入的明确缺陷。清理范围保持不变,新增屏障覆盖全部 CTA 线程且不与数学线程屏障冲突,FP8 与 paged 路径保持原有行为。当前环境未找到 CUDA 编译器或 GPU 工具,未复跑正确性和性能测试。

v5

The change moves the MXFP8/MXFP4 -inf logits cleanup out of the concurrent specialized-warp branch and runs it after the CTA's producer/consumer work and TMEM release, behind a new full-CTA named barrier (ID 1). FP8 keeps the concurrent cleaner. The cleaner body is hoisted verbatim into a clean_logits lambda guarded by if constexpr (kCleanLogits).

Verified (no action needed):

  • Correctness: the cleaner's ranges [0, kv_base) and [coverage_end, logits_stride) are disjoint from the math warps' [kv_base, coverage_end) stores, so serializing them cannot change output (consistent with the claimed bitwise-identical results). The lambda body is byte-for-byte the original cleaner loop (same grid-stride scheduler, same coverage_end clamp, same two fill_row ranges).
  • Barrier reachability: every branch of the warp-role if/else if chain (Q producer, KV producer, MMA warp, cleaner warp, math warps) falls through to the trailing if constexpr (kCleanLogits and kIsMXSF) block, so all kNumSpecializedThreads + kNumMathThreads (= 384 in the current launcher) threads arrive at barrier 1. Count is a multiple of 32 and <= 1024.
  • Barrier IDs: cutlass::arch::NamedBarrier(n, id) adds ReservedNamedBarrierCount to user IDs (checked in third-party/cutlass/include/cutlass/arch/barrier.h), so ID 0 (pre-existing math-only sync) and ID 1 (new full-CTA sync) are distinct from each other and from __syncthreads()/hardware barrier 0.
  • Paged path: SM100PagedMQALogitsScheduler has no make_cleaner; the paged entry passes kCleanLogits=false and the if constexpr inside the lambda discards the body, so the contiguous-only method is never instantiated. make_cleaner exists only in scheduler/sm100_mqa_logits.cuh.
  • Host launcher (csrc/jit_kernels/impls/sm100_mqa_logits.hpp) and the public Python API need no changes; template signature and block dims are untouched.
  • Lambda captures (lane_idx, sm_idx, make_scheduler, logits, logits_stride) are by reference to values that outlive the call; the cleaner warp still runs under the reduced 56-register budget as before.

No blocking issues found; approve. Two non-blocking suggestions are attached. Performance claims (1.6-2.3x on GB200 for MX clean_logits=True, FP8 within 0.2%) could not be re-run in this environment and are taken from the PR description; the change is low-risk with respect to output semantics.

v4

Reviewed the single-file scheduling change in deep_gemm/include/deep_gemm/impls/sm100_mqa_logits.cuh (30 insertions / 16 deletions vs baseline 66081d4). The patch is correct and semantics-preserving. It (1) hoists the existing logits cleaner into a local lambda clean_logits whose body is guarded by if constexpr (kCleanLogits), (2) keeps FP8 cleanup concurrent in the specialized cleaner warp kSpecWarpStart + 3 via if constexpr (kCleanLogits and not kIsMXSF), and (3) moves MXFP8/MXFP4 cleanup to after the producer/consumer branches and the math-only TMEM release, gated by a new full-CTA named barrier with ID 1. The barrier separation from the math-only barrier (ID 0) is correct: barrier 1 waits for all kNumSpecializedThreads + kNumMathThreads, so the MX cleaner can no longer overlap the main loop. The grid-stride cleanup ranges, in-coverage masks, metadata layout and public API are untouched. The lambda's inner if constexpr is load-bearing: it prevents the contiguous-only make_cleaner/LogitsCleaner path from being instantiated for paged schedulers, which pass kCleanLogits=false. I found no correctness bug, no deadlock path (there are no early returns before the new barrier), and the barrier thread count matches the launch bounds exactly. The reported 1.6-2.3x cleanup speedup and the clean_logits=False control (<0.2%) are consistent with the change removing cleaner/main-loop interference.

Files reviewed: 1
Issues found: 🔵 5 suggestion
Inline comments posted: 5

@OutWhite

Copy link
Copy Markdown
Author

Continued in #436 using the intended source branch name. The implementation and validation are unchanged; this PR retains the earlier review discussion.

@OutWhite OutWhite closed this Sep 10, 2026
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