Skip to content
Merged
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
47 changes: 47 additions & 0 deletions tests/v1/attention/test_b12x_sparse_mla_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -845,6 +845,53 @@ def test_glm_selector_metadata_builder_updates_draft_acceptance() -> None:
assert torch.equal(accepted, torch.ones(4, dtype=torch.int32))


def test_dsa_builder_refreshes_fused_dcp_lengths(monkeypatch) -> None:
builder = B12xMLASparseMetadataBuilder.__new__(B12xMLASparseMetadataBuilder)
builder.requires_glm_next_selector_metadata = False
builder.supports_draft_decode_metadata_update = True
builder.dcp_world_size = 4
builder.dcp_rank = 2
builder.cp_kv_cache_interleave_size = 1
global_seq_lens = torch.tensor([17, 9], dtype=torch.int32)
local_seq_lens = torch.zeros(2, dtype=torch.int32)
calls = []

def refresh(*args) -> None:
calls.append(args)

monkeypatch.setattr(b12x_mla_sparse, "refresh_dcp_local_seq_lens_", refresh)
metadata = SimpleNamespace(
dcp_global_seq_lens=global_seq_lens,
seq_lens=local_seq_lens,
num_reqs=2,
selector_num_accepted_tokens=None,
)

builder.update_draft_decode_metadata(metadata)

assert calls == [
(local_seq_lens, global_seq_lens, 2, 4, 2, 1),
]


def test_dsa_builder_rejects_missing_fused_dcp_lengths() -> None:
builder = B12xMLASparseMetadataBuilder.__new__(B12xMLASparseMetadataBuilder)
builder.requires_glm_next_selector_metadata = False
builder.supports_draft_decode_metadata_update = True
builder.dcp_world_size = 4
builder.dcp_rank = 0
builder.cp_kv_cache_interleave_size = 1
metadata = SimpleNamespace(
dcp_global_seq_lens=None,
seq_lens=torch.zeros(1, dtype=torch.int32),
num_reqs=1,
selector_num_accepted_tokens=None,
)

with pytest.raises(RuntimeError, match="global sequence lengths"):
builder.update_draft_decode_metadata(metadata)


def test_dsv4_metadata_builder_does_not_claim_glm_selector_state() -> None:
builder = B12xMLASparseMetadataBuilder.__new__(B12xMLASparseMetadataBuilder)
builder.requires_glm_next_selector_metadata = False
Expand Down
161 changes: 159 additions & 2 deletions tests/v1/attention/test_indexer_dcp_localize.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,13 +5,22 @@
import torch

import vllm.model_executor.layers.sparse_attn_indexer as sparse_indexer
import vllm.v1.attention.backends.mla.indexer as indexer_backend
from vllm.platforms import current_platform
from vllm.utils.import_utils import has_cutedsl
from vllm.v1.attention.backends.mla.indexer import build_prefill_chunk_metadata
from vllm.v1.attention.backends.mla.indexer import (
DeepSeekV32IndexerDecodeMetadata,
DeepseekV32IndexerMetadata,
DeepseekV32IndexerMetadataBuilder,
build_prefill_chunk_metadata,
)
from vllm.v1.attention.backends.mla.sparse_utils import (
triton_filter_and_convert_dcp_index,
)
from vllm.v1.attention.backends.utils import get_dcp_local_seq_lens
from vllm.v1.attention.backends.utils import (
get_dcp_local_seq_lens,
refresh_dcp_local_seq_lens_,
)
from vllm.v1.attention.ops.dcp import CPTritonContext, correct_attn_out


Expand Down Expand Up @@ -327,6 +336,154 @@ def test_get_dcp_local_seq_lens_must_run_after_decode_expansion():
)


@pytest.mark.skipif(not current_platform.is_cuda(), reason="This test requires CUDA")
@pytest.mark.parametrize("rank", [0, 2, 3])
def test_refresh_dcp_local_seq_lens_updates_storage_in_place(rank: int):
global_seq_lens = torch.arange(33, dtype=torch.int32, device="cuda")
local_seq_lens = torch.full((40,), -1, dtype=torch.int32, device="cuda")
storage_ptr = local_seq_lens.data_ptr()

refresh_dcp_local_seq_lens_(
local_seq_lens,
global_seq_lens,
global_seq_lens.numel(),
4,
rank,
2,
)

expected = get_dcp_local_seq_lens(global_seq_lens, 4, rank, 2)
torch.testing.assert_close(local_seq_lens[: global_seq_lens.numel()], expected)
assert torch.count_nonzero(local_seq_lens[global_seq_lens.numel() :]) == 0
assert local_seq_lens.data_ptr() == storage_ptr


def test_update_draft_decode_metadata_refreshes_dcp_lens_and_schedule(monkeypatch):
class _KVCacheSpec:
num_states = 64

builder = object.__new__(DeepseekV32IndexerMetadataBuilder)
builder.dcp_world_size = 4
builder.dcp_rank = 2
builder.cp_kv_cache_interleave_size = 1
builder.kv_cache_spec = _KVCacheSpec()
builder.num_sms = 8
builder.offsets_buffer = torch.arange(4, dtype=torch.int32)
builder.global_decode_seq_lens_buffer = torch.empty(12, dtype=torch.int32)

global_seq_lens = torch.tensor([11, 14, 0], dtype=torch.int32)
decode_seq_lens = torch.full((3, 3), -1, dtype=torch.int32)
schedule_metadata = torch.zeros((4, 2), dtype=torch.int32)
decode = DeepSeekV32IndexerDecodeMetadata(
block_table=torch.zeros((3, 1), dtype=torch.int32),
seq_lens=decode_seq_lens,
decode_lens=torch.ones(3, dtype=torch.int32),
requires_padding=False,
schedule_metadata=schedule_metadata,
global_seq_lens=global_seq_lens,
)
metadata = DeepseekV32IndexerMetadata(
seq_lens=global_seq_lens,
max_seq_len=14,
slot_mapping=torch.zeros(3, dtype=torch.int64),
num_decodes=3,
num_decode_tokens=9,
num_prefills=0,
num_prefill_tokens=0,
decode=decode,
)

planned_lens = []

def _fake_plan(context_lens, block_size, num_sms, indices=None):
assert block_size == 64
assert num_sms == 8
assert indices is None
planned_lens.append(context_lens.clone())
return torch.full_like(schedule_metadata, len(planned_lens))

monkeypatch.setattr(indexer_backend, "get_paged_mqa_logits_metadata", _fake_plan)
monkeypatch.setattr(
indexer_backend,
"refresh_dcp_local_seq_lens_",
lambda out, seq_lens, num_reqs, world, rank, interleave: out.copy_(
get_dcp_local_seq_lens(seq_lens[:num_reqs], world, rank, interleave)
),
)
seq_lens_ptr = decode_seq_lens.data_ptr()
schedule_ptr = schedule_metadata.data_ptr()

builder.update_draft_decode_metadata(metadata)
torch.testing.assert_close(
decode_seq_lens,
torch.tensor([[2, 2, 3], [3, 3, 3], [0, 0, 0]], dtype=torch.int32),
)
torch.testing.assert_close(planned_lens[-1], decode_seq_lens)
assert torch.all(schedule_metadata == 1)

global_seq_lens.copy_(torch.tensor([15, 18, 0], dtype=torch.int32))
builder.update_draft_decode_metadata(metadata)
torch.testing.assert_close(
decode_seq_lens,
torch.tensor([[3, 3, 4], [4, 4, 4], [0, 0, 0]], dtype=torch.int32),
)
torch.testing.assert_close(planned_lens[-1], decode_seq_lens)
assert torch.all(schedule_metadata == 2)
assert decode_seq_lens.data_ptr() == seq_lens_ptr
assert schedule_metadata.data_ptr() == schedule_ptr


def test_update_draft_decode_metadata_presents_rank_two_lens_to_native_planner(
monkeypatch,
):
class _KVCacheSpec:
num_states = 64

builder = object.__new__(DeepseekV32IndexerMetadataBuilder)
builder.dcp_world_size = 1
builder.dcp_rank = 0
builder.cp_kv_cache_interleave_size = 1
builder.kv_cache_spec = _KVCacheSpec()
builder.num_sms = 8
builder.offsets_buffer = torch.arange(4, dtype=torch.int32)
builder.global_decode_seq_lens_buffer = torch.empty(8, dtype=torch.int32)

global_seq_lens = torch.tensor([11, 14], dtype=torch.int32)
decode_seq_lens = torch.full((2,), -1, dtype=torch.int32)
schedule_metadata = torch.zeros((4, 2), dtype=torch.int32)
decode = DeepSeekV32IndexerDecodeMetadata(
block_table=torch.zeros((2, 1), dtype=torch.int32),
seq_lens=decode_seq_lens,
decode_lens=torch.ones(2, dtype=torch.int32),
requires_padding=False,
schedule_metadata=schedule_metadata,
global_seq_lens=None,
)
metadata = DeepseekV32IndexerMetadata(
seq_lens=global_seq_lens,
max_seq_len=14,
slot_mapping=torch.zeros(2, dtype=torch.int64),
num_decodes=2,
num_decode_tokens=2,
num_prefills=0,
num_prefill_tokens=0,
decode=decode,
)

planned_shapes = []

def _fake_plan(context_lens, block_size, num_sms, indices=None):
planned_shapes.append(context_lens.shape)
return torch.ones_like(schedule_metadata)

monkeypatch.setattr(indexer_backend, "get_paged_mqa_logits_metadata", _fake_plan)

builder.update_draft_decode_metadata(metadata)

assert planned_shapes == [torch.Size((2, 1))]
torch.testing.assert_close(decode_seq_lens, global_seq_lens)


@pytest.mark.parametrize("interleave", [1, 2])
def test_sparse_dcp_attention_matches_global_topk_attention(interleave: int):
torch.manual_seed(0)
Expand Down
41 changes: 32 additions & 9 deletions vllm/v1/attention/backends/mla/b12x_mla_sparse.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,10 @@
triton_convert_req_index_to_global_index,
triton_filter_and_convert_dcp_index,
)
from vllm.v1.attention.backends.utils import get_dcp_local_seq_lens
from vllm.v1.attention.backends.utils import (
get_dcp_local_seq_lens,
refresh_dcp_local_seq_lens_,
)
from vllm.v1.kv_cache_interface import AttentionSpec, MLAAttentionSpec
from vllm.v1.kv_cache_layout import KVCacheLayout
from vllm.v1.worker.workspace import (
Expand Down Expand Up @@ -554,6 +557,7 @@ class B12xMLASparseMetadata(AttentionMetadata):
num_decodes: int
num_prefills: int
num_decode_tokens: int
dcp_global_seq_lens: torch.Tensor | None = None
prefill_max_seq_len: int = 0
prefill: MLACommonPrefillMetadata | None = None
prefill_query_lens_cpu: torch.Tensor | None = None
Expand Down Expand Up @@ -599,9 +603,9 @@ def __init__(
):
raise ValueError(dcp_error)
super().__init__(kv_cache_spec, layer_names, vllm_config, device)
self.supports_draft_decode_metadata_update = (
self.requires_glm_next_selector_metadata
)
# All step-dependent state is persistent. Generic DSA DCP additionally
# refreshes its rank-local sequence lengths in place between steps.
self.supports_draft_decode_metadata_update = True
self.dcp_rank = get_dcp_group().rank_in_group if self.dcp_world_size > 1 else 0
scheduler_config = vllm_config.scheduler_config
max_tokens = scheduler_config.max_num_batched_tokens
Expand Down Expand Up @@ -779,6 +783,9 @@ def _build(
else common.seq_lens
)
metadata.seq_lens = seq_lens
metadata.dcp_global_seq_lens = (
common.seq_lens[: common.num_reqs] if use_dcp else None
)

if common.max_query_len <= 1 and num_tokens == common.num_reqs:
per_token_lens = seq_lens[:num_tokens]
Expand Down Expand Up @@ -958,12 +965,28 @@ def update_draft_decode_metadata(
self,
metadata: B12xMLASparseMetadata,
) -> None:
accepted = metadata.selector_num_accepted_tokens
if not self.requires_glm_next_selector_metadata or accepted is None:
raise RuntimeError(
"GLM5Next draft decode metadata requires accepted-token counts"
if self.dcp_world_size > 1:
global_seq_lens = metadata.dcp_global_seq_lens
if global_seq_lens is None:
raise RuntimeError(
"B12X fused DCP draft decode requires global sequence lengths"
)
refresh_dcp_local_seq_lens_(
metadata.seq_lens,
global_seq_lens,
metadata.num_reqs,
self.dcp_world_size,
self.dcp_rank,
self.cp_kv_cache_interleave_size,
)
accepted.fill_(1)

if self.requires_glm_next_selector_metadata:
accepted = metadata.selector_num_accepted_tokens
if accepted is None:
raise RuntimeError(
"GLM5Next draft decode metadata requires accepted-token counts"
)
accepted.fill_(1)


class B12xGLM5NextMLASparseMetadataBuilder(B12xMLASparseMetadataBuilder):
Expand Down
Loading
Loading