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
173 changes: 172 additions & 1 deletion tests/models/test_glm5next_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -252,7 +252,7 @@ def test_glm5next_kda_splits_mixed_decode_prefill_batch(monkeypatch) -> None:
layer.gate_lower_bound = -5.0
layer.A_log = torch.ones(1)
layer.dt_bias = torch.ones(1)
layer._b12x_kda_binding = None
layer._b12x_kda_plan = None
layer.conv1d = SimpleNamespace(
weight=torch.ones(3, 1, 3),
bias=torch.zeros(3),
Expand Down Expand Up @@ -339,6 +339,61 @@ def test_glm5next_alone_opts_into_b12x_kda_decode() -> None:
assert Glm5NextLinearAttention.b12x_kda_null_state_index == 0


@pytest.mark.parametrize(
("backend", "speculative", "uses_b12x"),
[
("auto", False, False),
("auto", True, True),
("b12x", False, True),
("b12x", True, True),
("triton", False, False),
("triton", True, False),
],
)
def test_glm5next_selects_configured_kda_decode_backend(
monkeypatch,
backend: str,
speculative: bool,
uses_b12x: bool,
) -> None:
monkeypatch.setattr(
KimiGatedDeltaNetAttention,
"__init__",
lambda self, config, vllm_config, prefix: None,
)
vllm_config = SimpleNamespace(
additional_config={"glm53_kda_decode_backend": backend},
speculative_config=object() if speculative else None,
)

layer = Glm5NextLinearAttention(object(), vllm_config)

assert layer.enable_b12x_kda_decode is uses_b12x


def test_glm5next_defaults_to_hybrid_kda_decode(monkeypatch) -> None:
monkeypatch.setattr(
KimiGatedDeltaNetAttention,
"__init__",
lambda self, config, vllm_config, prefix: None,
)
vllm_config = SimpleNamespace(additional_config={}, speculative_config=None)

layer = Glm5NextLinearAttention(object(), vllm_config)

assert layer._glm53_kda_decode_backend == "auto"
assert not layer.enable_b12x_kda_decode


def test_glm5next_rejects_unknown_kda_decode_backend() -> None:
vllm_config = SimpleNamespace(
additional_config={"glm53_kda_decode_backend": "unknown"}
)

with pytest.raises(ValueError, match="KDA decode backend"):
Glm5NextLinearAttention(object(), vllm_config)


def test_glm5next_b12x_mhc_builds_first_layer_broadcast_fn() -> None:
hidden_size = 4
hc_mult = 4
Expand Down Expand Up @@ -500,6 +555,122 @@ def plan(caps):
assert captured_caps["null_state_index"] == 0


def test_glm5next_b12x_kda_binds_spec_tensors_without_staging(monkeypatch) -> None:
scratch = torch.empty(32, dtype=torch.uint8)
bindings: list[dict[str, object]] = []
runs: list[object] = []

class FakePlan:
@staticmethod
def shapes_and_dtypes():
return (((32,), torch.uint8),)

class FakeApi:
@staticmethod
def bind_kda(plan, **kwargs):
binding = {"plan": plan, **kwargs}
bindings.append(binding)
return binding

@staticmethod
def run_kda(binding, **kwargs):
runs.append(binding)

workspace = SimpleNamespace(get_simultaneous=lambda *specs: (scratch,))
monkeypatch.setattr(
kimi_gdn_linear_attn,
"current_workspace_manager",
lambda: workspace,
)

layer = Glm5NextLinearAttention.__new__(Glm5NextLinearAttention)
torch.nn.Module.__init__(layer)
layer._b12x_kda_api = FakeApi()
layer._b12x_kda_plan = FakePlan()
layer._b12x_kda_max_tokens = 4
layer._b12x_kda_max_seqs = 2
layer._b12x_kda_state_index_columns = 2
layer.local_num_heads = 1
layer.head_dim = 2
layer.gate_lower_bound = -5.0
layer.A_log = torch.ones(1)
layer.dt_bias = torch.ones(2)
layer.o_norm = SimpleNamespace(weight=torch.ones(2), eps=1e-6)
layer.kv_cache = (torch.empty(0), torch.zeros(3, 1, 2, 2))
layer._b12x_kda_mixed_qkv = torch.empty(4, 6)
layer._b12x_kda_raw_g = torch.empty(4, 1, 2)
layer._b12x_kda_raw_beta = torch.empty(4, 1)
layer._b12x_kda_z = torch.empty(4, 1, 2)
layer._b12x_kda_output = torch.zeros(4, 1, 2)
layer._b12x_kda_query_start_loc = torch.zeros(3, dtype=torch.int32)
layer._b12x_kda_num_accepted_tokens = torch.ones(2, dtype=torch.int32)
layer._b12x_kda_state_indices = torch.zeros(2, 2, dtype=torch.int32)
layer._b12x_kda_num_seqs = torch.zeros(1, dtype=torch.int32)
layer._b12x_kda_num_tokens = torch.zeros(1, dtype=torch.int32)

kwargs = dict(
mixed_qkv=torch.ones(1, 6),
raw_g=torch.ones(1, 1, 2),
raw_beta=torch.ones(1, 1),
z=torch.ones(1, 1, 2),
output=torch.empty(1, 1, 2),
state_indices=torch.tensor([[1]], dtype=torch.int32),
query_start_loc=torch.tensor([0, 1], dtype=torch.int32),
num_accepted_tokens=None,
num_requests=1,
)
layer._run_b12x_kda_decode_post_conv(**kwargs)
spec_kwargs = dict(
mixed_qkv=torch.ones(4, 6),
raw_g=torch.ones(4, 1, 2),
raw_beta=torch.ones(4, 1),
z=torch.ones(4, 1, 2),
output=torch.empty(4, 1, 2),
state_indices=torch.tensor([[1, 2], [2, 1]], dtype=torch.int32),
query_start_loc=torch.tensor([0, 2, 4], dtype=torch.int32),
num_accepted_tokens=torch.tensor([1, 2], dtype=torch.int32),
num_requests=2,
)
layer._run_b12x_kda_decode_post_conv(**spec_kwargs)

assert len(bindings) == 2
assert bindings[0] is not bindings[1]
assert bindings[0]["scratch"] is scratch
assert bindings[1]["scratch"] is scratch
assert bindings[0]["mixed_qkv"] is layer._b12x_kda_mixed_qkv
assert bindings[0]["raw_g"] is layer._b12x_kda_raw_g
assert bindings[0]["raw_beta"] is layer._b12x_kda_raw_beta
assert bindings[0]["z"] is layer._b12x_kda_z
assert bindings[0]["output"] is layer._b12x_kda_output
assert bindings[0]["query_start_loc"] is layer._b12x_kda_query_start_loc
assert bindings[0]["state_indices"] is layer._b12x_kda_state_indices
assert bindings[0]["num_tokens"] is layer._b12x_kda_num_tokens
torch.testing.assert_close(bindings[0]["mixed_qkv"][0], kwargs["mixed_qkv"][0])
assert bindings[1]["mixed_qkv"] is spec_kwargs["mixed_qkv"]
assert (
bindings[1]["num_accepted_tokens"].data_ptr()
== spec_kwargs["num_accepted_tokens"].data_ptr()
)
assert (
bindings[1]["query_start_loc"].data_ptr()
== spec_kwargs["query_start_loc"].data_ptr()
)
assert (
bindings[1]["state_indices"].data_ptr()
== spec_kwargs["state_indices"].data_ptr()
)
assert bindings[1]["num_tokens"] is layer._b12x_kda_num_tokens
assert bindings[1]["num_tokens"].item() == spec_kwargs["mixed_qkv"].shape[0]
assert (
bindings[1]["num_tokens"].data_ptr()
!= spec_kwargs["query_start_loc"][2:].data_ptr()
)
assert bindings[1]["output"].data_ptr() == spec_kwargs["output"].data_ptr()
assert len(runs) == 2
assert runs[0] is bindings[0]
assert runs[1] is bindings[1]


def test_glm5next_sparse_mla_selects_b12x_backend(monkeypatch) -> None:
captured: dict[str, object] = {}

Expand Down
99 changes: 67 additions & 32 deletions vllm/model_executor/layers/mamba/gdn/kimi_gdn_linear_attn.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
from vllm.transformers_utils.configs.kimi_linear import KimiLinearConfig
from vllm.utils.b12x import get_b12x_gdn_decode
from vllm.v1.attention.backends.gdn_attn import GDNAttentionMetadata
from vllm.v1.worker.workspace import current_workspace_manager

from ...linear import (
ColumnParallelLinear,
Expand Down Expand Up @@ -308,7 +309,6 @@ def __init__(
self.o_norm = FusedRMSNormGated(self.head_dim, activation="sigmoid")
self._b12x_kda_api: Any | None = None
self._b12x_kda_plan = None
self._b12x_kda_binding = None
self._initialize_b12x_kda_decode(vllm_config)
self.o_proj = RowParallelLinear(
self.projection_size,
Expand Down Expand Up @@ -412,12 +412,6 @@ def _initialize_b12x_kda_decode(self, vllm_config: VllmConfig) -> None:
torch.zeros(1, dtype=torch.int32, device=device),
persistent=False,
)
scratch_spec = provisional.scratch_specs()[0]
self.register_buffer(
"_b12x_kda_scratch",
torch.empty(scratch_spec.shape, dtype=scratch_spec.dtype, device=device),
persistent=False,
)

def _make_b12x_kda_plan(self, max_state_slots: int):
api = self._b12x_kda_api
Expand Down Expand Up @@ -450,27 +444,8 @@ def bind_kv_cache(self, kv_cache: torch.Tensor) -> None:
recurrent_state = self.kv_cache[1]
plan = self._make_b12x_kda_plan(max_state_slots=recurrent_state.shape[0])
self._b12x_kda_plan = plan
self._b12x_kda_binding = api.bind_kda(
plan,
scratch=self._b12x_kda_scratch,
mixed_qkv=self._b12x_kda_mixed_qkv,
raw_g=self._b12x_kda_raw_g,
raw_beta=self._b12x_kda_raw_beta,
z=self._b12x_kda_z,
A_log=self.A_log,
dt_bias=self.dt_bias.view(self.local_num_heads, self.head_dim),
norm_weight=self.o_norm.weight,
recurrent_state=recurrent_state,
query_start_loc=self._b12x_kda_query_start_loc,
num_accepted_tokens=self._b12x_kda_num_accepted_tokens,
state_indices=self._b12x_kda_state_indices,
num_seqs=self._b12x_kda_num_seqs,
num_tokens=self._b12x_kda_num_tokens,
output=self._b12x_kda_output,
)

def unbind_kv_cache(self) -> None:
self._b12x_kda_binding = None
self._b12x_kda_plan = None
super().unbind_kv_cache()

Expand All @@ -486,7 +461,7 @@ def rearrange_mixed_qkv(

def _can_use_b12x_kda_decode(self, m: GDNAttentionMetadata) -> bool:
if (
self._b12x_kda_binding is None
self._b12x_kda_plan is None
or m.num_prefills != 0
or (m.num_decodes == 0 and m.num_spec_decodes == 0)
):
Expand Down Expand Up @@ -518,16 +493,17 @@ def _run_b12x_kda_decode_post_conv(
num_accepted_tokens: torch.Tensor | None,
num_requests: int,
) -> None:
binding = self._b12x_kda_binding
plan = self._b12x_kda_plan
api = self._b12x_kda_api
if binding is None or api is None:
if plan is None or api is None:
raise RuntimeError("b12x KDA KV cache was not bound before inference")
num_tokens = int(mixed_qkv.shape[0])
state_columns = int(state_indices.shape[1])
if (
num_tokens > self._b12x_kda_max_tokens
or num_requests > self._b12x_kda_max_seqs
or state_columns > self._b12x_kda_state_index_columns
or num_tokens > num_requests * state_columns
):
raise ValueError(
"b12x KDA capacity exceeded: "
Expand All @@ -537,6 +513,49 @@ def _run_b12x_kda_decode_post_conv(
f"{self._b12x_kda_state_index_columns}"
)

scratch_buffers = current_workspace_manager().get_simultaneous(
*plan.shapes_and_dtypes()
)
if not scratch_buffers:
raise RuntimeError("b12x KDA plan did not expose caller scratch")
scratch: torch.Tensor | tuple[torch.Tensor, ...]
scratch = (
scratch_buffers[0] if len(scratch_buffers) == 1 else tuple(scratch_buffers)
)

if (
num_accepted_tokens is not None
and state_columns == self._b12x_kda_state_index_columns
):
self._b12x_kda_num_seqs.fill_(num_requests)
self._b12x_kda_num_tokens.fill_(num_tokens)
query_start_loc = query_start_loc[: num_requests + 1]
binding = api.bind_kda(
plan,
scratch=scratch,
mixed_qkv=mixed_qkv,
raw_g=raw_g,
raw_beta=raw_beta,
z=z,
A_log=self.A_log,
dt_bias=self.dt_bias.view(self.local_num_heads, self.head_dim),
norm_weight=self.o_norm.weight,
recurrent_state=self.kv_cache[1],
query_start_loc=query_start_loc,
num_accepted_tokens=num_accepted_tokens[:num_requests],
state_indices=state_indices[:num_requests, :state_columns],
num_seqs=self._b12x_kda_num_seqs,
num_tokens=self._b12x_kda_num_tokens,
output=output[:num_tokens],
)
api.run_kda(
binding,
lower_bound=self.gate_lower_bound,
eps=self.o_norm.eps,
scale=self.head_dim**-0.5,
)
return

self._b12x_kda_mixed_qkv[:num_tokens].copy_(mixed_qkv)
self._b12x_kda_raw_g[:num_tokens].copy_(raw_g)
self._b12x_kda_raw_beta[:num_tokens].copy_(raw_beta)
Expand All @@ -555,10 +574,26 @@ def _run_b12x_kda_decode_post_conv(
state_indices[:num_requests]
)
self._b12x_kda_num_seqs.fill_(num_requests)
self._b12x_kda_num_tokens.copy_(
query_start_loc[num_requests : num_requests + 1]
)
self._b12x_kda_num_tokens.fill_(num_tokens)

binding = api.bind_kda(
plan,
scratch=scratch,
mixed_qkv=self._b12x_kda_mixed_qkv,
raw_g=self._b12x_kda_raw_g,
raw_beta=self._b12x_kda_raw_beta,
z=self._b12x_kda_z,
A_log=self.A_log,
dt_bias=self.dt_bias.view(self.local_num_heads, self.head_dim),
norm_weight=self.o_norm.weight,
recurrent_state=self.kv_cache[1],
query_start_loc=self._b12x_kda_query_start_loc,
num_accepted_tokens=self._b12x_kda_num_accepted_tokens,
state_indices=self._b12x_kda_state_indices,
num_seqs=self._b12x_kda_num_seqs,
num_tokens=self._b12x_kda_num_tokens,
output=self._b12x_kda_output,
)
api.run_kda(
binding,
lower_bound=self.gate_lower_bound,
Expand Down
Loading
Loading