Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
38 commits
Select commit Hold shift + click to select a range
933876c
feat: add GLM-5.3-Flash support
ZJY0516 Aug 26, 2026
54d2987
Merge branch 'main' into glm-release
ZJY0516 Aug 27, 2026
6f03690
Fix GLM merge regressions
ZJY0516 Aug 27, 2026
724f073
[ROCm] Support GLM-5.3-Flash on AMD GPUs (#5)
andyluo7 Aug 27, 2026
142062f
Fix video loader and NIXL test regressions
ZJY0516 Aug 27, 2026
7d464af
Merge branch 'main' into glm-release
ZJY0516 Aug 27, 2026
031c899
[Model] Use FlashInfer for GLM-5.3-Flash NoPE MLA
ZJY0516 Aug 27, 2026
0f82208
[Model] Fix GLM-5.3-Flash CI regressions
ZJY0516 Aug 27, 2026
2a5e8af
Merge branch 'main' into glm-release
ZJY0516 Aug 28, 2026
e34f7f7
[Kernel] Use FlashInfer rc10 for NoPE MLA
ZJY0516 Aug 28, 2026
3c01e04
[Model] Refine FlashInfer MLA integration
ZJY0516 Aug 28, 2026
878631b
[Model] Harden GLM-5.3 FlashInfer integration
ZJY0516 Aug 28, 2026
81ed4c3
[Bugfix][NIXL] Register compressed MLA transfer pages
izhuhaoran Aug 28, 2026
b7274aa
Merge branch 'main' into glm-release
ZJY0516 Aug 30, 2026
af282b7
Align kpool indexer block to a 64 wide pool page on sm120 (#10)
ima-helikoptaaa Aug 30, 2026
36bb379
[Bugfix] Ignore cache opt-outs in partial-hit gating
ZJY0516 Aug 30, 2026
7cf764c
Merge branch 'main' into glm-release
ZJY0516 Aug 31, 2026
f221389
[Kernel][Perf] Simplify GLM-5.3 short decode indexer path (#12)
ZJY0516 Aug 31, 2026
1f0369b
Merge branch 'main' into glm-release
ZJY0516 Aug 31, 2026
f3ba9c0
[Bugfix][KV Connector] Mooncake Store: fine-grained lookup + skip non…
JaredforReal Aug 31, 2026
ffbcee2
[Bugfix] Fix GLM release test regressions
ZJY0516 Sep 1, 2026
342c54d
Merge branch 'main' into glm-release
ZJY0516 Sep 1, 2026
ea506fd
[Bugfix] Preserve SWA KV cache groups in Mooncake
ZJY0516 Sep 1, 2026
4170d83
[Bugfix][ROCm] Narrow MLA page geometry types
ZJY0516 Sep 1, 2026
3dd9f06
[Model] Remove unused smart image background toggle
ZJY0516 Sep 1, 2026
010922c
[Bugfix] Fix GLM MTP and ROCm FlashInfer tests
ZJY0516 Sep 1, 2026
89b43ba
Merge branch 'main' into glm-release
ZJY0516 Sep 1, 2026
c4440e6
[Model] Move GLM KDA kernels under model directory
ZJY0516 Sep 1, 2026
9809b45
[Model] Remove unrelated changes from GLM support
ZJY0516 Sep 1, 2026
279bf7c
[Model] Vendor GLM KDA kernels for AMD
ZJY0516 Sep 1, 2026
22d6dfd
Merge branch 'main' into glm-release
ZJY0516 Sep 2, 2026
7e2d791
[Model] Restore Triton MLA non-causal decode support
ZJY0516 Sep 2, 2026
d89723d
[Multimodal] Remove stale PyAV GLM video test
ZJY0516 Sep 2, 2026
4f3d5a8
[Tests] Skip FlashInfer GDN prefill test on ROCm
ZJY0516 Sep 2, 2026
e91fb03
Merge branch 'main' into glm-release
ZJY0516 Sep 2, 2026
8f8cc41
[Bugfix] Make GLM-5.3 kpool metadata graph-safe without prefix cachin…
andyluo7 Sep 2, 2026
3ee3230
Merge branch 'main' into glm-release
ZJY0516 Sep 3, 2026
4500c80
Merge branch 'main' into glm-release
JaredforReal Sep 3, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion .buildkite/test_areas/kernels.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -147,7 +147,7 @@ steps:
commands:
- pytest -v -s kernels/attention/test_triton_unified_attention_diffkv.py

- label: ":nvidia: (H100) FlashMLA Kernels"
- label: ":nvidia: (H100) MLA Kernel Test"
key: kernels-flashmla-test-h100
timeout_in_minutes: 25
device: h100
Expand All @@ -157,12 +157,16 @@ steps:
- vllm/v1/attention/ops/flashmla.py
- vllm/v1/attention/backends/mla/flashmla.py
- vllm/v1/attention/backends/mla/flashmla_sparse.py
- vllm/v1/attention/backends/mla/flashinfer_mla_sparse_sm90.py
- vllm/utils/flashinfer.py
- tests/kernels/attention/test_flashinfer_mla_decode.py
- tests/kernels/attention/test_flashmla.py
- tests/kernels/attention/test_flashmla_sparse.py
- tests/kernels/attention/test_mla_cross_layer_kernel_equivalence.py
commands:
- pytest -v -s kernels/attention/test_flashmla.py
- pytest -v -s kernels/attention/test_flashmla_sparse.py
- pytest -v -s kernels/attention/test_flashinfer_mla_decode.py
- pytest -v -s kernels/attention/test_mla_cross_layer_kernel_equivalence.py

- label: ":nvidia: (L4) Quantization Kernels Shard %N"
Expand Down Expand Up @@ -318,6 +322,7 @@ steps:
- csrc/libtorch_stable/fused_minimax_m3_qknorm_rope_kv_insert_kernel.cu
- csrc/libtorch_stable/ops.h
- csrc/libtorch_stable/torch_bindings.cpp
- tests/kernels/attention/test_flashinfer_mla_decode.py
- tests/kernels/attention/test_minimax_m3_msa_cutlass_sparse_decode.py
- tests/kernels/test_fused_minimax_m3_qknorm_rope_kv_insert.py
- tests/kernels/mamba/test_gdn_prefill_cutedsl.py
Expand Down
13 changes: 9 additions & 4 deletions csrc/libtorch_stable/cache_kernels.cu
Original file line number Diff line number Diff line change
Expand Up @@ -1268,6 +1268,9 @@ __global__ void gather_and_maybe_dequant_cache_page(
#define CALL_GATHER_CACHE_576(SCALAR_T, CACHE_T, KV_DTYPE) \
CALL_GATHER_CACHE(SCALAR_T, CACHE_T, KV_DTYPE, 576)

#define CALL_GATHER_CACHE_512(SCALAR_T, CACHE_T, KV_DTYPE) \
CALL_GATHER_CACHE(SCALAR_T, CACHE_T, KV_DTYPE, 512)

#define CALL_GATHER_CACHE_320(SCALAR_T, CACHE_T, KV_DTYPE) \
CALL_GATHER_CACHE(SCALAR_T, CACHE_T, KV_DTYPE, 320)

Expand Down Expand Up @@ -1305,10 +1308,9 @@ void gather_and_maybe_dequant_cache(
seq_starts.value().scalar_type() == torch::headeronly::ScalarType::Int,
"seq_starts must be int32");
}
STD_TORCH_CHECK(
head_dim == 320 || head_dim == 576,
"gather_and_maybe_dequant_cache only support the head_dim to 320 or 576 "
"for better performance")
STD_TORCH_CHECK(head_dim == 320 || head_dim == 512 || head_dim == 576,
"gather_and_maybe_dequant_cache only support the head_dim to "
"320 or 512 or 576 for better performance")

STD_TORCH_CHECK(src_cache.device() == dst.device(),
"src_cache and dst must be on the same device");
Expand Down Expand Up @@ -1344,6 +1346,9 @@ void gather_and_maybe_dequant_cache(
if (head_dim == 576) {
DISPATCH_BY_KV_CACHE_DTYPE(dst.scalar_type(), kv_cache_dtype,
CALL_GATHER_CACHE_576);
} else if (head_dim == 512) {
DISPATCH_BY_KV_CACHE_DTYPE(dst.scalar_type(), kv_cache_dtype,
CALL_GATHER_CACHE_512);
} else {
DISPATCH_BY_KV_CACHE_DTYPE(dst.scalar_type(), kv_cache_dtype,
CALL_GATHER_CACHE_320);
Expand Down
170 changes: 168 additions & 2 deletions tests/kernels/attention/test_flashinfer_mla_decode.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,13 +9,23 @@

FLASHINFER_WORKSPACE_BUFFER_SIZE = 128 * 1024 * 1024

if not current_platform.has_device_capability(100):
if not current_platform.is_cuda() or not current_platform.has_device_capability(90):
pytest.skip(
reason="FlashInfer MLA Requires compute capability of 10 or above.",
reason="FlashInfer MLA requires CUDA compute capability 9.0 or above.",
allow_module_level=True,
)
else:
from flashinfer.decode import trtllm_batch_decode_with_kv_cache_mla
from flashinfer.mla import BatchMLAPagedAttentionWrapper

requires_sm90 = pytest.mark.skipif(
not current_platform.is_device_capability_family(90),
reason="This test requires an SM90 GPU.",
)
requires_sm10x = pytest.mark.skipif(
not current_platform.is_device_capability_family(100),
reason="This test requires an SM10x GPU.",
)

# Deepseek R1 MLA config.
NUM_HEADS = 128
Expand Down Expand Up @@ -82,6 +92,7 @@ def ref_mla(
@pytest.mark.parametrize("dtype", [torch.bfloat16])
@pytest.mark.parametrize("bs", [1, 2, 4, 16])
@pytest.mark.parametrize("block_size", [32, 64])
@requires_sm10x
def test_flashinfer_mla_decode(dtype: torch.dtype, bs: int, block_size: int):
torch.set_default_device("cuda")
torch.manual_seed(42)
Expand Down Expand Up @@ -119,6 +130,161 @@ def test_flashinfer_mla_decode(dtype: torch.dtype, bs: int, block_size: int):
torch.testing.assert_close(out_ans, out_ref, atol=1e-2, rtol=1e-2)


@requires_sm10x
def test_flashinfer_trtllm_sparse_mla_decode_without_rope():
"""The native sparse MLA path supports a zero-width rotary tail."""
torch.set_default_device("cuda")
torch.manual_seed(42)

batch_size = 2
block_size = 64
num_blocks = 4
sparse_topk = 128
valid_lens = torch.tensor([17, 73], dtype=torch.int32)

query = torch.randn(
batch_size,
1,
NUM_HEADS,
KV_LORA_RANK,
dtype=torch.bfloat16,
)
kv_cache = torch.randn(
num_blocks,
block_size,
KV_LORA_RANK,
dtype=torch.bfloat16,
)

num_slots = num_blocks * block_size
slot_tables = torch.stack(
[torch.randperm(num_slots)[:sparse_topk] for _ in range(batch_size)]
).to(torch.int32)
for row, valid_len in zip(slot_tables, valid_lens.tolist()):
row[valid_len:] = -1

workspace_buffer = torch.empty(
FLASHINFER_WORKSPACE_BUFFER_SIZE,
dtype=torch.int8,
)
out = trtllm_batch_decode_with_kv_cache_mla(
query=query,
kv_cache=kv_cache.unsqueeze(1),
workspace_buffer=workspace_buffer,
qk_nope_head_dim=QK_NOPE_HEAD_DIM,
kv_lora_rank=KV_LORA_RANK,
qk_rope_head_dim=0,
block_tables=slot_tables.unsqueeze(1),
seq_lens=valid_lens,
max_seq_len=sparse_topk,
sparse_mla_top_k=sparse_topk,
sparse_mla_top_k_lens=valid_lens,
bmm1_scale=QK_NOPE_HEAD_DIM**-0.5,
bmm2_scale=1.0,
).squeeze(1)

flat_cache = kv_cache.view(num_slots, KV_LORA_RANK).float()
refs = []
for batch_idx, valid_len in enumerate(valid_lens.tolist()):
selected_kv = flat_cache[slot_tables[batch_idx, :valid_len].long()]
scores = torch.einsum("hd,kd->hk", query[batch_idx, 0].float(), selected_kv)
probs = torch.softmax(scores * QK_NOPE_HEAD_DIM**-0.5, dim=-1)
refs.append(torch.einsum("hk,kd->hd", probs, selected_kv))
ref = torch.stack(refs).to(torch.bfloat16)

torch.testing.assert_close(out, ref, atol=2e-2, rtol=2e-2)


@requires_sm90
def test_flashinfer_sm90_fp8_mla_decode_without_rope():
"""Hopper FA3 supports BF16 queries over an FP8 cache without KPE."""
torch.manual_seed(42)
device = torch.device("cuda")
batch_size = 2
num_heads = 16
page_size = 16
num_pages = 6

q_nope = torch.randn(
batch_size,
num_heads,
KV_LORA_RANK,
dtype=torch.bfloat16,
device=device,
)
q_pe = torch.empty(
batch_size,
num_heads,
0,
dtype=torch.bfloat16,
device=device,
)

ckv = torch.randn(
num_pages,
page_size,
KV_LORA_RANK,
device=device,
)
fp8_max = torch.finfo(torch.float8_e4m3fn).max
ckv_scale = ckv.abs().max().item() / fp8_max
ckv_fp8 = (ckv / ckv_scale).clamp(-fp8_max, fp8_max).to(torch.float8_e4m3fn)
scale_bf16 = torch.tensor(ckv_scale, dtype=torch.bfloat16, device=device)
ckv_ref = ckv_fp8.to(torch.bfloat16) * scale_bf16
kpe_fp8 = torch.empty(
num_pages,
page_size,
0,
dtype=torch.float8_e4m3fn,
device=device,
)
kpe_ref = torch.empty(
num_pages,
page_size,
0,
dtype=torch.bfloat16,
device=device,
)

qo_indptr = torch.tensor([0, 1, 2], dtype=torch.int32, device=device)
kv_indptr = torch.tensor([0, 3, 5], dtype=torch.int32, device=device)
kv_indices = torch.tensor([4, 1, 3, 0, 5], dtype=torch.int32, device=device)
kv_lens = torch.tensor([45, 29], dtype=torch.int32, device=device)
sm_scale = QK_NOPE_HEAD_DIM**-0.5

def run(
ckv_cache: torch.Tensor,
kpe_cache: torch.Tensor,
**kwargs,
) -> torch.Tensor:
workspace = torch.empty(
FLASHINFER_WORKSPACE_BUFFER_SIZE,
dtype=torch.uint8,
device=device,
)
wrapper = BatchMLAPagedAttentionWrapper(workspace, backend="fa3")
wrapper.plan(
qo_indptr,
kv_indptr,
kv_indices,
kv_lens,
num_heads,
KV_LORA_RANK,
0,
page_size,
False,
sm_scale,
q_data_type=torch.bfloat16,
kv_data_type=ckv_cache.dtype,
)
return wrapper.run(q_nope, q_pe, ckv_cache, kpe_cache, **kwargs)

out_ref = run(ckv_ref, kpe_ref)
out = run(ckv_fp8, kpe_fp8, ckv_scale=ckv_scale, kpe_scale=1.0)
torch.testing.assert_close(out, out_ref, atol=2e-2, rtol=2e-2)


@requires_sm10x
def test_flashinfer_mla_decode_workspace_supports_autotune():
"""vLLM's FlashInfer MLA decode workspace must be int8 for autotuning.

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,8 @@ def _make_builder():
builder.device = torch.device("cpu")
builder.kv_cache_spec = SimpleNamespace(block_size=1)
builder.model_dtype = torch.bfloat16
builder.kv_cache_dtype = "fp8"
builder.mla_dims = SimpleNamespace(kv_lora_rank=512, qk_rope_head_dim=64)
builder.topk_tokens = topk_tokens
builder.req_id_per_token_buffer = torch.zeros(
max_num_batched_tokens, dtype=torch.int32, device="cpu"
Expand Down
53 changes: 53 additions & 0 deletions tests/kernels/mamba/test_gdn_prefill_flashinfer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,53 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

import pytest
import torch

from vllm.platforms import current_platform

if current_platform.is_rocm():
pytest.skip(
reason="FlashInfer GDN prefill is not supported on ROCm.",
allow_module_level=True,
)

import flashinfer.gdn_prefill # noqa: E402

from vllm.model_executor.layers.mamba.gdn.qwen_gdn_linear_attn import (
fi_chunk_gated_delta_rule,
) # noqa: E402


def test_flashinfer_gdn_prefill_uses_int64_cu_seqlens(monkeypatch):
captured_cu_seqlens = None

def fake_chunk_gated_delta_rule(**kwargs):
nonlocal captured_cu_seqlens
captured_cu_seqlens = kwargs["cu_seqlens"]
return kwargs["q"]

monkeypatch.setattr(
flashinfer.gdn_prefill,
"chunk_gated_delta_rule",
fake_chunk_gated_delta_rule,
)
q = torch.zeros(1, 2, 1, 2)
cu_seqlens = torch.tensor([0, 2], dtype=torch.int32)

output, final_state = fi_chunk_gated_delta_rule(
q=q,
k=q,
v=q,
g=torch.zeros(1, 2, 1),
beta=torch.zeros(1, 2, 1),
initial_state=torch.zeros(1, 1, 2, 2),
output_final_state=False,
cu_seqlens=cu_seqlens,
use_qk_l2norm_in_kernel=False,
)

assert captured_cu_seqlens is not None
assert captured_cu_seqlens.dtype == torch.int64
assert output.shape == q.shape
assert final_state is None
Loading
Loading