Skip to content
Open
Show file tree
Hide file tree
Changes from 5 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)
180 changes: 180 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,180 @@
# 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",
}


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",
}
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()
9 changes: 9 additions & 0 deletions vllm_fl/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -152,6 +152,15 @@ def register_model():
register_quant_linear()
register_router()

# Replace only the DSV4 registry entry. The lazy thin model subclasses the
# vLLM 0.24 NVIDIA implementation and preserves its non-INT8 behavior.
from vllm import ModelRegistry

ModelRegistry.register_model(
"DeepseekV4ForCausalLM",
"vllm_fl.models.deepseek_v4:DeepseekV4FLForCausalLM",
)

# Register GLM-5 (GlmMoeDsa) — config not yet upstream
try:
from vllm.transformers_utils.config import _CONFIG_REGISTRY
Expand Down
1 change: 1 addition & 0 deletions vllm_fl/dispatch/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -525,6 +525,7 @@ Currently supported operators:
| Operator | Description | FlagGems | Reference | Vendor |
|----------|-------------|----------|-----------|--------|
| `dynamic_per_token_quant_int8` | vLLM-compatible symmetric dynamic per-token INT8 quantization | ✓ | ✓ | - |
| `deepseek_v4_inv_rope_quant_int8` | DSV4 inverse-RoPE with group-major INT8 activation quantization | ✓ | ✓ | ✓ |
| `silu_and_mul` | SiLU activation + element-wise multiplication | ✓ | ✓ | ✓ |
| `rms_norm` | RMS normalization | ✓ | ✓ | ✓ |
| `rotary_embedding` | Rotary position embedding | ✓ | ✓ | ✓ |
Expand Down
Loading
Loading