Skip to content
Open
Show file tree
Hide file tree
Changes from 7 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
39 changes: 39 additions & 0 deletions tests/unit_tests/dispatch/test_deepseek_v4_attention.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
# Copyright (c) 2026 BAAI. All rights reserved.

"""Tests for the DeepSeek-V4 attention compile boundary."""

from types import SimpleNamespace

import torch

from vllm_fl.models import deepseek_v4


def test_deepseek_v4_fl_attention_writes_preallocated_output(monkeypatch):
calls = []

class Layer:
def attention_impl(self, *args):
calls.append(args)
args[-1].fill_(7)

layer = Layer()
monkeypatch.setattr(
deepseek_v4,
"get_forward_context",
lambda: SimpleNamespace(no_compile_layers={"layer": layer}),
)

tensors = [torch.empty(1) for _ in range(7)]
out = torch.empty(2, 3, 4)
result = deepseek_v4._deepseek_v4_fl_attention(*tensors, out, "layer")

assert result is None
assert len(calls) == 1
assert calls[0] == (*tensors, out)
assert calls[0][-1] is out
assert torch.equal(out, torch.full_like(out, 7))
schema = torch._C._dispatch_find_schema_or_throw(
"vllm::deepseek_v4_fl_attention", ""
).schema()
assert "Tensor(a7!) out" in str(schema)
194 changes: 194 additions & 0 deletions tests/unit_tests/dispatch/test_deepseek_v4_ops.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,194 @@
# Copyright (c) 2026 BAAI. All rights reserved.

"""Tests for DeepSeek-V4 operator dispatch."""

from unittest.mock import Mock

import torch

from vllm_fl.dispatch.backends.reference.impl.deepseek_v4 import (
deepseek_v4_hc_head_torch,
deepseek_v4_int8_scaled_mm_torch,
deepseek_v4_inv_rope_quant_int8_torch,
deepseek_v4_mhc_post_torch,
)
from vllm_fl.dispatch.types import BackendImplKind
from vllm_fl.ops import deepseek_v4_int8_woa

DSV4_OPS = {
"deepseek_v4_inv_rope_quant_int8",
"deepseek_v4_inv_rope_quant_fp8",
"deepseek_v4_int8_scaled_mm",
"deepseek_v4_mhc_pre",
"deepseek_v4_mhc_fused_post_pre",
"deepseek_v4_mhc_post",
"deepseek_v4_hc_head",
"deepseek_v4_fused_q_kv_rmsnorm",
"deepseek_v4_qnorm_rope_kv_quant_insert",
"deepseek_v4_qnorm_rope_kv_bf16_insert",
"deepseek_v4_qnorm_rope_kv_fp8_insert",
"deepseek_v4_compute_global_topk_indices_and_lens",
"deepseek_v4_flash_mla_with_kvcache",
"deepseek_v4_dequantize_and_gather_k_cache",
"deepseek_v4_combine_topk_swa_indices",
"deepseek_v4_flash_mla_sparse_fwd",
"deepseek_v4_fused_indexer_q_rope_quant",
"deepseek_v4_fused_indexer_q_rope_quant_int8",
"deepseek_v4_compress_int8_indexer_k_cache",
"deepseek_v4_int8_mqa_logits",
"deepseek_v4_int8_paged_mqa_logits",
}


def test_reference_inv_rope_quant_int8():
o = torch.tensor(
[[[1, 2, 3, 4], [-1, -2, 5, 6]]],
dtype=torch.bfloat16,
)
positions = torch.tensor([0], dtype=torch.int32)
cos_sin_cache = torch.tensor([[0, 1]], dtype=torch.float32)

quantized, scales = deepseek_v4_inv_rope_quant_int8_torch(
o,
positions,
cos_sin_cache,
n_groups=1,
heads_per_group=2,
nope_dim=2,
rope_dim=2,
)

expected = torch.tensor(
[[[21, 42, 85, -64, -21, -42, 127, -106]]],
dtype=torch.int8,
)
assert torch.equal(quantized, expected)
torch.testing.assert_close(
scales,
torch.tensor([[[6 / 127]]], dtype=torch.float32),
)


def test_frontend_dispatches_through_cached_op(monkeypatch):
expected = (Mock(), Mock())
dispatch = Mock(return_value=expected)
monkeypatch.setattr(
deepseek_v4_int8_woa,
"_dispatch_inv_rope_quant_int8",
dispatch,
)
args = (
Mock(),
Mock(),
Mock(),
2,
4,
64,
64,
)

actual = deepseek_v4_int8_woa.fused_inv_rope_quant_int8(*args)

assert actual is expected
dispatch.assert_called_once_with(*args)


def test_all_backends_register_deepseek_v4_op(monkeypatch):
from vllm_fl.dispatch.backends.flaggems import register_ops as flaggems_ops
from vllm_fl.dispatch.backends.reference import register_ops as reference_ops
from vllm_fl.dispatch.backends.vendor.cuda import register_ops as cuda_ops

registered = []

class Registry:
def register_many(self, impls):
registered.extend(impls)

monkeypatch.setattr(
flaggems_ops,
"use_flaggems_op",
lambda op_name: op_name == deepseek_v4_int8_woa.DSV4_INV_ROPE_QUANT_INT8_OP,
)
registry = Registry()
flaggems_ops.register_builtins(registry)
cuda_ops.register_builtins(registry)
reference_ops.register_builtins(registry)

implementations = [
impl
for impl in registered
if impl.op_name == deepseek_v4_int8_woa.DSV4_INV_ROPE_QUANT_INT8_OP
]
assert {impl.impl_id for impl in implementations} == {
"default.flagos",
"vendor.cuda",
"reference.torch",
}
assert {impl.kind for impl in implementations} == {
BackendImplKind.DEFAULT,
BackendImplKind.VENDOR,
BackendImplKind.REFERENCE,
}


def test_reference_scaled_mm_and_mhc_ops():
x_q = torch.tensor([[1, -2]], dtype=torch.int8)
weight = torch.tensor([[3, 4], [5, 6]], dtype=torch.int8)
actual = deepseek_v4_int8_scaled_mm_torch(
x_q,
weight,
torch.tensor([[0.5]]),
torch.tensor([0.25, 0.5]),
torch.float32,
)
torch.testing.assert_close(actual, torch.tensor([[-0.875, -2.0]]))

residual = torch.tensor([[[1, 2], [3, 4]]], dtype=torch.bfloat16)
layer = torch.tensor([[2, -1]], dtype=torch.bfloat16)
post = torch.tensor([[[0.5], [1.0]]], dtype=torch.float32)
comb = torch.eye(2, dtype=torch.float32).unsqueeze(0)
torch.testing.assert_close(
deepseek_v4_mhc_post_torch(layer, residual, post, comb),
torch.tensor([[[2, 1.5], [5, 3]]], dtype=torch.bfloat16),
)

fn = torch.zeros((2, 4), dtype=torch.float32)
head = deepseek_v4_hc_head_torch(
residual,
fn,
torch.ones(1),
torch.zeros(2),
1e-6,
0.0,
)
torch.testing.assert_close(head, residual.float().mean(dim=1).to(torch.bfloat16))


def test_all_backends_register_all_deepseek_v4_ops(monkeypatch):
from vllm_fl.dispatch.backends.flaggems import register_ops as flaggems_ops
from vllm_fl.dispatch.backends.reference import register_ops as reference_ops
from vllm_fl.dispatch.backends.vendor.cuda import register_ops as cuda_ops

registered = []

class Registry:
def register_many(self, impls):
registered.extend(impls)

monkeypatch.setattr(
flaggems_ops,
"use_flaggems_op",
lambda op_name: op_name in DSV4_OPS,
)
registry = Registry()
flaggems_ops.register_builtins(registry)
cuda_ops.register_builtins(registry)
reference_ops.register_builtins(registry)

for op_name in DSV4_OPS:
implementations = [impl for impl in registered if impl.op_name == op_name]
assert {impl.impl_id for impl in implementations} == {
"default.flagos",
"vendor.cuda",
"reference.torch",
}
109 changes: 109 additions & 0 deletions tests/unit_tests/ops/test_deepseek_v4_int8_indexer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,109 @@
# Copyright (c) 2026 BAAI. All rights reserved.

"""CUDA correctness tests for the DeepSeek-V4 INT8 indexer kernels."""

import pytest
import torch

from vllm_fl.ops.deepseek_v4_int8_indexer import (
int8_mqa_logits,
int8_paged_mqa_logits,
)

pytestmark = pytest.mark.skipif(
not torch.cuda.is_available(), reason="requires CUDA"
)


def test_int8_mqa_logits_matches_torch():
torch.manual_seed(2)
num_queries, num_keys, num_heads, head_dim = 1, 97, 64, 128
q = torch.randint(
-127,
128,
(num_queries, num_heads, head_dim),
dtype=torch.int8,
)
k = torch.randint(-127, 128, (num_keys, head_dim), dtype=torch.int8)
k_scale = torch.rand(num_keys, dtype=torch.float32) * 0.02
weights = torch.randn(num_queries, num_heads, dtype=torch.float32)
cu_ks = torch.tensor([0], dtype=torch.int32)
cu_ke = torch.tensor([num_keys], dtype=torch.int32)

actual = int8_mqa_logits(
q.cuda(),
k.cuda(),
k_scale.cuda(),
weights.cuda(),
cu_ks.cuda(),
cu_ke.cuda(),
).cpu()
dots = torch.einsum("mhd,nd->mhn", q.float(), k.float())
expected = (dots * k_scale[None, None, :]).relu()
expected = (expected * weights[:, :, None]).sum(dim=1)

torch.testing.assert_close(actual, expected, atol=2e-3, rtol=2e-3)


def test_int8_paged_mqa_logits_matches_torch():
torch.manual_seed(3)
batch, next_n, num_heads, head_dim = 1, 1, 64, 128
block_size, num_blocks, context_len = 64, 2, 100
q = torch.randint(
-127,
128,
(batch, next_n, num_heads, head_dim),
dtype=torch.int8,
)
k = torch.randint(
-127,
128,
(num_blocks, block_size, head_dim),
dtype=torch.int8,
)
k_scale = torch.rand(num_blocks, block_size, dtype=torch.float32) * 0.02
weights = torch.randn(batch * next_n, num_heads, dtype=torch.float32)

# The compressor stores one packed page as all INT8 K bytes followed by
# all fp32 scales. The logical tensor shape only reserves 132 bytes/token;
# its final dimension must not be interpreted as an interleaved layout.
cache = torch.empty(
num_blocks,
block_size,
head_dim + torch.tensor([], dtype=torch.float32).element_size(),
dtype=torch.uint8,
)
flat_cache = cache.view(-1)
for block in range(num_blocks):
page_base = block * cache.stride(0)
k_bytes = k[block].contiguous().view(torch.uint8).reshape(-1)
scale_bytes = (
k_scale[block].contiguous().view(torch.uint8).reshape(-1)
)
flat_cache[page_base : page_base + k_bytes.numel()].copy_(k_bytes)
scale_start = page_base + k_bytes.numel()
flat_cache[scale_start : scale_start + scale_bytes.numel()].copy_(
scale_bytes
)

context_lens = torch.tensor([[context_len]], dtype=torch.int32)
block_table = torch.tensor([[0, 1]], dtype=torch.int32)
actual = int8_paged_mqa_logits(
q.cuda(),
cache.cuda(),
weights.cuda(),
context_lens.cuda(),
block_table.cuda(),
num_blocks * block_size,
).cpu()

flat_k = k.reshape(-1, head_dim)[:context_len]
flat_scale = k_scale.reshape(-1)[:context_len]
dots = torch.einsum("hd,nd->hn", q[0, 0].float(), flat_k.float())
expected = (dots * flat_scale[None, :]).relu()
expected = (expected * weights[0, :, None]).sum(dim=0)

torch.testing.assert_close(
actual[0, :context_len], expected, atol=2e-3, rtol=2e-3
)
assert torch.isfinite(actual[0, :context_len]).all()
48 changes: 48 additions & 0 deletions tests/unit_tests/quantization/test_w8a8_linear.py
Original file line number Diff line number Diff line change
Expand Up @@ -254,3 +254,51 @@ def fake_scaled_mm(

assert actual.shape == (1, 2, 3)
assert torch.equal(actual, expected)


def test_w8a8_grouped_linear_prepares_group_major_weights():
checkpoint_weight = torch.arange(32, dtype=torch.int8).reshape(8, 4)
checkpoint_scale = torch.arange(1, 9, dtype=torch.float32).reshape(8, 1)
layer = torch.nn.Module()
layer.is_bmm = True
layer.bmm_batch_size = 2
layer.register_parameter(
"weight",
torch.nn.Parameter(checkpoint_weight.clone(), requires_grad=False),
)
layer.register_parameter(
"weight_scale",
torch.nn.Parameter(checkpoint_scale.clone(), requires_grad=False),
)
layer.register_parameter("input_scale", None)
layer.register_parameter("input_zero_point", None)
layer.register_parameter("azp_adj", None)

kernel = object.__new__(linear.FLW8A8DynamicLinearKernel)
kernel.layer_param_names = [
"weight",
"weight_scale",
"input_scale",
"input_zero_point",
"azp_adj",
]
kernel.process_weights_after_loading(layer)

grouped_weight = layer._fl_w8a8_grouped_weight
assert grouped_weight.shape == (2, 4, 4)
assert grouped_weight.is_contiguous()
assert torch.equal(
grouped_weight[0],
checkpoint_weight[:4].contiguous(),
)
assert torch.equal(
grouped_weight[1],
checkpoint_weight[4:].contiguous(),
)
assert grouped_weight[0].transpose(0, 1).stride() == (1, 4)
assert torch.equal(
layer._fl_w8a8_grouped_weight_scale,
checkpoint_scale.reshape(2, 4),
)
assert "_fl_w8a8_grouped_weight" not in layer.state_dict()
assert "_fl_w8a8_grouped_weight_scale" not in layer.state_dict()
Loading
Loading