Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
130 changes: 130 additions & 0 deletions tests/models/test_deepseek_v4_dspark_rocm.py
Original file line number Diff line number Diff line change
@@ -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)
207 changes: 207 additions & 0 deletions tests/v1/attention/test_deepseek_v4_rocm_adaptive.py
Original file line number Diff line number Diff line change
@@ -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(),
)
Loading
Loading