diff --git a/.github/workflows/configs/nightly_config.yaml b/.github/workflows/configs/nightly_config.yaml index a2da55967cf2..7064b25bb825 100644 --- a/.github/workflows/configs/nightly_config.yaml +++ b/.github/workflows/configs/nightly_config.yaml @@ -237,6 +237,10 @@ a3: multi_card: test_config: # pytest-driven tests + - name: kimi-k3-execution-parity + os: linux-aarch64-nightly-a3-16 + tests: tests/e2e/nightly/single_node/models/test_kimi_k3_execution_parity.py + testcase_timeout: 180 - name: qwen3-30b-acc os: linux-aarch64-nightly-a3-4 tests: tests/e2e/weekly/single_node/models/test_qwen3_30b_acc.py diff --git a/.github/workflows/scripts/test_config.yaml b/.github/workflows/scripts/test_config.yaml index 002709e81e1f..f36d6f3e4a48 100644 --- a/.github/workflows/scripts/test_config.yaml +++ b/.github/workflows/scripts/test_config.yaml @@ -311,6 +311,7 @@ - tests/ut/models/test_deepseek_v4_compressor.py - tests/ut/models/test_deepseek_v4_indexer.py - tests/ut/models/test_deepseek_v4_moe.py + - tests/ut/models/test_kimi_k3_adapter.py - tests/e2e/pull_request/four_card/test_deepseek_v4.py - name: models_minimax_m3 diff --git a/tests/e2e/nightly/single_node/models/fixtures/kimi_k3_5layers_16experts/config.json b/tests/e2e/nightly/single_node/models/fixtures/kimi_k3_5layers_16experts/config.json new file mode 100644 index 000000000000..0dd1a5e42c0e --- /dev/null +++ b/tests/e2e/nightly/single_node/models/fixtures/kimi_k3_5layers_16experts/config.json @@ -0,0 +1,63 @@ +{ + "activation_situ_beta": 4, + "activation_situ_linear_beta": 25, + "architectures": [ + "KimiLinearForCausalLM" + ], + "attn_res_block_size": 12, + "bos_token_id": 163584, + "dtype": "bfloat16", + "eos_token_id": 163586, + "first_k_dense_replace": 1, + "hidden_act": "situ", + "hidden_size": 7168, + "intermediate_size": 33792, + "kv_lora_rank": 512, + "latent_moe_use_norm": true, + "linear_attn_config": { + "full_attn_layers": [ + 4 + ], + "gate_lower_bound": -5, + "head_dim": 128, + "kda_layers": [ + 1, + 2, + 3, + 5 + ], + "num_heads": 96, + "short_conv_kernel_size": 4, + "use_full_rank_gate": true + }, + "max_position_embeddings": 1048576, + "mla_use_nope": true, + "mla_use_output_gate": true, + "model_type": "kimi_linear", + "moe_intermediate_size": 3072, + "moe_layer_freq": 1, + "moe_renormalize": true, + "moe_router_activation_func": "sigmoid", + "num_attention_heads": 96, + "num_expert_group": 1, + "num_experts": 16, + "num_experts_per_token": 16, + "num_hidden_layers": 5, + "num_key_value_heads": 96, + "num_nextn_predict_layers": 0, + "num_shared_experts": 2, + "pad_token_id": 163839, + "q_lora_rank": 1536, + "qk_nope_head_dim": 128, + "qk_rope_head_dim": 64, + "rms_norm_eps": 1e-05, + "routed_expert_hidden_size": 3584, + "routed_scaling_factor": 1, + "tie_word_embeddings": false, + "topk_group": 1, + "topk_method": "noaux_tc", + "use_cache": true, + "use_grouped_topk": true, + "v_head_dim": 128, + "vocab_size": 163840 +} diff --git a/tests/e2e/nightly/single_node/models/test_kimi_k3_execution_parity.py b/tests/e2e/nightly/single_node/models/test_kimi_k3_execution_parity.py new file mode 100644 index 000000000000..a28c793923b5 --- /dev/null +++ b/tests/e2e/nightly/single_node/models/test_kimi_k3_execution_parity.py @@ -0,0 +1,110 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved. +"""Storage-light Kimi K3 execution parity guard. + +The committed fixture keeps Kimi K3's production dimensions and its mixed +KDA/MLA layout, but limits the model to five layers and sixteen experts. Dummy +weights deliberately make this an execution-parity test, not a semantic +accuracy test. Full-checkpoint GPQA remains a separate release gate. +""" + +from pathlib import Path + +import pytest +import torch +from vllm import SamplingParams +from vllm.inputs import TokensPrompt + +from tests.e2e.conftest import VllmRunner, wait_until_npu_memory_free + +MODEL_CONFIG = Path(__file__).parent / "fixtures" / "kimi_k3_5layers_16experts" +SCHEDULER_BLOCK_SIZE = 16 +PROMPT_TOKEN_IDS = [163584, *range(100, 100 + SCHEDULER_BLOCK_SIZE)] +MAX_TOKENS = 4 + + +def _assert_complete_output(request_output): + assert request_output is not None + assert request_output.finished + assert request_output.outputs is not None + assert len(request_output.outputs) == 1 + + completion = request_output.outputs[0] + assert completion is not None + assert completion.token_ids is not None + assert len(completion.token_ids) == MAX_TOKENS + assert completion.logprobs is not None + assert len(completion.logprobs) == MAX_TOKENS + + chosen_logprobs = [] + for token_id, step_logprobs in zip(completion.token_ids, completion.logprobs): + assert step_logprobs is not None + assert token_id in step_logprobs + logprob = step_logprobs[token_id].logprob + assert logprob is not None + assert torch.isfinite(torch.tensor(logprob)) + chosen_logprobs.append(logprob) + + return list(completion.token_ids), torch.tensor(chosen_logprobs, dtype=torch.float32) + + +@pytest.mark.e2e_model("sgl-npu/Kimi-K3-W4A8") +@pytest.mark.e2e_coverage( + arch="moe", + feature="aclgraph,prefix_caching,logprobs", + parallel="TP,EP", + deploy="pd_mix", + hardware="A3", + quantization="BF16", + graph_mode="full_decode_only", +) +@wait_until_npu_memory_free() +def test_kimi_k3_dummy_prefix_cache_one_token_prefill_parity(): + """Compare cold prefill with the cached block-size-plus-one path.""" + sampling_params = SamplingParams( + temperature=0, + max_tokens=MAX_TOKENS, + logprobs=1, + ignore_eos=True, + seed=0, + ) + prompt = TokensPrompt(prompt_token_ids=PROMPT_TOKEN_IDS) + + with VllmRunner( + str(MODEL_CONFIG), + skip_tokenizer_init=True, + load_format="dummy", + dtype="bfloat16", + seed=0, + block_size=SCHEDULER_BLOCK_SIZE, + max_model_len=64, + max_num_seqs=1, + max_num_batched_tokens=64, + tensor_parallel_size=16, + enable_expert_parallel=True, + enable_prefix_caching=True, + gpu_memory_utilization=0.75, + compilation_config={ + "cudagraph_mode": "FULL_DECODE_ONLY", + "cudagraph_capture_sizes": [1], + }, + ) as vllm_model: + cold = vllm_model.model.generate([prompt], sampling_params, use_tqdm=False)[0] + hit = vllm_model.model.generate([prompt], sampling_params, use_tqdm=False)[0] + + assert cold.num_cached_tokens in (None, 0) + assert hit.num_cached_tokens == SCHEDULER_BLOCK_SIZE + assert len(PROMPT_TOKEN_IDS) - hit.num_cached_tokens == 1 + + cold_tokens, cold_logprobs = _assert_complete_output(cold) + hit_tokens, hit_logprobs = _assert_complete_output(hit) + assert hit_tokens == cold_tokens + torch.testing.assert_close(hit_logprobs, cold_logprobs, rtol=5e-3, atol=5e-3) + + assert vllm_model.model.reset_prefix_cache() + reset = vllm_model.model.generate([prompt], sampling_params, use_tqdm=False)[0] + assert reset.num_cached_tokens in (None, 0) + reset_tokens, reset_logprobs = _assert_complete_output(reset) + assert reset_tokens == cold_tokens + torch.testing.assert_close(reset_logprobs, cold_logprobs, rtol=5e-3, atol=5e-3) diff --git a/tests/ut/models/test_kimi_k3_adapter.py b/tests/ut/models/test_kimi_k3_adapter.py new file mode 100644 index 000000000000..2e86d78211b3 --- /dev/null +++ b/tests/ut/models/test_kimi_k3_adapter.py @@ -0,0 +1,564 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest +import torch +from torch import nn +from vllm.config import VllmConfig, set_current_vllm_config + +from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec +from vllm_ascend.models import kimi_k3 +from vllm_ascend.models.kimi_k3 import ( + AscendKimiK3ForConditionalGeneration, + AscendKimiK3MultiModalProjector, + AscendKimiLinearForCausalLM, + AscendKimiLinearModel, + AscendKimiMLAAttention, + AscendKimiMLP, + AscendKimiMoE, +) +from vllm_ascend.models.kimi_k3_dspark import ( + AscendK3DSparkDecoderLayer, + AscendK3DSparkForCausalLM, + AscendK3DSparkModel, +) + + +def test_ascend_attn_res_matches_canonical_k3_math(monkeypatch): + prefix_sum = torch.tensor([[1.0, 2.0], [3.0, 4.0]]) + block_residual = torch.tensor( + [ + [[0.5, 1.5], [2.5, 3.5], [1000.0, 1000.0]], + [[1.0, 0.0], [0.0, 1.0], [1000.0, 1000.0]], + ] + ) + norm = SimpleNamespace(weight=torch.tensor([1.0, 1.5]), variance_epsilon=1e-5) + proj = SimpleNamespace(weight=torch.tensor([[0.25, -0.5]])) + + monkeypatch.setattr( + kimi_k3, + "_EXTRA_CTX", + SimpleNamespace(flash_comm_v1_enabled=False), + ) + + output = kimi_k3._apply_ascend_attn_res( + prefix_sum, + block_residual, + proj, + norm, + num_valid_blocks=2, + ) + + values = torch.cat( + (block_residual[:, :2], prefix_sum.unsqueeze(1)), + dim=1, + ).float() + inverse_rms = torch.rsqrt(values.square().mean(-1, keepdim=True) + norm.variance_epsilon) + normalized_without_gamma = values * inverse_rms + score_weight = norm.weight.float() * proj.weight.squeeze(0).float() + probabilities = (normalized_without_gamma * score_weight).sum(-1).softmax(-1).unsqueeze(1) + expected = torch.matmul(probabilities, values).squeeze(1).to(prefix_sum.dtype) + torch.testing.assert_close(output, expected) + + +def test_ascend_attn_res_avoids_broadcast_score_product(monkeypatch): + prefix_sum = torch.ones(2, 4) + block_residual = torch.ones(2, 3, 4) + norm = SimpleNamespace(weight=torch.ones(4), variance_epsilon=1e-5) + proj = SimpleNamespace(weight=torch.ones(1, 4)) + original_matmul = torch.matmul + score_matmul_shapes = [] + + def record_matmul(left, right, *args, **kwargs): + if left.shape == (2, 3, 4) and right.shape == (4,): + score_matmul_shapes.append((left.shape, right.shape)) + return original_matmul(left, right, *args, **kwargs) + + monkeypatch.setattr( + kimi_k3, + "_EXTRA_CTX", + SimpleNamespace(flash_comm_v1_enabled=False), + ) + monkeypatch.setattr(torch, "matmul", record_matmul) + + kimi_k3._apply_ascend_attn_res( + prefix_sum, + block_residual, + proj, + norm, + num_valid_blocks=2, + ) + + assert score_matmul_shapes == [((2, 3, 4), (4,))] + + +def test_ascend_kimi_moe_delegates_padding_to_routed_experts(monkeypatch): + config = SimpleNamespace(min_moe_intermediate_per_partition=256) + delegated = {} + + def fake_init(self, *, config, **kwargs): + nn.Module.__init__(self) + self.use_latent_moe = False + delegated["config"] = config + delegated["kwargs"] = kwargs + + monkeypatch.setattr(kimi_k3.UpstreamKimiMoE, "__init__", fake_init) + + AscendKimiMoE( + config=config, + prefix="model.layers.1.block_sparse_moe", + layer_idx=1, + ) + + assert delegated["config"] is not config + assert delegated["config"].min_moe_intermediate_per_partition == 0 + assert config.min_moe_intermediate_per_partition == 256 + assert delegated["kwargs"] == { + "quant_config": None, + "prefix": "model.layers.1.block_sparse_moe", + "layer_idx": 1, + } + + +def test_ascend_kimi_mlp_forwards_explicit_situ_parameters( + monkeypatch, +): + delegated: dict[str, object] = {} + + def fake_init( + self, + *args, + hidden_act, + activation_situ_beta, + activation_situ_linear_beta, + **kwargs, + ): + nn.Module.__init__(self) + delegated.update( + hidden_act=hidden_act, + activation_situ_beta=activation_situ_beta, + activation_situ_linear_beta=activation_situ_linear_beta, + ) + + monkeypatch.setattr(kimi_k3.KimiMLP, "__init__", fake_init) + + with set_current_vllm_config(VllmConfig()): + mlp = AscendKimiMLP( + hidden_act="situ", + activation_situ_beta=4.0, + activation_situ_linear_beta=25.0, + ) + + assert delegated == { + "hidden_act": "situ", + "activation_situ_beta": 4.0, + "activation_situ_linear_beta": 25.0, + } + assert mlp.act_fn.beta == 4.0 + assert mlp.act_fn.linear_beta == 25.0 + + +def test_dspark_decoder_uses_upstream_mlp_activation_contract( + monkeypatch, +): + config = SimpleNamespace( + hidden_size=8, + num_attention_heads=2, + qk_nope_head_dim=2, + qk_rope_head_dim=2, + v_head_dim=2, + q_lora_rank=4, + kv_lora_rank=4, + intermediate_size=16, + hidden_act="silu", + rms_norm_eps=1e-6, + ) + vllm_config = SimpleNamespace(cache_config=None) + mlp_factory = MagicMock(return_value=nn.Identity()) + monkeypatch.setattr( + "vllm_ascend.models.kimi_k3_dspark.get_draft_quant_config", + lambda _: None, + ) + monkeypatch.setattr( + "vllm_ascend.models.kimi_k3_dspark.AscendKimiMLAAttention", + lambda **_: nn.Identity(), + ) + monkeypatch.setattr( + "vllm_ascend.models.kimi_k3_dspark.AscendKimiMLP", + mlp_factory, + ) + + with set_current_vllm_config(VllmConfig()): + AscendK3DSparkDecoderLayer( + vllm_config=vllm_config, + config=config, + layer_idx=0, + start_layer_id=4, + prefix="model", + ) + + assert mlp_factory.call_args.kwargs["hidden_act"] == "silu" + assert "activation_situ_beta" not in mlp_factory.call_args.kwargs + assert "activation_situ_linear_beta" not in mlp_factory.call_args.kwargs + + +def test_ascend_kimi_moe_quantizes_modelslim_latent_projections(monkeypatch): + class FakeLinear(nn.Module): + def __init__(self, input_size, output_size, **kwargs): + super().__init__() + self.input_size = input_size + self.output_size = output_size + self.kwargs = kwargs + + class FakeRunner(nn.Module): + def __init__(self): + super().__init__() + self.routed_input_transform = nn.Identity() + self.routed_output_transform = nn.Identity() + + config = SimpleNamespace( + hidden_size=16, + min_moe_intermediate_per_partition=256, + ) + quant_config = MagicMock() + quant_config.get_name.return_value = "ascend" + norm = nn.Identity() + + def fake_init(self, *, config, **kwargs): + nn.Module.__init__(self) + self.use_latent_moe = True + self.moe_hidden_size = 8 + self.routed_expert_norm = norm + self.routed_expert_down_proj = nn.Identity() + self.routed_expert_up_proj = nn.Identity() + self.routed_output_transform = nn.Identity() + self.experts = FakeRunner() + + monkeypatch.setattr(kimi_k3.UpstreamKimiMoE, "__init__", fake_init) + monkeypatch.setattr(kimi_k3, "ReplicatedLinear", FakeLinear) + + moe = AscendKimiMoE( + config=config, + quant_config=quant_config, + prefix="model.layers.1.block_sparse_moe", + layer_idx=1, + ) + + assert moe.routed_expert_down_proj.input_size == 16 + assert moe.routed_expert_down_proj.output_size == 8 + assert moe.routed_expert_down_proj.kwargs == { + "bias": False, + "quant_config": quant_config, + "prefix": "model.layers.1.block_sparse_moe.routed_expert_down_proj", + } + assert moe.routed_expert_up_proj.input_size == 8 + assert moe.routed_expert_up_proj.output_size == 16 + assert moe.routed_expert_up_proj.kwargs == { + "bias": False, + "quant_config": quant_config, + "prefix": "model.layers.1.block_sparse_moe.routed_expert_up_proj", + } + assert moe.experts.routed_input_transform is moe.routed_expert_down_proj + assert moe.experts.routed_output_transform is moe.routed_output_transform + + +def test_kimi_text_model_retains_upstream_checkpoint_packing(): + assert AscendKimiLinearForCausalLM.packed_modules_mapping == { + "gate_up_proj": ["gate_proj", "up_proj"], + "in_proj_qkvgfab": [ + "q_proj", + "k_proj", + "v_proj", + "b_proj", + "f_a_proj", + ], + "conv1d": ["q_conv1d", "k_conv1d", "v_conv1d"], + "fused_qkv_a_proj": ["q_a_proj", "kv_a_proj_with_mqa"], + } + + +def test_kimi_mixed_kda_gate_weights_load_into_float_packed_projection(monkeypatch): + model = AscendKimiLinearModel.__new__(AscendKimiLinearModel) + nn.Module.__init__(model) + layer = nn.Module() + layer.self_attn = nn.Module() + layer.self_attn.in_proj_gfab = nn.Module() + layer.self_attn.in_proj_gfab.load_shard_weight = MagicMock() + packed_weight = nn.Parameter(torch.empty(6, 4)) + layer.self_attn.in_proj_gfab.register_parameter("weight", packed_weight) + layer.router = nn.Linear(4, 1, bias=False) + model.layers = nn.ModuleList([layer]) + + remaining = [] + + def fake_upstream_load_weights(_self, weights): + remaining.extend(weights) + return {name for name, *_ in remaining} + + monkeypatch.setattr( + kimi_k3.UpstreamKimiLinearModel, + "load_weights", + fake_upstream_load_weights, + ) + source_weights = [ + ("layers.0.router.weight", torch.full((1, 4), 0.5)), + ("layers.0.self_attn.g_proj.weight", torch.full((1,), 1.0)), + ("layers.0.self_attn.f_a_proj.weight", torch.full((1,), 2.0)), + ("layers.0.self_attn.b_proj.weight", torch.full((1,), 3.0)), + ("layers.0.self_attn.o_proj.weight", torch.full((1,), 4.0)), + ] + + loaded = model.load_weights(iter(source_weights)) + + weight_loader = layer.self_attn.in_proj_gfab.load_shard_weight + assert [call.args[2] for call in weight_loader.call_args_list] == [0, 1, 2] + assert [call.args[1].item() for call in weight_loader.call_args_list] == [1.0, 2.0, 3.0] + assert remaining == [source_weights[0], source_weights[-1]] + assert loaded == { + "layers.0.self_attn.in_proj_gfab.weight", + "layers.0.router.weight", + "layers.0.self_attn.o_proj.weight", + } + + +def test_kimi_text_model_layer_factory_accepts_prefix_keyword(monkeypatch): + config = SimpleNamespace( + vocab_size=64, + hidden_size=16, + num_hidden_layers=1, + rms_norm_eps=1e-5, + attn_res_block_size=None, + num_attention_heads=1, + ) + vllm_config = MagicMock() + vllm_config.model_config.hf_text_config = config + pp_group = SimpleNamespace(is_first_rank=False, is_last_rank=False) + decoder_layer = nn.Identity() + decoder_layer_factory = MagicMock(return_value=decoder_layer) + + def fake_make_layers(num_hidden_layers, layer_fn, *, prefix): + assert num_hidden_layers == 1 + assert layer_fn(prefix=f"{prefix}.0") is decoder_layer + return 0, 1, nn.ModuleList([decoder_layer]) + + monkeypatch.setattr(kimi_k3, "get_pp_group", lambda: pp_group) + monkeypatch.setattr(kimi_k3, "get_tensor_model_parallel_world_size", lambda: 1) + monkeypatch.setattr(kimi_k3, "AscendKimiDecoderLayer", decoder_layer_factory) + monkeypatch.setattr(kimi_k3, "make_layers", fake_make_layers) + + model = AscendKimiLinearModel(vllm_config=vllm_config, prefix="model") + + assert model.start_layer == 0 + assert model.end_layer == 1 + decoder_layer_factory.assert_called_once_with( + config, + vllm_config, + "model.layers.0", + ) + + +def test_kimi_mla_cache_spec_preserves_hybrid_page_padding(): + real_page_size = 128 * 576 * torch.bfloat16.itemsize + padded_page_size = real_page_size + 128 + spec = AscendMLAAttentionSpec( + block_size=128, + num_kv_heads=1, + head_size=576, + dtype=torch.bfloat16, + page_size_padded=padded_page_size, + ) + + assert spec.real_page_size_bytes == real_page_size + assert spec.page_size_bytes == padded_page_size + assert AscendMLAAttentionSpec.merge([spec, spec]).page_size_bytes == padded_page_size + + +def test_ascend_mla_exposes_layer_and_cache_contract(): + attention = AscendKimiMLAAttention.__new__(AscendKimiMLAAttention) + layer = MagicMock() + layer.layer_name = "model.layers.1.self_attn.attn" + layer.impl = object() + layer.kv_cache = (object(), object()) + layer.kv_cache_dtype = "auto" + layer._k_scale = 1.0 + attention.mla_attn = MagicMock() + attention.mla_attn.mla_attn = layer + attention.mla_attn.is_vl_first_layer = True + + assert attention.layer_name == layer.layer_name + assert attention.impl is layer.impl + assert attention.kv_cache is layer.kv_cache + assert attention.kv_cache_dtype == layer.kv_cache_dtype + assert attention._k_scale == layer._k_scale + assert attention.is_vl_first_layer is True + + +def test_projector_applies_optional_modelslim_rotation(): + class ScaleLinear(nn.Module): + def forward(self, hidden_states): + return hidden_states * 2, None + + projector = AscendKimiK3MultiModalProjector.__new__(AscendKimiK3MultiModalProjector) + nn.Module.__init__(projector) + image_features = torch.tensor([[1.0, 2.0]]) + + with patch.object( + kimi_k3.KimiK25MultiModalProjector, + "forward", + lambda self, hidden_states: hidden_states, + ): + projector.rot_proj = ScaleLinear() + torch.testing.assert_close( + projector(image_features), + image_features * 2, + ) + projector.rot_proj = None + torch.testing.assert_close(projector(image_features), image_features) + + +def test_projector_rotation_is_removed_when_checkpoint_omits_it(monkeypatch): + wrapper = AscendKimiK3ForConditionalGeneration.__new__(AscendKimiK3ForConditionalGeneration) + nn.Module.__init__(wrapper) + wrapper.mm_projector = nn.Module() + wrapper.mm_projector.rot_proj = nn.Linear(1, 1, bias=False) + + loader = MagicMock() + loader.load_weights.return_value = {"mm_projector.linear_1.weight"} + monkeypatch.setattr(kimi_k3, "AutoWeightsLoader", lambda model: loader) + + loaded = wrapper.load_weights(iter(())) + + assert loaded == {"mm_projector.linear_1.weight"} + assert wrapper.mm_projector.rot_proj is None + + +def test_projector_rotation_is_kept_when_checkpoint_provides_it(monkeypatch): + wrapper = AscendKimiK3ForConditionalGeneration.__new__(AscendKimiK3ForConditionalGeneration) + nn.Module.__init__(wrapper) + wrapper.mm_projector = nn.Module() + wrapper.mm_projector.rot_proj = nn.Linear(1, 1, bias=False) + + loader = MagicMock() + loader.load_weights.return_value = {"mm_projector.rot_proj.weight"} + monkeypatch.setattr(kimi_k3, "AutoWeightsLoader", lambda model: loader) + + wrapper.load_weights(iter(())) + + assert wrapper.mm_projector.rot_proj is not None + + +class _DraftTokenEmbedder(nn.Module): + def __init__(self) -> None: + super().__init__() + self.embedding = nn.Embedding.from_pretrained( + torch.tensor( + [ + [0.0, 0.0], + [1.0, 2.0], + [3.0, 4.0], + ] + ) + ) + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.embedding(input_ids) + + +def _make_k3_dspark_for_embedding_test() -> AscendK3DSparkForCausalLM: + model = AscendK3DSparkForCausalLM.__new__(AscendK3DSparkForCausalLM) + nn.Module.__init__(model) + model.model = _DraftTokenEmbedder() + return model + + +def test_k3_dspark_load_weights_keeps_per_layer_context_kv(monkeypatch): + model = AscendK3DSparkForCausalLM.__new__(AscendK3DSparkForCausalLM) + nn.Module.__init__(model) + source_weights = [ + ( + "layers.0.self_attn.kv_a_proj_with_mqa.weight", + torch.ones(1, 1), + ) + ] + seen_names: list[str] = [] + + class CapturingLoader: + def __init__(self, loaded_model, *, skip_substrs): + assert loaded_model is model + assert skip_substrs == list(model.checkpoint_skip_substrs) + + def load_weights(self, weights, *, mapper): + assert mapper is model.hf_to_vllm_mapper + seen_names.extend(name for name, _ in weights) + return {"model.layers.0.self_attn.fused_qkv_a_proj.weight"} + + monkeypatch.setattr( + "vllm_ascend.models.kimi_k3_dspark.AutoWeightsLoader", + CapturingLoader, + ) + + loaded = model.load_weights(iter(source_weights)) + + assert seen_names == [source_weights[0][0]] + assert loaded == {"model.layers.0.self_attn.fused_qkv_a_proj.weight"} + + +def test_k3_dspark_embed_input_ids_keeps_text_only_path(): + model = _make_k3_dspark_for_embedding_test() + + output = model.embed_input_ids(torch.tensor([1, 2])) + + torch.testing.assert_close( + output, + torch.tensor([[1.0, 2.0], [3.0, 4.0]]), + ) + + +def test_k3_dspark_embed_input_ids_merges_multimodal_embeddings(): + model = _make_k3_dspark_for_embedding_test() + input_ids = torch.tensor([1, 999, 2]) + is_multimodal = torch.tensor([False, True, False]) + image_embedding = torch.tensor([[9.0, 10.0]]) + + output = model.embed_input_ids( + input_ids, + multimodal_embeddings=(image_embedding,), + is_multimodal=is_multimodal, + ) + + torch.testing.assert_close( + output, + torch.tensor( + [ + [1.0, 2.0], + [9.0, 10.0], + [3.0, 4.0], + ] + ), + ) + + +def test_k3_dspark_embed_input_ids_requires_multimodal_mask(): + model = _make_k3_dspark_for_embedding_test() + + with pytest.raises(ValueError, match="is_multimodal"): + model.embed_input_ids( + torch.tensor([1]), + multimodal_embeddings=(torch.tensor([[9.0, 10.0]]),), + ) + + +def test_k3_dspark_rejects_incomplete_context_slot_mappings(): + model = AscendK3DSparkModel.__new__(AscendK3DSparkModel) + nn.Module.__init__(model) + model.layers = nn.ModuleList([nn.Identity(), nn.Identity()]) + + with pytest.raises(ValueError, match="one entry per draft layer"): + model.precompute_and_store_context_kv( + torch.ones(1, 1), + torch.zeros(1, dtype=torch.int64), + [torch.zeros(1, dtype=torch.int32)], + ) diff --git a/vllm_ascend/models/__init__.py b/vllm_ascend/models/__init__.py index 54d1d393c146..2f994de14c41 100644 --- a/vllm_ascend/models/__init__.py +++ b/vllm_ascend/models/__init__.py @@ -2,6 +2,28 @@ def register_model(): + ModelRegistry.register_model( + "KimiLinearForCausalLM", + "vllm_ascend.models.kimi_k3:AscendKimiLinearForCausalLM", + ) + # Keep the release-branch text architecture as a compatibility alias for + # checkpoints whose config predates vLLM's KimiLinear rename. + ModelRegistry.register_model( + "KimiK3ForCausalLM", + "vllm_ascend.models.kimi_k3:AscendKimiLinearForCausalLM", + ) + ModelRegistry.register_model( + "KimiK3ForConditionalGeneration", + "vllm_ascend.models.kimi_k3:AscendKimiK3ForConditionalGeneration", + ) + ModelRegistry.register_model( + "KimiK3MTPModel", + "vllm_ascend.models.kimi_k3_mtp:AscendKimiK3MTP", + ) + ModelRegistry.register_model( + "K3DSparkModel", + "vllm_ascend.models.kimi_k3_dspark:AscendK3DSparkForCausalLM", + ) ModelRegistry.register_model( "DeepseekV4ForCausalLM", "vllm_ascend.models.deepseek_v4.model:AscendDeepseekV4ForCausalLM" ) diff --git a/vllm_ascend/models/kimi_k3.py b/vllm_ascend/models/kimi_k3.py new file mode 100644 index 000000000000..c5bb95c803d9 --- /dev/null +++ b/vllm_ascend/models/kimi_k3.py @@ -0,0 +1,891 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Kimi K3 model adapters for vLLM 0.27 on Ascend. + +vLLM owns Kimi's configuration, multimodal processor, weight mappings, and +model-level forward contract. This module composes those upstream pieces with +the generic MLA/MoE implementation and the Ascend KDA backend. +""" + +import math +from collections.abc import Iterable +from copy import copy + +import torch +from torch import nn +from vllm.config import CacheConfig, VllmConfig +from vllm.distributed import ( + get_pp_group, + get_tensor_model_parallel_world_size, +) +from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.model_executor.layers.linear import ( + ColumnParallelLinear, + MergedColumnParallelLinear, + ReplicatedLinear, + RowParallelLinear, +) +from vllm.model_executor.layers.logits_processor import LogitsProcessor +from vllm.model_executor.layers.mla import ( + MLAModules, + MultiHeadLatentAttentionWrapper, +) +from vllm.model_executor.layers.quantization import QuantizationConfig +from vllm.model_executor.layers.quantization.compressed_tensors import ( + compressed_tensors, +) +from vllm.model_executor.layers.rotary_embedding import get_rope +from vllm.model_executor.layers.vocab_parallel_embedding import ( + ParallelLMHead, + VocabParallelEmbedding, +) +from vllm.model_executor.models.kimi_k25_vit import ( + KimiK25MultiModalProjector, + MoonViT3dPretrainedModel, +) +from vllm.model_executor.models.utils import ( + AutoWeightsLoader, + PPMissingLayer, + init_vllm_registered_model, + make_layers, + maybe_prefix, +) +from vllm.model_executor.models.vision import is_vit_use_data_parallel +from vllm.models.kimi_k3.amd.linear import ( + KimiDecoderLayer as UpstreamKimiDecoderLayer, +) +from vllm.models.kimi_k3.amd.linear import KimiLinearForCausalLM as UpstreamKimiLinearForCausalLM +from vllm.models.kimi_k3.amd.linear import KimiLinearModel as UpstreamKimiLinearModel +from vllm.models.kimi_k3.amd.linear import ( + KimiMLP, + KimiRoutedOutputTransform, +) +from vllm.models.kimi_k3.amd.linear import ( + KimiMoE as UpstreamKimiMoE, +) +from vllm.models.kimi_k3.amd.model import ( + KimiK3ForConditionalGeneration as UpstreamKimiK3ForConditionalGeneration, +) +from vllm.models.kimi_k3.common.mm_preprocess import ( + KimiK3DummyInputsBuilder, + KimiK3MultiModalProcessor, + KimiK3ProcessingInfo, +) +from vllm.multimodal import MULTIMODAL_REGISTRY +from vllm.platforms import current_platform +from vllm.sequence import IntermediateTensors +from vllm.utils.math_utils import cdiv + +from vllm_ascend.ascend_forward_context import _EXTRA_CTX +from vllm_ascend.ops.activation import AscendSituAndMul # type: ignore[attr-defined] +from vllm_ascend.ops.kimi_kda import AscendKimiK3DeltaAttention # type: ignore[import-untyped] + + +def _apply_ascend_attn_res( + prefix_sum: torch.Tensor, + block_residual: torch.Tensor, + proj: ReplicatedLinear, + norm: RMSNorm, + num_valid_blocks: int, +) -> torch.Tensor: + """Apply Kimi's canonical learned residual mixture with native ops.""" + if num_valid_blocks <= 0: + return prefix_sum + + values = torch.cat( + ( + block_residual[:, :num_valid_blocks, :], + prefix_sum.unsqueeze(1), + ), + dim=1, + ) + values_fp32 = values.float() + inverse_rms = torch.rsqrt(values_fp32.square().mean(-1, keepdim=True) + norm.variance_epsilon) + normalized_without_gamma = values_fp32 * inverse_rms + score_weight = norm.weight.float() * proj.weight.squeeze(0).float() + # Avoid materializing a broadcasted FP32 tensor as large as the entire + # normalized residual stack. + scores = torch.matmul(normalized_without_gamma, score_weight) + probabilities = scores.softmax(-1).unsqueeze(1) + mixed = torch.matmul(probabilities, values_fp32).squeeze(1).to(values.dtype) + if _EXTRA_CTX.flash_comm_v1_enabled: + mixed = torch.ops.vllm.maybe_chunk_residual(prefix_sum, mixed) + return mixed + + +class AscendKimiMLP(KimiMLP): + """Use the Ascend SiTU module for dense and shared-expert MLPs.""" + + def __init__( + self, + *args, + hidden_act: str, + activation_situ_beta: float | None = None, + activation_situ_linear_beta: float | None = None, + **kwargs, + ) -> None: + super().__init__( + *args, + hidden_act=hidden_act, + activation_situ_beta=activation_situ_beta, + activation_situ_linear_beta=activation_situ_linear_beta, + **kwargs, + ) + if hidden_act == "situ": + self.act_fn = AscendSituAndMul( + beta=activation_situ_beta or 1.0, + linear_beta=activation_situ_linear_beta, + ) + + +class AscendKimiMoE(UpstreamKimiMoE): + """Adapt Kimi K3 latent MoE construction to Ascend. + + The upstream AMD implementation pads small expert partitions at the model + layer before the backend MoE factory resolves TP versus EP. Under EP this + applies a TP-sized pad even though each rank owns complete experts. Ascend + routed experts already perform any backend-required size rounding, so keep + Kimi's checkpoint intermediate size here and leave padding to that layer. + + Native ModelSlim checkpoints also quantize the latent down/up projections, + while the upstream implementation always creates them unquantized. Rebind + those projections and the runner transforms to their Ascend quantized + modules after the common Kimi structure has been assembled. + """ + + def __init__( + self, + *, + config, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + **kwargs, + ) -> None: + ascend_config = copy(config) + ascend_config.min_moe_intermediate_per_partition = 0 + super().__init__( + config=ascend_config, + quant_config=quant_config, + prefix=prefix, + **kwargs, + ) + + latent_quant_config = quant_config if quant_config is not None and quant_config.get_name() == "ascend" else None + if not self.use_latent_moe or latent_quant_config is None: + return + + self.routed_expert_down_proj = ReplicatedLinear( + config.hidden_size, + self.moe_hidden_size, + bias=False, + quant_config=latent_quant_config, + prefix=f"{prefix}.routed_expert_down_proj", + ) + self.routed_expert_up_proj = ReplicatedLinear( + self.moe_hidden_size, + config.hidden_size, + bias=False, + quant_config=latent_quant_config, + prefix=f"{prefix}.routed_expert_up_proj", + ) + self.routed_output_transform = KimiRoutedOutputTransform( + self.routed_expert_norm, + self.routed_expert_up_proj, + ) + self.experts.routed_input_transform = self.routed_expert_down_proj + self.experts.routed_output_transform = self.routed_output_transform + + +class AscendKimiMLAAttention(nn.Module): + """Generic vLLM MLA composition backed by Ascend's pluggable wrapper.""" + + def __init__( + self, + config, + hidden_size: int, + num_heads: int, + qk_nope_head_dim: int, + qk_rope_head_dim: int, + v_head_dim: int, + q_lora_rank: int | None, + kv_lora_rank: int, + use_output_gate: bool, + use_rope: bool, + cache_config: CacheConfig | None = None, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + non_causal_multi_token_decode: bool = False, + ) -> None: + """Assemble Kimi's projections, optional RoPE, and generic MLA wrapper.""" + super().__init__() + self.hidden_size = hidden_size + self.qk_nope_head_dim = qk_nope_head_dim + self.qk_rope_head_dim = qk_rope_head_dim + self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim + self.v_head_dim = v_head_dim + self.q_lora_rank = q_lora_rank + self.kv_lora_rank = kv_lora_rank + self.num_heads = num_heads + tp_size = get_tensor_model_parallel_world_size() + assert num_heads % tp_size == 0 + self.num_local_heads = num_heads // tp_size + self.scaling = self.qk_head_dim**-0.5 + + self.fused_qkv_a_proj = None + self.kv_a_proj_with_mqa = None + self.q_a_layernorm = None + self.q_b_proj = None + self.q_proj = None + if q_lora_rank is not None: + self.fused_qkv_a_proj = MergedColumnParallelLinear( + hidden_size, + [q_lora_rank, kv_lora_rank + qk_rope_head_dim], + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.fused_qkv_a_proj", + disable_tp=True, + ) + self.q_a_layernorm = RMSNorm( + q_lora_rank, + eps=config.rms_norm_eps, + ) + self.q_b_proj = ColumnParallelLinear( + q_lora_rank, + num_heads * self.qk_head_dim, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.q_b_proj", + ) + else: + self.q_proj = ColumnParallelLinear( + hidden_size, + num_heads * self.qk_head_dim, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.q_proj", + ) + self.kv_a_proj_with_mqa = ReplicatedLinear( + hidden_size, + kv_lora_rank + qk_rope_head_dim, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.kv_a_proj_with_mqa", + ) + + self.kv_a_layernorm = RMSNorm( + kv_lora_rank, + eps=config.rms_norm_eps, + ) + self.kv_b_proj = ColumnParallelLinear( + kv_lora_rank, + num_heads * (qk_nope_head_dim + v_head_dim), + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.kv_b_proj", + ) + self.g_proj = ( + ColumnParallelLinear( + hidden_size, + num_heads * v_head_dim, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.g_proj", + ) + if use_output_gate + else None + ) + self.o_proj = RowParallelLinear( + num_heads * v_head_dim, + hidden_size, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.o_proj", + ) + + self.rotary_emb = None + if use_rope: + rope_parameters = dict(config.rope_parameters) + if rope_parameters["rope_type"] != "default": + rope_parameters["rope_type"] = ( + "deepseek_yarn" if rope_parameters.get("apply_yarn_scaling", True) else "deepseek_llama_scaling" + ) + self.rotary_emb = get_rope( + qk_rope_head_dim, + max_position=config.max_position_embeddings, + rope_parameters=rope_parameters, + is_neox_style=False, + ) + if rope_parameters["rope_type"] == "deepseek_yarn": + scaling_factor = float(rope_parameters["factor"]) + mscale_all_dim = float(rope_parameters.get("mscale_all_dim", 0.0)) + if scaling_factor > 1 and mscale_all_dim: + mscale = 0.1 * mscale_all_dim * math.log(scaling_factor) + 1.0 + self.scaling *= mscale * mscale + + mla_modules = MLAModules( + kv_a_layernorm=self.kv_a_layernorm, + kv_b_proj=self.kv_b_proj, + rotary_emb=self.rotary_emb, + o_proj=self.o_proj, + fused_qkv_a_proj=self.fused_qkv_a_proj, + kv_a_proj_with_mqa=self.kv_a_proj_with_mqa, + q_a_layernorm=self.q_a_layernorm, + q_b_proj=self.q_b_proj, + q_proj=self.q_proj, + indexer=None, + is_sparse=False, + topk_indices_buffer=None, + g_proj=self.g_proj, + ) + self.mla_attn = MultiHeadLatentAttentionWrapper( + hidden_size, + self.num_local_heads, + self.scaling, + qk_nope_head_dim, + qk_rope_head_dim, + v_head_dim, + q_lora_rank, + kv_lora_rank, + mla_modules, + cache_config, + quant_config, + prefix, + non_causal_multi_token_decode=non_causal_multi_token_decode, + ) + + @property + def _attention_layer(self): + return self.mla_attn.mla_attn + + @property + def is_vl_first_layer(self) -> bool: + return self.mla_attn.is_vl_first_layer + + @property + def layer_name(self) -> str: + return self._attention_layer.layer_name + + @property + def impl(self): + return self._attention_layer.impl + + @property + def kv_cache(self): + return self._attention_layer.kv_cache + + @property + def kv_cache_dtype(self): + return self._attention_layer.kv_cache_dtype + + @property + def _k_scale(self): + return self._attention_layer._k_scale + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + ) -> torch.Tensor: + return self.mla_attn(positions, hidden_states) + + +class AscendKimiDecoderLayer(UpstreamKimiDecoderLayer): + """Upstream Kimi decoder structure with Ascend attention backends.""" + + def __init__( + self, + config, + vllm_config: VllmConfig, + prefix: str = "", + ) -> None: + """Select KDA or no-RoPE MLA and configure the layer residual path.""" + nn.Module.__init__(self) + self.hidden_size = config.hidden_size + self.layer_idx = int(prefix.rsplit(".", 1)[1]) + self.is_moe = config.is_moe + layer_idx = self.layer_idx + cache_config = vllm_config.cache_config + quant_config = vllm_config.quant_config + + if config.is_kda_layer(layer_idx): + self.self_attn = AscendKimiK3DeltaAttention( + config, + vllm_config, + prefix=f"{prefix}.self_attn", + ) + self._self_attn_writes_output = False + else: + qk_nope_head_dim = config.qk_nope_head_dim + qk_rope_head_dim = config.qk_rope_head_dim + v_head_dim = config.v_head_dim + kv_lora_rank = config.kv_lora_rank + assert qk_nope_head_dim is not None + assert qk_rope_head_dim is not None + assert v_head_dim is not None + assert kv_lora_rank is not None + assert config.mla_use_nope is True + self.self_attn = AscendKimiMLAAttention( + config=config, + hidden_size=self.hidden_size, + num_heads=config.num_attention_heads, + qk_nope_head_dim=qk_nope_head_dim, + qk_rope_head_dim=qk_rope_head_dim, + v_head_dim=v_head_dim, + q_lora_rank=config.q_lora_rank, + kv_lora_rank=kv_lora_rank, + use_output_gate=bool(config.mla_use_output_gate), + use_rope=False, + cache_config=cache_config, + quant_config=quant_config, + prefix=f"{prefix}.self_attn", + ) + self._self_attn_writes_output = False + + self.is_moe_layer = ( + self.is_moe + and config.num_experts is not None + and layer_idx >= config.first_k_dense_replace + and layer_idx % config.moe_layer_freq == 0 + ) + if self.is_moe_layer: + self.block_sparse_moe = AscendKimiMoE( + config=config, + quant_config=quant_config, + prefix=f"{prefix}.block_sparse_moe", + layer_idx=layer_idx, + ) + self.mlp = self.block_sparse_moe + else: + self.mlp = AscendKimiMLP( + hidden_size=self.hidden_size, + intermediate_size=config.intermediate_size, + hidden_act=config.hidden_act, + quant_config=quant_config, + prefix=f"{prefix}.mlp", + activation_situ_beta=config.activation_situ_beta, + activation_situ_linear_beta=config.activation_situ_linear_beta, + ) + self.input_layernorm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + self.post_attention_layernorm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + + attn_res_block_size = config.attn_res_block_size + self.use_attn_residuals = attn_res_block_size is not None + if attn_res_block_size is not None: + self.attn_res_block_size = attn_res_block_size + self.is_block_write_layer = layer_idx % attn_res_block_size == 0 + self.block_write_idx = layer_idx // attn_res_block_size + self.prev_valid_blocks = cdiv(layer_idx, attn_res_block_size) + self.self_attention_res_norm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + self.mlp_res_norm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + self.self_attention_res_proj = ReplicatedLinear( + config.hidden_size, + 1, + bias=False, + quant_config=None, + prefix=f"{prefix}.self_attention_res_proj", + ) + self.mlp_res_proj = ReplicatedLinear( + config.hidden_size, + 1, + bias=False, + quant_config=None, + prefix=f"{prefix}.mlp_res_proj", + ) + + self.is_vl_first_layer = self.self_attn.is_vl_first_layer + + def _run_self_attn( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + ) -> torch.Tensor: + if not self._self_attn_writes_output: + return self.self_attn( + hidden_states=hidden_states, + positions=positions, + ) + output = torch.empty_like(hidden_states) + self.self_attn( + hidden_states=hidden_states, + positions=positions, + output=output, + ) + return output + + def forward_attn_residual( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + block_residual: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Run upstream attn-res with the multimodal FlashComm transition.""" + prefix_sum: torch.Tensor | None = hidden_states + hidden_states = _apply_ascend_attn_res( + prefix_sum, + block_residual, + self.self_attention_res_proj, + self.self_attention_res_norm, + self.prev_valid_blocks, + ) + if self.is_block_write_layer: + assert prefix_sum is not None + block_residual[:, self.block_write_idx, :].copy_(prefix_sum) + prefix_sum = None + + hidden_states = self.input_layernorm(hidden_states) + hidden_states = self._run_self_attn(positions, hidden_states) + + if self.is_vl_first_layer and _EXTRA_CTX.flash_comm_v1_enabled: + block_residual = torch.ops.vllm.maybe_chunk_residual( + hidden_states.unsqueeze(1), + block_residual, + ) + + prefix_sum = hidden_states if prefix_sum is None else prefix_sum + hidden_states + mlp_valid_blocks = self.prev_valid_blocks + (1 if self.is_block_write_layer else 0) + hidden_states = _apply_ascend_attn_res( + prefix_sum, + block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + mlp_valid_blocks, + ) + hidden_states = self.post_attention_layernorm(hidden_states) + hidden_states = self.mlp(hidden_states) + hidden_states = prefix_sum + hidden_states + return hidden_states, block_residual + + +class AscendKimiLinearModel(UpstreamKimiLinearModel): + """Kimi text model assembled from the Ascend decoder layer.""" + + packed_modules_mapping = { + "gate_up_proj": ["gate_proj", "up_proj"], + "in_proj_qkvgfab": [ + "q_proj", + "k_proj", + "v_proj", + "b_proj", + "f_a_proj", + ], + "conv1d": ["q_conv1d", "k_conv1d", "v_conv1d"], + "fused_qkv_a_proj": ["q_a_proj", "kv_a_proj_with_mqa"], + } + + def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: + nn.Module.__init__(self) + config = vllm_config.model_config.hf_text_config + self.config = config + self.vocab_size = config.vocab_size + + if get_pp_group().is_first_rank: + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + prefix=f"{prefix}.embed_tokens", + ) + else: + self.embed_tokens = PPMissingLayer() + + def get_layer(prefix: str): + return AscendKimiDecoderLayer( + config, + vllm_config, + prefix, + ) + + self.start_layer, self.end_layer, self.layers = make_layers( + config.num_hidden_layers, + get_layer, + prefix=f"{prefix}.layers", + ) + + if get_pp_group().is_last_rank: + self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + if config.attn_res_block_size is not None: + self.output_attn_res_norm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + self.output_attn_res_proj = ReplicatedLinear( + config.hidden_size, + 1, + bias=False, + quant_config=None, + prefix=f"{prefix}.output_attn_res_proj", + ) + else: + self.norm = PPMissingLayer() + if config.attn_res_block_size is not None: + self.output_attn_res_norm = PPMissingLayer() + self.output_attn_res_proj = PPMissingLayer() + + world_size = get_tensor_model_parallel_world_size() + assert config.num_attention_heads % world_size == 0, "num_attention_heads must be divisible by world_size" + + def load_weights(self, weights): + """Load mixed-precision KDA gates into the FLOAT packed module.""" + params_dict = dict(self.named_parameters()) + gate_mapping = ( + (".g_proj", ".in_proj_gfab", 0), + (".f_a_proj", ".in_proj_gfab", 1), + (".b_proj", ".in_proj_gfab", 2), + ) + loaded_gate_params = set() + + def load_non_gate_weights(): + for args in weights: + name, loaded_weight = args[:2] + for source, target, shard_id in gate_mapping: + if source not in name: + continue + mapped_name = name.replace(source, target) + if mapped_name in params_dict: + param = params_dict[mapped_name] + module_name = mapped_name.rsplit(".", 1)[0] + module = self.get_submodule(module_name) + module.load_shard_weight(param, loaded_weight, shard_id) + loaded_gate_params.add(mapped_name) + break + else: + yield args + + return super().load_weights(load_non_gate_weights()) | loaded_gate_params + + def forward( + self, + input_ids: torch.Tensor | None, + positions: torch.Tensor, + intermediate_tensors: IntermediateTensors | None, + inputs_embeds: torch.Tensor | None = None, + **kwargs, + ) -> torch.Tensor | IntermediateTensors | tuple[torch.Tensor, list[torch.Tensor]]: + if self.config.attn_res_block_size is None: + return super().forward( + input_ids=input_ids, + positions=positions, + intermediate_tensors=intermediate_tensors, + inputs_embeds=inputs_embeds, + **kwargs, + ) + + if get_pp_group().is_first_rank: + hidden_states = inputs_embeds if inputs_embeds is not None else self.embed_input_ids(input_ids) + residual = None + else: + assert intermediate_tensors is not None + hidden_states = intermediate_tensors["hidden_states"] + residual = intermediate_tensors["residual"] + + aux_hidden_states = self._maybe_add_hidden_state( + [], + self.start_layer, + hidden_states, + residual, + ) + attn_res_block_num = cdiv( + self.end_layer, + self.config.attn_res_block_size, + ) + block_residual = hidden_states.new_empty( + hidden_states.size(0), + attn_res_block_num, + hidden_states.size(1), + ) + if residual is not None: + block_residual[:, : residual.size(1), :].copy_(residual) + residual = block_residual + + for layer_idx, layer in enumerate( + self.layers[self.start_layer : self.end_layer], + start=self.start_layer, + ): + hidden_states, residual = layer( + positions=positions, + hidden_states=hidden_states, + residual=residual, + ) + if (layer_idx + 1) in self.aux_hidden_state_layers: + self._maybe_add_hidden_state( + aux_hidden_states, + layer_idx + 1, + hidden_states, + residual, + ) + + if not get_pp_group().is_last_rank: + return IntermediateTensors( + { + "hidden_states": hidden_states, + "residual": residual, + } + ) + + hidden_states = _apply_ascend_attn_res( + hidden_states, + residual, + self.output_attn_res_proj, + self.output_attn_res_norm, + attn_res_block_num, + ) + if aux_hidden_states: + return hidden_states, aux_hidden_states + return hidden_states + + +class AscendKimiLinearForCausalLM(UpstreamKimiLinearForCausalLM): + """Causal-LM wrapper retaining vLLM 0.27 state/cache interfaces.""" + + packed_modules_mapping = AscendKimiLinearModel.packed_modules_mapping + + def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: + nn.Module.__init__(self) + self.model_config = vllm_config.model_config + self.vllm_config = vllm_config + self.config = self.model_config.hf_config + self.quant_config = vllm_config.quant_config + self.model = AscendKimiLinearModel( + vllm_config=vllm_config, + prefix=maybe_prefix(prefix, "model"), + ) + if get_pp_group().is_last_rank: + self.lm_head = ParallelLMHead( + self.config.vocab_size, + self.config.hidden_size, + quant_config=self.quant_config, + prefix=maybe_prefix(prefix, "lm_head"), + ) + else: + self.lm_head = PPMissingLayer() + self.logits_processor = LogitsProcessor( + self.config.vocab_size, + scale=getattr(self.config, "logit_scale", 1.0), + ) + + +class AscendKimiK3MultiModalProjector(KimiK25MultiModalProjector): + """Kimi projector with the optional ModelSlim output rotation.""" + + def __init__(self, config, *args, prefix: str = "", **kwargs) -> None: + super().__init__(config, *args, prefix=prefix, **kwargs) + output_size = config.text_hidden_size + self.rot_proj: ReplicatedLinear | None = ReplicatedLinear( + output_size, + output_size, + bias=False, + quant_config=None, + prefix=f"{prefix}.rot_proj", + ) + + def forward(self, image_features: torch.Tensor) -> torch.Tensor: + hidden_states = super().forward(image_features) + rot_proj = self.rot_proj + if rot_proj is not None: + hidden_states = rot_proj(hidden_states)[0] + return hidden_states + + +@MULTIMODAL_REGISTRY.register_processor( + KimiK3MultiModalProcessor, + info=KimiK3ProcessingInfo, + dummy_inputs=KimiK3DummyInputsBuilder, +) +class AscendKimiK3ForConditionalGeneration(UpstreamKimiK3ForConditionalGeneration): + """Upstream Kimi K3 multimodal wrapper with Ascend text/projector layers.""" + + def __init__(self, vllm_config: VllmConfig, prefix: str = "") -> None: + nn.Module.__init__(self) + model_config = vllm_config.model_config + self.config = model_config.hf_config + self.quant_config = vllm_config.quant_config + multimodal_config = model_config.multimodal_config + assert multimodal_config is not None + + self.use_data_parallel = is_vit_use_data_parallel( + self.config.vision_config.num_attention_heads, + ) + self.hidden_size = self.config.text_config.hidden_size + self.device = current_platform.current_device() + vision_quant_config = self._maybe_ignore_quant_config(self.quant_config) + + with self._mark_tower_model(vllm_config, "image"): + self.vision_tower = MoonViT3dPretrainedModel( + self.config.vision_config, + quant_config=vision_quant_config, + prefix=maybe_prefix(prefix, "vision_tower"), + ) + if vision_quant_config is not None: + self.vision_tower = self.vision_tower.to(device=self.device) + else: + self.vision_tower = self.vision_tower.to( + device=self.device, + dtype=model_config.dtype, + ) + + self.mm_projector = AscendKimiK3MultiModalProjector( + self.config.vision_config, + use_data_parallel=self.use_data_parallel, + quant_config=vision_quant_config, + prefix=maybe_prefix(prefix, "mm_projector"), + ) + if vision_quant_config is not None: + self.mm_projector = self.mm_projector.to(device=self.device) + else: + self.mm_projector = self.mm_projector.to( + device=self.device, + dtype=model_config.dtype, + ) + + with self._mark_language_model(vllm_config): + self.language_model = init_vllm_registered_model( + vllm_config=vllm_config, + hf_config=self.config.text_config, + prefix=maybe_prefix(prefix, "language_model"), + architectures=["KimiLinearForCausalLM"], + ) + self.make_empty_intermediate_tensors = ( # type: ignore[method-assign] + self.language_model.make_empty_intermediate_tensors + ) + self.media_placeholder = self.config.media_placeholder_token_id + + def _maybe_ignore_quant_config( + self, + quant_config: QuantizationConfig | None, + ) -> QuantizationConfig | None: + if isinstance( + quant_config, + compressed_tensors.CompressedTensorsConfig, + ): + return None + return quant_config + + def load_weights( + self, + weights: Iterable[tuple[str, torch.Tensor]], + ) -> set[str]: + rot_proj = self.mm_projector.rot_proj + loader = AutoWeightsLoader(self) + rot_proj_weight_names = ( + {name for name, _ in rot_proj.named_parameters(prefix="mm_projector.rot_proj")} + if rot_proj is not None + else set() + ) + loaded_weights = loader.load_weights( + weights, + mapper=self.hf_to_vllm_mapper, + ) + if rot_proj is not None and rot_proj_weight_names.isdisjoint(loaded_weights): + self.mm_projector.rot_proj = None + return loaded_weights diff --git a/vllm_ascend/models/kimi_k3_dspark.py b/vllm_ascend/models/kimi_k3_dspark.py new file mode 100644 index 000000000000..7013bf4377be --- /dev/null +++ b/vllm_ascend/models/kimi_k3_dspark.py @@ -0,0 +1,292 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Kimi K3 MLA DSpark draft model for Ascend.""" + +from collections.abc import Iterable + +import torch +from torch import nn +from vllm.config import VllmConfig +from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.model_executor.layers.linear import ReplicatedLinear +from vllm.model_executor.layers.logits_processor import LogitsProcessor +from vllm.model_executor.models.interfaces import MultiModalEmbeddings +from vllm.model_executor.models.qwen3_dspark import DSparkMarkovHead +from vllm.model_executor.models.utils import ( + AutoWeightsLoader, + _merge_multimodal_embeddings, + get_draft_quant_config, + maybe_prefix, +) +from vllm.models.kimi_k3.nvidia.dspark_mla import ( + K3DSparkForCausalLM as UpstreamK3DSparkForCausalLM, +) + +from vllm_ascend.models.kimi_k3 import ( + AscendKimiMLAAttention, + AscendKimiMLP, +) +from vllm_ascend.ops.rotary_embedding import get_cos_and_sin_mla + + +class AscendK3DSparkDecoderLayer(nn.Module): + def __init__( + self, + *, + vllm_config: VllmConfig, + config, + layer_idx: int, + start_layer_id: int, + prefix: str, + ) -> None: + super().__init__() + quant_config = get_draft_quant_config(vllm_config) + layer_prefix = maybe_prefix( + prefix, + f"layers.{start_layer_id + layer_idx}", + ) + self.self_attn = AscendKimiMLAAttention( + config=config, + hidden_size=config.hidden_size, + num_heads=config.num_attention_heads, + qk_nope_head_dim=config.qk_nope_head_dim, + qk_rope_head_dim=config.qk_rope_head_dim, + v_head_dim=config.v_head_dim, + q_lora_rank=config.q_lora_rank, + kv_lora_rank=config.kv_lora_rank, + use_output_gate=False, + use_rope=True, + cache_config=vllm_config.cache_config, + quant_config=quant_config, + prefix=f"{layer_prefix}.self_attn", + non_causal_multi_token_decode=True, + ) + self.mlp = AscendKimiMLP( + hidden_size=config.hidden_size, + intermediate_size=config.intermediate_size, + hidden_act=config.hidden_act, + quant_config=quant_config, + prefix=f"{layer_prefix}.mlp", + ) + self.input_layernorm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + self.post_attention_layernorm = RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + residual: torch.Tensor | None, + ) -> tuple[torch.Tensor, torch.Tensor]: + if residual is None: + residual = hidden_states + hidden_states = self.input_layernorm(hidden_states) + else: + hidden_states, residual = self.input_layernorm( + hidden_states, + residual, + ) + hidden_states = self.self_attn( + positions=positions, + hidden_states=hidden_states, + ) + hidden_states, residual = self.post_attention_layernorm( + hidden_states, + residual, + ) + hidden_states = self.mlp(hidden_states) + return hidden_states, residual + + +class AscendK3DSparkModel(nn.Module): + def __init__( + self, + *, + vllm_config: VllmConfig, + start_layer_id: int, + prefix: str, + ) -> None: + super().__init__() + assert vllm_config.speculative_config is not None + draft_model_config = vllm_config.speculative_config.draft_model_config + assert draft_model_config is not None + self.config = draft_model_config.hf_config + self.quant_config = get_draft_quant_config(vllm_config) + self.embed_tokens: nn.Module | None = None + + self.context_proj = ReplicatedLinear( + self.config.target_hidden_size * self.config.num_target_layers, + self.config.hidden_size, + bias=False, + return_bias=False, + quant_config=self.quant_config, + prefix=maybe_prefix(prefix, "context_proj"), + ) + self.context_norm = RMSNorm( + self.config.hidden_size, + eps=self.config.rms_norm_eps, + ) + self.layers = nn.ModuleList( + [ + AscendK3DSparkDecoderLayer( + vllm_config=vllm_config, + config=self.config, + layer_idx=layer_idx, + start_layer_id=start_layer_id, + prefix=prefix, + ) + for layer_idx in range(self.config.num_hidden_layers) + ] + ) + self.final_norm = RMSNorm( + self.config.hidden_size, + eps=self.config.rms_norm_eps, + ) + self.markov_head = DSparkMarkovHead( + self.config.vocab_size, + self.config.draft_vocab_size, + self.config.markov_rank, + prefix=maybe_prefix(prefix, "markov_head"), + ) + self._context_kv_fusion_available = False + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + assert self.embed_tokens is not None + return self.embed_tokens(input_ids) + + def combine_hidden_states(self, hidden_states: torch.Tensor) -> torch.Tensor: + return self.context_norm(self.context_proj(hidden_states)) + + @torch.inference_mode() + def precompute_and_store_context_kv( + self, + context_states: torch.Tensor, + context_positions: torch.Tensor, + context_slot_mapping: ( + torch.Tensor | list[torch.Tensor | None] | tuple[torch.Tensor | None, ...] | None + ) = None, + ) -> None: + if context_slot_mapping is None or context_states.numel() == 0: + return + per_layer_slot_mapping = isinstance(context_slot_mapping, (list, tuple)) + if per_layer_slot_mapping and len(context_slot_mapping) != len(self.layers): + raise ValueError( + "context_slot_mapping must contain one entry per draft layer: " + f"got {len(context_slot_mapping)} entries for " + f"{len(self.layers)} layers" + ) + cos, sin = get_cos_and_sin_mla(context_positions) + for layer_idx, layer in enumerate(self.layers): + attn = layer.self_attn + assert attn.fused_qkv_a_proj is not None + assert attn.q_lora_rank is not None + qkv_lora = attn.fused_qkv_a_proj(context_states)[0] + kv_no_split = qkv_lora[..., attn.q_lora_rank :].contiguous() + slots = context_slot_mapping[layer_idx] if per_layer_slot_mapping else context_slot_mapping + if slots is None: + continue + attn.impl.exec_kv_prefill( + kv_no_split, + cos, + sin, + attn.kv_cache, + slots, + ) + + def _build_fused_context_kv_buffers(self) -> None: + # The Ascend path keeps the quantization-aware per-layer projections. + self._context_kv_fusion_available = False + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + inputs_embeds: torch.Tensor | None = None, + ) -> torch.Tensor: + if inputs_embeds is None: + inputs_embeds = self.embed_input_ids(input_ids) + hidden_states = inputs_embeds + residual = None + for layer in self.layers: + hidden_states, residual = layer( + positions=positions, + hidden_states=hidden_states, + residual=residual, + ) + hidden_states, _ = self.final_norm(hidden_states, residual) + return hidden_states + + +class AscendK3DSparkForCausalLM(UpstreamK3DSparkForCausalLM): + def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: + nn.Module.__init__(self) + assert vllm_config.speculative_config is not None + self.draft_model_config = vllm_config.speculative_config.draft_model_config + assert self.draft_model_config is not None + self.config = self.draft_model_config.hf_config + target_layer_num = vllm_config.model_config.get_num_layers(vllm_config.parallel_config) + self.model = AscendK3DSparkModel( + vllm_config=vllm_config, + start_layer_id=target_layer_num, + prefix=maybe_prefix(prefix, "model"), + ) + self.lm_head: nn.Module | None = None + self.logits_processor = LogitsProcessor( + self.config.draft_vocab_size, + scale=getattr(self.config, "logit_scale", 1.0), + ) + + def load_weights( + self, + weights: Iterable[tuple[str, torch.Tensor]], + ) -> set[str]: + """Load the per-layer KV projections used by the Ascend draft model. + + Upstream additionally duplicates these weights into a CUDA-specific + cross-layer ``context_kv_proj``. Ascend deliberately retains the + quantization-aware per-layer projections, so use vLLM's public loader + interface without creating that extra packed parameter. + """ + loader = AutoWeightsLoader( + self, + skip_substrs=list(self.checkpoint_skip_substrs), + ) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) + + def embed_input_ids( + self, + input_ids: torch.Tensor, + multimodal_embeddings: MultiModalEmbeddings | None = None, + *, + is_multimodal: torch.Tensor | None = None, + ) -> torch.Tensor: + """Embed draft tokens and replace multimodal placeholder positions. + + vLLM 0.27 passes the target model's precomputed multimodal embeddings + through the speculative proposer. K3 DSpark shares the target token + embedding but upstream still exposes the older text-only method + signature, so adapt that interface without duplicating the vision + tower in the draft model. + """ + if multimodal_embeddings is None or len(multimodal_embeddings) == 0: + return self.model.embed_input_ids(input_ids) + if is_multimodal is None: + raise ValueError("is_multimodal is required when multimodal_embeddings are provided") + + # Placeholder ids are overwritten below. Mask them before the shared + # vocabulary lookup so out-of-vocabulary multimodal ids are safe too. + text_input_ids = input_ids.masked_fill( + is_multimodal.to(device=input_ids.device, non_blocking=True), + 0, + ) + inputs_embeds = self.model.embed_input_ids(text_input_ids) + return _merge_multimodal_embeddings( + inputs_embeds=inputs_embeds, + multimodal_embeddings=multimodal_embeddings, + is_multimodal=is_multimodal, + ) diff --git a/vllm_ascend/models/kimi_k3_mtp.py b/vllm_ascend/models/kimi_k3_mtp.py new file mode 100644 index 000000000000..145e46688d48 --- /dev/null +++ b/vllm_ascend/models/kimi_k3_mtp.py @@ -0,0 +1,146 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Kimi K3 MTP draft model for Ascend.""" + +import copy + +import torch +from torch import nn +from vllm.config import VllmConfig +from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.model_executor.layers.logits_processor import LogitsProcessor +from vllm.model_executor.layers.vocab_parallel_embedding import VocabParallelEmbedding +from vllm.model_executor.models.utils import maybe_prefix +from vllm.models.kimi_k3.amd.mtp import ( + KimiK3MTP as UpstreamKimiK3MTP, +) +from vllm.models.kimi_k3.amd.mtp import SharedHead +from vllm.models.kimi_k3.common.mtp import fused_mtp_input + +from vllm_ascend.models.kimi_k3 import AscendKimiDecoderLayer + + +class AscendKimiK3MultiTokenPredictorLayer(nn.Module): + def __init__(self, config, vllm_config: VllmConfig, prefix: str) -> None: + super().__init__() + self.config = config + self.enorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.hnorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.eh_proj = nn.Linear( + config.hidden_size * 2, + config.hidden_size, + bias=False, + ) + self.shared_head = SharedHead( + config=config, + prefix=prefix, + quant_config=vllm_config.quant_config, + ) + block_config = copy.copy(config) + block_config.attn_res_block_size = None + self.mtp_block = AscendKimiDecoderLayer( + block_config, + vllm_config, + prefix=prefix, + ) + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + previous_hidden_states: torch.Tensor, + inputs_embeds: torch.Tensor | None = None, + spec_step_index: int = 0, + ) -> tuple[torch.Tensor, torch.Tensor]: + del input_ids, spec_step_index + assert inputs_embeds is not None + hidden_states = self.eh_proj( + fused_mtp_input( + positions, + inputs_embeds, + previous_hidden_states, + self.enorm.weight, + self.hnorm.weight, + self.enorm.variance_epsilon, + ) + ) + hidden_states, residual = self.mtp_block( + positions=positions, + hidden_states=hidden_states, + residual=None, + ) + logits_hidden_states, hidden_states = self.shared_head.norm( + hidden_states, + residual, + ) + return logits_hidden_states, hidden_states + + +class AscendKimiK3MultiTokenPredictor(nn.Module): + def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: + super().__init__() + config = vllm_config.model_config.hf_text_config + self.config = config + self.mtp_start_layer_idx = config.num_hidden_layers + self.num_mtp_layers = config.num_nextn_predict_layers + self.layers = nn.ModuleDict( + { + str(idx): AscendKimiK3MultiTokenPredictorLayer( + config, + vllm_config, + f"{prefix}.layers.{idx}", + ) + for idx in range( + self.mtp_start_layer_idx, + self.mtp_start_layer_idx + self.num_mtp_layers, + ) + } + ) + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + prefix=maybe_prefix(prefix, "embed_tokens"), + ) + self.logits_processor = LogitsProcessor(config.vocab_size) + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.embed_tokens(input_ids) + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + previous_hidden_states: torch.Tensor, + inputs_embeds: torch.Tensor | None = None, + spec_step_idx: int = 0, + ) -> tuple[torch.Tensor, torch.Tensor]: + if inputs_embeds is None: + inputs_embeds = self.embed_tokens(input_ids) + current_step_idx = spec_step_idx % self.num_mtp_layers + return self.layers[str(self.mtp_start_layer_idx + current_step_idx)]( + input_ids, + positions, + previous_hidden_states, + inputs_embeds, + current_step_idx, + ) + + def compute_logits( + self, + hidden_states: torch.Tensor, + spec_step_idx: int = 0, + ) -> torch.Tensor: + current_step_idx = spec_step_idx % self.num_mtp_layers + mtp_layer = self.layers[str(self.mtp_start_layer_idx + current_step_idx)] + return self.logits_processor(mtp_layer.shared_head.head, hidden_states) + + +class AscendKimiK3MTP(UpstreamKimiK3MTP): + def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: + nn.Module.__init__(self) + self.config = vllm_config.model_config.hf_text_config + self.quant_config = vllm_config.quant_config + self.model = AscendKimiK3MultiTokenPredictor( + vllm_config=vllm_config, + prefix=maybe_prefix(prefix, "model"), + )