Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
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
2 changes: 2 additions & 0 deletions cmake/external_projects/flashmla.cmake
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,8 @@ if(FLASH_MLA_ARCHS)
${flashmla_SOURCE_DIR}/csrc/kernels/sm90/decode/sparse/instantiations/v4_persistent_h128.cu
${flashmla_SOURCE_DIR}/csrc/kernels/sm90/decode/sparse/instantiations/v32_persistent_h64.cu
${flashmla_SOURCE_DIR}/csrc/kernels/sm90/decode/sparse/instantiations/v32_persistent_h128.cu
${flashmla_SOURCE_DIR}/csrc/kernels/sm90/decode/sparse/instantiations/v32_persistent_h64_topklen.cu
${flashmla_SOURCE_DIR}/csrc/kernels/sm90/decode/sparse/instantiations/v32_persistent_h128_topklen.cu

# sm90 sparse prefill
${flashmla_SOURCE_DIR}/csrc/kernels/sm90/prefill/sparse/instantiations/phase1_k512.cu
Expand Down
167 changes: 157 additions & 10 deletions tests/v1/attention/test_sparse_mla_backends.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,7 @@
FlashMLASparseMetadata,
FlashMLASparseMetadataBuilder,
triton_convert_req_index_to_global_index,
triton_filter_and_convert_dcp_index,
)
from vllm.v1.attention.backends.mla.index_group import (
HiSparseMLAIndexGroup,
Expand Down Expand Up @@ -3629,25 +3630,25 @@ def test_fp8_mixed_batch_dcp_neutralizes_empty_rows(monkeypatch):
num_tokens, num_heads, head_dim = 3, 2, 3
q = torch.empty(num_tokens, num_heads, head_dim, device=DEVICE_TYPE)
local_indices = torch.tensor(
[[0, 1, -1, -1], [-1, -1, -1, -1], [2, -1, 3, -1]],
[[0, 1, -1, -1], [-1, -1, -1, -1], [2, 3, -1, -1]],
dtype=torch.int32,
device=DEVICE_TYPE,
)
local_lengths = torch.tensor([2, 0, 2], dtype=torch.int32, device=DEVICE_TYPE)

monkeypatch.setattr(
"vllm.v1.attention.backends.mla.flashmla_sparse."
"triton_filter_and_convert_dcp_index",
lambda *args, **kwargs: local_indices,
lambda *args, **kwargs: (local_indices, local_lengths),
)

def run_kernel(**kwargs):
out = torch.full(
(1, num_tokens, num_heads, 1), float("nan"), device=DEVICE_TYPE
)
lse = torch.full((1, num_heads, num_tokens), float("nan"), device=DEVICE_TYPE)
def run_kernel(q, kv_c_and_k_pe_cache, topk_indices, topk_length):
assert torch.equal(topk_length, local_lengths)
out = torch.full((num_tokens, num_heads, 1), float("nan"), device=DEVICE_TYPE)
lse = torch.full((num_tokens, num_heads), float("nan"), device=DEVICE_TYPE)
for token_id in (0, 2): # rows with local candidates get real values
out[0, token_id] = float(token_id + 1)
lse[0, :, token_id] = float(token_id + 1)
out[token_id] = float(token_id + 1)
lse[token_id] = float(token_id + 1)
return out, lse

metadata = SimpleNamespace(
Expand All @@ -3667,7 +3668,7 @@ def run_kernel(**kwargs):
dcp_world_size=2,
dcp_rank=0,
need_to_return_lse_for_decode=True,
_fp8_flash_mla_kernel=run_kernel,
_fp8_flash_mla_kernel_per_token=run_kernel,
)

out, lse = FlashMLASparseImpl._forward_fp8_kv_mixed_batch(
Expand All @@ -3684,6 +3685,152 @@ def run_kernel(**kwargs):
assert not lse.isnan().any()


@pytest.mark.parametrize("dcp_world_size,dcp_rank", [(2, 1), (4, 3)])
@pytest.mark.parametrize("interleave", [1, 16])
def test_fp8_dcp_topk_length_matches_masked_rows(
monkeypatch, dcp_world_size, dcp_rank, interleave
):
"""Handing the kernel each row's owned-slot count as topk_length must match
the call that masks the slots other DCP ranks own, for rows with fewer than
topk candidates, rows this rank owns nothing of, and batches split across
several schedule plans."""
ok, reason = flashmla.is_flashmla_sparse_supported()
if not ok:
pytest.skip(reason)
torch.manual_seed(0)
num_heads, head_dim, block_size, num_topk = 64, 576, 64, 2048
monkeypatch.setattr(
"vllm.v1.attention.backends.mla.flashmla_sparse._MAX_SPLIT_ACCUM_BYTES",
16 * num_heads * 512 * 4, # 16 tokens per call
)
seq_lens = [9000, 2100, 700]
# With interleave > 1 a rank's local slots can run up to interleave past
# seq_len / dcp.
local_blocks = [
cdiv(cdiv(s, dcp_world_size) + interleave, block_size) for s in seq_lens
]
block_table = torch.zeros(
len(seq_lens), max(local_blocks), dtype=torch.int32, device=DEVICE_TYPE
)
perm = torch.randperm(sum(local_blocks), device=DEVICE_TYPE).to(torch.int32)
start = 0
for r, n in enumerate(local_blocks):
block_table[r, :n] = perm[start : start + n]
start += n

# Full decode rows per request, then the first positions of a fresh
# prefill of request 0, whose rows see only a few candidates.
rows, req_ids = [], []
for r, s in enumerate(seq_lens):
for _ in range(4):
row = torch.full((num_topk,), -1, dtype=torch.int32, device=DEVICE_TYPE)
cand = torch.randperm(s, device=DEVICE_TYPE)[:num_topk]
row[: cand.numel()] = cand.to(torch.int32)
rows.append(row)
req_ids.append(r)
for pos in range(24):
row = torch.full((num_topk,), -1, dtype=torch.int32, device=DEVICE_TYPE)
row[: pos + 1] = torch.arange(pos + 1, dtype=torch.int32, device=DEVICE_TYPE)
rows.append(row)
req_ids.append(0)
global_topk = torch.stack(rows)
req_id = torch.tensor(req_ids, dtype=torch.int32, device=DEVICE_TYPE)
num_tokens = global_topk.shape[0]

num_slots = sum(local_blocks) * block_size
cache = torch.zeros(
sum(local_blocks), block_size, 656, dtype=torch.uint8, device=DEVICE_TYPE
)
ops.concat_and_cache_mla(
(torch.randn(num_slots, 512, device=DEVICE_TYPE) * 0.5).bfloat16(),
(torch.randn(num_slots, 64, device=DEVICE_TYPE) * 0.5).bfloat16(),
cache,
torch.arange(num_slots, device=DEVICE_TYPE),
kv_cache_dtype="fp8_ds_mla",
scale=torch.ones(1, device=DEVICE_TYPE),
)
q = (
torch.randn(num_tokens, num_heads, head_dim, device=DEVICE_TYPE) * 0.3
).bfloat16()
scale = head_dim**-0.5
filter_args = dict(
dcp_size=dcp_world_size,
dcp_rank=dcp_rank,
cp_kv_cache_interleave_size=interleave,
BLOCK_SIZE=block_size,
NUM_TOPK_TOKENS=num_topk,
)

masked = triton_filter_and_convert_dcp_index(
req_id, block_table, global_topk, compact_valid_to_front=False, **filter_args
)
ref_out, ref_lse = flashmla.flash_mla_with_kvcache(
q=q.unsqueeze(0),
k_cache=cache.unsqueeze(-2),
block_table=None,
cache_seqlens=None,
head_dim_v=512,
tile_scheduler_metadata=flashmla.get_mla_metadata()[0],
is_fp8_kvcache=True,
indices=masked.unsqueeze(0),
softmax_scale=scale,
)
ref_out, ref_lse = ref_out.squeeze(0), ref_lse.squeeze(0).transpose(0, 1)

compacted, lengths = triton_filter_and_convert_dcp_index(
req_id, block_table, global_topk, return_valid_counts=True, **filter_args
)
impl = SimpleNamespace(
fp8_decode_padded_heads=num_heads, kv_lora_rank=512, softmax_scale=scale
)
impl._pad_q_heads_for_fp8_kernel = MethodType(
FlashMLASparseImpl._pad_q_heads_for_fp8_kernel, impl
)
out, lse = FlashMLASparseImpl._fp8_flash_mla_kernel_per_token(
impl, q, cache, compacted, lengths
)

owned = lengths > 0
assert not owned.all(), "expected rows this rank owns nothing of"
torch.testing.assert_close(out[owned], ref_out[owned], atol=2e-2, rtol=2e-2)
torch.testing.assert_close(lse[owned], ref_lse[owned], atol=1e-4, rtol=1e-5)


def test_hyv4_fp8_per_token_kernel_passes_sink(monkeypatch):
"""HY V4 threads its sink into every FlashMLA call, including the
per-token one the DCP path uses."""
from vllm.models.hy_v4.nvidia.flashmla_sparse import HYV4FlashMLASparseImpl

num_tokens, num_heads, head_dim = 3, 32, 576
impl = object.__new__(HYV4FlashMLASparseImpl)
impl.sinks = torch.randn(num_heads, device=DEVICE_TYPE)
impl.fp8_decode_padded_heads = 64
impl.kv_lora_rank = 512
impl.softmax_scale = head_dim**-0.5
seen = {}

def fake_kernel(**kwargs):
seen.update(kwargs)
lse = torch.zeros(kwargs["q"].shape[0], 64, 1, device=DEVICE_TYPE)
return kwargs["out"], lse

backend = "vllm.v1.attention.backends.mla.flashmla_sparse"
monkeypatch.setattr(f"{backend}.flash_mla_with_kvcache", fake_kernel)
monkeypatch.setattr(f"{backend}.get_mla_metadata", lambda: (None, None))
q = torch.zeros(num_tokens, num_heads, head_dim, device=DEVICE_TYPE)
indices = torch.zeros(num_tokens, 64, dtype=torch.int32, device=DEVICE_TYPE)
lengths = torch.ones(num_tokens, dtype=torch.int32, device=DEVICE_TYPE)
out, lse = impl._fp8_flash_mla_kernel_per_token(
q, torch.empty(1, 64, 656, dtype=torch.uint8), indices, lengths
)

sink = seen["attn_sink"]
assert torch.equal(sink[:num_heads], impl.sinks)
assert torch.isneginf(sink[num_heads:]).all()
assert out.shape == (num_tokens, num_heads, 512)
assert lse.shape == (num_tokens, num_heads)


def test_hisparse_prefill_reuses_builder_staging_plan():
"""Every layer must reuse the batch plan instead of synchronizing to dedupe."""
plan = SimpleNamespace(
Expand Down
18 changes: 17 additions & 1 deletion vllm/models/hy_v4/nvidia/flashmla_sparse.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@
`_canonicalize_sparse_mla_kv_cache_dtype` promotes a quantized KV cache to
``fp8_ds_mla`` for it, and `FlashMLASparseImpl` asserts that layout — so a new
name would silently change KV cache behaviour. Only``supports_sink`` and the
two kernel wrappers differ from the parent.
kernel wrappers differ from the parent.
"""

from typing import TYPE_CHECKING
Expand Down Expand Up @@ -214,6 +214,22 @@ def _fp8_flash_mla_kernel(

return out, lse

def _fp8_flash_mla_kernel_per_token(
self,
q: torch.Tensor,
kv_c_and_k_pe_cache: torch.Tensor,
topk_indices: torch.Tensor,
topk_length: torch.Tensor,
attn_sink: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
# q shape: (num_tokens, num_heads, head_dim)
attn_sink = self._sinks_for_query(
q, head_dim=1, kernel_heads=self.fp8_decode_padded_heads
)
return super()._fp8_flash_mla_kernel_per_token(
q, kv_c_and_k_pe_cache, topk_indices, topk_length, attn_sink=attn_sink
)

def _bf16_flash_mla_kernel(
self,
q: torch.Tensor,
Expand Down
Loading
Loading