diff --git a/tests/models/test_deepseek_v4_dspark_rocm.py b/tests/models/test_deepseek_v4_dspark_rocm.py new file mode 100644 index 000000000000..efa97e10d7f9 --- /dev/null +++ b/tests/models/test_deepseek_v4_dspark_rocm.py @@ -0,0 +1,130 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from types import SimpleNamespace + +import pytest +import torch + +from vllm.models.deepseek_v4.amd import dspark as dspark_module +from vllm.models.deepseek_v4.amd.dspark import DSparkDeepseekV4ForCausalLM + + +def _make_uninitialized_model(confidence_head): + model = DSparkDeepseekV4ForCausalLM.__new__(DSparkDeepseekV4ForCausalLM) + object.__setattr__( + model, + "model", + SimpleNamespace(confidence_head=confidence_head), + ) + return model + + +def _prepare_loader_model(model, named_parameters): + object.__setattr__( + model, + "config", + SimpleNamespace( + n_routed_experts=1, + expert_dtype="fp4", + num_attention_heads=1, + ), + ) + object.__setattr__(model, "named_parameters", lambda: named_parameters) + + +def _disable_distributed_loader_paths(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr( + dspark_module, + "fused_moe_make_expert_params_mapping", + lambda *args, **kwargs: [], + ) + monkeypatch.setattr( + dspark_module, "get_tensor_model_parallel_world_size", lambda: 1 + ) + monkeypatch.setattr(dspark_module, "get_tensor_model_parallel_rank", lambda: 0) + + +@pytest.mark.cpu_test +def test_deepseek_v4_rocm_dspark_maps_enabled_confidence_head(): + model = _make_uninitialized_model(object()) + + assert ( + model._remap_dspark_name("mtp.2.confidence_head.proj.weight") + == "model.confidence_head.proj.weight" + ) + + +@pytest.mark.cpu_test +def test_deepseek_v4_rocm_dspark_skips_disabled_confidence_head(): + model = _make_uninitialized_model(None) + + assert model._remap_dspark_name("mtp.2.confidence_head.proj.weight") is None + + +@pytest.mark.cpu_test +def test_deepseek_v4_rocm_dspark_disables_unloaded_confidence_head( + monkeypatch: pytest.MonkeyPatch, +): + model = _make_uninitialized_model(object()) + _prepare_loader_model(model, []) + _disable_distributed_loader_paths(monkeypatch) + + assert model.load_weights([]) == set() + assert model.model.confidence_head is None + + +@pytest.mark.cpu_test +def test_deepseek_v4_rocm_dspark_loads_complete_confidence_head( + monkeypatch: pytest.MonkeyPatch, +): + class FakeParameter: + loaded_weight = None + + def weight_loader(self, param, loaded_weight): + assert param is self + self.loaded_weight = loaded_weight + + confidence_head = object() + parameter = FakeParameter() + model = _make_uninitialized_model(confidence_head) + _prepare_loader_model( + model, + [("model.confidence_head.proj.weight", parameter)], + ) + _disable_distributed_loader_paths(monkeypatch) + loaded_weight = torch.tensor([[1.0, 2.0]]) + + loaded = model.load_weights([("mtp.2.confidence_head.proj.weight", loaded_weight)]) + + assert loaded == {"model.confidence_head.proj.weight"} + assert parameter.loaded_weight is loaded_weight + assert model.model.confidence_head is confidence_head + + +@pytest.mark.cpu_test +def test_deepseek_v4_rocm_dspark_confidence_is_probability(): + class ConfidenceHead: + def __call__(self, head_hidden, markov_embed): + return (head_hidden[:, 0] + markov_embed[:, 0]).float() + + model = _make_uninitialized_model(ConfidenceHead()) + head_hidden = torch.tensor([[0.0], [1.0]], dtype=torch.bfloat16) + markov_embed = torch.tensor([[0.0], [-2.0]], dtype=torch.bfloat16) + + confidence = model.compute_confidence(head_hidden, markov_embed) + + torch.testing.assert_close( + confidence, + torch.sigmoid(torch.tensor([0.0, -1.0])), + ) + assert torch.all((confidence >= 0) & (confidence <= 1)) + + +@pytest.mark.cpu_test +def test_deepseek_v4_rocm_dspark_confidence_requires_a_head(): + model = _make_uninitialized_model(None) + empty = torch.zeros((1, 1)) + + with pytest.raises(RuntimeError, match="confidence_head"): + model.compute_confidence(empty, empty) diff --git a/tests/v1/attention/test_deepseek_v4_rocm_adaptive.py b/tests/v1/attention/test_deepseek_v4_rocm_adaptive.py new file mode 100644 index 000000000000..3549ddf06f25 --- /dev/null +++ b/tests/v1/attention/test_deepseek_v4_rocm_adaptive.py @@ -0,0 +1,207 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from types import SimpleNamespace + +import pytest +import torch + +from vllm.models.deepseek_v4.amd.rocm import ( + DeepseekV4ROCMAiterMLASparseMetadataBuilder, + DeepseekV4ROCMAiterSparseSWAMetadataBuilder, +) +from vllm.v1.attention.backend import AttentionCGSupport +from vllm.v1.attention.backends.mla import indexer +from vllm.v1.attention.backends.mla.indexer import ( + DeepseekV4IndexerBackend, + DeepseekV32IndexerMetadataBuilder, +) + + +def _make_indexer_builder(*, adaptive: bool, capacity: int = 12): + builder = DeepseekV32IndexerMetadataBuilder.__new__( + DeepseekV32IndexerMetadataBuilder + ) + builder.vllm_config = SimpleNamespace( + speculative_config=SimpleNamespace(enable_adaptive_verification=adaptive) + ) + builder.supports_varlen = False + builder.decode_seq_lens_buffer = torch.zeros(capacity, dtype=torch.int32) + builder.expanded_block_table_buffer = torch.zeros((capacity, 2), dtype=torch.int32) + builder.decode_lens_buffer = torch.zeros(capacity, dtype=torch.int32) + builder.arange_buffer = torch.arange(capacity, dtype=torch.int32) + return builder + + +@pytest.mark.cpu_test +def test_deepseek_v4_rocm_adaptive_builders_support_varlen_full_graphs(): + adaptive_config = SimpleNamespace( + speculative_config=SimpleNamespace(enable_adaptive_verification=True) + ) + fixed_config = SimpleNamespace( + speculative_config=SimpleNamespace(enable_adaptive_verification=False) + ) + + for builder_cls in ( + DeepseekV4ROCMAiterMLASparseMetadataBuilder, + DeepseekV4ROCMAiterSparseSWAMetadataBuilder, + ): + assert ( + builder_cls.get_cudagraph_support(adaptive_config, SimpleNamespace()) + == AttentionCGSupport.ALWAYS + ) + assert ( + builder_cls.get_cudagraph_support(fixed_config, SimpleNamespace()) + == AttentionCGSupport.UNIFORM_BATCH + ) + + +@pytest.mark.cpu_test +def test_deepseek_v4_rocm_adaptive_indexer_support(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(indexer.current_platform, "is_rocm", lambda: True) + adaptive_config = SimpleNamespace( + num_speculative_tokens=1, + speculative_config=SimpleNamespace(enable_adaptive_verification=True), + ) + fixed_config = SimpleNamespace( + num_speculative_tokens=1, + speculative_config=SimpleNamespace(enable_adaptive_verification=False), + ) + + assert DeepseekV4IndexerBackend.supports_device_cpu_query_lens_mismatch() + assert ( + DeepseekV4IndexerBackend.get_builder_cls() is DeepseekV32IndexerMetadataBuilder + ) + assert indexer._use_flattening(adaptive_config) + assert ( + DeepseekV32IndexerMetadataBuilder.get_cudagraph_support( + adaptive_config, SimpleNamespace() + ) + == AttentionCGSupport.ALWAYS + ) + assert not indexer._use_flattening(fixed_config) + assert ( + DeepseekV32IndexerMetadataBuilder.get_cudagraph_support( + fixed_config, SimpleNamespace() + ) + == AttentionCGSupport.UNIFORM_BATCH + ) + + +@pytest.mark.cpu_test +def test_rocm_adaptive_indexer_preserves_single_request_uniform_path( + monkeypatch: pytest.MonkeyPatch, +): + builder = _make_indexer_builder(adaptive=True, capacity=8) + + class FakeUniformKernel: + called = False + + def __call__( + self, + seq_lens, + decode_seq_lens, + block_table, + expanded_block_table, + decode_lens, + num_decode_tokens, + max_decode_len, + ): + self.called = True + assert num_decode_tokens == max_decode_len == 4 + decode_seq_lens[:num_decode_tokens] = torch.arange( + seq_lens[0] - max_decode_len + 1, + seq_lens[0] + 1, + dtype=torch.int32, + ) + expanded_block_table[:num_decode_tokens] = block_table[0] + decode_lens[:num_decode_tokens] = 1 + + fake_kernel = FakeUniformKernel() + monkeypatch.setattr(indexer, "_PREPARE_UNIFORM_DECODE_KERNEL", fake_kernel) + + seq_lens, block_table, decode_lens, batch_size, requires_padding = ( + builder._prepare_decode_tensors( + seq_lens=torch.tensor([10], dtype=torch.int32), + block_table=torch.tensor([[1, 2]], dtype=torch.int32), + decode_lens=torch.tensor([4], dtype=torch.int32), + decode_lens_cpu=torch.tensor([4], dtype=torch.int32), + query_start_loc=torch.tensor([0], dtype=torch.int32), + num_decodes=1, + num_decode_tokens=4, + use_native=False, + next_n=8, + max_decode_len=4, + ) + ) + + assert fake_kernel.called + torch.testing.assert_close(seq_lens, torch.tensor([7, 8, 9, 10], dtype=torch.int32)) + torch.testing.assert_close( + block_table, + torch.tensor([[1, 2], [1, 2], [1, 2], [1, 2]], dtype=torch.int32), + ) + torch.testing.assert_close(decode_lens, torch.ones(4, dtype=torch.int32)) + assert batch_size == 4 + assert not requires_padding + + +@pytest.mark.cpu_test +def test_rocm_adaptive_indexer_replays_changed_allocations_with_stable_buffers(): + builder = _make_indexer_builder(adaptive=True) + buffer_ptrs = ( + builder.decode_seq_lens_buffer.data_ptr(), + builder.expanded_block_table_buffer.data_ptr(), + builder.decode_lens_buffer.data_ptr(), + ) + block_table = torch.tensor([[1, 2], [3, 4], [5, 6]], dtype=torch.int32) + decode_lens_cpu = torch.tensor([4, 4, 0], dtype=torch.int32) + + first = builder._prepare_decode_tensors( + seq_lens=torch.tensor([20, 20, 0], dtype=torch.int32), + block_table=block_table, + decode_lens=torch.tensor([7, 1, 0], dtype=torch.int32), + decode_lens_cpu=decode_lens_cpu, + query_start_loc=torch.tensor([0, 7, 8], dtype=torch.int32), + num_decodes=3, + num_decode_tokens=10, + use_native=False, + next_n=8, + max_decode_len=4, + ) + torch.testing.assert_close( + first[0], + torch.tensor([14, 15, 16, 17, 18, 19, 20, 20, 0, 0], dtype=torch.int32), + ) + torch.testing.assert_close( + first[1][:, 0], + torch.tensor([1, 1, 1, 1, 1, 1, 1, 3, 0, 0], dtype=torch.int32), + ) + + second = builder._prepare_decode_tensors( + seq_lens=torch.tensor([20, 20, 0], dtype=torch.int32), + block_table=block_table, + decode_lens=torch.tensor([1, 7, 0], dtype=torch.int32), + decode_lens_cpu=decode_lens_cpu, + query_start_loc=torch.tensor([0, 1, 8], dtype=torch.int32), + num_decodes=3, + num_decode_tokens=10, + use_native=False, + next_n=8, + max_decode_len=4, + ) + torch.testing.assert_close( + second[0], + torch.tensor([20, 14, 15, 16, 17, 18, 19, 20, 0, 0], dtype=torch.int32), + ) + torch.testing.assert_close( + second[1][:, 0], + torch.tensor([1, 3, 3, 3, 3, 3, 3, 3, 0, 0], dtype=torch.int32), + ) + torch.testing.assert_close(second[2], torch.ones(10, dtype=torch.int32)) + assert second[3:] == (10, False) + assert buffer_ptrs == ( + builder.decode_seq_lens_buffer.data_ptr(), + builder.expanded_block_table_buffer.data_ptr(), + builder.decode_lens_buffer.data_ptr(), + ) diff --git a/vllm/models/deepseek_v4/amd/dspark.py b/vllm/models/deepseek_v4/amd/dspark.py index edca59032aa3..c3f249fa80d5 100644 --- a/vllm/models/deepseek_v4/amd/dspark.py +++ b/vllm/models/deepseek_v4/amd/dspark.py @@ -44,6 +44,7 @@ ) from vllm.model_executor.model_loader.weight_utils import default_weight_loader from vllm.model_executor.models.qwen3_dspark import ( + DSparkConfidenceHead, DSparkMarkovHead, ) from vllm.model_executor.models.utils import maybe_prefix @@ -102,7 +103,7 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: ] ) - # Heads: final norm + hc_head, and the Markov head + # Heads: final norm + hc_head, and the Markov + confidence heads # Loaded from the "final" MTP layer weights (mtp.*) in the target checkpoint self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) hc_dim = self.hc_mult * config.hidden_size @@ -125,6 +126,12 @@ def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: config.dspark_markov_rank, prefix=maybe_prefix(prefix, "markov_head"), ) + self.confidence_head: DSparkConfidenceHead | None = None + if getattr(config, "enable_confidence_head", True): + self.confidence_head = DSparkConfidenceHead( + config.hidden_size + config.dspark_markov_rank, + prefix=maybe_prefix(prefix, "confidence_head"), + ) # MHC head CustomOp dispatcher (aiter / tilelang / triton / torch), # replacing the direct nvidia tilelang kernel call. @@ -358,6 +365,17 @@ def markov_embed(self, token_ids: torch.Tensor) -> torch.Tensor: def markov_bias(self, markov_embed: torch.Tensor) -> torch.Tensor: return self.model.markov_head.bias(markov_embed, self.logits_processor) + def compute_confidence( + self, head_hidden: torch.Tensor, markov_embed: torch.Tensor + ) -> torch.Tensor: + """Per-position acceptance probability for each drafted token.""" + if self.model.confidence_head is None: + raise RuntimeError( + "compute_confidence() requires a confidence head, but the " + "checkpoint did not provide confidence_head weights." + ) + return torch.sigmoid(self.model.confidence_head(head_hidden, markov_embed)) + # --- Weight loading ---------------------------------------------------- def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: @@ -391,6 +409,7 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: params_dict = dict(self.named_parameters()) loaded_params: set[str] = set() + loaded_confidence_head = False tp_size = get_tensor_model_parallel_world_size() tp_rank = get_tensor_model_parallel_rank() @@ -403,6 +422,8 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: if mapped is None: continue name = mapped + if "confidence_head." in name: + loaded_confidence_head = True # ``.scale`` -> per-method scale suffix. if name.endswith(".scale"): @@ -439,8 +460,8 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: continue # Stacked rules only apply to decoder-layer weights. Head-stack params - # (main_proj/norm/hc_head/markov_head) load directly — otherwise e.g. - # "markov_w1" would collide with the "w1" shard rule. + # (main_proj/norm/hc_head/markov_head/confidence_head) load directly — + # otherwise e.g. "markov_w1" would collide with the "w1" shard rule. is_layer_param = name.startswith("model.layers.") for param_name, weight_name, stacked_shard_id in stacked_params_mapping: if not is_layer_param or weight_name not in name: @@ -469,6 +490,8 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: weight_loader(param, loaded_weight) loaded_params.add(name) + if self.model.confidence_head is not None and not loaded_confidence_head: + self.model.confidence_head = None logger.info_once("DSpark draft model loaded: %d params", len(loaded_params)) return loaded_params @@ -482,8 +505,7 @@ def _remap_dspark_name(self, name: str) -> str | None: return None stage = int(m.group(1)) rest = m.group(2) - # The confidence head is not wired into inference yet; drop its weights. - if rest.startswith("confidence_head."): + if rest.startswith("confidence_head.") and self.model.confidence_head is None: return None # Head-stack params live at model level (mtp.last), context combiner at # model level (mtp.0); everything else is a per-layer decoder block. @@ -493,6 +515,7 @@ def _remap_dspark_name(self, name: str) -> str | None: "hc_head_base", "hc_head_scale", "markov_head.", + "confidence_head.", ) if rest.startswith(("main_proj.", "main_norm.")) or rest.startswith( head_prefixes diff --git a/vllm/models/deepseek_v4/amd/rocm.py b/vllm/models/deepseek_v4/amd/rocm.py index 3cc9827c1fbe..5d884083de58 100644 --- a/vllm/models/deepseek_v4/amd/rocm.py +++ b/vllm/models/deepseek_v4/amd/rocm.py @@ -7,6 +7,7 @@ import torch +from vllm.config import VllmConfig from vllm.distributed import ( get_tensor_model_parallel_world_size, tensor_model_parallel_all_reduce, @@ -25,6 +26,7 @@ from vllm.triton_utils import tl, triton from vllm.utils.multi_stream_utils import execute_in_parallel from vllm.v1.attention.backend import ( + AttentionCGSupport, CommonAttentionMetadata, ) from vllm.v1.attention.backends.mla.sparse_swa import ( @@ -37,6 +39,7 @@ rocm_sparse_attn_decode, rocm_sparse_attn_prefill, ) +from vllm.v1.kv_cache_interface import KVCacheSpec from vllm.v1.worker.workspace import current_workspace_manager logger = init_logger(__name__) @@ -416,6 +419,20 @@ class DeepseekV4ROCMAiterSparseSWAMetadata(DeepseekSparseSWAMetadata): class DeepseekV4ROCMAiterMLASparseMetadataBuilder(DeepseekV4SparseMLAMetadataBuilder): + @classmethod + def get_cudagraph_support( + cls, + vllm_config: VllmConfig, + kv_cache_spec: KVCacheSpec, + ) -> AttentionCGSupport: + spec_config = vllm_config.speculative_config + if spec_config is not None and spec_config.enable_adaptive_verification: + # All per-token metadata is built from device query boundaries into + # persistent buffers, so adaptive verification can replay varlen + # FULL decode graphs after reallocating drafts across requests. + return AttentionCGSupport.ALWAYS + return super().get_cudagraph_support(vllm_config, kv_cache_spec) + def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self.c128a_decode_topk_ragged_indices_buffer: torch.Tensor | None = None @@ -483,6 +500,20 @@ def build_for_cudagraph_capture( class DeepseekV4ROCMAiterSparseSWAMetadataBuilder(DeepseekSparseSWAMetadataBuilder): + @classmethod + def get_cudagraph_support( + cls, + vllm_config: VllmConfig, + kv_cache_spec: KVCacheSpec, + ) -> AttentionCGSupport: + spec_config = vllm_config.speculative_config + if spec_config is not None and spec_config.enable_adaptive_verification: + # SWA indices, lengths, and token-to-request mappings are built from + # device boundaries into persistent buffers, so adaptive verification + # can replay varlen FULL decode graphs safely. + return AttentionCGSupport.ALWAYS + return super().get_cudagraph_support(vllm_config, kv_cache_spec) + # Keep fused multi-step decode disabled until update_draft_decode_metadata() # also refreshes the ROCm-specific ragged SWA indices and indptrs. supports_draft_decode_metadata_update = False diff --git a/vllm/v1/attention/backends/mla/indexer.py b/vllm/v1/attention/backends/mla/indexer.py index 2321cf8d13d4..863909308949 100644 --- a/vllm/v1/attention/backends/mla/indexer.py +++ b/vllm/v1/attention/backends/mla/indexer.py @@ -238,6 +238,15 @@ def get_builder_cls() -> type["KpoolTailMetadataBuilder"]: # type: ignore[overr class DeepseekV4IndexerBackend(DeepseekV32IndexerBackend): + @classmethod + def supports_device_cpu_query_lens_mismatch(cls) -> bool: + # ROCm runs adaptive verification through the per-token flattened + # indexer path, which derives row ownership from device decode lengths. + return ( + _rocm_supports_flattened_device_query_lens() + or super().supports_device_cpu_query_lens_mismatch() + ) + @staticmethod def get_name() -> str: return "DEEPSEEK_V4_INDEXER" @@ -722,6 +731,10 @@ def _supports_flattened_device_query_lens() -> bool: ) +def _rocm_supports_flattened_device_query_lens() -> bool: + return current_platform.is_rocm() + + def _supports_native_decode(next_n: int) -> bool: """Whether decode can pass `next_n` Q rows per request to the kernel instead of flattening to one single-token row per query, which re-reads @@ -742,7 +755,10 @@ def _use_flattening(vllm_config: VllmConfig) -> bool: return not _supports_native_decode(next_n) or ( speculative_config is not None and speculative_config.enable_adaptive_verification - and _supports_flattened_device_query_lens() + and ( + _supports_flattened_device_query_lens() + or _rocm_supports_flattened_device_query_lens() + ) )