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
126 changes: 118 additions & 8 deletions tests/v1/attention/test_b12x_sparse_mla_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,9 +35,12 @@
B12xMLASparseImpl,
B12xMLASparseMetadata,
B12xMLASparseMetadataBuilder,
_global_causal_lens_for_ckv_gather,
_is_glm_next_ckv_source_layout,
_is_speculative_decode_batch,
_max_speculative_decode_query_len,
_selected_index_block_stride_rows,
_use_b12x_full_ckv_gather,
)
from vllm.v1.attention.backends.mla.sparse_utils import _remap_tiling
from vllm.v1.attention.backends.registry import AttentionBackendEnum
Expand Down Expand Up @@ -172,10 +175,12 @@ def test_b12x_glm5_next_cache_spec_and_layout(monkeypatch) -> None:
assert invalid_reasons == []
assert unidentified == probe
assert packed_by_glm_backend.state_content_bytes == 528
assert packed_by_glm_backend.page_size_padded == 64 * 528 + 16 * 128 * 2
assert packed_by_glm_backend.page_size_padded is None
assert packed_by_glm_backend.page_size_bytes == 64 * (528 + 33)
assert packed_by_glm_backend.model_version == "glm5_next"
assert packed.state_content_bytes == 528
assert packed.page_size_padded == 64 * 528 + 16 * 128 * 2
assert packed.page_size_padded is None
assert packed.page_size_bytes == 64 * (528 + 33)
assert packed.model_version == "glm5_next"
assert packed_without_config_context == packed
assert layouts == (KVCacheLayout.BLHNC,)
Expand Down Expand Up @@ -258,6 +263,52 @@ def test_b12x_glm5_next_accepts_dcp_with_speculation(monkeypatch) -> None:
assert invalid_reasons == []


@pytest.mark.parametrize(
("max_query_len", "num_decode_tokens", "num_tokens", "expected"),
[
(1, 0, 32, False),
(6, 192, 192, False),
(6, 0, 192, True),
(128, 0, 8192, True),
(128, 0, 600000, False),
],
)
def test_b12x_full_ckv_gather_excludes_decode_and_mtp_batches(
max_query_len: int,
num_decode_tokens: int,
num_tokens: int,
expected: bool,
) -> None:
assert (
_use_b12x_full_ckv_gather(
enabled=True,
is_glm_next=True,
dcp_world_size=4,
max_query_len=max_query_len,
num_tokens=num_tokens,
num_decode_tokens=num_decode_tokens,
min_tokens=16,
max_tokens=524288,
)
is expected
)


def test_b12x_full_ckv_gather_uses_global_causal_lengths() -> None:
global_seq_lens = torch.tensor([5, 12], dtype=torch.int32)
query_start_loc = torch.tensor([0, 2, 5], dtype=torch.int32)
req_id_per_token = torch.tensor([0, 0, 1, 1, 1], dtype=torch.int32)

actual = _global_causal_lens_for_ckv_gather(
global_seq_lens,
query_start_loc,
req_id_per_token,
num_actual_tokens=5,
)

assert actual.tolist() == [4, 5, 10, 11, 12]


def test_b12x_glm5_next_accepts_dcp_with_prefix_caching(monkeypatch) -> None:
monkeypatch.setattr(b12x_mla_sparse, "get_b12x_sparse_mla", lambda: object())
with set_current_vllm_config(
Expand Down Expand Up @@ -341,6 +392,8 @@ def test_b12x_glm5_next_selected_indices_use_physical_slots() -> None:
)
== 64
)
assert _is_glm_next_ckv_source_layout(cache, page_size=64)
assert not _is_glm_next_ckv_source_layout(cache[:, :, ::2], page_size=64)


def test_sparse_index_remap_tiling_covers_glm5_next_width() -> None:
Expand Down Expand Up @@ -369,19 +422,29 @@ def test_b12x_glm5_next_cache_writer_ignores_empty_rope() -> None:
assert calls == [(kv_c, kv_cache, slots)]


def test_b12x_glm5_next_cache_bind_replans_aligned_manager_page(monkeypatch) -> None:
def test_b12x_glm5_next_cache_geometry_is_finalized_before_bind(monkeypatch) -> None:
planned: list[SimpleNamespace] = []
reservations: list[tuple[tuple[tuple[int, ...], torch.dtype], ...]] = []
monkeypatch.setattr(torch.accelerator, "current_device_index", lambda: 0)
monkeypatch.setattr(
b12x_mla_sparse,
"is_workspace_manager_initialized",
lambda: False,
)
monkeypatch.setattr(
b12x_mla_sparse,
"current_workspace_manager",
lambda: SimpleNamespace(reserve_all=lambda *specs: reservations.append(specs)),
)

class FakeCaps(SimpleNamespace):
def __init__(self, **kwargs):
super().__init__(**kwargs)

@staticmethod
def shapes_and_dtypes():
return ()

class FakeModule:
Caps = FakeCaps

Expand All @@ -394,7 +457,10 @@ def plan(caps):
impl._is_glm_next = True
impl._module = FakeModule
impl._kernel_page_size = 64
impl._kernel_page_size_finalized = False
impl._input_num_heads = 64
impl.num_heads = 16
impl.dcp_world_size = 4
impl._max_tokens = 4096
impl._max_seqs = 4
impl._max_speculative_decode_query_len = 6
Expand All @@ -404,23 +470,60 @@ def plan(caps):
impl._q_head_dim = 512
impl.kv_lora_rank = 512
impl._model_type = 1
impl._ckv_gather_enabled = True
impl._ckv_capacity_tokens = 131200
impl._ckv_local_capacity = 131200
impl._decode_plan = SimpleNamespace()
impl._extend_plan = SimpleNamespace()
impl._ckv_extend_plan = SimpleNamespace()
owner = SimpleNamespace(impl=impl, indexer=None)
cache = torch.empty((2, 1, 2304, 528), dtype=torch.uint8)

MLAAttention.finalize_kv_cache_geometry(
owner,
SimpleNamespace(cache_config=SimpleNamespace(block_size=2304)),
)
MLAAttention.bind_kv_cache(owner, cache)

assert owner.kv_cache.shape == (2, 2304, 528)
assert impl._kernel_page_size == 2304
assert impl._kernel_page_size_finalized
assert [(caps.mode, caps.page_size) for caps in planned] == [
("decode", 2304),
("extend", 2304),
("extend", 2304),
]
plan_geometry = [
(caps.num_q_heads, caps.max_q_rows, caps.max_batch) for caps in planned
]
assert [(caps.max_q_rows, caps.max_batch) for caps in planned] == [
(24, 24),
(4096, 4096),
assert plan_geometry == [
(64, 24, 24),
(64, 4096, 4096),
(16, 4096, 4096),
]
assert len(reservations) == 3
assert reservations[0] == (((4096, 64, 512), torch.bfloat16),)
assert reservations[1] == (((4096, 64, 512), torch.bfloat16),)
assert reservations[2] == (
((4096, 16, 512), torch.bfloat16),
((131328, 528), torch.uint8),
((525312, 528), torch.uint8),
)

with pytest.raises(RuntimeError, match="immutable after finalization"):
impl.finalize_kv_cache_geometry(64)
with pytest.raises(RuntimeError, match="does not match the finalized"):
impl.bind_kv_cache(torch.empty((2, 64, 528), dtype=torch.uint8))


def test_b12x_glm5_next_full_ckv_bind_requires_geometry_finalization() -> None:
impl = object.__new__(B12xMLASparseImpl)
impl._is_glm_next = True
impl._ckv_gather_enabled = True
impl._kernel_page_size_finalized = False

with pytest.raises(RuntimeError, match="before KV-cache memory profiling"):
impl.bind_kv_cache(torch.empty((2, 2304, 528), dtype=torch.uint8))


@pytest.mark.parametrize(
Expand Down Expand Up @@ -533,6 +636,7 @@ def _bare_glm_selector_metadata_builder() -> B12xMLASparseMetadataBuilder:
builder = B12xMLASparseMetadataBuilder.__new__(B12xMLASparseMetadataBuilder)
builder.requires_glm_next_selector_metadata = True
builder.supports_draft_decode_metadata_update = True
builder._ckv_gather_requested = False
builder.dcp_world_size = 1
builder._max_speculative_decode_query_len = 6
builder._capture_default_state_slot_ids = torch.arange(4, dtype=torch.int32)
Expand Down Expand Up @@ -631,7 +735,10 @@ def test_glm_selector_metadata_builder_stages_padded_rows_and_capture(
monkeypatch.setattr(
SparseMLACommonMetadataBuilder,
"build",
lambda *args, **kwargs: SimpleNamespace(),
lambda *args, **kwargs: SimpleNamespace(
num_prefills=0,
num_decode_tokens=0,
),
)
builder = _bare_glm_selector_metadata_builder()
common = SimpleNamespace(
Expand Down Expand Up @@ -707,7 +814,10 @@ def test_glm_selector_metadata_builder_requires_complete_runtime_state(
monkeypatch.setattr(
SparseMLACommonMetadataBuilder,
"build",
lambda *args, **kwargs: SimpleNamespace(),
lambda *args, **kwargs: SimpleNamespace(
num_prefills=0,
num_decode_tokens=0,
),
)
builder = _bare_glm_selector_metadata_builder()
common = SimpleNamespace(
Expand Down
27 changes: 27 additions & 0 deletions tests/v1/attention/test_mla_backends.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
MLAAttention,
QueryLenSupport,
_DecodeConcatQuantFP8,
_select_mqa_query,
build_mla_chunked_context_metadata,
)
from vllm.model_executor.layers.quantization.utils.quant_utils import GroupShape
Expand Down Expand Up @@ -68,6 +69,32 @@
DEVICE_TYPE = current_platform.device_type


@pytest.mark.cpu_test
def test_full_ckv_dcp_prefers_local_query_geometry() -> None:
q = torch.zeros((2, 2, 6))
q_dcp_replicated = torch.ones((2, 8, 6))

selected, replicated = _select_mqa_query(
q,
q_dcp_replicated,
num_mqa_tokens=1,
full_ckv_dcp=True,
)
assert selected.shape == (1, 2, 6)
assert not replicated
assert torch.equal(selected, q[:1])

selected, replicated = _select_mqa_query(
q,
q_dcp_replicated,
num_mqa_tokens=1,
full_ckv_dcp=False,
)
assert selected.shape == (1, 8, 6)
assert replicated
assert torch.equal(selected, q_dcp_replicated[:1])


@pytest.mark.parametrize(
("cache_dtype", "expected_quant_mode"),
[
Expand Down
28 changes: 28 additions & 0 deletions tests/v1/worker/test_workspace.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,34 @@ def test_workspace_lanes_compose_with_ubatches(monkeypatch) -> None:
assert len(pointers) == 4


def test_workspace_reservation_covers_every_execution_slot(monkeypatch) -> None:
active_ubatch = [0]
monkeypatch.setattr(workspace, "dbo_current_ubatch_id", lambda: active_ubatch[0])
manager = workspace.WorkspaceManager(
torch.device("cpu"), num_ubatches=2, num_lanes=2
)

manager.reserve_all(((257,), torch.uint8), ((1,), torch.float32))

assert [
buffer.numel() if buffer is not None else 0
for buffer in manager._current_workspaces
] == [768, 768, 768, 768]
for ubatch_id in range(2):
active_ubatch[0] = ubatch_id
for lane in range(2):
workspace_id = ubatch_id * 2 + lane
with workspace.use_workspace_lane(lane):
(view,) = manager.get_simultaneous(((8,), torch.uint8))
reserved = manager._current_workspaces[workspace_id]
assert reserved is not None
assert view.data_ptr() == reserved.data_ptr()

manager.lock()
with pytest.raises(AssertionError, match="reserve_all"):
manager.reserve_all(((1024,), torch.uint8))


def test_workspace_lane_validation(monkeypatch) -> None:
monkeypatch.setattr(workspace, "dbo_current_ubatch_id", lambda: 0)
manager = workspace.WorkspaceManager(torch.device("cpu"), num_lanes=1)
Expand Down
15 changes: 15 additions & 0 deletions vllm/envs.py
Original file line number Diff line number Diff line change
Expand Up @@ -190,6 +190,9 @@
VLLM_HUMMING_USE_F16_ACCUM: bool = False
VLLM_HUMMING_MOE_GEMM_TYPE: Literal["indexed", "grouped", "auto"] | None = None
VLLM_B12X_MOE_FP4_FORCE_A16: bool = False
VLLM_B12X_MLA_CKV_GATHER: bool = False
VLLM_B12X_MLA_CKV_GATHER_MIN_TOKENS: int = 16
VLLM_B12X_MLA_CKV_GATHER_MAX_TOKENS: int = 524288
VLLM_PLE_CPU_OFFLOAD: bool = False
VLLM_DEEPEPLL_NVFP4_DISPATCH: bool = False
VLLM_V1_USE_OUTLINES_CACHE: bool = False
Expand Down Expand Up @@ -1627,6 +1630,18 @@ def _resolve_rust_cli_path() -> str | None:
"VLLM_B12X_MOE_FP4_FORCE_A16": lambda: bool(
int(os.getenv("VLLM_B12X_MOE_FP4_FORCE_A16", "0"))
),
# Gather DCP-sharded C4 records before B12X sparse-MLA prefill. This avoids
# query replication plus the per-rank LSE combine and is opt-in while the
# path is being qualified on GLM5Next.
"VLLM_B12X_MLA_CKV_GATHER": lambda: (
os.getenv("VLLM_B12X_MLA_CKV_GATHER", "0").lower() in ("1", "true", "yes", "on")
),
"VLLM_B12X_MLA_CKV_GATHER_MIN_TOKENS": lambda: int(
os.getenv("VLLM_B12X_MLA_CKV_GATHER_MIN_TOKENS", "16")
),
"VLLM_B12X_MLA_CKV_GATHER_MAX_TOKENS": lambda: int(
os.getenv("VLLM_B12X_MLA_CKV_GATHER_MAX_TOKENS", "524288")
),
# Qwen3.8-Flash-Next only. Store PLE table payloads in CUDA-mapped host
# memory unless additional_config.ple_table_memory is explicitly set.
"VLLM_PLE_CPU_OFFLOAD": lambda: bool(int(os.getenv("VLLM_PLE_CPU_OFFLOAD", "0"))),
Expand Down
36 changes: 28 additions & 8 deletions vllm/model_executor/layers/attention/mla_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -346,6 +346,19 @@ def _detect_output_quant_key(
return kFp8StaticTensorSym


def _select_mqa_query(
q: torch.Tensor,
q_dcp_replicated: torch.Tensor | None,
*,
num_mqa_tokens: int,
full_ckv_dcp: bool,
) -> tuple[torch.Tensor, bool]:
"""Select local or replicated query geometry for MLA decode/prefill."""
if q_dcp_replicated is not None and not full_ckv_dcp:
return q_dcp_replicated[:num_mqa_tokens], True
return q[:num_mqa_tokens], False


def _canonicalize_sparse_mla_kv_cache_dtype(
attn_backend: type[AttentionBackend],
kv_cache_dtype: CacheDType,
Expand Down Expand Up @@ -903,12 +916,19 @@ def forward_impl(
)

if num_mqa_tokens > 0:
if q_dcp_replicated is not None:
mqa_q = q_dcp_replicated[:num_mqa_tokens]
qrep_decode = True
else:
mqa_q = q[:num_mqa_tokens]
qrep_decode = False
full_ckv_dcp = self.impl.uses_full_ckv_dcp( # type: ignore[attr-defined]
attn_metadata, num_mqa_tokens
)
# Full-CKV prefill already makes every rank's cache visible to
# its local query heads. Prefer the local projection even when
# dcp_q_replicate retained a global query for ordinary DCP decode;
# the replicated query does not fit the local-head CKV plan.
mqa_q, qrep_decode = _select_mqa_query(
q,
q_dcp_replicated,
num_mqa_tokens=num_mqa_tokens,
full_ckv_dcp=full_ckv_dcp,
)
mqa_output_slice = output[:num_mqa_tokens]

mqa_q_nope, mqa_q_pe = mqa_q.split(
Expand Down Expand Up @@ -998,7 +1018,7 @@ def forward_impl(
if isinstance(mqa_q, tuple):
# concatenate mqa_ql_nope and mqa_q_pe -> (B, N, L + P)
mqa_q = torch.cat(mqa_q, dim=-1)
if not qrep_decode:
if not qrep_decode and not full_ckv_dcp:
assert self.dcp_manager.query_gather is not None
mqa_q = self.dcp_manager.query_gather(mqa_q)

Expand All @@ -1008,7 +1028,7 @@ def forward_impl(
attn_out, lse = self.impl.forward_mqa(mqa_q, kv_cache, attn_metadata, self) # type: ignore[attr-defined]

# correct dcp attn_out with lse.
if self.impl.dcp_world_size > 1:
if self.impl.dcp_world_size > 1 and not full_ckv_dcp:
assert lse is not None
assert self.dcp_manager is not None
decode_metadata = getattr(attn_metadata, "decode", None)
Expand Down
Loading
Loading