[NVIDIA] [GDN] Enable FlashInfer MTP verify on SM100+ (Blackwell) (#23273)

Co-authored-by: Yangmin Li <yangminl@nvidia.com>
Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
Shu Wang
2026-06-01 18:56:42 -07:00
committed by GitHub
co-authored by Yangmin Li Claude Opus 4.7
parent 54143264bf
commit 0574d2b8a5
4 changed files with 163 additions and 86 deletions
@@ -123,9 +123,9 @@ class GDNKernelDispatcher:
else:
raise ValueError(f"Unsupported GDN prefill backend: {prefill_backend}")
# Verify kernel: use FlashInfer only when the selected FlashInfer kernel
# supports MTP verify. On SM100+ FlashInfer GDN decode is supported, but
# its MTP verify path is not, so keep Triton as the verify fallback.
# Verify kernel: use FlashInfer when the selected FlashInfer kernel
# supports MTP verify. SM90 uses the fp32-state path; SM100 uses the
# bf16-state adapter in FlashInferGDNKernel.
if (
decode_backend.is_flashinfer() or prefill_backend.is_flashinfer()
) and flashinfer_kernel.supports_target_verify:
@@ -1,11 +1,11 @@
"""FlashInfer-based kernels for GDN (Gated Delta Network) linear attention.
Both SM90 and SM100+ use the same pool layout: [pool, HV, V, K] (K-last).
Both SM90 and SM100 use the same pool layout: [pool, HV, V, K] (K-last).
SM90 (Hopper): full support — decode, prefill, MTP. State dtype: fp32.
SM100+ (Blackwell+): decode and prefill with bf16 state. MTP verify on the way.
SM100 (Blackwell): full support — decode, prefill, MTP.
Requires flashinfer >= 0.6.4 (SM90) or >= 0.6.5 (SM100+).
Requires flashinfer >= 0.6.7.
"""
import logging
@@ -27,14 +27,15 @@ _flashinfer_gdn_available: Optional[bool] = None
_flashinfer_chunk_gated_delta_rule = None
_flashinfer_gated_delta_rule_mtp = None
_flashinfer_gated_delta_rule_decode = None
_flashinfer_gated_delta_rule_mtp_bf16 = None
def _get_flashinfer_gdn_kernels():
"""Lazy import for FlashInfer GDN prefill, decode and verify (MTP) kernels.
Returns (available, prefill_fn, mtp_fn, decode_fn).
Returns (available, prefill_fn, mtp_fn, decode_fn, mtp_bf16_fn).
"""
global _flashinfer_gdn_available, _flashinfer_chunk_gated_delta_rule, _flashinfer_gated_delta_rule_mtp, _flashinfer_gated_delta_rule_decode
global _flashinfer_gdn_available, _flashinfer_chunk_gated_delta_rule, _flashinfer_gated_delta_rule_mtp, _flashinfer_gated_delta_rule_decode, _flashinfer_gated_delta_rule_mtp_bf16
if _flashinfer_gdn_available is None:
try:
os.environ.setdefault("FLASHINFER_DISABLE_VERSION_CHECK", "1")
@@ -43,10 +44,14 @@ def _get_flashinfer_gdn_kernels():
gated_delta_rule_decode_pretranspose,
gated_delta_rule_mtp,
)
from flashinfer.gdn_kernels.gdn_decode_bf16_state import (
gated_delta_rule_mtp as gated_delta_rule_mtp_bf16,
)
from flashinfer.gdn_prefill import chunk_gated_delta_rule
_flashinfer_chunk_gated_delta_rule = chunk_gated_delta_rule
_flashinfer_gated_delta_rule_mtp = gated_delta_rule_mtp
_flashinfer_gated_delta_rule_mtp_bf16 = gated_delta_rule_mtp_bf16
_flashinfer_gated_delta_rule_decode = gated_delta_rule_decode_pretranspose
_flashinfer_gdn_available = (
torch.cuda.is_available() and torch.cuda.get_device_capability()[0] >= 9
@@ -62,6 +67,7 @@ def _get_flashinfer_gdn_kernels():
_flashinfer_chunk_gated_delta_rule,
_flashinfer_gated_delta_rule_mtp,
_flashinfer_gated_delta_rule_decode,
_flashinfer_gated_delta_rule_mtp_bf16,
)
@@ -74,9 +80,9 @@ class FlashInferGDNKernel(LinearAttnKernelBase):
"""FlashInfer kernel for GDN with K-last SSM state layout.
SM90 (Hopper): decode uses gather/scatter; prefill and MTP verify supported.
SM100+ (Blackwell+): decode and prefill supported; MTP verify not yet supported.
SM100 (Blackwell): decode uses gather/scatter; prefill and MTP verify supported.
Requires flashinfer >= 0.6.4 (SM90) or >= 0.6.5 (SM100+).
Requires flashinfer >= 0.6.7.
"""
def __init__(self):
@@ -85,6 +91,7 @@ class FlashInferGDNKernel(LinearAttnKernelBase):
self._prefill_fn,
self._mtp_fn,
self._decode_fn,
mtp_bf16_fn,
) = _get_flashinfer_gdn_kernels()
if not available:
@@ -97,13 +104,46 @@ class FlashInferGDNKernel(LinearAttnKernelBase):
sm_major = torch.cuda.get_device_capability()[0]
self.use_state_pool = sm_major >= 10
self.supports_target_verify = sm_major == 9
self.supports_target_verify = sm_major in (9, 10)
if sm_major == 9:
if self._prefill_fn is None:
raise RuntimeError("FlashInfer GDN prefill kernel is unavailable.")
if self._mtp_fn is None:
raise RuntimeError("FlashInfer GDN MTP (verify) kernel is unavailable.")
if sm_major == 9 and self._prefill_fn is None:
raise RuntimeError("FlashInfer GDN prefill kernel is unavailable.")
if self._mtp_fn is None:
raise RuntimeError("FlashInfer GDN MTP (verify) kernel is unavailable.")
if self.use_state_pool and mtp_bf16_fn is not None:
# Adapt bf16 kernel to fp32 kernel interface so target_verify needs no branching.
def _mtp_bf16_adapted(
q,
k,
v,
initial_state,
initial_state_indices,
A_log,
a,
dt_bias,
b,
use_qk_l2norm=True,
**kw,
):
out = mtp_bf16_fn(
A_log=A_log.float(),
a=a,
dt_bias=dt_bias,
softplus_beta=1.0,
softplus_threshold=20.0,
q=q,
k=k,
v=v,
b=b,
initial_state_source=initial_state,
initial_state_indices=initial_state_indices,
use_qk_l2norm_in_kernel=use_qk_l2norm,
**kw,
)
return out, None
self._mtp_fn = _mtp_bf16_adapted
logger.info("Using FlashInfer GDN kernels")
@@ -280,12 +320,7 @@ class FlashInferGDNKernel(LinearAttnKernelBase):
retrieve_parent_token: torch.Tensor,
**kwargs,
) -> torch.Tensor:
if self.use_state_pool:
raise NotImplementedError(
"FlashInfer GDN MTP verify is not yet supported on SM100+."
)
# SM90: MTP verify using FlashInfer gated_delta_rule_mtp kernel.
# MTP verify using FlashInfer gated_delta_rule_mtp kernel (SM90 + SM100+).
if retrieve_parent_token is not None:
raise RuntimeError(
"FlashInfer GDN verify kernel only supports topk=1 "
@@ -313,6 +348,13 @@ class FlashInferGDNKernel(LinearAttnKernelBase):
a_mtp = a.view(batch_size, draft_token_num, num_v_heads)
b_mtp = b.view(batch_size, draft_token_num, num_v_heads)
intermediate_states_buffer_mtp = intermediate_states_buffer
if self.use_state_pool and intermediate_states_buffer is not None:
# The SM100 bf16 MTP kernel indexes this scratch buffer by the
# per-call batch id, while SGLang's speculative state cache is
# pool-scoped and may include an extra dummy slot.
intermediate_states_buffer_mtp = intermediate_states_buffer[:batch_size]
output_fi, _ = self._mtp_fn(
q=query_mtp,
k=key_mtp,
@@ -325,7 +367,7 @@ class FlashInferGDNKernel(LinearAttnKernelBase):
b=b_mtp,
scale=None,
output=None,
intermediate_states_buffer=intermediate_states_buffer,
intermediate_states_buffer=intermediate_states_buffer_mtp,
disable_state_update=True,
use_qk_l2norm=True,
)
+3 -6
View File
@@ -3141,17 +3141,14 @@ class ServerArgs:
def _handle_linear_attn_backend(self):
import torch
# SM100+: default to FlashInfer GDN decode when the user hasn't
# explicitly chosen a decode backend and mamba-ssm-dtype is bf16
# (required by FlashInfer GDN on SM100+).
# SM100+: default to FlashInfer GDN decode (and MTP verify, via pool API)
# when the user hasn't explicitly chosen a decode backend and
# mamba-ssm-dtype is bf16 (required by FlashInfer GDN on SM100+).
# Fixed in FlashInfer v0.6.7: flashinfer-ai/flashinfer#2810
# Excluded when MTP speculative decoding is enabled because
# FlashInfer GDN MTP verify is not yet supported on SM100+.
if (
self.linear_attn_decode_backend is None
and is_sm100_supported()
and self.mamba_ssm_dtype == "bfloat16"
and self.speculative_algorithm is None
):
self.linear_attn_decode_backend = "flashinfer"
logger.info(
@@ -15,11 +15,74 @@ from sglang.test.test_utils import (
popen_launch_server,
)
register_cuda_ci(est_time=340, stage="base-c", runner_config="4-gpu-b200")
register_cuda_ci(est_time=740, stage="base-c", runner_config="4-gpu-b200")
QWEN35_FP4_MODEL = "nvidia/Qwen3.5-397B-A17B-NVFP4"
ACC_THRESHOLDS = {QWEN35_FP4_MODEL: {"gsm8k": 0.95}}
MTP_BASE_ARGS = [
"--tp-size",
"4",
"--chunked-prefill-size",
"2048",
"--mamba-scheduler-strategy",
"extra_buffer",
"--mamba-track-interval",
"128",
"--mamba-ssm-dtype",
"bfloat16",
"--max-running-requests",
"128",
"--reasoning-parser",
"qwen3",
"--attention-backend",
"trtllm_mha",
"--quantization",
"modelopt_fp4",
"--speculative-algorithm",
"NEXTN",
"--speculative-num-steps",
"3",
"--speculative-eagle-topk",
"1",
"--speculative-num-draft-tokens",
"4",
"--mem-fraction-static",
"0.8",
"--model-loader-extra-config",
'{"enable_multithread_load": true,"num_threads": 64}',
]
def _run_mtp_gsm8k(test_case):
args = SimpleNamespace(
model=test_case.model,
eval_name="gsm8k",
num_shots=5,
num_examples=200,
max_tokens=16000,
num_threads=128,
repeat=1,
temperature=0.6,
top_p=0.95,
top_k=20,
base_url=test_case.base_url,
host="http://127.0.0.1",
port=int(test_case.base_url.split(":")[-1]),
)
metrics = run_eval(args)
print(f"{metrics=}")
test_case.assertGreaterEqual(
metrics["score"], ACC_THRESHOLDS[test_case.model]["gsm8k"]
)
server_info = requests.get(test_case.base_url + "/server_info")
avg_spec_accept_length = server_info.json()["internal_states"][0][
"avg_spec_accept_length"
]
print(f"{avg_spec_accept_length=}")
test_case.assertGreater(avg_spec_accept_length, 3.3)
class TestQwen35FP4MTP(ReasoningTokenUsageMixin, CustomTestCase):
reasoning_parser_name = "qwen3"
@@ -34,37 +97,36 @@ class TestQwen35FP4MTP(ReasoningTokenUsageMixin, CustomTestCase):
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--tp-size",
"4",
"--chunked-prefill-size",
"2048",
"--mamba-scheduler-strategy",
"extra_buffer",
"--mamba-track-interval",
"128",
"--mamba-ssm-dtype",
"bfloat16",
"--max-running-requests",
"128",
"--reasoning-parser",
"qwen3",
"--attention-backend",
"trtllm_mha",
"--quantization",
"modelopt_fp4",
"--speculative-algorithm",
"NEXTN",
"--speculative-num-steps",
"3",
"--speculative-eagle-topk",
"1",
"--speculative-num-draft-tokens",
"4",
"--mem-fraction-static",
"0.8",
"--model-loader-extra-config",
'{"enable_multithread_load": true,"num_threads": 64}',
other_args=MTP_BASE_ARGS,
)
@classmethod
def tearDownClass(cls):
envs.SGLANG_ENABLE_SPEC_V2.set(False)
kill_process_tree(cls.process.pid)
def test_gsm8k(self):
_run_mtp_gsm8k(self)
class TestQwen35FP4MTPFlashInfer(ReasoningTokenUsageMixin, CustomTestCase):
reasoning_parser_name = "qwen3"
@classmethod
def setUpClass(cls):
cls.model = QWEN35_FP4_MODEL
cls.base_url = DEFAULT_URL_FOR_TEST
cls.init_reasoning_token_verifier()
envs.SGLANG_ENABLE_SPEC_V2.set(True)
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=MTP_BASE_ARGS
+ [
"--linear-attn-decode-backend",
"flashinfer",
"--enforce-disable-flashinfer-allreduce-fusion",
],
)
@@ -74,31 +136,7 @@ class TestQwen35FP4MTP(ReasoningTokenUsageMixin, CustomTestCase):
kill_process_tree(cls.process.pid)
def test_gsm8k(self):
args = SimpleNamespace(
model=self.model,
eval_name="gsm8k",
num_shots=5,
num_examples=200,
max_tokens=16000,
num_threads=128,
repeat=1,
temperature=0.6,
top_p=0.95,
top_k=20,
base_url=self.base_url,
host="http://127.0.0.1",
port=int(self.base_url.split(":")[-1]),
)
metrics = run_eval(args)
print(f"{metrics=}")
self.assertGreaterEqual(metrics["score"], ACC_THRESHOLDS[self.model]["gsm8k"])
server_info = requests.get(self.base_url + "/server_info")
avg_spec_accept_length = server_info.json()["internal_states"][0][
"avg_spec_accept_length"
]
print(f"{avg_spec_accept_length=}")
self.assertGreater(avg_spec_accept_length, 3.3)
_run_mtp_gsm8k(self)
if __name__ == "__main__":