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
Original file line number Diff line number Diff line change
@@ -0,0 +1,159 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

import pytest
import torch
import torch_npu # noqa: F401

from vllm_ascend.ops.triton.glm5_next_kpool_state_compress import glm5_next_kpool_state_compress_and_write_cache_triton

HEAD_DIM = 128
BF16_RTOL = 1e-2
BF16_ATOL = 1e-2


@pytest.mark.parametrize("pool,capacity", [(4, 4), (4, 16), (8, 8)])
@pytest.mark.parametrize("use_graph", [False, True])
@torch.inference_mode()
def test_paged_state_long_prefill_padding_and_rollback(pool, capacity, use_graph):
generator = torch.Generator().manual_seed(13)
dim = HEAD_DIM
# Physical pages deliberately have padding and request IDs are reordered.
state_backing = torch.full((32, capacity + 2, 2 * dim), -7.0, device="npu")
cache_backing = torch.full((3, 18, 1, dim), -7.0, dtype=torch.bfloat16, device="npu")
state, cache = state_backing[:, :capacity], cache_backing[:, :16]
expected_state_backing, expected_cache_backing = state_backing.cpu(), cache_backing.cpu()
expected_state, expected_cache = expected_state_backing[:, :capacity], expected_cache_backing[:, :16]
ape = torch.randn(pool, dim, generator=generator) * 0.1
state_table = torch.tensor([[5, 9, 3, 12, 7, 16, 2, 20], [11, 8, 1, 18, 4, 22, 15, 23]], dtype=torch.int32)
cache_blocks = [1, 2]
history: list[dict[int, tuple[torch.Tensor, torch.Tensor]]] = [{}, {}]
graph, captured_args = None, None
graph_capacity = 32

def run(starts, lengths, invalidate=False):
nonlocal graph, captured_args
positions = torch.cat([torch.arange(s, s + n) for s, n in zip(starts, lengths)])
num_tokens = positions.numel()
keys = torch.randn(num_tokens, dim, generator=generator)
gates = torch.randn(num_tokens, dim, generator=generator) * 0.1
ends = torch.tensor(lengths, dtype=torch.int32).cumsum(0).to(torch.int32)
seq_lens = torch.tensor([s + n for s, n in zip(starts, lengths)], dtype=torch.int32)
state_slots, indexer_slots = [], []
cursor = 0
for req, (start, length) in enumerate(zip(starts, lengths)):
for local in range(length):
row, pos = cursor + local, start + local
history[req][pos] = (keys[row], gates[row])
state_slots.append(int(state_table[req, pos // capacity]) * capacity + pos % capacity)
slot = cache_blocks[req] * 16 + pos // pool if (pos + 1) % pool == 0 else -1
indexer_slots.append(slot)
if slot >= 0:
window = [history[req][p] for p in range(pos - pool + 1, pos + 1)]
pooled = (
torch.softmax(torch.stack([g for _, g in window]) + ape, dim=0)
* torch.stack([k for k, _ in window])
).sum(0)
expected_cache[cache_blocks[req], pos // pool, 0] = pooled.bfloat16()
cursor += length
if invalidate:
state_slots[0] = -1
state_slots[1] = state.shape[0] * capacity
indexer_slots[0] = cache.shape[0] * cache.shape[1]
# Each logical position has its own scheduler-provided physical slot.
cursor = 0
for req, (start, length) in enumerate(zip(starts, lengths)):
for local in range(length):
row, pos = cursor + local, start + local
if 0 <= state_slots[row] < state.shape[0] * capacity:
expected_state[state_slots[row] // capacity, pos % capacity] = torch.cat((keys[row], gates[row]))
cursor += length
# Graph-capacity rows have no owning request, even if a stale slot is positive.
padded_keys = torch.nn.functional.pad(keys, (0, 0, 0, graph_capacity - num_tokens)).npu()
padded_gates = torch.nn.functional.pad(gates, (0, 0, 0, graph_capacity - num_tokens)).npu()
args = (
state,
cache,
padded_keys,
padded_gates,
ape.npu(),
torch.cat((positions, torch.zeros(graph_capacity - num_tokens))).long().npu(),
ends.npu(),
seq_lens.npu(),
torch.tensor(state_slots + [0] + [-1] * (graph_capacity - num_tokens - 1), device="npu"),
state_table.npu(),
torch.tensor(indexer_slots + [0] + [-1] * (graph_capacity - num_tokens - 1), device="npu"),
pool,
)
if use_graph:
if graph is None:
captured_args = args
glm5_next_kpool_state_compress_and_write_cache_triton(*args)
torch.npu.synchronize()
graph = torch.npu.NPUGraph()
with torch.npu.graph(graph):
glm5_next_kpool_state_compress_and_write_cache_triton(*captured_args)
else:
for target, value in zip(captured_args[2:-1], args[2:-1]):
target.copy_(value)
graph.replay()
else:
glm5_next_kpool_state_compress_and_write_cache_triton(*args)
torch.testing.assert_close(state_backing.cpu(), expected_state_backing, rtol=0, atol=0)
torch.testing.assert_close(cache_backing.cpu(), expected_cache_backing, rtol=BF16_RTOL, atol=BF16_ATOL)

run([0, 0], [2, 3], invalidate=True)
run([0, 3], [2, 13]) # Rewrite invalid req0 slots; req1 reads history before a long prefill.
run([2, 16], [3, 3]) # Verification writes future candidates into separate paged slots.
run([3, 17], [1, 3]) # Reject candidates and overwrite their positions.
state_table[1, : 16 // capacity] = -1 # Evicted old pages must not be consulted.
run([4, 20], [pool, pool])


@pytest.mark.parametrize("empty", ["tokens", "requests"])
def test_empty_compression_preserves_caches(empty):
state = torch.ones(1, 4, 2 * HEAD_DIM, device="npu")
cache = torch.ones(1, 16, 1, HEAD_DIM, dtype=torch.bfloat16, device="npu")
count = 0 if empty == "tokens" else 1
keys = torch.zeros(count, HEAD_DIM, device="npu")
ends = torch.tensor([count] if empty == "tokens" else [], dtype=torch.int32, device="npu")
glm5_next_kpool_state_compress_and_write_cache_triton(
state,
cache,
keys,
keys,
torch.zeros(4, HEAD_DIM, device="npu"),
torch.zeros(count, dtype=torch.int64, device="npu"),
ends,
ends,
torch.full((count,), -1, device="npu"),
torch.zeros(1, 1, dtype=torch.int32, device="npu"),
torch.full((count,), -1, device="npu"),
4,
)
torch.testing.assert_close(state.cpu(), torch.ones_like(state, device="cpu"), rtol=0, atol=0)
torch.testing.assert_close(cache.cpu(), torch.ones_like(cache, device="cpu"), rtol=0, atol=0)


@pytest.mark.parametrize("block_size,pages", [(2, 1), (4, 0)])
def test_invalid_state_layout_raises(block_size, pages):
state = torch.zeros(1, block_size, 2 * HEAD_DIM, device="npu")
cache = torch.zeros(1, 16, 1, HEAD_DIM, dtype=torch.bfloat16, device="npu")
keys = torch.zeros(1, HEAD_DIM, device="npu")
ends = torch.ones(1, dtype=torch.int32, device="npu")
slots = torch.zeros(1, dtype=torch.int64, device="npu")
with pytest.raises(ValueError, match="nonempty state page table and block size >= pool size"):
glm5_next_kpool_state_compress_and_write_cache_triton(
state,
cache,
keys,
keys,
torch.zeros(4, HEAD_DIM, device="npu"),
slots,
ends,
ends,
slots,
torch.zeros(1, pages, dtype=torch.int32, device="npu"),
slots,
4,
)
Original file line number Diff line number Diff line change
@@ -0,0 +1,114 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

import pytest
import torch
import torch_npu # noqa: F401

from vllm_ascend.ops.triton import glm5_next_lightning_indexer as indexer

# GLM-5.3-Flash indexer parameters, including its actual top-k width.
HEAD_DIM = 128
NUM_HEADS = 32
POOL_SIZE = 4
INDEX_TOPK = 2048
CACHE_BLOCK_SIZE = 16
OUTPUT_WIDTH = INDEX_TOPK + POOL_SIZE - 1


def _assert_selection(result, query, cache, weights, ends, pool_lens, table, positions):
"""CPU reference: score each query head before applying the head weights."""
result, query, cache, weights = (x.cpu() for x in (result, query, cache, weights))
start = 0
for req, end in enumerate(ends.tolist()):
for row in range(start, end):
pos = int(positions[row])
count = min((pos + 1) // POOL_SIZE, int(pool_lens[req]))
ids = torch.arange(count)
keys = cache[table[req, ids // CACHE_BLOCK_SIZE].long(), ids % CACHE_BLOCK_SIZE, 0].float()
per_head_scores = query[row].float() @ keys.T
scores = (per_head_scores * weights[row].float()[:, None]).sum(0)
selected = torch.topk(scores, min(INDEX_TOPK // POOL_SIZE, count)).indices
history = (selected[:, None] * POOL_SIZE + torch.arange(POOL_SIZE)).flatten()
expected = torch.full((OUTPUT_WIDTH,), -1, dtype=torch.int32)
expected[: history.numel()] = history.to(torch.int32)
tail = torch.arange((pos + 1) // POOL_SIZE * POOL_SIZE, pos + 1, dtype=torch.int32)
expected[INDEX_TOPK : INDEX_TOPK + tail.numel()] = tail
actual = result[row, 0]
# Top-k ordering may differ for tied scores; token membership,
# multiplicity, fixed tail columns, and every padding lane must match.
torch.testing.assert_close(
actual[:INDEX_TOPK].sort().values,
expected[:INDEX_TOPK].sort().values,
rtol=0,
atol=0,
)
torch.testing.assert_close(actual[INDEX_TOPK:], expected[INDEX_TOPK:], rtol=0, atol=0)
start = end


@pytest.mark.parametrize("max_pool_seq_len", [0, 4, 512, 2050])
@pytest.mark.parametrize("use_graph", [False, True])
@torch.inference_mode()
def test_pool_selection_real_shape_paging_and_causal_tail(max_pool_seq_len, use_graph, monkeypatch):
generator = torch.Generator().manual_seed(19)
# Three requests exercise non-power-of-two bucketization. Two extra rows
# model graph padding; the caller is responsible for ignoring their output.
ends = torch.tensor([3, 5, 8], dtype=torch.int32)
num_tokens = 10
pages = (max_pool_seq_len + CACHE_BLOCK_SIZE - 1) // CACHE_BLOCK_SIZE
num_blocks = max(1, 3 * pages)
table = torch.randperm(num_blocks, generator=generator)[: 3 * pages].reshape(3, pages).to(torch.int32)
backing = torch.randn(num_blocks, CACHE_BLOCK_SIZE + 2, 1, HEAD_DIM, generator=generator).bfloat16()
cache = backing.npu()[:, :CACHE_BLOCK_SIZE]
query = torch.randn(num_tokens, NUM_HEADS, HEAD_DIM, generator=generator).bfloat16().npu()
weights = torch.randn(num_tokens, NUM_HEADS, generator=generator).bfloat16().npu()
pool_lens = torch.tensor([0, min(4, max_pool_seq_len), max_pool_seq_len], dtype=torch.int32)
last_pos = max(2, max_pool_seq_len * POOL_SIZE + 2)
positions = torch.tensor([0, 1, 2, 1, 2, last_pos - 2, last_pos - 1, last_pos, 0, 0])
device_lens, device_positions = pool_lens.npu(), positions.npu()
args = (query, cache, weights, ends.npu(), device_lens, table.npu(), device_positions)
kwargs = dict(index_topk=INDEX_TOPK, index_kpool=POOL_SIZE, max_pool_seq_len=max_pool_seq_len)
# Force several token chunks without a large scratch allocation. This also
# checks that request lookup uses the batch-global token offset.
monkeypatch.setattr(indexer, "TRITON_SCORES_CHUNK_BYTES", max(1, max_pool_seq_len) * 4 * 3)

if use_graph:
indexer.glm5_next_lightning_indexer_triton(*args, **kwargs)
torch.npu.synchronize()
graph = torch.npu.NPUGraph()
with torch.npu.graph(graph):
result = indexer.glm5_next_lightning_indexer_triton(*args, **kwargs)
for step in range(2):
if step:
query.copy_(torch.randn(num_tokens, NUM_HEADS, HEAD_DIM, generator=generator).bfloat16().npu())
weights.copy_(torch.randn(num_tokens, NUM_HEADS, generator=generator).bfloat16().npu())
pool_lens[2] //= 2
positions[3:5] = torch.tensor([15, 16])
device_lens.copy_(pool_lens)
device_positions.copy_(positions)
if use_graph:
graph.replay()
else:
result = indexer.glm5_next_lightning_indexer_triton(*args, **kwargs)
assert result.shape == (num_tokens, 1, OUTPUT_WIDTH)
assert result.dtype == torch.int32
_assert_selection(result, query, cache, weights, ends, pool_lens, table, positions)
torch.testing.assert_close(cache.cpu(), backing[:, :CACHE_BLOCK_SIZE], rtol=0, atol=0)


def test_empty_query():
result = indexer.glm5_next_lightning_indexer_triton(
torch.empty(0, NUM_HEADS, HEAD_DIM, dtype=torch.bfloat16, device="npu"),
torch.empty(1, CACHE_BLOCK_SIZE, 1, HEAD_DIM, dtype=torch.bfloat16, device="npu"),
torch.empty(0, NUM_HEADS, dtype=torch.bfloat16, device="npu"),
torch.empty(0, dtype=torch.int32, device="npu"),
torch.empty(0, dtype=torch.int32, device="npu"),
torch.empty(0, 0, dtype=torch.int32, device="npu"),
torch.empty(0, dtype=torch.int64, device="npu"),
index_topk=INDEX_TOPK,
index_kpool=POOL_SIZE,
max_pool_seq_len=0,
)
assert result.shape == (0, 1, OUTPUT_WIDTH)
assert result.dtype == torch.int32
57 changes: 57 additions & 0 deletions tests/ut/models/test_glm5next_indexer_backend.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,57 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

from types import SimpleNamespace

import torch
from torch import nn

from vllm_ascend.ops.mla import IndexerWrapper


class _FakeKPoolBackend(nn.Module):
def __init__(self, source, rope_dim):
super().__init__()
del rope_dim
self.k_cache = source.k_cache
self.head_dim = source.head_dim
self.enable_sparse_li_c8 = False
self.num_cache_tensors = 1

def process_weights_after_loading(self):
return None


class _FakeIndexer(nn.Module):
def __init__(self):
super().__init__()
self.n_head = 2
self.head_dim = 4
self.topk_tokens = 8
self.q_lora_rank = 3
self.wq_b = nn.Linear(3, 8, bias=False)
self.wk_weights_proj = nn.Linear(5, 6, bias=False)
self.k_norm = nn.LayerNorm(4)
self.softmax_scale = 0.5
self.index_kpool_compress_ape = nn.Parameter(torch.zeros(4, 4))
self.index_kpool_compress_gate = nn.Parameter(torch.zeros(4, 5))
self.k_cache = SimpleNamespace(prefix="indexer.k_cache")

def get_ascend_indexer_backend_cls(self):
return _FakeKPoolBackend


def test_wrapper_selects_model_backend_and_preserves_weight_names() -> None:
wrapper = IndexerWrapper(_FakeIndexer(), qk_rope_head_dim=0)

assert isinstance(wrapper.impl, _FakeKPoolBackend)
names = set(dict(wrapper.named_parameters()))
assert {
"wq_b.weight",
"wk_weights_proj.weight",
"k_norm.weight",
"k_norm.bias",
"index_kpool_compress_ape",
"index_kpool_compress_gate",
}.issubset(names)
assert not any(name.startswith("impl.") for name in names)
18 changes: 12 additions & 6 deletions vllm_ascend/ops/mla.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,10 +32,7 @@
from vllm.utils.torch_utils import direct_register_custom_op
from vllm.v1.attention.backend import AttentionMetadata # type: ignore

from vllm_ascend.attention.indexer import (
AscendSFAIndexerBackend,
AscendSFAIndexerMetadata,
)
from vllm_ascend.attention.indexer import AscendSFAIndexerBackend


class IndexerWrapper(nn.Module):
Expand All @@ -60,7 +57,16 @@ def __init__(self, vllm_indexer: nn.Module, qk_rope_head_dim: int) -> None:
self.wk_weights_proj = vllm_indexer.wk_weights_proj
self.k_norm = vllm_indexer.k_norm
self.softmax_scale = vllm_indexer.softmax_scale
self.impl = AscendSFAIndexerBackend(vllm_indexer, qk_rope_head_dim)
# Preserve checkpoint-visible direct Parameters for every indexer
# family. Registering them here keeps paths at ``...indexer.<name>``
# rather than adding an implementation segment.
if isinstance(vllm_indexer, nn.Module):
for name, parameter in vllm_indexer.named_parameters(recurse=False):
self.register_parameter(name, parameter)

backend_factory = getattr(type(vllm_indexer), "get_ascend_indexer_backend_cls", None)
backend_cls = backend_factory(vllm_indexer) if backend_factory is not None else AscendSFAIndexerBackend
self.impl = backend_cls(vllm_indexer, qk_rope_head_dim)

# Interface consumed by the SFA impl - delegated to the backend impl.
@property
Expand Down Expand Up @@ -89,7 +95,7 @@ def forward(
cos: torch.Tensor,
sin: torch.Tensor,
k_hidden_states: torch.Tensor,
indexer_metadata: AscendSFAIndexerMetadata,
indexer_metadata: AttentionMetadata,
compute_topk: bool = True,
) -> torch.Tensor | None:
return self.impl(hidden_states, q_c, cos, sin, k_hidden_states, indexer_metadata, compute_topk)
Expand Down
Loading
Loading