Skip to content

[SM90] Support dynamic top-k for V3.2 sparse decode - #216

Open
Livinfly wants to merge 1 commit into
deepseek-ai:mainfrom
Livinfly:test/sm90-v32-dynamic-topk-upstream
Open

Livinfly wants to merge 1 commit into
deepseek-ai:mainfrom
Livinfly:test/sm90-v32-dynamic-topk-upstream

Conversation

@Livinfly

Copy link
Copy Markdown

Summary

This PR adds optional per-request topk_length support to the SM90
DeepSeek-V3.2 FP8 sparse decode kernel (d_qk=576). SM90 V3.2 can now use
dynamic top-k instead of always processing the full allocated top-k capacity.

When topk_length is provided, the kernel:

  • loads the valid length for each batch row;
  • schedules only max(ceil_div(topk_length, TOPK_BLOCK_SIZE), 1) blocks;
  • masks a partial tail before its indices can be used for KV addressing; and
  • includes the absolute top-k position in the softmax validity predicate.

When topk_length == nullptr, the existing fixed top-k behavior is preserved.
There is no public API change, and MODEL1 and SM100 behavior are unchanged.

Why template<bool DYNAMIC_TOPK, ...>

I also tested a single-device-kernel implementation that checks
params.topk_length != nullptr at runtime. The condition is grid-uniform, but
the compiler cannot remove the length load, tail predicates, extra control
flow, or their value lifetimes from the fixed path.

template<bool DYNAMIC_TOPK> moves the choice to host launch dispatch:

  • false: the dynamic-length logic is compiled out, preserving the fixed path;
  • true: per-row block counts and partial-tail safety checks are enabled.

The comparison implementation is available at
7b7c831 on
test/sm90-v32-dynamic-topk-runtime-device. It regresses the representative
B64-B128 fixed cases by 8.4%-9.4%, while the specialized implementation stays
within 0.5% of the baseline. Both builds report 168 registers, 12 barriers,
and no spills, so the generated hot-path instructions—not an occupancy
change—explain the observed difference.

Correctness

python tests/test_flash_mla_sparse_decoding.py

Focused reference checks passed for:

Case Result
V3.2 dynamic H64 / H128 pass
Existing MODEL1 dynamic H64 / H128 pass
V3.2 tail lengths 0,1,17,63,64,65,127,128 pass

For the tail-safety check, every suffix entry after topk_length was replaced
with INT_MAX. H64 and H128 matched the corresponding safe-tail results, with
no illegal memory access.

Kernel performance

Environment: NVIDIA H100 80GB HBM3 (SM90), driver 580.126.09, nvcc 12.8.93,
and PyTorch 2.13.0+cu129. Each number is the median of seven group means, with
10 warmups and 20 timed calls per group using CUDA events. Scheduler metadata
is reused and excluded from timing.

These are standalone FlashMLA sparse-decode measurements; they exclude input
generation, index compaction, framework integration, collectives, and model
execution. No end-to-end speedup is claimed.

Compared revisions:

Build Revision
Upstream baseline 15f13e5
Runtime device check 7b7c831
Final specialization cfc99c2

Fixed top-k regression check

Shape: S_q=2, H=128, d_qk=576, s_kv=32768, page block 64,
variable KV lengths, topk=2048. Positive deltas are regressions.

Batch Baseline (us) Runtime check (us) Delta Specialized (us) Delta
2 25.10 24.99 -0.46% 25.11 +0.03%
64 172.93 189.21 +9.41% 172.90 -0.02%
74 214.30 232.31 +8.41% 213.79 -0.24%
128 339.07 369.87 +9.08% 337.39 -0.50%
Geomean - - +6.53% - -0.18%

For B64-B128, the runtime-check version regresses by 8.97% geomean. The final
specialized version is -0.25% versus baseline, showing no fixed-path
regression within measurement variation.

Dynamic top-k

Shape: [B=256, S_q=1, H=128, d_qk=576], s_kv=4096, page block 64,
allocated top-k 2048, effective per-row top-k 1024. The fixed input uses an
invalid suffix; the dynamic input passes topk_length=1024 and poisons the
ignored suffix with INT_MAX.

Implementation/path Median (us) Group range (us) Delta vs own fixed Speedup
Baseline fixed 319.76 319.21-320.08 - -
Runtime-check fixed 346.96 346.59-348.54 - -
Runtime-check dynamic 202.78 202.44-203.22 -41.56% 1.71x
Specialized fixed 319.68 319.48-320.52 - -
Specialized dynamic 188.80 188.47-189.02 -40.94% 1.69x

The final dynamic path saves 130.89 us per call versus its fixed path. It is
also 6.90% faster than the shared runtime-check dynamic path, while
specialization avoids the latter's 8.51% fixed-path regression in this shape.

Reproduction

Build each worktree with:

FLASH_MLA_DISABLE_SM100=1 \
MAX_JOBS=32 \
NVCC_THREADS=4 \
python setup.py build_ext --inplace

Save the folded harness below as bench_sm90_v32_dynamic_topk.py, then
run it in a fresh process for each build:

# Fixed production sweep: run on baseline, runtime-check, and final worktrees.
python bench_sm90_v32_dynamic_topk.py \
  --repo <worktree> --label <build-label> --suite fixed

# Same-layout fixed/dynamic comparison: run on runtime-check and final worktrees.
python bench_sm90_v32_dynamic_topk.py \
  --repo <worktree> --label <build-label> --suite row-batch --mode both
Benchmark harness (click to expand)
#!/usr/bin/env python
import argparse
import gc
import json
import statistics
import sys
from pathlib import Path
from typing import Callable


parser = argparse.ArgumentParser()
parser.add_argument("--repo", type=Path, required=True)
parser.add_argument("--label", required=True)
parser.add_argument("--suite", choices=("fixed", "row-batch"), required=True)
parser.add_argument("--mode", choices=("fixed", "dynamic", "both"), default="fixed")
parser.add_argument("--warmup", type=int, default=10)
parser.add_argument("--groups", type=int, default=7)
parser.add_argument("--repeats", type=int, default=20)
args = parser.parse_args()

repo = args.repo.resolve()
sys.path[:0] = [str(repo), str(repo / "tests")]

import flash_mla  # noqa: E402
import lib  # noqa: E402
import torch  # noqa: E402
from lib import RawTestParamForDecode as RawTestParam  # noqa: E402

torch.set_default_dtype(torch.bfloat16)
torch.set_default_device("cuda:0")
torch.cuda.set_device(0)
torch.set_float32_matmul_precision("high")
torch.set_num_threads(32)


def measure(fn: Callable) -> list[float]:
    for _ in range(args.warmup):
        fn()
    torch.cuda.synchronize()
    samples = []
    for _ in range(args.groups):
        start = torch.cuda.Event(enable_timing=True)
        end = torch.cuda.Event(enable_timing=True)
        start.record()
        for _ in range(args.repeats):
            fn()
        end.record()
        end.synchronize()
        samples.append(start.elapsed_time(end) * 1000 / args.repeats)
    return samples


def run(p, testcase, case: str, mode: str, expected=None):
    metadata, _ = flash_mla.get_mla_metadata()

    def call():
        return lib.run_flash_mla_decode(p, testcase, metadata, None)

    output = call()
    torch.cuda.synchronize()
    if expected is not None:
        torch.testing.assert_close(output[0], expected[0], atol=1e-3, rtol=2.01 / 128)
        torch.testing.assert_close(output[1], expected[1], atol=1e-6, rtol=8.01 / 65536)
        print("fixed/dynamic output and LSE: pass", flush=True)

    samples = measure(call)
    print(json.dumps({
        "label": args.label,
        "case": case,
        "mode": mode,
        "median_us": statistics.median(samples),
        "min_us": min(samples),
        "max_us": max(samples),
        "samples_us": samples,
    }, sort_keys=True), flush=True)
    return output


def fixed_suite():
    for batch in (2, 64, 74, 128):
        torch.cuda.empty_cache()
        p = RawTestParam(
            batch, 128, 2, 1, 32768, True, topk=2048, d_qk=576,
            check_correctness=False, num_runs=0, seed=20260830,
        ).to_test_param()
        testcase = lib.generate_testcase_for_decode(p)
        run(p, testcase, f"B{batch}_Sq2_H128_K2048", "fixed")
        del testcase
        gc.collect()
        torch.cuda.empty_cache()


def configure(testcase, original_indices, mode: str):
    testcase.kv_scope.indices_in_kvcache.copy_(original_indices)
    if mode == "fixed":
        testcase.kv_scope.indices_in_kvcache[..., 1024:] = -1
        testcase.kv_scope.topk_length = None
    else:
        testcase.kv_scope.indices_in_kvcache[..., 1024:] = 2147483647
        testcase.kv_scope.topk_length = torch.full(
            (256,), 1024, dtype=torch.int32, device="cuda"
        )


def row_batch_suite():
    p = RawTestParam(
        256, 128, 1, 1, 4096, False, topk=2048, d_qk=576,
        check_correctness=False, num_runs=0, seed=20260830,
    ).to_test_param()
    testcase = lib.generate_testcase_for_decode(p)
    original_indices = testcase.kv_scope.indices_in_kvcache.clone()
    case = "B256_Sq1_H128_K2048_effective1024"

    fixed_output = None
    if args.mode in ("fixed", "both"):
        configure(testcase, original_indices, "fixed")
        fixed_output = run(p, testcase, case, "fixed")
    elif args.mode == "dynamic":
        configure(testcase, original_indices, "fixed")
        metadata, _ = flash_mla.get_mla_metadata()
        fixed_output = lib.run_flash_mla_decode(p, testcase, metadata, None)

    if args.mode in ("dynamic", "both"):
        configure(testcase, original_indices, "dynamic")
        run(p, testcase, case, "dynamic", fixed_output)


if args.suite == "fixed":
    fixed_suite()
else:
    row_batch_suite()

Scope

This change only enables dynamic top-k for the existing SM90 V3.2 sparse
decode specialization. It does not compact selected indices; callers that
provide topk_length must place valid indices in a prefix. Framework-side
layout conversion and integration should be benchmarked separately.

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.

1 participant