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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 15 additions & 0 deletions tests/config/test_config_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -221,6 +221,21 @@ def test_cache_config_hash_ignores_prefix_cache_retention_interval():
assert CacheConfig(prefix_cache_retention_interval=64).compute_hash() == base_hash


def test_swa_page_sizes_have_distinct_compile_cache_hashes():
configs = (
CacheConfig(swa_block_size=32),
CacheConfig(swa_block_size=64),
CacheConfig(swa_block_size=128),
)
assert len({config.compute_hash() for config in configs}) == 3


@pytest.mark.parametrize("block_size", [0, -32, 16, 96, 256])
def test_swa_page_size_rejects_unsupported_kernel_geometry(block_size):
with pytest.raises(ValueError, match="swa_block_size"):
CacheConfig(swa_block_size=block_size)


def test_envs_compile_factors_relocation_invariant(tmp_path):
"""Relocating HOME or the XDG roots must not change the compile-cache
env hash.
Expand Down
16 changes: 16 additions & 0 deletions tests/engine/test_arg_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -490,6 +490,22 @@ def test_attention_config():
engine_args.create_engine_config()


def test_swa_block_size_cli_choices(tmp_path):
parser = EngineArgs.add_cli_args(FlexibleArgumentParser())
model_args = ["--model", str(tmp_path)]
assert (
EngineArgs.from_cli_args(parser.parse_args(model_args)).swa_block_size is None
)
args = parser.parse_args([*model_args, "--swa-block-size", "None"])
assert EngineArgs.from_cli_args(args).swa_block_size is None
for size in (32, 64, 128):
args = parser.parse_args([*model_args, "--swa-block-size", str(size)])
assert EngineArgs.from_cli_args(args).swa_block_size == size
parser.exit_on_error = False
with pytest.raises(ArgumentError):
parser.parse_args(["--swa-block-size", "96"])


def test_prefix_cache_default():
parser = EngineArgs.add_cli_args(FlexibleArgumentParser())
args = parser.parse_args([])
Expand Down
7 changes: 5 additions & 2 deletions tests/models/test_deepseek_v4_1_ced_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -181,6 +181,7 @@ def test_indexer_short_scan_reservation_preserves_long_scan_and_decode(
capacity=4096,
n_local_heads=16,
swa_width=128,
swa_cache_layer=SimpleNamespace(block_size=64),
is_ced_decoder=ced,
layer_id=2,
candidate_source_layer=20,
Expand Down Expand Up @@ -306,7 +307,8 @@ def test_ced_global_context_and_full_row_output_abi(monkeypatch, compact):
positions = torch.arange(100, 108)
inputs = torch.arange(16, dtype=torch.float32).reshape(8, 2) / 8
indices = torch.tensor([2, 3, 6, 7, -1, -1]) if compact else None
shared, row_work = {}, []
shared: dict[str, torch.Tensor | int] = {}
row_work: list[tuple[str, int, int]] = []
buffer = torch.empty(8, 8)
target = SimpleNamespace(
embed_input_ids=lambda _: inputs,
Expand Down Expand Up @@ -345,6 +347,7 @@ def test_ced_global_context_and_full_row_output_abi(monkeypatch, compact):
if index >= 2:
expected_aux.append(expected.clone())
if compact:
assert indices is not None
keep = torch.zeros(8, dtype=torch.bool)
keep[indices[indices >= 0]] = True
expected[~keep] = 0
Expand All @@ -360,7 +363,7 @@ def test_ced_global_context_and_full_row_output_abi(monkeypatch, compact):
atol=1e-6,
)
assert shared["global_rows"] == 8
decoder_rows = len(indices) if compact else 8
decoder_rows = len(indices) if indices is not None else 8
assert row_work == [
(kind, i, 8 if i < 2 else decoder_rows)
for i in range(4)
Expand Down
34 changes: 34 additions & 0 deletions tests/test_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
import vllm.envs as envs
from vllm.compilation.backends import VllmBackend
from vllm.config import (
CacheConfig,
CompilationConfig,
KernelConfig,
ModelConfig,
Expand All @@ -41,6 +42,39 @@
DEVICE_TYPE = current_platform.device_type


@pytest.mark.parametrize("swa_size,prefix_unit", [(None, 128), (32, 64), (128, 96)])
def test_swa_page_size_rejects_incompatible_prefix_matching(swa_size, prefix_unit):
config = SimpleNamespace(
model_config=SimpleNamespace(architecture="DeepseekV41ForCausalLM"),
speculative_config=None,
cache_config=CacheConfig(
swa_block_size=swa_size, prefix_match_unit=prefix_unit
),
)
with pytest.raises(ValueError, match="must be divisible by --prefix-match-unit"):
VllmConfig.validate_swa_block_size(config)
config.cache_config.prefix_match_unit = 32
VllmConfig.validate_swa_block_size(config)


def test_swa_page_size_is_scoped_to_v41_target_and_draft():
target = SimpleNamespace(architecture="DeepseekV41ForCausalLM")
draft = SimpleNamespace(architecture="DSparkDeepseekV4ForCausalLM")
config = SimpleNamespace(
model_config=draft,
speculative_config=SimpleNamespace(
target_model_config=target, draft_model_config=draft
),
cache_config=CacheConfig(swa_block_size=128, prefix_match_unit=32),
)
VllmConfig.validate_swa_block_size(config)
target.architecture = "DeepseekV4ForCausalLM"
with pytest.raises(ValueError, match="only supported by native DeepSeek V4.1"):
VllmConfig.validate_swa_block_size(config)
config.cache_config.swa_block_size = None
VllmConfig.validate_swa_block_size(config)


@pytest.mark.parametrize(
"configured,expected",
[(None, "auto"), ("b12x", "b12x"), ("FLASHINFER-CUTLASS", "flashinfer_cutlass")],
Expand Down
118 changes: 92 additions & 26 deletions tests/v1/attention/test_b12x_v41_workspace.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,7 +76,7 @@ def test_indexer_loads_complete_projections_on_tp_rank(native_workspace, monkeyp
torch.testing.assert_close(projection.weight, source, rtol=0, atol=0)


def _layer(attention, layer_id=0):
def _layer(attention, layer_id=0, swa_page=32):
# Avoid checkpoint/model construction: exercise the real planning and
# attention methods with the same TP4 head geometry and serving capacity.
layer = attention.DeepseekV4Attention.__new__(attention.DeepseekV4Attention)
Expand Down Expand Up @@ -105,13 +105,63 @@ def _layer(attention, layer_id=0):
layer.indexer = SimpleNamespace(
heads=32, k_cache=SimpleNamespace(prefix=layer.prefix + ".indexer.k_cache")
)
layer.swa_cache_layer = SimpleNamespace(prefix=layer.prefix + ".swa_cache")
layer.swa_cache_layer = SimpleNamespace(
prefix=layer.prefix + ".swa_cache", block_size=swa_page
)
layer.compressor = None
layer._context = {layer.prefix: layer}
layer._ready = False
return layer


@pytest.mark.parametrize("swa_page", [32, 64, 128])
def test_context_cache_write_preserves_page_boundaries_and_padding(
native_workspace, swa_page
):
from b12x.attention._shared.mla.compressed_reference import (
pack_deepseek_v41_cache_reference,
)

attention, _, _ = native_workspace
device = torch.device("cuda")
torch.manual_seed(145)
layer = _layer(attention, swa_page=swa_page)
rows = 2 * swa_page + 3
layer.rotary_emb = SimpleNamespace(
cos_sin_cache=torch.cat(
(
torch.ones((rows, 32), device=device),
torch.zeros((rows, 32), device=device),
),
dim=-1,
)
)
page_bytes = attention.mla.page_nbytes(
swa_page, cache_kind="swa", cache_format="deepseek_v41"
)
storage = torch.full((6, 113920), 0xA5, dtype=torch.uint8, device=device)
layer.swa_cache_layer.kv_cache = storage[:, :page_bytes]
expected = storage.clone()
logical = torch.arange(rows, device=device) + swa_page - 1
pages = torch.tensor([3, 1, 4, 2], device=device)
slots = pages[logical // swa_page] * swa_page + logical % swa_page
slots[-2:] = -1
positions = torch.arange(rows, device=device)
kv = torch.randn((rows, 512), device=device, dtype=torch.bfloat16)
records = pack_deepseek_v41_cache_reference(
kv, page_size=swa_page, cache_kind="swa"
).view(-1, 528)
valid = slots >= 0
columns = (slots[valid] % swa_page)[:, None] * 528 + torch.arange(
528, device=device
)
expected[(slots[valid] // swa_page)[:, None], columns] = records[:rows][valid]

layer.insert_context_kv(kv, positions, slots)

torch.testing.assert_close(storage, expected, rtol=0, atol=0)


def test_prepare_memory_is_metadata_not_capacity_activations(native_workspace):
attention, manager, _ = native_workspace
device = torch.device("cuda", torch.accelerator.current_device_index())
Expand Down Expand Up @@ -287,6 +337,7 @@ def record(device, capacity, hidden):
assert bool(torch.isfinite(y).all())


@pytest.mark.parametrize("main_page,swa_page", [(64, 32), (128, 64), (256, 128)])
@pytest.mark.parametrize(
"is_decode,rows,live_rows",
[
Expand All @@ -299,14 +350,15 @@ def record(device, capacity, hidden):
],
)
def test_attention_shared_scratch_graph_replay(
native_workspace, monkeypatch, is_decode, rows, live_rows
native_workspace, monkeypatch, is_decode, rows, live_rows, main_page, swa_page
):
attention, manager, workspace = native_workspace
# A TP4 rank must score all replicated heads without obtaining a TP group.
monkeypatch.setattr(attention, "get_tensor_model_parallel_world_size", lambda: 4)
device = torch.device("cuda", torch.accelerator.current_device_index())
torch.manual_seed(142)
layer = _layer(attention)
layer = _layer(attention, swa_page=swa_page)
layer.config.cache_config.block_size = main_page
layer.max_model_len = 4096
if rows == 36:
layer.config.speculative_config.parallel_drafting = True
Expand Down Expand Up @@ -346,14 +398,14 @@ def metadata(page):
context = SimpleNamespace(
cudagraph_runtime_mode=CUDAGraphMode.NONE,
attn_metadata={
layer.swa_cache_layer.prefix: metadata(32),
layer.prefix: metadata(64),
layer.indexer.k_cache.prefix: metadata(64),
layer.swa_cache_layer.prefix: metadata(swa_page),
layer.prefix: metadata(main_page),
layer.indexer.k_cache.prefix: metadata(main_page),
},
)
monkeypatch.setattr(attention, "get_forward_context", lambda: context)
kv = torch.randn((length, 512), device=device, dtype=torch.bfloat16)
for kind, page in (("swa", 32), ("indexed", 64)):
for kind, page in (("swa", swa_page), ("indexed", main_page)):
cache = torch.empty(
(
length // page + 1,
Expand All @@ -377,7 +429,10 @@ def metadata(page):
else:
layer.kv_cache = cache
layer.indexer.k_cache.kv_cache = torch.empty(
(length // 64 + 1, attention.dsa_indexer.MXFP4_INDEX_PAGE_BYTES),
(
length // main_page + 1,
attention.dsa_indexer.index_mxfp4_page_bytes(main_page),
),
dtype=torch.uint8,
device=device,
)
Expand All @@ -391,7 +446,8 @@ def metadata(page):
attention.dsa_indexer.quantize_write_index_k_mxfp4(
index_keys,
index_k_cache=layer.indexer.k_cache.kv_cache,
slot_mapping=torch.arange(length, device=device) + 64,
slot_mapping=torch.arange(length, device=device) + main_page,
page_size=main_page,
)
layer.attn_sink = torch.zeros(layer.n_local_heads, device=device)
q = torch.randn(
Expand Down Expand Up @@ -780,9 +836,10 @@ def test_ced_global_preparation_preserves_full_row_cache_bytes(
)


@pytest.mark.parametrize("main_page,swa_page", [(64, 32), (128, 64), (256, 128)])
@torch.inference_mode()
def test_ced_compact_attention_bounded_oracle_and_frozen_replay(
native_workspace, monkeypatch
native_workspace, monkeypatch, main_page, swa_page
):
from b12x import freeze_kernel_resolution, unfreeze_kernel_resolution
from b12x.attention._shared.mla.compressed_reference import (
Expand All @@ -792,7 +849,8 @@ def test_ced_compact_attention_bounded_oracle_and_frozen_replay(
attention, manager, workspace = native_workspace
torch.manual_seed(713)
device = torch.device("cuda")
layer = _layer(attention, 20)
layer = _layer(attention, 20, swa_page=swa_page)
layer.config.cache_config.block_size = main_page
layer.is_ced_decoder = True
layer._prepare(device)
rows, length, boundary = 128, 256, 128
Expand Down Expand Up @@ -825,14 +883,14 @@ def metadata(page):

context = SimpleNamespace(
attn_metadata={
layer.swa_cache_layer.prefix: metadata(32),
layer.prefix: metadata(64),
layer.indexer.k_cache.prefix: metadata(64),
layer.swa_cache_layer.prefix: metadata(swa_page),
layer.prefix: metadata(main_page),
layer.indexer.k_cache.prefix: metadata(main_page),
}
)
monkeypatch.setattr(attention, "get_forward_context", lambda: context)
kv = torch.randn((length, 512), device=device, dtype=torch.bfloat16)
for kind, page in (("swa", 32), ("indexed", 64)):
for kind, page in (("swa", swa_page), ("indexed", main_page)):
cache = torch.zeros(
(
length // page + 1,
Expand Down Expand Up @@ -865,7 +923,10 @@ def metadata(page):
cache, page_size=page, cache_kind=kind
)[page : page + length].clone()
layer.indexer.k_cache.kv_cache = torch.zeros(
(length // 64 + 1, attention.dsa_indexer.index_mxfp4_page_bytes(64)),
(
length // main_page + 1,
attention.dsa_indexer.index_mxfp4_page_bytes(main_page),
),
device=device,
dtype=torch.uint8,
)
Expand All @@ -875,8 +936,8 @@ def metadata(page):
attention.dsa_indexer.quantize_write_index_k_mxfp4(
keys,
index_k_cache=layer.indexer.k_cache.kv_cache,
slot_mapping=torch.arange(length, device=device) + 64,
page_size=64,
slot_mapping=torch.arange(length, device=device) + main_page,
page_size=main_page,
)
q = torch.randn((rows, 16, 512), device=device, dtype=torch.bfloat16)
heads = layer.indexer.heads
Expand Down Expand Up @@ -955,17 +1016,18 @@ def oracle(live):
del graph, resources


@pytest.mark.parametrize("swa_page", [32, 64, 128])
@torch.inference_mode()
def test_ced_replay_window_high_page_stride_and_invalid_rows(
native_workspace, monkeypatch
native_workspace, monkeypatch, swa_page
):
from b12x.attention._shared.mla.compressed_reference import (
unpack_deepseek_v41_cache_reference,
)

attention, _, _ = native_workspace
device = torch.device("cuda")
layer = _layer(attention, 20)
layer = _layer(attention, 20, swa_page=swa_page)
layer.is_ced_decoder = True
layer.compress_ratio = 0
layer.indexer = None
Expand All @@ -975,14 +1037,16 @@ def test_ced_replay_window_high_page_stride_and_invalid_rows(
storage = torch.empty((high_pid + 1, stride), device=device, dtype=torch.uint8)
cache = storage[
:,
: attention.mla.page_nbytes(32, cache_kind="swa", cache_format="deepseek_v41"),
: attention.mla.page_nbytes(
swa_page, cache_kind="swa", cache_format="deepseek_v41"
),
]
kv = torch.randn((1, 512), device=device, dtype=torch.bfloat16)
attention.mla.write_cache(
kv,
cache,
torch.tensor([high_pid * 32], device=device),
page_size=32,
torch.tensor([high_pid * swa_page], device=device),
page_size=swa_page,
cache_kind="swa",
cache_format="deepseek_v41",
)
Expand All @@ -992,7 +1056,7 @@ def test_ced_replay_window_high_page_stride_and_invalid_rows(
positions=positions,
req_id_per_token=torch.tensor([0, 0, -1], device=device, dtype=torch.int32),
block_table=torch.tensor(
[[-1, -1, -1, -1, high_pid]], device=device, dtype=torch.int32
[[-1] * (128 // swa_page) + [high_pid]], device=device, dtype=torch.int32
),
query_start_loc=torch.tensor([0, 1], device=device, dtype=torch.int32),
request_positions=positions[:1],
Expand All @@ -1012,7 +1076,9 @@ def test_ced_replay_window_high_page_stride_and_invalid_rows(
layer.attn_sink = torch.zeros(16, device=device)
layer.forward_mqa(q, None, positions, out)
value = unpack_deepseek_v41_cache_reference(
cache[high_pid : high_pid + 1].contiguous(), page_size=32, cache_kind="swa"
cache[high_pid : high_pid + 1].contiguous(),
page_size=swa_page,
cache_kind="swa",
)[0]
score = torch.einsum("hd,d->h", q[0].float(), value) * 512**-0.5
expected = score.sigmoid()[:, None] * value
Expand Down
Loading
Loading