diff --git a/tests/models/kimi_k3/test_mla_padding.py b/tests/models/kimi_k3/test_mla_padding.py index 007556801152..cf137217e41a 100644 --- a/tests/models/kimi_k3/test_mla_padding.py +++ b/tests/models/kimi_k3/test_mla_padding.py @@ -88,6 +88,117 @@ def safe_bmm(q, weight, output, *, use_safe_op): torch.testing.assert_close(result, expected.transpose(0, 1)) +def test_kimi_mla_decode_query_uses_replicated_absorbed_weight(monkeypatch): + from vllm.models.kimi_k3.nvidia import mla + + attention = object.__new__(mla.MultiHeadLatentAttention) + torch.nn.Module.__init__(attention) + attention.kv_lora_rank = 5 + attention.dcp_q_replicate = True + attention.W_UK_T = torch.nn.Parameter( + torch.randn((2, 3, 5), dtype=torch.bfloat16), requires_grad=False + ) + attention.W_UK_T_dcp_qrep = torch.randn((4, 3, 5), dtype=torch.bfloat16) + query = torch.randn((3, 4, 3), dtype=torch.bfloat16) + captured = {} + + def safe_bmm(q, weight, output, *, use_safe_op): + captured["weight"] = weight + torch.bmm(q.contiguous(), weight, out=output) + + monkeypatch.setattr(mla, "_run_mla_query_bmm", safe_bmm) + + result = attention._absorb_decode_query(query) + + assert captured["weight"] is attention.W_UK_T_dcp_qrep + expected = torch.bmm( + query.transpose(0, 1).contiguous(), + attention.W_UK_T_dcp_qrep, + ) + torch.testing.assert_close(result, expected.transpose(0, 1)) + + +def test_kimi_mla_qrep_layer_allowlist(monkeypatch): + from vllm.models.kimi_k3.nvidia import mla + + config = SimpleNamespace( + parallel_config=SimpleNamespace( + decode_context_parallel_size=8, + prefill_context_parallel_size=1, + ) + ) + monkeypatch.setattr(mla.envs, "VLLM_DCP_Q_REPLICATE", True) + monkeypatch.setattr( + mla.envs, + "VLLM_K3_DCP_Q_REPLICATE_LAYERS", + "4,12-16,92", + ) + + assert mla._k3_dcp_qrep_enabled("model.layers.4.self_attn", config) + assert mla._k3_dcp_qrep_enabled("model.layers.14.self_attn", config) + assert not mla._k3_dcp_qrep_enabled("model.layers.20.self_attn", config) + assert mla._k3_dcp_qrep_enabled("model.layers.92.self_attn", config) + + +def test_kimi_mla_qrep_requires_explicit_memory_policy(monkeypatch): + from vllm.models.kimi_k3.nvidia import mla + + config = SimpleNamespace( + parallel_config=SimpleNamespace( + decode_context_parallel_size=8, + prefill_context_parallel_size=1, + ) + ) + monkeypatch.setattr(mla.envs, "VLLM_DCP_Q_REPLICATE", True) + monkeypatch.setattr(mla.envs, "VLLM_K3_DCP_Q_REPLICATE_LAYERS", None) + + with pytest.raises(ValueError, match="explicit layer list"): + mla._k3_dcp_qrep_enabled("model.layers.4.self_attn", config) + + +def test_kimi_mla_qrep_explicit_all(monkeypatch): + from vllm.models.kimi_k3.nvidia import mla + + config = SimpleNamespace( + parallel_config=SimpleNamespace( + decode_context_parallel_size=8, + prefill_context_parallel_size=1, + ) + ) + monkeypatch.setattr(mla.envs, "VLLM_DCP_Q_REPLICATE", True) + monkeypatch.setattr(mla.envs, "VLLM_K3_DCP_Q_REPLICATE_LAYERS", "all") + + assert mla._k3_dcp_qrep_enabled("model.layers.4.self_attn", config) + + +def test_kimi_mla_qrep_uses_dcp_group_head_width(): + from vllm.models.kimi_k3.nvidia import mla + + assert mla._k3_projected_query_heads(12, 8, True) == 96 + assert mla._k3_projected_query_heads(6, 8, True) == 48 + assert mla._k3_projected_query_heads(6, 8, False) == 6 + + +def test_kimi_mla_qrep_rejects_invalid_layer_range(monkeypatch): + from vllm.models.kimi_k3.nvidia import mla + + config = SimpleNamespace( + parallel_config=SimpleNamespace( + decode_context_parallel_size=8, + prefill_context_parallel_size=1, + ) + ) + monkeypatch.setattr(mla.envs, "VLLM_DCP_Q_REPLICATE", True) + monkeypatch.setattr( + mla.envs, + "VLLM_K3_DCP_Q_REPLICATE_LAYERS", + "16-12", + ) + + with pytest.raises(ValueError, match="Invalid K3 qrep layer range"): + mla._k3_dcp_qrep_enabled("model.layers.12.self_attn", config) + + def test_kimi_mla_defines_graph_padding_before_output_projection(monkeypatch): from vllm.models.kimi_k3.nvidia import mla diff --git a/tests/v1/attention/test_b12x_mla.py b/tests/v1/attention/test_b12x_mla.py index 8d2dc5d28b74..8f41cac1fbf9 100644 --- a/tests/v1/attention/test_b12x_mla.py +++ b/tests/v1/attention/test_b12x_mla.py @@ -63,6 +63,20 @@ def test_b12x_mla_limits_active_cache_splits( assert b12x_mla._active_dense_mla_splits(plan, max_seq_len) == expected +def test_b12x_mla_row_caps_cover_full_graph_shapes() -> None: + assert b12x_mla._dense_mla_plan_row_caps(28) == (1, 2, 4, 8, 16, 28) + + +def test_b12x_mla_selects_smallest_covering_plan() -> None: + plans = {1: "b1", 2: "b2", 4: "b4", 8: "b8"} + + assert b12x_mla._select_dense_mla_plan(plans, 1) == "b1" + assert b12x_mla._select_dense_mla_plan(plans, 3) == "b4" + assert b12x_mla._select_dense_mla_plan(plans, 8) == "b8" + with pytest.raises(ValueError, match="exceed the planned capacities"): + b12x_mla._select_dense_mla_plan(plans, 9) + + def test_b12x_mla_plans_local_interleaved_dcp_cache() -> None: config = SimpleNamespace( parallel_config=SimpleNamespace( @@ -76,15 +90,15 @@ def test_b12x_mla_plans_local_interleaved_dcp_cache() -> None: def test_mla_uses_one_kv_shard_for_replicated_dcp_cache() -> None: - replicated = SimpleNamespace(dcp_replicated=True, dcp_kv_shard_count=None) - sharded = SimpleNamespace(dcp_replicated=False, dcp_kv_shard_count=None) + replicated = SimpleNamespace(get_num_dcp_kv_shards=lambda _: 1) + sharded = SimpleNamespace(get_num_dcp_kv_shards=lambda dcp_size: dcp_size) assert mla_attention._get_mla_kv_dcp_world_size(replicated, 16) == 1 assert mla_attention._get_mla_kv_dcp_world_size(sharded, 16) == 16 def test_mla_rejects_partial_dcp_cache_without_matching_subgroup() -> None: - partial = SimpleNamespace(dcp_replicated=False, dcp_kv_shard_count=4) + partial = SimpleNamespace(get_num_dcp_kv_shards=lambda _: 4) with pytest.raises(NotImplementedError, match="partial DCP KV sharding"): mla_attention._get_mla_kv_dcp_world_size(partial, 16) @@ -195,6 +209,7 @@ def _fake_impl(monkeypatch, *, num_heads: int = 8) -> tuple[B12xMLAImpl, _FakeDe impl.dcp_world_size = 1 impl._dcp_comm_backend = "a2a" impl._dcp_max_batch_size = 16 + impl.dcp_q_replicate = False impl._compiled_bindings = set() dense_mla = _FakeDenseMLA() impl._dense_mla = dense_mla @@ -450,6 +465,59 @@ def test_b12x_mla_builder_maps_causal_verification_lengths_to_dcp_rank( ) +def test_b12x_mla_builder_preserves_tiled_q4_dcp_verification( + monkeypatch, +) -> None: + builder = object.__new__(B12xMLAMetadataBuilder) + builder._dense_mla_plan = _FakePlan() + builder._dense_mla_plans = {4: _FakePlan()} + verify_plan = SimpleNamespace(caps=SimpleNamespace(max_page_table_width=4)) + builder._dense_mla_verify_plans = {1: verify_plan} + builder._dense_mla_scratch = torch.empty(256, dtype=torch.uint8) + builder._dense_mla_padded_q = None + builder._dense_mla_padded_output = None + builder._max_dense_mla_rows = 8 + builder._dense_mla_flat_block_table = torch.zeros(8, 4, dtype=torch.int32) + builder._dense_mla_flat_seq_lens = torch.empty(8, dtype=torch.int32) + builder._dense_mla_flat_query_start_loc = torch.arange(9, dtype=torch.int32) + builder._dense_mla_causal_offsets = torch.arange(-3, 1, dtype=torch.int32) + builder._dense_mla_flat_global_seq_lens = torch.empty(8, dtype=torch.int32) + builder._dense_mla_flat_dcp_remainder = torch.empty(8, dtype=torch.int32) + builder.dcp_world_size = 4 + builder._dcp_rank = 3 + builder.cp_kv_cache_interleave_size = 2 + + source_table = torch.tensor([[3, 4, 5, 6, 90]], dtype=torch.int32) + metadata = SimpleNamespace( + causal=True, + num_decodes=1, + num_decode_tokens=4, + decode=SimpleNamespace( + block_table=source_table, + seq_lens=torch.tensor([4], dtype=torch.int32), + dcp_tot_seq_lens=torch.tensor([17], dtype=torch.int32), + ), + ) + monkeypatch.setattr( + b12x_mla.MLACommonMetadataBuilder, + "build", + lambda *args, **kwargs: metadata, + ) + + result = builder.build(0, SimpleNamespace()) + + assert result.dense_mla_plan is verify_plan + torch.testing.assert_close( + result.dense_mla_verify_block_table, + source_table[:, :4], + ) + torch.testing.assert_close( + result.dense_mla_query_cache_seq_lens, + torch.tensor([2, 3, 4, 4], dtype=torch.int32), + ) + assert getattr(result, "dense_mla_flat_block_table", None) is None + + def test_b12x_mla_builder_bounds_single_token_draft_table(monkeypatch) -> None: builder = object.__new__(B12xMLAMetadataBuilder) builder._dense_mla_plan = _FakePlan() @@ -533,6 +601,42 @@ def test_b12x_mla_adapter_uses_flattened_non_causal_rows(monkeypatch) -> None: assert lse is not None and lse.shape == (query_rows, 6) +def test_b12x_mla_adapter_uses_tiled_query_visibility(monkeypatch) -> None: + impl, dense_mla = _fake_impl(monkeypatch) + query_rows = 4 + q = torch.randn(query_rows, 8, 576, dtype=torch.bfloat16) + cache = torch.randn(4, 16, 576, dtype=torch.bfloat16) + source_table = torch.tensor([[0, 1, 2, 3, 90]], dtype=torch.int32) + verify_table = source_table[:, :4].contiguous() + query_cache_seq_lens = torch.tensor([29, 30, 31, 32], dtype=torch.int32) + query_start_loc = torch.tensor([0, 4], dtype=torch.int32) + metadata = SimpleNamespace( + dense_mla_plan=_FakePlan(), + dense_mla_scratch=torch.empty(256, dtype=torch.uint8), + dense_mla_verify_block_table=verify_table, + dense_mla_query_cache_seq_lens=query_cache_seq_lens, + query_start_loc=query_start_loc, + decode=SimpleNamespace( + block_table=source_table, + seq_lens=torch.tensor([32], dtype=torch.int32), + ), + ) + + output, lse = impl.forward_mqa( + q, + cache, + metadata, + SimpleNamespace(_q_scale=None, _k_scale=None), + ) + + binding = dense_mla.bindings[0] + assert binding.page_table is verify_table + assert binding.cu_seqlens_q.data_ptr() == query_start_loc.data_ptr() + assert binding.query_cache_seqlens is query_cache_seq_lens + assert output.shape == (query_rows, 8, 512) + assert lse is not None and lse.shape == (query_rows, 8) + + def test_b12x_mla_adapter_passes_fp8_scales(monkeypatch) -> None: impl, dense_mla = _fake_impl(monkeypatch) q = torch.empty(1, 8, 576, dtype=torch.float8_e4m3fn) @@ -616,6 +720,53 @@ def reduce(output, lse, actual_group, **kwargs): assert lse is None +def test_b12x_mla_adapter_skips_query_gather_for_qrep(monkeypatch) -> None: + impl, dense_mla = _fake_impl(monkeypatch, num_heads=6) + impl.dcp_world_size = 8 + impl.dcp_q_replicate = True + batch = 2 + q = torch.randn(batch, 48, 576, dtype=torch.bfloat16) + cache = torch.randn(4, 16, 576, dtype=torch.bfloat16) + group = SimpleNamespace(world_size=8) + calls: list[str] = [] + + monkeypatch.setattr(b12x_mla, "get_dcp_group", lambda: group) + + def unexpected_gather(*args, **kwargs): + raise AssertionError("qrep must skip the query all-gather") + + def reduce(output, lse, actual_group, **kwargs): + assert actual_group is group + calls.append("reduce") + return output[:, :6] + + monkeypatch.setattr(b12x_mla, "dcp_b12x_all_gather_heads", unexpected_gather) + monkeypatch.setattr(b12x_mla, "dcp_a2a_lse_reduce", reduce) + metadata = SimpleNamespace( + dense_mla_plan=_FakePlan(), + dense_mla_scratch=torch.empty(256, dtype=torch.uint8), + dense_mla_padded_q=torch.empty(batch, 48, 576, dtype=torch.bfloat16), + dense_mla_padded_output=torch.zeros(batch, 48, 512, dtype=torch.bfloat16), + query_start_loc=torch.tensor([0, 1, 2], dtype=torch.int32), + decode=SimpleNamespace( + block_table=torch.tensor([[0, 1], [2, 3]], dtype=torch.int32), + seq_lens=torch.tensor([17, 31], dtype=torch.int32), + ), + ) + + output, lse = impl.forward_mqa( + q, + cache, + metadata, + SimpleNamespace(_q_scale=None, _k_scale=None), + ) + + assert calls == ["reduce"] + assert dense_mla.bindings[0].q.data_ptr() == q.data_ptr() + assert output.shape == (batch, 6, 512) + assert lse is None + + def test_b12x_mla_adapter_skips_dcp_for_replicated_cache(monkeypatch) -> None: impl, dense_mla = _fake_impl(monkeypatch, num_heads=6) impl.dcp_world_size = 8 diff --git a/vllm/envs.py b/vllm/envs.py index 92d8c3e0098a..0ca61813913c 100755 --- a/vllm/envs.py +++ b/vllm/envs.py @@ -236,6 +236,12 @@ VLLM_USE_DEEP_GEMM_E8M0: bool = True VLLM_USE_DEEP_GEMM_TMA_ALIGNED_SCALES: bool = True VLLM_DCP_Q_REPLICATE: bool = False + VLLM_K3_DCP_Q_REPLICATE_LAYERS: str | None = None + VLLM_K3_DYNAMIC_SPARSE_STRIDE: int = 1 + VLLM_K3_DYNAMIC_SPARSE_MIN_TOKENS: int = 0 + VLLM_K3_DYNAMIC_SPARSE_SINK_TOKENS: int = 4096 + VLLM_K3_DYNAMIC_SPARSE_RECENT_TOKENS: int = 32768 + VLLM_K3_DYNAMIC_SPARSE_REFRESH_INTERVAL: int = 128 VLLM_USE_DIRECT_DCP_A2A: bool | None = None VLLM_USE_DIRECT_DCP_Q_GATHER: bool | None = None VLLM_USE_DIRECT_DCP_KV_GATHER: bool | None = None @@ -1805,6 +1811,28 @@ def _resolve_rust_cli_path() -> str | None: ), # Opt-in MLA DCP query replication: skip the decode query all-gather. "VLLM_DCP_Q_REPLICATE": lambda: bool(int(os.getenv("VLLM_DCP_Q_REPLICATE", "0"))), + # Required Kimi-K3 layer allow-list when DCP query replication is enabled. + # Use "all" only after explicitly qualifying the persistent VRAM cost. + "VLLM_K3_DCP_Q_REPLICATE_LAYERS": lambda: os.getenv( + "VLLM_K3_DCP_Q_REPLICATE_LAYERS" + ), + # Experimental Kimi-K3 dense-MLA sparsity. Stride 1 is exact and leaves + # the production path unchanged. + "VLLM_K3_DYNAMIC_SPARSE_STRIDE": lambda: int( + os.getenv("VLLM_K3_DYNAMIC_SPARSE_STRIDE", "1") + ), + "VLLM_K3_DYNAMIC_SPARSE_MIN_TOKENS": lambda: int( + os.getenv("VLLM_K3_DYNAMIC_SPARSE_MIN_TOKENS", "0") + ), + "VLLM_K3_DYNAMIC_SPARSE_SINK_TOKENS": lambda: int( + os.getenv("VLLM_K3_DYNAMIC_SPARSE_SINK_TOKENS", "4096") + ), + "VLLM_K3_DYNAMIC_SPARSE_RECENT_TOKENS": lambda: int( + os.getenv("VLLM_K3_DYNAMIC_SPARSE_RECENT_TOKENS", "32768") + ), + "VLLM_K3_DYNAMIC_SPARSE_REFRESH_INTERVAL": lambda: int( + os.getenv("VLLM_K3_DYNAMIC_SPARSE_REFRESH_INTERVAL", "128") + ), # DeepGemm JITs the kernels on-demand. The warmup attempts to make DeepGemm # JIT all the required kernels before model execution so there is no # JIT'ing in the hot-path. However, this warmup increases the engine diff --git a/vllm/models/kimi_k3/nvidia/mla.py b/vllm/models/kimi_k3/nvidia/mla.py index 766bd1570bfc..330ebfdf2bb6 100644 --- a/vllm/models/kimi_k3/nvidia/mla.py +++ b/vllm/models/kimi_k3/nvidia/mla.py @@ -29,6 +29,7 @@ """ import math +import re from typing import TYPE_CHECKING, cast import torch @@ -43,6 +44,7 @@ get_current_vllm_config, ) from vllm.distributed import ( + get_dcp_group, get_tensor_model_parallel_world_size, ) from vllm.forward_context import get_forward_context @@ -60,6 +62,7 @@ from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import ( ColumnParallelLinear, + DCPGroupColumnParallelLinear, MergedColumnParallelLinear, ReplicatedLinear, RowParallelLinear, @@ -117,6 +120,64 @@ _MLA_CALLER_OUTPUT_MIN_TOKENS = 1024 +def _parse_k3_qrep_layers(spec: str) -> frozenset[int] | None: + if spec.strip().lower() == "all": + return None + layers: set[int] = set() + for item in spec.split(","): + item = item.strip() + if not item: + continue + if "-" in item: + start_text, end_text = item.split("-", 1) + start = int(start_text) + end = int(end_text) + if start < 0 or end < start: + raise ValueError(f"Invalid K3 qrep layer range: {item!r}") + layers.update(range(start, end + 1)) + else: + layer = int(item) + if layer < 0: + raise ValueError(f"Invalid K3 qrep layer: {item!r}") + layers.add(layer) + return frozenset(layers) + + +def _k3_dcp_qrep_enabled(prefix: str, vllm_config: VllmConfig) -> bool: + parallel_config = vllm_config.parallel_config + if ( + not envs.VLLM_DCP_Q_REPLICATE + or parallel_config.decode_context_parallel_size <= 1 + or parallel_config.prefill_context_parallel_size > 1 + ): + return False + layer_spec = envs.VLLM_K3_DCP_Q_REPLICATE_LAYERS + if layer_spec is None: + raise ValueError( + "Kimi-K3 DCP query replication duplicates query and absorbed " + "projection weights. Set VLLM_K3_DCP_Q_REPLICATE_LAYERS to an " + "explicit layer list/range, or to 'all' after verifying the VRAM " + "budget." + ) + match = re.search(r"(?:^|\.)layers\.(\d+)(?:\.|$)", prefix) + if match is None: + raise ValueError( + "VLLM_K3_DCP_Q_REPLICATE_LAYERS requires a layer-qualified prefix, " + f"got {prefix!r}" + ) + layers = _parse_k3_qrep_layers(layer_spec) + return layers is None or int(match.group(1)) in layers + + +def _k3_projected_query_heads( + num_local_heads: int, + dcp_world_size: int, + dcp_q_replicate: bool, +) -> int: + """Return the DCP-group head width emitted by the query projection.""" + return int(num_local_heads) * (int(dcp_world_size) if dcp_q_replicate else 1) + + @torch.compile(backend=current_platform.simple_compile_backend) def _gate_sigmoid_mul(attn_out: torch.Tensor, gate: torch.Tensor) -> torch.Tensor: """Apply the sigmoid output gate to a precomputed ``g_proj`` projection.""" @@ -294,6 +355,13 @@ def __init__( assert num_heads % tp_size == 0 self.num_heads = num_heads self.num_local_heads = num_heads // tp_size + vllm_config = get_current_vllm_config() + self.dcp_q_replicate = _k3_dcp_qrep_enabled(prefix, vllm_config) + q_proj_cls = ( + DCPGroupColumnParallelLinear + if self.dcp_q_replicate + else ColumnParallelLinear + ) # ---- Pre-attention projections (fusable front-end) ---- # Two query variants: a low-rank q-LoRA path (Kimi-K3) fused with the @@ -325,7 +393,7 @@ def __init__( disable_tp=True, ) self.q_a_layernorm = RMSNorm(self.q_lora_rank, eps=config.rms_norm_eps) - self.q_b_proj = ColumnParallelLinear( + self.q_b_proj = q_proj_cls( self.q_lora_rank, self.num_heads * self.qk_head_dim, bias=False, @@ -335,7 +403,7 @@ def __init__( else: # Uncompressed query: full-rank q_proj (TP-split over heads) plus a # replicated kv-down projection (shared latent across TP ranks). - self.q_proj = ColumnParallelLinear( + self.q_proj = q_proj_cls( self.hidden_size, self.num_heads * self.qk_head_dim, bias=False, @@ -440,7 +508,6 @@ def __init__( self.impl.dcp_rank = 0 self.q_pad_num_heads = getattr(self.impl, "q_pad_num_heads", None) - vllm_config = get_current_vllm_config() parallel_config = vllm_config.parallel_config assert parallel_config.prefill_context_parallel_size == 1, ( "Kimi-K3 MultiHeadLatentAttention does not support prefill context " @@ -450,6 +517,13 @@ def __init__( self.backend_owns_decode_dcp = _backend_owns_decode_dcp( self.impl, self.dcp_world_size ) + if self.dcp_q_replicate: + if not self.backend_owns_decode_dcp: + raise NotImplementedError( + "Kimi-K3 DCP query replication requires a backend-owned " + "decode DCP path." + ) + self.impl.dcp_q_replicate = True assert ( self.dcp_world_size <= 1 or self.rotary_emb is None @@ -639,6 +713,12 @@ def process_weights_after_loading(self, act_dtype: torch.dtype) -> None: pre_w_uk_t.copy_(w_uk_t) w_uk_t = pre_w_uk_t replace_parameter(self, "W_UK_T", w_uk_t, prefer_copy=True) + self.W_UK_T_dcp_qrep: torch.Tensor | None = None + if self.dcp_q_replicate: + self.W_UK_T_dcp_qrep = get_dcp_group().all_gather( + self.W_UK_T.contiguous(), + dim=0, + ) quant_method = ( self.quant_config.get_quant_method(self, prefix=self.layer_name) @@ -678,9 +758,15 @@ def _absorb_decode_query(self, q_nope: torch.Tensor) -> torch.Tensor: """ query = q_nope.transpose(0, 1).contiguous() output = query.new_empty((query.shape[0], query.shape[1], self.kv_lora_rank)) + weight = ( + self.W_UK_T_dcp_qrep + if getattr(self, "dcp_q_replicate", False) + else self.W_UK_T + ) + assert weight is not None _run_mla_query_bmm( query, - self.W_UK_T, + weight, output, use_safe_op=True, ) @@ -730,13 +816,21 @@ def _forward_attn( self.kv_a_layernorm.weight.data, self.rms_norm_eps, ) - q = self.q_b_proj(q_c)[0].view(-1, self.num_local_heads, self.qk_head_dim) + q_heads = _k3_projected_query_heads( + self.num_local_heads, + self.dcp_world_size, + self.dcp_q_replicate, + ) + q = self.q_b_proj(q_c)[0].view(-1, q_heads, self.qk_head_dim) else: # Uncompressed query: project directly (no q-LoRA, no q norm) and # normalize only the kv latent. - q = self.q_proj(hidden_states)[0].view( - -1, self.num_local_heads, self.qk_head_dim + q_heads = _k3_projected_query_heads( + self.num_local_heads, + self.dcp_world_size, + self.dcp_q_replicate, ) + q = self.q_proj(hidden_states)[0].view(-1, q_heads, self.qk_head_dim) kv_lora = self.kv_a_proj_with_mqa(hidden_states)[0] kv_c, k_pe = kv_lora.split( [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1 @@ -855,8 +949,12 @@ def _attention( # ---- Prefill: fused key-concat + cache-insert + attention ---- if num_mha_tokens > 0: + prefill_q = q[num_mqa_tokens:] + if getattr(self, "dcp_q_replicate", False): + q_proj = self.q_b_proj if self.q_lora_rank is not None else self.q_proj + prefill_q = q_proj._local_view(prefill_q) self._forward_prefill_fused( - q[num_mqa_tokens:], + prefill_q, kv_c_normed[num_mqa_tokens:], k_pe[num_mqa_tokens:], rope_positions[num_mqa_tokens:] if rope_positions is not None else None, diff --git a/vllm/v1/attention/backends/mla/b12x_mla.py b/vllm/v1/attention/backends/mla/b12x_mla.py index b22c4c253673..98c017fd5afc 100644 --- a/vllm/v1/attention/backends/mla/b12x_mla.py +++ b/vllm/v1/attention/backends/mla/b12x_mla.py @@ -4,11 +4,13 @@ from __future__ import annotations +from bisect import bisect_left from dataclasses import dataclass from typing import Any, ClassVar, cast import torch +from vllm import envs from vllm.config import VllmConfig, get_current_vllm_config from vllm.config.cache import CacheDType from vllm.distributed.parallel_state import get_dcp_group @@ -141,6 +143,34 @@ def _active_dense_mla_splits(plan: Any, max_seq_len: int | None) -> int: ) +def _dense_mla_plan_row_caps(max_rows: int) -> tuple[int, ...]: + """Return CUDA-graph-friendly row capacities through ``max_rows``.""" + if max_rows <= 0: + raise ValueError("dense MLA row capacity must be positive") + caps: list[int] = [] + row_cap = 1 + while row_cap < max_rows: + caps.append(row_cap) + row_cap *= 2 + caps.append(max_rows) + return tuple(caps) + + +def _select_dense_mla_plan( + plans: dict[int, Any], + total_rows: int, +) -> Any: + """Select the smallest launch plan that covers the live query rows.""" + row_caps = tuple(sorted(plans)) + index = bisect_left(row_caps, total_rows) + if total_rows <= 0 or index >= len(row_caps): + raise ValueError( + "B12X_MLA query rows exceed the planned capacities: " + f"rows={total_rows}, capacities={row_caps}" + ) + return plans[row_caps[index]] + + def _create_dense_mla_plan( vllm_config: VllmConfig, device: torch.device, @@ -148,6 +178,9 @@ def _create_dense_mla_plan( page_size: int, num_q_heads: int, max_total_q: int | None = None, + max_batch: int | None = None, + mode: str = "decode", + uses_query_cache_seqlens: bool = False, dcp_size: int | None = None, max_cache_tokens: int | None = None, ) -> Any: @@ -162,6 +195,40 @@ def _create_dense_mla_plan( if max_cache_tokens is not None else _max_dcp_local_cache_tokens(vllm_config, dcp_size=dcp_size) ) + max_batch = int(max_total_q if max_batch is None else max_batch) + dcp_size = int( + vllm_config.parallel_config.decode_context_parallel_size + if dcp_size is None + else dcp_size + ) + + def local_tokens(global_tokens: int) -> int: + return (max(int(global_tokens), 0) + dcp_size - 1) // dcp_size + + sparse_stride = int(envs.VLLM_K3_DYNAMIC_SPARSE_STRIDE) + if sparse_stride > 1: + logger.warning_once( + "Kimi-K3 dynamic sparse MLA is enabled with stride=%d. This " + "changes attention semantics and requires model-quality " + "qualification; stride=1 is the exact production default.", + sparse_stride, + ) + sparse_min_tokens = local_tokens(envs.VLLM_K3_DYNAMIC_SPARSE_MIN_TOKENS) + sparse_sink_chunks = ( + local_tokens(envs.VLLM_K3_DYNAMIC_SPARSE_SINK_TOKENS) + 63 + ) // 64 + sparse_recent_chunks = ( + local_tokens(envs.VLLM_K3_DYNAMIC_SPARSE_RECENT_TOKENS) + 63 + ) // 64 + # A local DCP shard length can differ from its peers at an interleave + # boundary. Until the kernel accepts a global refresh clock, disabling the + # periodic dense refresh under DCP avoids mixing dense and sparse shards in + # one exact LSE reduction. The sparse policy itself remains rank-consistent. + sparse_refresh_interval = ( + local_tokens(envs.VLLM_K3_DYNAMIC_SPARSE_REFRESH_INTERVAL) + if dcp_size == 1 + else 0 + ) if max_total_q > _MAX_B12X_QUERY_ROWS: raise ValueError( "B12X_MLA supports at most " @@ -175,17 +242,23 @@ def _create_dense_mla_plan( caps = dense_mla.Caps( device=device, - mode="decode", + mode=mode, dtype=torch.bfloat16, kv_dtype=_planned_kv_dtype(vllm_config), num_q_heads=num_q_heads, page_size=page_size, max_total_q=max_total_q, - max_batch=max_total_q, + max_batch=max_batch, max_cache_tokens=max_cache_tokens, max_page_table_width=_page_table_width(max_cache_tokens, page_size), num_cache_pages=_MAX_I32, use_cuda_graph=True, + uses_query_cache_seqlens=uses_query_cache_seqlens, + sparse_stride=sparse_stride, + sparse_min_tokens=sparse_min_tokens, + sparse_sink_chunks=sparse_sink_chunks, + sparse_recent_chunks=sparse_recent_chunks, + sparse_refresh_interval=sparse_refresh_interval, ) return dense_mla.plan(caps) @@ -201,6 +274,8 @@ class B12xMLAMetadata(MLACommonMetadata): dense_mla_flat_block_table: torch.Tensor | None = None dense_mla_flat_seq_lens: torch.Tensor | None = None dense_mla_flat_query_start_loc: torch.Tensor | None = None + dense_mla_verify_block_table: torch.Tensor | None = None + dense_mla_query_cache_seq_lens: torch.Tensor | None = None dense_mla_dcp_world_size: int = 1 @@ -245,19 +320,55 @@ def __init__( sliding_window = getattr(kv_cache_spec, "sliding_window", None) if sliding_window is not None: max_cache_tokens = min(max_cache_tokens, int(sliding_window)) - self._dense_mla_plan = _create_dense_mla_plan( - vllm_config, - device, - page_size=self.page_size, - num_q_heads=self._kernel_heads, - max_total_q=max_dense_mla_rows, - dcp_size=self.dcp_world_size, - max_cache_tokens=max_cache_tokens, - ) - self._workspace_specs = self._dense_mla_plan.shapes_and_dtypes() - if len(self._workspace_specs) != 1: - raise RuntimeError("B12X_MLA expected exactly one scratch buffer.") - scratch_shape, scratch_dtype = self._workspace_specs[0] + self._dense_mla_plans = { + rows: _create_dense_mla_plan( + vllm_config, + device, + page_size=self.page_size, + num_q_heads=self._kernel_heads, + max_total_q=rows, + dcp_size=self.dcp_world_size, + max_cache_tokens=max_cache_tokens, + ) + for rows in _dense_mla_plan_row_caps(max_dense_mla_rows) + } + self._dense_mla_verify_plans: dict[int, Any] = {} + if _planned_kv_dtype(vllm_config) == torch.float8_e4m3fn: + self._dense_mla_verify_plans = { + batch: _create_dense_mla_plan( + vllm_config, + device, + page_size=self.page_size, + num_q_heads=self._kernel_heads, + max_total_q=batch * 4, + max_batch=batch, + mode="verify", + uses_query_cache_seqlens=True, + dcp_size=self.dcp_world_size, + max_cache_tokens=max_cache_tokens, + ) + for batch in range( + 1, + int(vllm_config.scheduler_config.max_num_seqs) + 1, + ) + } + self._dense_mla_plan = self._dense_mla_plans[max_dense_mla_rows] + workspace_specs = [ + plan.shapes_and_dtypes() + for plan in ( + *self._dense_mla_plans.values(), + *self._dense_mla_verify_plans.values(), + ) + ] + if any(len(specs) != 1 for specs in workspace_specs): + raise RuntimeError("B12X_MLA expected exactly one scratch buffer per plan.") + scratch_dtype = workspace_specs[0][0][1] + if any(specs[0][1] != scratch_dtype for specs in workspace_specs): + raise RuntimeError("B12X_MLA plan scratch dtypes do not match.") + scratch_shape = max( + (specs[0][0] for specs in workspace_specs), + key=lambda shape: shape[0], + ) # Every attention layer represented by this builder executes serially # on the model stream. One builder-owned buffer therefore gives each # eager bind a stable caller-owned address without a backend workspace @@ -321,18 +432,83 @@ def __init__( else None ) logger.info_once( - "B12X dense K3 MLA plan: local_heads=%d, effective_heads=%d, " + "B12X dense K3 MLA plans: local_heads=%d, effective_heads=%d, " "kernel_heads=%d, page_size=%d, " - "max_decode_rows=%d, max_cache_tokens=%d, splits=%d", + "max_decode_rows=%d, max_cache_tokens=%d, rows/splits=%s, " + "verify_batch/splits=%s", self.num_heads, self._effective_heads, self._kernel_heads, self.page_size, max_dense_mla_rows, max_cache_tokens, - self._dense_mla_plan.num_splits, + ",".join( + f"{rows}/{plan.num_splits}" + for rows, plan in self._dense_mla_plans.items() + ), + ",".join( + f"{batch}/{plan.num_splits}" + for batch, plan in self._dense_mla_verify_plans.items() + ) + or "disabled", ) + def _materialize_query_cache_seq_lens( + self, + metadata: B12xMLAMetadata, + decode_metadata: Any, + *, + query_len: int, + total_q: int, + ) -> torch.Tensor: + flat_lens = self._dense_mla_flat_seq_lens[:total_q] + if not metadata.causal: + flat_lens.copy_( + decode_metadata.seq_lens[:, None].expand(-1, query_len).reshape(total_q) + ) + return flat_lens + + offsets = self._dense_mla_causal_offsets[-query_len:] + if self.dcp_world_size == 1: + torch.add( + decode_metadata.seq_lens[:, None], + offsets, + out=flat_lens.view(metadata.num_decodes, query_len), + ) + return flat_lens + + global_source_lens = decode_metadata.dcp_tot_seq_lens + if global_source_lens is None: + raise RuntimeError( + "B12X_MLA causal DCP verification requires global decode " + "sequence lengths." + ) + assert self._dense_mla_flat_global_seq_lens is not None + assert self._dense_mla_flat_dcp_remainder is not None + global_flat_lens = self._dense_mla_flat_global_seq_lens[:total_q] + torch.add( + global_source_lens[:, None], + offsets, + out=global_flat_lens.view(metadata.num_decodes, query_len), + ) + virtual_block = self.dcp_world_size * self.cp_kv_cache_interleave_size + torch.div( + global_flat_lens, + virtual_block, + rounding_mode="floor", + out=flat_lens, + ) + flat_lens.mul_(self.cp_kv_cache_interleave_size) + remainder = self._dense_mla_flat_dcp_remainder[:total_q] + torch.remainder(global_flat_lens, virtual_block, out=remainder) + remainder.sub_(self._dcp_rank * self.cp_kv_cache_interleave_size) + remainder.clamp_( + min=0, + max=self.cp_kv_cache_interleave_size, + ) + flat_lens.add_(remainder) + return flat_lens + def build( self, common_prefix_len: int, @@ -347,95 +523,74 @@ def build( fast_build=fast_build, ), ) - metadata.dense_mla_plan = self._dense_mla_plan + live_rows = max(1, int(metadata.num_decode_tokens)) + plans = getattr( + self, + "_dense_mla_plans", + {self._max_dense_mla_rows: self._dense_mla_plan}, + ) + metadata.dense_mla_plan = _select_dense_mla_plan(plans, live_rows) metadata.dense_mla_scratch = self._dense_mla_scratch metadata.dense_mla_padded_q = self._dense_mla_padded_q metadata.dense_mla_padded_output = self._dense_mla_padded_output metadata.dense_mla_dcp_world_size = self.dcp_world_size decode_metadata = metadata.decode - flatten_decode = False - if decode_metadata is not None and metadata.num_decodes > 0: - flatten_decode = metadata.num_decode_tokens > metadata.num_decodes or int( - decode_metadata.block_table.shape[1] - ) > int(self._dense_mla_plan.caps.max_page_table_width) - if flatten_decode: - assert decode_metadata is not None - total_q = int(metadata.num_decode_tokens) - if total_q > self._max_dense_mla_rows: - raise ValueError( - "B12X_MLA query block exceeds its flattened capacity: " - f"rows={total_q}, capacity={self._max_dense_mla_rows}." - ) - if total_q % metadata.num_decodes: - raise ValueError( - "B12X_MLA requires a uniform query block, got " - f"tokens={total_q}, requests={metadata.num_decodes}." - ) - query_len = total_q // metadata.num_decodes - source_table = decode_metadata.block_table - flat_table = self._dense_mla_flat_block_table[:total_q] - # A bounded speculative cache can retain a position-indexed worker - # table wider than the resident cache. Sequence lengths make the - # omitted suffix unreachable by the dense-MLA kernel. - source_width = min(int(source_table.shape[1]), int(flat_table.shape[1])) - flat_table[:, :source_width].copy_( - source_table[:, None, :source_width] - .expand(-1, query_len, -1) - .reshape(total_q, source_width) - ) - flat_lens = self._dense_mla_flat_seq_lens[:total_q] - if metadata.causal: - offsets = self._dense_mla_causal_offsets[-query_len:] - if self.dcp_world_size > 1: - global_source_lens = decode_metadata.dcp_tot_seq_lens - if global_source_lens is None: - raise RuntimeError( - "B12X_MLA causal DCP verification requires global " - "decode sequence lengths." - ) - assert self._dense_mla_flat_global_seq_lens is not None - assert self._dense_mla_flat_dcp_remainder is not None - global_flat_lens = self._dense_mla_flat_global_seq_lens[:total_q] - torch.add( - global_source_lens[:, None], - offsets, - out=global_flat_lens.view(metadata.num_decodes, query_len), - ) - virtual_block = ( - self.dcp_world_size * self.cp_kv_cache_interleave_size - ) - torch.div( - global_flat_lens, - virtual_block, - rounding_mode="floor", - out=flat_lens, - ) - flat_lens.mul_(self.cp_kv_cache_interleave_size) - remainder = self._dense_mla_flat_dcp_remainder[:total_q] - torch.remainder(global_flat_lens, virtual_block, out=remainder) - remainder.sub_(self._dcp_rank * self.cp_kv_cache_interleave_size) - remainder.clamp_( - min=0, - max=self.cp_kv_cache_interleave_size, - ) - flat_lens.add_(remainder) - else: - torch.add( - decode_metadata.seq_lens[:, None], - offsets, - out=flat_lens.view(metadata.num_decodes, query_len), - ) - else: - flat_lens.copy_( - decode_metadata.seq_lens[:, None] - .expand(-1, query_len) - .reshape(total_q) - ) - metadata.dense_mla_flat_block_table = flat_table - metadata.dense_mla_flat_seq_lens = flat_lens - metadata.dense_mla_flat_query_start_loc = ( - self._dense_mla_flat_query_start_loc[: total_q + 1] + if decode_metadata is None or metadata.num_decodes <= 0: + return metadata + multi_query = metadata.num_decode_tokens > metadata.num_decodes + table_too_wide = int(decode_metadata.block_table.shape[1]) > int( + self._dense_mla_plan.caps.max_page_table_width + ) + if not (multi_query or table_too_wide): + return metadata + + total_q = int(metadata.num_decode_tokens) + if total_q > self._max_dense_mla_rows: + raise ValueError( + "B12X_MLA query block exceeds its flattened capacity: " + f"rows={total_q}, capacity={self._max_dense_mla_rows}." + ) + if total_q % metadata.num_decodes: + raise ValueError( + "B12X_MLA requires a uniform query block, got " + f"tokens={total_q}, requests={metadata.num_decodes}." ) + query_len = total_q // metadata.num_decodes + source_table = decode_metadata.block_table + flat_lens = self._materialize_query_cache_seq_lens( + metadata, + decode_metadata, + query_len=query_len, + total_q=total_q, + ) + verify_plans = getattr(self, "_dense_mla_verify_plans", {}) + tiled_verify = ( + metadata.causal and query_len == 4 and metadata.num_decodes in verify_plans + ) + if tiled_verify: + verify_table = self._dense_mla_flat_block_table[: metadata.num_decodes] + source_width = min( + int(source_table.shape[1]), + int(verify_table.shape[1]), + ) + verify_table[:, :source_width].copy_(source_table[:, :source_width]) + metadata.dense_mla_plan = verify_plans[metadata.num_decodes] + metadata.dense_mla_verify_block_table = verify_table + metadata.dense_mla_query_cache_seq_lens = flat_lens + return metadata + + flat_table = self._dense_mla_flat_block_table[:total_q] + source_width = min(int(source_table.shape[1]), int(flat_table.shape[1])) + flat_table[:, :source_width].copy_( + source_table[:, None, :source_width] + .expand(-1, query_len, -1) + .reshape(total_q, source_width) + ) + metadata.dense_mla_flat_block_table = flat_table + metadata.dense_mla_flat_seq_lens = flat_lens + metadata.dense_mla_flat_query_start_loc = self._dense_mla_flat_query_start_loc[ + : total_q + 1 + ] return metadata @@ -638,6 +793,7 @@ def __init__( self._dense_mla = _load_dense_mla() self._dcp_comm_backend = vllm_config.parallel_config.dcp_comm_backend self._dcp_max_batch_size = vllm_config.scheduler_config.max_num_batched_tokens + self.dcp_q_replicate = False self._compiled_bindings: set[tuple[object, ...]] = set() def forward_mqa( @@ -663,6 +819,18 @@ def forward_mqa( block_table = attn_metadata.decode.block_table seq_lens = attn_metadata.decode.seq_lens query_start_loc = attn_metadata.query_start_loc + query_cache_seq_lens = getattr( + attn_metadata, + "dense_mla_query_cache_seq_lens", + None, + ) + verify_block_table = getattr( + attn_metadata, + "dense_mla_verify_block_table", + None, + ) + if verify_block_table is not None: + block_table = verify_block_table flat_block_table = getattr(attn_metadata, "dense_mla_flat_block_table", None) if flat_block_table is not None: block_table = flat_block_table @@ -677,16 +845,11 @@ def forward_mqa( batch = int(seq_lens.shape[0]) total_q = int(q.shape[0]) - if total_q != batch: + if query_cache_seq_lens is None and total_q != batch: raise ValueError( "B12X_MLA requires one query row per prepared decode sequence, " f"got {total_q} rows for {batch} sequences." ) - if int(q.shape[1]) != self.num_heads: - raise ValueError( - f"B12X_MLA expected {self.num_heads} query heads, got {q.shape[1]}." - ) - metadata_dcp_world_size = int( getattr(attn_metadata, "dense_mla_dcp_world_size", self.dcp_world_size) ) @@ -697,33 +860,41 @@ def forward_mqa( ) effective_heads = self.num_heads * metadata_dcp_world_size kernel_heads = _kernel_query_heads(self.num_heads, metadata_dcp_world_size) + qrep_decode = self.dcp_q_replicate and metadata_dcp_world_size > 1 + expected_input_heads = effective_heads if qrep_decode else self.num_heads + if int(q.shape[1]) != expected_input_heads: + raise ValueError( + f"B12X_MLA expected {expected_input_heads} query heads, " + f"got {q.shape[1]}." + ) dcp_group = None if metadata_dcp_world_size > 1: dcp_group = get_dcp_group() - gathered_q = getattr(attn_metadata, "dense_mla_padded_q", None) - if gathered_q is None: - raise RuntimeError( - "B12X_MLA DCP metadata is missing caller-owned query storage." - ) - if int(gathered_q.shape[0]) < total_q: - raise ValueError( - "B12X_MLA DCP query capacity is smaller than the decode " - f"batch: capacity={gathered_q.shape[0]}, required={total_q}." - ) - if gathered_q.dtype != q.dtype: - raise TypeError( - "B12X_MLA DCP query storage does not match the live query: " - f"buffer={gathered_q.dtype}, query={q.dtype}." + if not qrep_decode: + gathered_q = getattr(attn_metadata, "dense_mla_padded_q", None) + if gathered_q is None: + raise RuntimeError( + "B12X_MLA DCP metadata is missing caller-owned query storage." + ) + if int(gathered_q.shape[0]) < total_q: + raise ValueError( + "B12X_MLA DCP query capacity is smaller than the decode " + f"batch: capacity={gathered_q.shape[0]}, required={total_q}." + ) + if gathered_q.dtype != q.dtype: + raise TypeError( + "B12X_MLA DCP query storage does not match the live query: " + f"buffer={gathered_q.dtype}, query={q.dtype}." + ) + gathered_q = gathered_q[:total_q, :effective_heads] + q = dcp_b12x_all_gather_heads( + q, + dcp_group, + max_batch_size=self._dcp_max_batch_size, + output_head_dim=self.kv_lora_rank, + out=gathered_q, ) - gathered_q = gathered_q[:total_q, :effective_heads] - q = dcp_b12x_all_gather_heads( - q, - dcp_group, - max_batch_size=self._dcp_max_batch_size, - output_head_dim=self.kv_lora_rank, - out=gathered_q, - ) actual_heads = int(q.shape[1]) if actual_heads != effective_heads: @@ -790,6 +961,7 @@ def forward_mqa( output=output, page_table=block_table, cache_seqlens=seq_lens, + query_cache_seqlens=query_cache_seq_lens, cu_seqlens_q=query_start_loc[: batch + 1], q_scale=layer._q_scale if quantized else None, kv_scale=layer._k_scale if quantized else None,