init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -0,0 +1,102 @@
# SPDX-License-Identifier: Apache-2.0
# Compare vllm_ascend.sample.penalties.apply_all_penalties (Triton-Ascend) with
# vllm.v1.sample.ops.penalties.apply_all_penalties (PyTorch via model_executor).
# Requires NPU and Triton-Ascend.
import gc
import pytest
import torch
from vllm.v1.sample.ops.penalties import apply_all_penalties as v1_apply_all_penalties
from vllm_ascend.sample.penalties import apply_all_penalties as ascend_apply_all_penalties
# Same scenario grid as test_apply_penalties_model_executor (equivalence + boundaries).
APPLY_PENALTY_CASES = [
pytest.param(0, 0, "mixed", id="empty-both"),
pytest.param(0, 16, "mixed", id="empty-prompt"),
pytest.param(32, 0, "mixed", id="empty-output"),
pytest.param(1, 1, "mixed", id="single-token-each"),
pytest.param(32, 16, "mixed", id="typical-small"),
pytest.param(128, 64, "mixed", id="typical-large"),
pytest.param(128, 64, "all_padding", id="all-padding"),
]
def _make_tokens(
num_seqs: int,
seq_len: int,
vocab_size: int,
mode: str,
device: str,
) -> torch.Tensor:
if mode == "all_padding":
return torch.full((num_seqs, seq_len), vocab_size, device=device, dtype=torch.int64)
if seq_len == 0:
return torch.empty((num_seqs, 0), device=device, dtype=torch.int64)
tokens = torch.randint(0, vocab_size, (num_seqs, seq_len), device=device, dtype=torch.int64)
pad_mask = torch.rand(num_seqs, seq_len, device=device) > 0.7
tokens[pad_mask] = vocab_size
return tokens
@pytest.mark.skip("Probabilistic failure, need zengtian after fix")
@pytest.mark.parametrize("num_seqs", [1, 8, 32, 128])
@pytest.mark.parametrize("vocab_size", [5120, 151936])
@pytest.mark.parametrize(
"max_prompt_len,max_output_len,token_mode",
APPLY_PENALTY_CASES,
)
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
@torch.inference_mode()
def test_apply_all_penalties_v1_vs_ascend(
num_seqs,
vocab_size,
max_prompt_len,
max_output_len,
token_mode,
dtype,
device="npu",
seed=42,
):
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
init_device_properties_triton()
torch.manual_seed(seed)
logits_v1 = torch.randn(num_seqs, vocab_size, device=device, dtype=dtype)
logits_ascend = logits_v1.clone()
prompt_tokens = _make_tokens(num_seqs, max_prompt_len, vocab_size, token_mode, device)
output_tokens = _make_tokens(num_seqs, max_output_len, vocab_size, token_mode, device)
output_token_ids = [row.tolist() for row in output_tokens.cpu()]
presence_penalties = torch.rand(num_seqs, device=device, dtype=torch.float32) * 0.2
frequency_penalties = torch.rand(num_seqs, device=device, dtype=torch.float32) * 0.2
repetition_penalties = torch.rand(num_seqs, device=device, dtype=torch.float32) * 0.4 + 1.0
v1_apply_all_penalties(
logits_v1,
prompt_tokens,
presence_penalties,
frequency_penalties,
repetition_penalties,
output_token_ids,
)
ascend_apply_all_penalties(
logits_ascend,
prompt_tokens,
presence_penalties,
frequency_penalties,
repetition_penalties,
output_token_ids,
)
atol = 1e-2 if dtype == torch.bfloat16 else 1e-3
rtol = 1e-2 if dtype == torch.bfloat16 else 1e-3
assert torch.allclose(logits_ascend.float(), logits_v1.float(), atol=atol, rtol=rtol), (
f"Max diff: {(logits_ascend.float() - logits_v1.float()).abs().max().item()}"
)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,202 @@
# SPDX-License-Identifier: Apache-2.0
# Test vllm_ascend.worker.v2.sample.bad_words.apply_bad_words (Triton-Ascend).
# Requires NPU and Triton-Ascend.
import pytest
import torch
from vllm_ascend.worker.v2.sample.bad_words import apply_bad_words
# Test cases for different input shapes
BAD_WORDS_TEST_CASES = [
pytest.param(512, 50257, 16, 3, 2, id="small-case"),
pytest.param(1024, 50257, 32, 5, 3, id="medium-case"),
pytest.param(2048, 50257, 64, 8, 4, id="large-case"),
]
def create_test_data(num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length, device):
"""Create test data for testing"""
# Create logits
logits = torch.randn(num_tokens, vocab_size, dtype=torch.float32, device=device)
# Create expanded_idx_mapping (map each token to a request)
expanded_idx_mapping = torch.randint(0, num_requests, (num_tokens,), dtype=torch.int32, device=device)
# Create bad_word_token_ids and bad_word_offsets
MAX_BAD_WORDS_TOTAL_TOKENS = 1024
MAX_NUM_BAD_WORDS = 128
bad_word_token_ids = torch.zeros((num_requests, MAX_BAD_WORDS_TOTAL_TOKENS), dtype=torch.int32, device=device)
bad_word_offsets = torch.zeros((num_requests, MAX_NUM_BAD_WORDS + 1), dtype=torch.int32, device=device)
num_bad_words = torch.zeros(num_requests, dtype=torch.int32, device=device)
# Fill bad words data
for req_idx in range(num_requests):
offset = 0
actual_bad_words = 0
for bw_idx in range(num_bad_words_per_req):
# Check if adding this bad word would exceed the token limit
if offset + bad_word_length > MAX_BAD_WORDS_TOTAL_TOKENS:
break
# Create a bad word with specific tokens
bad_word = torch.tensor([100 + req_idx * 10 + bw_idx] * bad_word_length, dtype=torch.int32, device=device)
bad_word_token_ids[req_idx, offset : offset + bad_word_length] = bad_word
bad_word_offsets[req_idx, bw_idx] = offset
offset += bad_word_length
actual_bad_words += 1
bad_word_offsets[req_idx, actual_bad_words] = offset
num_bad_words[req_idx] = actual_bad_words
# Create all_token_ids with some matching bad words
max_seq_len = 1024
all_token_ids = torch.randint(0, vocab_size, (num_requests, max_seq_len), dtype=torch.int32, device=device)
# Create prompt_len and total_len
prompt_len = torch.tensor([50] * num_requests, dtype=torch.int32, device=device)
total_len = torch.tensor([max_seq_len] * num_requests, dtype=torch.int32, device=device)
# Create input_ids with the same bad words, so they can be detected
input_ids = torch.randint(0, vocab_size, (num_tokens,), dtype=torch.int32, device=device)
# For each token, set input_ids to match the bad word for its request
for token_idx in range(num_tokens):
req_idx = expanded_idx_mapping[token_idx].item()
if num_bad_words[req_idx] > 0:
# Set input_ids to match the first bad word
bad_word = bad_word_token_ids[req_idx, :bad_word_length]
# For each position in the bad word, set input_ids accordingly
for i in range(bad_word_length):
if token_idx - i >= 0:
input_ids[token_idx - i] = bad_word[bad_word_length - 1 - i]
# Create expanded_local_pos - set to bad_word_length - 1 so that effective_len = output_len + (bad_word_length - 1)
# This ensures that we're checking the current token as the end of a bad word
expanded_local_pos = torch.full((num_tokens,), bad_word_length - 1, dtype=torch.int32, device=device)
return (
logits,
expanded_idx_mapping,
bad_word_token_ids,
bad_word_offsets,
num_bad_words,
all_token_ids,
prompt_len,
total_len,
input_ids,
expanded_local_pos,
)
@pytest.mark.parametrize(
"num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length", BAD_WORDS_TEST_CASES
)
@torch.inference_mode()
def test_apply_bad_words_different_shapes(
num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length, device="npu"
):
"""Test apply_bad_words with different input shapes"""
test_data = create_test_data(num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length, device)
# Make a copy of logits to compare
logits_before = test_data[0].clone()
logits_after = test_data[0].clone()
# Apply bad words
apply_bad_words(logits_after, *test_data[1:], num_bad_words_per_req)
# Verify that logits were modified
assert not torch.allclose(logits_before, logits_after), "Logits should be modified when bad words are present"
print(f"Test passed: tokens={num_tokens}, requests={num_requests}")
@torch.inference_mode()
def test_apply_bad_words_no_bad_words(device="npu"):
"""Test apply_bad_words with no bad words"""
num_tokens = 1024
vocab_size = 50257
num_requests = 32
num_bad_words_per_req = 0
bad_word_length = 3
test_data = create_test_data(num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length, device)
# Make a copy of logits to compare
logits_before = test_data[0].clone()
logits_after = test_data[0].clone()
# Apply bad words
apply_bad_words(logits_after, *test_data[1:], num_bad_words_per_req)
# Verify that logits were not modified
assert torch.allclose(logits_before, logits_after), "Logits should not be modified when no bad words are present"
print("No bad words test passed")
@torch.inference_mode()
def test_apply_bad_words_edge_cases(device="npu"):
"""Test apply_bad_words with edge cases"""
# Test with maximum bad words
num_tokens = 1024
vocab_size = 50257
num_requests = 16
num_bad_words_per_req = 128 # Maximum allowed
bad_word_length = 2
print("\nTesting edge case: maximum bad words")
test_data = create_test_data(num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length, device)
# Make a copy of logits to compare
logits_before = test_data[0].clone()
logits_after = test_data[0].clone()
# Apply bad words
apply_bad_words(logits_after, *test_data[1:], num_bad_words_per_req)
# Verify that logits were modified
assert not torch.allclose(logits_before, logits_after), (
"Logits should be modified when maximum bad words are present"
)
print("Maximum bad words test passed")
@torch.inference_mode()
def test_apply_bad_words_token_limit(device="npu"):
"""Test apply_bad_words with token limit cases"""
num_tokens = 1024
vocab_size = 50257
num_requests = 16
# Test case 1: Total tokens within limit
print("\nTesting case: total tokens within limit")
num_bad_words_per_req = 32
bad_word_length = 32 # 32 * 32 = 1024 tokens (exactly at limit)
test_data = create_test_data(num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length, device)
# Make a copy of logits to compare
logits_before = test_data[0].clone()
logits_after = test_data[0].clone()
# Apply bad words
apply_bad_words(logits_after, *test_data[1:], num_bad_words_per_req)
# Verify that logits were modified
assert not torch.allclose(logits_before, logits_after), (
"Logits should be modified when total tokens are within limit"
)
print("Total tokens within limit test passed")
# Test case 2: Total tokens exceeding limit (this should still work but only process up to limit)
print("\nTesting case: total tokens exceeding limit")
num_bad_words_per_req = 33
bad_word_length = 32 # 33 * 32 = 1056 tokens (exceeding limit)
test_data = create_test_data(num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length, device)
# Make a copy of logits to compare
logits_before = test_data[0].clone()
logits_after = test_data[0].clone()
# Apply bad words
apply_bad_words(logits_after, *test_data[1:], num_bad_words_per_req)
# Verify that logits were modified (even though we exceed the limit)
assert not torch.allclose(logits_before, logits_after), "Logits should be modified when total tokens exceed limit"
print("Total tokens exceeding limit test passed")

View File

@@ -0,0 +1,35 @@
import pytest
import torch
from vllm_ascend.ops.triton.batch_memcpy import batch_memcpy_kernel
@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float32])
def test_batch_memcpy(dtype):
element_size = 2 if dtype == torch.bfloat16 else 4
device = "npu:0"
# this is a typical case when used in mamba states copy.
sizes = torch.tensor([24576, 262144, 24576, 262144], device=device, dtype=torch.int32)
src_tensors_list = []
src_addr_list = []
dst_tensors_list = []
dst_addr_list = []
for i in range(len(sizes)):
src_tensors_list.append(torch.rand(sizes[i].item() // element_size, dtype=dtype, device=device))
src_addr_list.append(src_tensors_list[-1].data_ptr())
dst_tensors_list.append(torch.empty(sizes[i].item() // element_size, dtype=dtype, device=device))
dst_addr_list.append(dst_tensors_list[-1].data_ptr())
src_addr_list = torch.tensor(src_addr_list, dtype=torch.int64, device=device)
dst_addr_list = torch.tensor(dst_addr_list, dtype=torch.int64, device=device)
batch = sizes.shape[0]
grid = (batch,)
# using larger block_size to accelerate copy.
BLOCK_SIZE = 8192
batch_memcpy_kernel[grid](src_addr_list, dst_addr_list, sizes, BLOCK_SIZE=BLOCK_SIZE)
for i in range(len(sizes)):
torch.testing.assert_close(src_tensors_list[i], dst_tensors_list[i], rtol=0, atol=0)

View File

@@ -0,0 +1,125 @@
import pytest
import torch
from vllm.triton_utils import triton
from vllm_ascend.worker.v2.sample.penalties import _bincount_kernel
def torch_bincount(
expanded_idx_mapping: torch.Tensor,
all_token_ids: torch.Tensor,
prompt_len: torch.Tensor,
prefill_len: torch.Tensor,
prompt_bin_mask: torch.Tensor,
output_bin_counts: torch.Tensor,
):
req_indices = expanded_idx_mapping
prompt_bin_mask[req_indices] = 0
output_bin_counts[req_indices] = 0
for token_idx in range(expanded_idx_mapping.shape[0]):
req_idx = expanded_idx_mapping[token_idx].item()
p_len = prompt_len[req_idx].item()
pref_len = prefill_len[req_idx].item()
tokens = all_token_ids[req_idx]
for pos in range(p_len):
token = tokens[pos].item()
bin_idx = token // 32
bit_idx = token % 32
prompt_bin_mask[req_idx, bin_idx] |= 1 << bit_idx
for pos in range(p_len, pref_len):
token = tokens[pos].item()
output_bin_counts[req_idx, token] += 1
@pytest.mark.skip(reason="atomic_or operator hangs in current npu_ir version")
def test_bincount_kernel():
"""
Compute the prompt binary mask and token bincount using the Triton kernel.
Args:
expanded_idx_mapping: Tensor containing the indices of requests to process.
all_token_ids: Batch of input token IDs for all requests.
prompt_len: Tensor storing the prompt length for each request.
prefill_len: Tensor storing the prefill length for each request.
prompt_bin_mask: Output binary mask tensor to mark prompt tokens.
output_bin_counts: Output tensor to store token frequency counts.
max_prefill_len: Maximum prefill length to limit kernel processing.
"""
torch.manual_seed(42)
expanded_idx_mapping = torch.tensor([63], dtype=torch.int32).npu()
all_token_ids = torch.randint(
low=0,
high=10,
size=(64, 40960),
dtype=torch.int32,
).npu()
prompt_len = torch.randint(
low=0,
high=10,
size=(64,),
dtype=torch.int32,
).npu()
prefill_len = torch.randint(
low=0,
high=10,
size=(64,),
dtype=torch.int32,
).npu()
prompt_bin_mask = torch.zeros(size=(64, 4748), dtype=torch.int32).npu()
output_bin_counts = torch.zeros(size=(64, 151936), dtype=torch.int32).npu()
ref_prompt_bin_mask = torch.zeros(size=(64, 4748), dtype=torch.int32).npu()
ref_output_bin_counts = torch.zeros(size=(64, 151936), dtype=torch.int32).npu()
max_prefill_len = 10
prompt_bin_mask[expanded_idx_mapping] = 0
output_bin_counts[expanded_idx_mapping] = 0
num_tokens = expanded_idx_mapping.shape[0]
BLOCK_SIZE = 1024
num_blocks = triton.cdiv(max_prefill_len, BLOCK_SIZE)
_bincount_kernel[(num_tokens, num_blocks)](
expanded_idx_mapping,
all_token_ids,
all_token_ids.stride(0),
prompt_len,
prefill_len,
prompt_bin_mask,
prompt_bin_mask.stride(0),
output_bin_counts,
output_bin_counts.stride(0),
BLOCK_SIZE=BLOCK_SIZE,
)
torch_bincount(
expanded_idx_mapping,
all_token_ids,
prompt_len,
prefill_len,
ref_prompt_bin_mask,
ref_output_bin_counts,
)
# ========== Verify results ==========
assert torch.equal(prompt_bin_mask, ref_prompt_bin_mask), (
f"prompt_bin_mask triton output differs from torch reference.\n"
f"Max diff: {torch.max(torch.abs(prompt_bin_mask - ref_prompt_bin_mask))}\n"
f"Mean diff: {torch.mean(torch.abs(prompt_bin_mask - ref_prompt_bin_mask))}"
)
assert torch.equal(output_bin_counts, ref_output_bin_counts), (
f"output_bin_counts triton output differs from torch reference.\n"
f"Max diff: {torch.max(torch.abs(output_bin_counts - ref_output_bin_counts))}\n"
f"Mean diff: {torch.mean(torch.abs(output_bin_counts - ref_output_bin_counts))}"
)

View File

@@ -0,0 +1,315 @@
import gc
import pytest
import torch
from vllm_ascend._310p.ops.causal_conv1d import causal_conv1d_fn as causal_conv1d_fn_ref
from vllm_ascend._310p.ops.causal_conv1d import causal_conv1d_update as causal_conv1d_update_ref
from vllm_ascend.ops.triton.mamba.causal_conv1d import PAD_SLOT_ID, causal_conv1d_fn
from vllm_ascend.ops.triton.mamba.causal_conv1d import causal_conv1d_update_npu as causal_conv1d_update
from vllm_ascend.utils import enable_custom_op
def validate_cmp(y_cal, y_ref, dtype, device="npu"):
y_cal = y_cal.to(device)
y_ref = y_ref.to(device)
if dtype == torch.float16:
torch.testing.assert_close(y_ref, y_cal, rtol=3e-03, atol=1e-02, equal_nan=True)
elif dtype == torch.bfloat16:
torch.testing.assert_close(y_ref, y_cal, rtol=1e-02, atol=1e-02, equal_nan=True)
elif dtype == torch.float32:
torch.testing.assert_close(y_ref, y_cal, rtol=1e-03, atol=4e-03, equal_nan=True)
elif (
dtype == torch.int32
or dtype == torch.int64
or dtype == torch.int16
or dtype == torch.int8
or dtype == torch.uint32
or dtype == torch.bool
):
assert torch.equal(y_cal, y_ref)
else:
raise ValueError('Invalid parameter "dtype" is found : {}'.format(dtype))
@pytest.mark.parametrize("has_initial_state", [False, True])
@pytest.mark.parametrize("itype", [torch.bfloat16])
@pytest.mark.parametrize("silu_activation", [True])
@pytest.mark.parametrize("has_bias", [True])
@pytest.mark.parametrize("seq_len", [[128, 1024, 2048, 4096]])
@pytest.mark.parametrize("extra_state_len", [0, 2])
@pytest.mark.parametrize("width", [4])
@pytest.mark.parametrize("dim", [2048])
def test_ascend_causal_conv1d(
dim, width, extra_state_len, seq_len, has_bias, silu_activation, itype, has_initial_state
):
torch.random.manual_seed(0)
enable_custom_op()
device = "npu"
cu_seqlen, num_seq = sum(seq_len), len(seq_len)
state_len = width - 1 + extra_state_len
x = torch.randn(cu_seqlen, dim, device=device, dtype=itype).transpose(0, 1)
weight = torch.randn(dim, width, device=device, dtype=itype) #
query_start_loc = torch.cumsum(torch.tensor([0] + seq_len, device=device, dtype=torch.int32), dim=0).to(
dtype=torch.int32
)
cache_indices = torch.arange(num_seq, device=device, dtype=torch.int32)
has_initial_state_tensor = torch.tensor([has_initial_state] * num_seq, device=device, dtype=torch.bool)
activation = None if not silu_activation else "silu"
if has_initial_state:
conv_states = torch.randn((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2)
conv_states_ref = (
torch.randn((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2).copy_(conv_states)
)
else:
conv_states = torch.zeros((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2)
conv_states_ref = torch.zeros((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2)
if has_bias:
bias = torch.randn(dim, device=device, dtype=itype)
else:
bias = None
out_ref = causal_conv1d_fn_ref(
x,
weight,
bias=bias,
activation=activation,
conv_states=conv_states_ref,
has_initial_state=has_initial_state_tensor,
cache_indices=cache_indices,
query_start_loc=query_start_loc,
)
# out = causal_conv1d_fn(x,
# weight,
# bias=bias,
# activation=activation,
# conv_states=conv_states,
# has_initial_state=has_initial_state_tensor,
# cache_indices=cache_indices,
# query_start_loc=query_start_loc)
x_origin = x.transpose(-1, -2)
weight_origin = weight.transpose(-1, -2)
conv_states_origin = conv_states.transpose(-1, -2)
activation_num = 1 if activation else 0
out = torch.empty_like(x_origin)
torch.ops._C_ascend.npu_causal_conv1d_custom(
out,
x_origin,
weight_origin,
conv_state=conv_states_origin,
bias_opt=bias,
query_start_loc_opt=query_start_loc,
cache_indices_opt=cache_indices,
initial_state_mode_opt=has_initial_state_tensor,
num_accepted_tokens_opt=None,
activation_mode=activation_num,
pad_slot_id=PAD_SLOT_ID,
run_mode=0,
)
out = out.transpose(-1, -2)
validate_cmp(out, out_ref, itype)
validate_cmp(conv_states, conv_states_ref, itype)
@pytest.mark.skip(
reason="To use this tirton ops:causal_conv1d_fn, you need to set `get_forward_context`. After\
the model side dumps the data, Zeng Tian has made the necessary fixes."
)
@pytest.mark.parametrize("has_initial_state", [False, True])
@pytest.mark.parametrize("itype", [torch.bfloat16])
@pytest.mark.parametrize("silu_activation", [True])
@pytest.mark.parametrize("has_bias", [True])
@pytest.mark.parametrize("seq_len", [[128, 1024, 2048, 4096]])
@pytest.mark.parametrize("extra_state_len", [0, 2])
@pytest.mark.parametrize("width", [2, 4])
@pytest.mark.parametrize("dim", [4160])
def test_causal_conv1d(dim, width, extra_state_len, seq_len, has_bias, silu_activation, itype, has_initial_state):
torch.random.manual_seed(0)
device = "npu"
cu_seqlen, num_seq = sum(seq_len), len(seq_len)
state_len = width - 1 + extra_state_len
x = torch.randn(cu_seqlen, dim, device=device, dtype=itype).transpose(0, 1)
weight = torch.randn(dim, width, device=device, dtype=itype)
query_start_loc = torch.cumsum(torch.tensor([0] + seq_len, device=device, dtype=torch.int32), dim=0)
cache_indices = torch.arange(num_seq, device=device, dtype=torch.int32)
has_initial_state_tensor = torch.tensor([has_initial_state] * num_seq, device=device, dtype=torch.bool)
activation = None if not silu_activation else "silu"
if has_initial_state:
conv_states = torch.randn((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2)
conv_states_ref = (
torch.randn((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2).copy_(conv_states)
)
else:
conv_states = torch.zeros((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2)
conv_states_ref = torch.zeros((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2)
if has_bias:
bias = torch.randn(dim, device=device, dtype=itype)
else:
bias = None
out_ref = causal_conv1d_fn_ref(
x,
weight,
bias=bias,
activation=activation,
conv_states=conv_states_ref,
has_initial_state=has_initial_state_tensor,
cache_indices=cache_indices,
query_start_loc=query_start_loc,
)
out = causal_conv1d_fn(
x,
weight,
bias=bias,
activation=activation,
conv_states=conv_states,
has_initial_state=has_initial_state_tensor,
cache_indices=cache_indices,
query_start_loc=query_start_loc,
)
validate_cmp(out, out_ref, itype)
validate_cmp(conv_states, conv_states_ref, itype)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.skip(
reason="In this scenario, using tirton ops:causal_conv1d_update will cause an overflow. \
Later, Zeng Tian was responsible for fixing this issue."
)
@pytest.mark.parametrize("itype", [torch.bfloat16])
@pytest.mark.parametrize("silu_activation", [True])
@pytest.mark.parametrize("has_bias", [False, True])
@pytest.mark.parametrize("seqlen", [1, 3])
@pytest.mark.parametrize("width", [3, 4])
@pytest.mark.parametrize("dim", [2048 + 16, 4096])
# tests correctness in case subset of the sequences are padded
@pytest.mark.parametrize("with_padding", [True, False])
@pytest.mark.parametrize("batch_size", [3, 64])
def test_causal_conv1d_update_with_batch_gather(
batch_size, with_padding, dim, width, seqlen, has_bias, silu_activation, itype
):
device = "npu"
rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (3e-3, 5e-3)
if itype == torch.bfloat16:
rtol, atol = 1e-2, 5e-2
padding = 5 if with_padding else 0
padded_batch_size = batch_size + padding
# total_entries = number of cache line
total_entries = 10 * batch_size
# x will be (batch, dim, seqlen) with contiguous along dim-axis
x = torch.randn(padded_batch_size, seqlen, dim, device=device, dtype=itype).transpose(1, 2)
x_ref = x.clone()
conv_state_indices = torch.randperm(total_entries)[:batch_size].to(dtype=torch.int32, device=device)
unused_states_bool = torch.ones(total_entries, dtype=torch.bool, device=device)
unused_states_bool[conv_state_indices] = False
padded_state_indices = torch.concat(
[
conv_state_indices,
torch.as_tensor([PAD_SLOT_ID] * padding, dtype=torch.int32, device=device),
],
dim=0,
)
# conv_state will be (cache_lines, dim, state_len)
# with contiguous along dim-axis
conv_state = torch.randn(total_entries, width - 1, dim, device=device, dtype=itype).transpose(1, 2)
conv_state_for_padding_test = conv_state.clone()
weight = torch.randn(dim, width, device=device, dtype=itype)
bias = torch.randn(dim, device=device, dtype=itype) if has_bias else None
conv_state_ref = conv_state[conv_state_indices, :].detach().clone()
activation = None if not silu_activation else "silu"
out = causal_conv1d_update(
x,
conv_state,
weight,
bias,
activation=activation,
conv_state_indices=padded_state_indices,
pad_slot_id=PAD_SLOT_ID,
)
out_ref = causal_conv1d_update_ref(
x_ref[:batch_size].transpose(1, 2), conv_state_ref, weight, bias, activation=activation
).transpose(1, 2)
assert torch.equal(conv_state[conv_state_indices, :], conv_state_ref)
assert torch.equal(conv_state[unused_states_bool], conv_state_for_padding_test[unused_states_bool])
assert torch.allclose(out[:batch_size], out_ref, rtol=rtol, atol=atol)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.skip("Probabilistic failure, need zengtian after fix")
def test_causal_conv1d_update_qwen3_next_shape():
device = "npu"
itype = torch.bfloat16
rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (3e-3, 5e-3)
if itype == torch.bfloat16:
rtol, atol = 1e-2, 5e-2
total_tokens = 192
dim = 4096
kernel_size = 4
batch_size = 96
num_states = 929
x = torch.randn(total_tokens, dim, dtype=itype, device=device)
conv_state = torch.randn(num_states, dim, kernel_size, dtype=itype, device=device)
weight = torch.randn(dim, kernel_size, dtype=itype, device=device)
bias = None
conv_state_indices = torch.randint(0, num_states, (batch_size,), dtype=torch.int32, device=device)
num_accepted_tokens = torch.ones(total_tokens, dtype=torch.int32, device=device)
query_start_loc = torch.arange(0, total_tokens + 1, dtype=torch.int32, device=device)
activation = "silu"
max_query_len = 2
pad_slot_id = -1
validate_data = False
block_idx_last_scheduled_token = None
initial_state_idx = None
out = causal_conv1d_update(
x,
conv_state,
weight,
bias,
activation,
conv_state_indices,
num_accepted_tokens,
query_start_loc,
max_query_len,
pad_slot_id,
block_idx_last_scheduled_token,
initial_state_idx,
validate_data,
)
x_ref = x.clone()
conv_state_ref = conv_state[conv_state_indices, :].detach().clone()
out_ref = causal_conv1d_update_ref(
x_ref[:batch_size].transpose(1, 2), conv_state_ref, weight, bias, activation=activation
).transpose(1, 2)
assert torch.allclose(out[:batch_size], out_ref, rtol=rtol, atol=atol)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,83 @@
from unittest.mock import MagicMock, patch
import torch
from tests.ut.base import PytestBase
from vllm_ascend._310p.ops.fla.chunk_gated_delta_rule import chunk_gated_delta_rule_pytorch
from vllm_ascend.ops.triton.fla.chunk import chunk_gated_delta_rule
class TestChunkGatedDeltaRule(PytestBase):
def test_triton_fusion_ops(self):
mock_attn_metadata = MagicMock()
mock_attn_metadata.num_decodes = 1
mock_forward_context = MagicMock()
mock_forward_context.attn_metadata = mock_attn_metadata
q = torch.randn(1, 17, 4, 128, dtype=torch.bfloat16).npu()
k = torch.randn(1, 17, 4, 128, dtype=torch.bfloat16).npu()
v = torch.randn(1, 17, 8, 128, dtype=torch.bfloat16).npu()
g = torch.randn(1, 17, 8, dtype=torch.float32).npu()
beta = torch.randn(1, 17, 8, dtype=torch.bfloat16).npu()
initial_state = torch.randn(3, 8, 128, 128, dtype=torch.bfloat16).npu()
q_start_loc = torch.range(0, 3, dtype=torch.int).npu()
mock_pcp_group = MagicMock()
mock_pcp_group.world_size = 1
with (
patch("vllm_ascend.ops.triton.fla.chunk.get_forward_context", return_value=mock_forward_context),
patch("vllm_ascend.ops.triton.fla.chunk.get_pcp_group", return_value=mock_pcp_group),
):
(
core_attn_out_non_spec,
last_recurrent_state,
) = chunk_gated_delta_rule(
q=q,
k=k,
v=v,
g=g,
beta=beta,
initial_state=initial_state,
output_final_state=True,
cu_seqlens=q_start_loc,
head_first=False,
use_qk_l2norm_in_kernel=True,
)
assert core_attn_out_non_spec.shape == (1, 17, 8, 128)
assert last_recurrent_state.shape == (3, 8, 128, 128)
def test_chunk_gated_delta_rule_310_state_layout_matches_vllm():
q = torch.tensor([[[[1.0, 0.0]]]], dtype=torch.float32)
k = torch.tensor([[[[1.0, 0.0]]]], dtype=torch.float32)
v = torch.tensor([[[[10.0, 20.0, 30.0]]]], dtype=torch.float32)
g = torch.zeros(1, 1, 1, dtype=torch.float32)
beta = torch.ones(1, 1, 1, dtype=torch.float32)
initial_state = torch.tensor(
[[[[1.0, 2.0], [4.0, 8.0], [16.0, 32.0]]]],
dtype=torch.float32,
)
out, final_state = chunk_gated_delta_rule_pytorch(
q=q,
k=k,
v=v,
g=g,
beta=beta,
initial_state=initial_state,
output_final_state=True,
cu_seqlens=None,
head_first=False,
use_qk_l2norm_in_kernel=False,
)
expected_out = torch.tensor([[[[10.0, 20.0, 30.0]]]], dtype=torch.float32) / (2.0**0.5)
expected_state = torch.tensor(
[[[[10.0, 2.0], [20.0, 8.0], [30.0, 32.0]]]],
dtype=torch.float32,
)
torch.testing.assert_close(out, expected_out, rtol=1e-5, atol=1e-5)
assert final_state is not None
torch.testing.assert_close(final_state, expected_state, rtol=1e-5, atol=1e-5)

View File

@@ -0,0 +1,162 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# mypy: ignore-errors
"""Precision tests for vllm's chunk_kda Triton operator on NPU.
Compares chunk_kda against a naive recurrent reference (float32).
"""
import pytest
import torch
import torch.nn.functional as F
import torch_npu # noqa: F401
from vllm_ascend.ops.triton.kda.kda import chunk_kda
DEVICE = "npu"
NPU_RMSE_RATIO_O = 0.005
NPU_RMSE_RATIO_HT = 0.005
def reference_l2norm(x: torch.Tensor, eps: float = 1e-6) -> torch.Tensor:
dtype = x.dtype
x = x.to(torch.float32)
return (x * torch.rsqrt(torch.sum(x * x, dim=-1, keepdim=True) + eps)).to(dtype)
def naive_recurrent_kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""Naive recurrent KDA reference, ported from FLA's naive.py."""
dtype = v.dtype
B, T, H, K, V = *q.shape, v.shape[-1]
if scale is None:
scale = K**-0.5
q, k, v, g, beta = map(lambda x: x.to(torch.float), [q, k, v, g, beta])
q = q * scale
S = k.new_zeros(B, H, K, V).to(q)
if initial_state is not None:
S += initial_state
o = torch.zeros_like(v)
for i in range(T):
q_i, k_i, v_i, g_i, b_i = q[:, i], k[:, i], v[:, i], g[:, i], beta[:, i]
S = S * g_i[..., None].exp()
S = S + torch.einsum(
"bhk,bhv->bhkv",
b_i[..., None] * k_i,
v_i - (k_i[..., None] * S).sum(-2),
)
o[:, i] = torch.einsum("bhk,bhkv->bhv", q_i, S)
if not output_final_state:
S = None
return o.to(dtype), S
def assert_close(
name: str,
ref: torch.Tensor,
tri: torch.Tensor,
ratio: float,
err_atol: float = 1e-6,
):
"""RMSE-based relative error comparison."""
abs_err = (ref.detach() - tri.detach()).flatten().abs().max().item()
rmse_diff = (ref.detach() - tri.detach()).flatten().square().mean().sqrt().item()
rmse_base = ref.detach().flatten().square().mean().sqrt().item()
rel_err = rmse_diff / (rmse_base + 1e-8)
print(f"{name:>4} | abs={abs_err:.6f} | rmse={rel_err:.6f} | thr={ratio}")
if abs_err <= err_atol:
return
assert not torch.isnan(ref).any(), f"{name}: NaN detected in ref"
assert not torch.isnan(tri).any(), f"{name}: NaN detected in tri"
assert rel_err < ratio, f"{name}: max abs err {abs_err:.6f}, rmse ratio {rel_err:.6f} >= {ratio}"
@pytest.mark.parametrize(
("H", "D", "cu_seqlens", "dtype"),
[
pytest.param(
*test,
id="H{}-D{}-cu{}-{}".format(*test),
)
for test in [
(32, 128, [0, 64], torch.float16),
(32, 128, [0, 1024], torch.float16),
(32, 128, [0, 15], torch.float16),
(32, 128, [0, 256, 512, 768, 1024], torch.float16),
(32, 128, [0, 15, 100, 300, 1200], torch.float16),
(64, 128, [0, 256, 500, 1000], torch.float16),
(32, 128, [0, 8192], torch.float16),
(32, 128, [0, 256, 500, 1000], torch.bfloat16),
(32, 128, [0, 4096], torch.float16),
]
],
)
@pytest.mark.skip_global_cleanup
@torch.inference_mode()
def test_chunk_kda(
H: int,
D: int,
cu_seqlens: list[int],
dtype: torch.dtype,
):
T = cu_seqlens[-1]
torch.manual_seed(42)
B = 1
cu_seqlens_t = torch.LongTensor(cu_seqlens).to(DEVICE)
N = len(cu_seqlens) - 1
q = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
k = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
v = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
g = F.logsigmoid(torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE)).to(dtype)
beta = torch.rand(B, T, H, dtype=dtype, device=DEVICE).sigmoid()
h0 = torch.randn(N, H, D, D, dtype=torch.float32, device=DEVICE)
ref_outputs = []
ref_states = []
for i in range(N):
s, e = cu_seqlens[i], cu_seqlens[i + 1]
q_i = reference_l2norm(q[:, s:e].contiguous())
k_i = reference_l2norm(k[:, s:e].contiguous())
o_i, ht_i = naive_recurrent_kda(
q_i,
k_i,
v[:, s:e],
g[:, s:e],
beta[:, s:e],
initial_state=h0[i],
output_final_state=True,
)
ref_outputs.append(o_i)
ref_states.append(ht_i)
ref_o = torch.cat(ref_outputs, dim=1)
ref_ht = torch.cat(ref_states, dim=0)
# h0 transposed to (V, K) layout for the kernel; naive uses (K, V)
tri_o, tri_ht = chunk_kda(
q=q.clone(),
k=k.clone(),
v=v.clone(),
g=g.clone(),
beta=beta.clone(),
initial_state=h0.transpose(-1, -2).contiguous().clone(),
output_final_state=True,
cu_seqlens=cu_seqlens_t,
use_qk_l2norm_in_kernel=True,
)
assert not torch.isnan(tri_o).any(), "Triton output o contains NaN"
assert not torch.isnan(tri_ht).any(), "Triton output ht contains NaN"
assert_close("o", ref_o, tri_o, NPU_RMSE_RATIO_O)
assert_close("ht", ref_ht, tri_ht.transpose(-1, -2).contiguous(), NPU_RMSE_RATIO_HT)

View File

@@ -0,0 +1,37 @@
import gc
import pytest
import torch
from vllm_ascend.ops.triton.fla.utils import clear_ssm_states
@pytest.mark.parametrize("dtype", [torch.bfloat16])
@pytest.mark.parametrize(
"state_shape",
[
(6, 3, 5, 7),
(4, 5, 25, 41),
],
)
def test_clear_ssm_states_ref_parity(state_shape, dtype):
torch.manual_seed(0)
device = "npu"
ssm_states = torch.randn(*state_shape, device=device, dtype=dtype)
has_initial_state = torch.tensor(
[True, False, True, False, False, True][: state_shape[0]],
device=device,
dtype=torch.bool,
)
ssm_states_ref = ssm_states.clone()
ssm_states_ref[~has_initial_state, ...] = 0
clear_ssm_states(ssm_states, has_initial_state)
torch.testing.assert_close(ssm_states, ssm_states_ref)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,109 @@
import torch
from vllm.v1.worker.gpu.block_table import _compute_slot_mappings_kernel as ref_compute_slot_mappings_kernel
from vllm_ascend.worker.v2.block_table import _compute_slot_mappings_kernel as ascend_compute_slot_mappings_kernel
def test_compute_slot_mapping_npu_kernel():
"""
Computes the physical slot IDs in KV cache for each token in the current batch.
This function maps the logical positions of tokens to their actual storage locations
in the block-managed KV cache, which is critical for efficient memory access in LLM inference.
Input:
- max_num_batched_tokens (int): Maximum preallocated batched tokens in KV cache (memory limit)
- idx_mapping (torch.Tensor): [num_reqs], int32 → Virtual-to-actual request index mapping
- query_start_loc (torch.Tensor): [num_reqs+1], int32 → Batch-level token start positions per request
- positions (torch.Tensor): [num_tokens], int64 → Per-token logical sequence positions in requests
- block_table_ptrs (torch.Tensor): [num_kv_cache_groups], int32 → Pointers to block tables (virtual→physical)
- block_table_strides (torch.Tensor): [num_kv_cache_groups], int32 → Stride for block table addressing
- block_sizes_tensor (torch.Tensor): [num_kv_cache_groups], int32 → Token capacity per KV cache block
- slot_mappings (torch.Tensor): [num_kv_cache_groups, max_num_batched_tokens], int32 → Output slot ID tensor
- slot_mappings_stride0 (int): Stride of the first dimension of slot_mappings (memory layout)
- cp_rank (int): Current device rank in column-parallel (CP) group
- CP_SIZE (int): Total devices in CP parallel group
- CP_INTERLEAVE (bool): Enable interleaved CP computation (memory access optimization)
- PAD_ID (int): Padding value for invalid slot IDs (-1)
- TRITON_BLOCK_SIZE (int): Block size for Triton kernel execution (hardware optimization),
'TOTAL_BLOCK_SIZE' must be greater than the 'position / (block_size * CP_SIZE) + 1024'
Output:
- slot_mappings (torch.Tensor): [num_kv_cache_groups, max_num_batched_tokens], int32 → Output slot ID tensor
"""
torch.manual_seed(42)
device = "npu" if torch.npu.is_available() else "cuda" if torch.cuda.is_available() else "cpu"
max_num_batched_tokens = 8192
idx_mapping = torch.tensor([63], dtype=torch.int32, device=device)
query_start_loc = torch.tensor([0, 5], dtype=torch.int32, device=device)
positions = torch.tensor([0, 1, 2, 3, 4, 0, 0, 0], dtype=torch.int64, device=device)
num_kv_cache_groups = 1
max_num_reqs = 64
max_num_blocks = 320
block_tables: list[torch.Tensor] = []
for i in range(num_kv_cache_groups):
block_table = torch.randint(0, 320, (max_num_reqs, max_num_blocks), dtype=torch.int32, device=device)
block_tables.append(block_table)
block_table_ptrs = torch.tensor([t.data_ptr() for t in block_table], dtype=torch.uint64, device=device)
block_table_strides = torch.tensor([320], dtype=torch.int32, device=device)
block_sizes_tensor = torch.tensor([128], dtype=torch.int32, device=device)
slot_mappings = torch.zeros(size=(1, 8192), dtype=torch.int64, device=device)
ref_slot_mappings = torch.zeros(size=(1, 8192), dtype=torch.int64, device=device)
cp_rank = 0
cp_size = 1
cp_interleave = 1
num_reqs = query_start_loc.shape[0] - 1
num_groups = num_kv_cache_groups
try:
ascend_compute_slot_mappings_kernel[(num_groups, num_reqs + 1)](
max_num_batched_tokens,
idx_mapping,
query_start_loc,
positions,
block_table_ptrs,
block_table_strides,
block_sizes_tensor,
slot_mappings,
slot_mappings.stride(0),
cp_rank,
CP_SIZE=cp_size,
CP_INTERLEAVE=cp_interleave,
PAD_ID=-1,
TRITON_BLOCK_SIZE=1024, # type: ignore
TOTAL_BLOCK_SIZE=4096,
)
ref_compute_slot_mappings_kernel[(num_groups, num_reqs + 1)](
max_num_batched_tokens,
idx_mapping,
query_start_loc,
positions,
block_table_ptrs,
block_table_strides,
block_sizes_tensor,
ref_slot_mappings,
ref_slot_mappings.stride(0),
cp_rank,
CP_SIZE=cp_size,
CP_INTERLEAVE=cp_interleave,
PAD_ID=-1,
TRITON_BLOCK_SIZE=1024, # type: ignore
)
# ========== Verify results ==========
assert torch.equal(slot_mappings, ref_slot_mappings), (
f"ascend output differs from gpu reference.\n"
f"Max diff: {torch.max(torch.abs(slot_mappings - ref_slot_mappings))}\n"
f"Mean diff: {torch.mean(torch.abs(slot_mappings - ref_slot_mappings).float())}"
)
except Exception as e:
print(f"Error during executionm: {e}")
import traceback
traceback.print_exc()

View File

@@ -0,0 +1,227 @@
import random
import pytest
import torch
from vllm_ascend.worker.v2.sample.logprob import compute_token_logprobs
def torch_compute_token_logprobs(logits: torch.Tensor, token_ids: torch.Tensor) -> torch.Tensor:
"""Pure PyTorch reference implementation of topk log softmax.
Computes log_softmax for the entire logits tensor, then gathers
the values at the specified token_ids positions.
Args:
logits: Tensor of shape (batch_size, vocab_size) containing the logits.
token_ids: Tensor of shape (batch_size, topk) containing the token indices.
Returns:
Tensor of shape (batch_size, topk) containing log probabilities.
"""
# Compute log_softmax along the vocab dimension
# log_softmax(x) = x - log(sum(exp(x))) = x - max(x) - log(sum(exp(x - max(x))))
log_probs = torch.nn.functional.log_softmax(logits.float(), dim=-1)
# Gather the log probabilities at the specified token positions
token_ids = token_ids.to(torch.int64)
result = torch.gather(log_probs, dim=1, index=token_ids)
return result.to(torch.float32)
# Common vocab sizes from mainstream models
VOCAB_SIZES = [
32000, # LLaMA / LLaMA2 / Mistral
50257, # GPT-2
65024, # ChatGLM
128256, # LLaMA3
151936, # Qwen2
]
# Different topk values to test
TOPK_VALUES = [1, 2, 5, 10, 32, 64]
@pytest.mark.skip("UB overflow, zengtian needs to fix it later")
@pytest.mark.parametrize(
"batch_size, vocab_size, topk",
[(random.randint(1, 64), vocab_size, topk) for vocab_size in VOCAB_SIZES for topk in TOPK_VALUES],
)
def test_topk_log_softmax_kernel(batch_size, vocab_size, topk):
"""
Test the Triton _topk_log_softmax_kernel against a pure PyTorch reference.
The kernel computes log_softmax and gathers values at specified token positions.
Args:
batch_size: Number of requests (rows) in the logits tensor.
vocab_size: Vocabulary size (columns) in the logits tensor.
topk: Number of token positions to compute log probabilities for.
"""
torch.manual_seed(42)
device = "npu"
# Build input tensors
logits = torch.randn((batch_size, vocab_size), dtype=torch.float32, device=device)
# Generate random token indices within vocab_size
token_ids = torch.randint(0, vocab_size, (batch_size, topk), dtype=torch.int64, device=device)
# ========== Run Triton kernel ==========
logprobs_triton = compute_token_logprobs(logits, token_ids)
# ========== Run PyTorch reference ==========
logprobs_ref = torch_compute_token_logprobs(logits, token_ids)
# ========== Verify results ==========
max_diff = torch.max(torch.abs(logprobs_triton - logprobs_ref)).item()
mean_diff = torch.mean(torch.abs(logprobs_triton - logprobs_ref)).item()
assert torch.allclose(logprobs_triton, logprobs_ref, atol=1e-4, rtol=1e-5), (
f"Triton topk_log_softmax kernel output differs from torch reference.\n"
f"batch_size={batch_size}, vocab_size={vocab_size}, topk={topk}\n"
f"Max diff: {max_diff}\n"
f"Mean diff: {mean_diff}"
)
@pytest.mark.skip("UB overflow, zengtian needs to fix it later")
@pytest.mark.parametrize("vocab_size", VOCAB_SIZES)
def test_topk_log_softmax_edge_cases(vocab_size):
"""
Test edge cases for the topk_log_softmax kernel.
Args:
vocab_size: Vocabulary size to test.
"""
torch.manual_seed(42)
device = "npu"
# Test case 1: Single batch, single topk
logits = torch.randn((1, vocab_size), dtype=torch.float32, device=device)
token_ids = torch.randint(0, vocab_size, (1, 1), dtype=torch.int64, device=device)
logprobs_triton = compute_token_logprobs(logits, token_ids)
logprobs_ref = torch_compute_token_logprobs(logits, token_ids)
assert torch.allclose(logprobs_triton, logprobs_ref, atol=1e-4, rtol=1e-5), (
f"Edge case (1,1) failed for vocab_size={vocab_size}"
)
# Test case 2: Logits with extreme values
logits_extreme = torch.randn((4, vocab_size), dtype=torch.float32, device=device)
logits_extreme[0, 0] = 100.0 # Very large positive
logits_extreme[1, 0] = -100.0 # Very large negative
logits_extreme[2, :] = 0.0 # All zeros
logits_extreme[3, :] = 1.0 # All ones
token_ids = torch.zeros((4, 5), dtype=torch.int64, device=device)
token_ids[:, 0] = 0 # Include the extreme value position
for i in range(1, 5):
token_ids[:, i] = torch.randint(1, vocab_size, (4,))
logprobs_triton = compute_token_logprobs(logits_extreme, token_ids)
logprobs_ref = torch_compute_token_logprobs(logits_extreme, token_ids)
assert torch.allclose(logprobs_triton, logprobs_ref, atol=1e-4, rtol=1e-5), (
f"Extreme values test failed for vocab_size={vocab_size}"
)
@pytest.mark.skip("UB overflow, zengtian needs to fix it later")
@pytest.mark.parametrize(
"batch_size, vocab_size, topk",
[
(16, 32000, 10),
(32, 50257, 5),
(64, 128256, 20),
],
)
def test_topk_log_softmax_deterministic(batch_size, vocab_size, topk):
"""
Test that the kernel produces deterministic results across multiple runs.
Args:
batch_size: Number of requests.
vocab_size: Vocabulary size.
topk: Number of token positions.
"""
torch.manual_seed(42)
device = "npu"
logits = torch.randn((batch_size, vocab_size), dtype=torch.float32, device=device)
token_ids = torch.randint(0, vocab_size, (batch_size, topk), dtype=torch.int64, device=device)
# Run multiple times and check consistency
results = []
for _ in range(3):
result = compute_token_logprobs(logits, token_ids)
results.append(result.clone())
for i in range(1, len(results)):
assert torch.equal(results[0], results[i]), f"Non-deterministic results detected in run {i}"
@pytest.mark.skip("UB overflow, zengtian needs to fix it later")
@pytest.mark.parametrize("dtype", [torch.float32, torch.float16])
def test_topk_log_softmax_dtypes(dtype):
"""
Test the kernel with different input dtypes.
Args:
dtype: Input tensor dtype.
"""
torch.manual_seed(42)
device = "npu"
batch_size = 8
vocab_size = 32000
topk = 10
logits = torch.randn((batch_size, vocab_size), dtype=dtype, device=device)
token_ids = torch.randint(0, vocab_size, (batch_size, topk), dtype=torch.int64, device=device)
logprobs_triton = compute_token_logprobs(logits, token_ids)
logprobs_ref = torch_compute_token_logprobs(logits.float(), token_ids)
# Use slightly larger tolerance for float16 due to precision loss
atol = 1e-3 if dtype == torch.float16 else 1e-4
assert torch.allclose(logprobs_triton, logprobs_ref, atol=atol, rtol=1e-4), f"dtype {dtype} test failed"
if __name__ == "__main__":
# Run a quick sanity check
print("Running quick sanity check...")
device = "npu"
print(f"Using device: {device}")
torch.manual_seed(42)
batch_size = 4
vocab_size = 32000
topk = 5
logits = torch.randn((batch_size, vocab_size), dtype=torch.float32, device=device)
token_ids = torch.randint(0, vocab_size, (batch_size, topk), dtype=torch.int64, device=device)
logprobs_triton = compute_token_logprobs(logits, token_ids)
logprobs_ref = torch_compute_token_logprobs(logits, token_ids)
max_diff = torch.max(torch.abs(logprobs_triton - logprobs_ref)).item()
mean_diff = torch.mean(torch.abs(logprobs_triton - logprobs_ref)).item()
print(f"Max diff: {max_diff}")
print(f"Mean diff: {mean_diff}")
print(f"All close (atol=1e-4): {torch.allclose(logprobs_triton, logprobs_ref, atol=1e-4, rtol=1e-5)}")
print("\nTriton output (first row):", logprobs_triton[0])
print("PyTorch output (first row):", logprobs_ref[0])
print("\nSanity check passed!")

View File

@@ -0,0 +1,61 @@
import pytest
import torch
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
from vllm_ascend.worker.v2.sample.logprob import compute_topk_logprobs
@pytest.mark.parametrize(
"batch_size,vocab_size,num_logprobs",
[
(48, 1024, 5),
(96, 1024, 0),
(24, 1519, 1),
(1, 320, 10),
],
)
def test_compute_topk_logprobs(batch_size, vocab_size, num_logprobs):
"""Test compute_topk_logprobs for correctness of IDs, logprobs, and ranks.
Args:
batch_size: Number of sequences in the batch
vocab_size: Size of the vocabulary
num_logprobs: Number of top-k logprobs to return (excluding the sampled token)
"""
init_device_properties_triton()
# ========== 1. Setup test data ==========
torch.manual_seed(42)
device = "npu"
logits = torch.randn(batch_size, vocab_size, device=device, dtype=torch.float32)
sampled_token_ids = torch.randint(0, vocab_size, (batch_size,), device=device, dtype=torch.int64)
# ========== 2. Execute Triton implementation ==========
triton_output = compute_topk_logprobs(logits, num_logprobs, sampled_token_ids)
torch.npu.synchronize()
# ========== 3. Compute reference values using PyTorch ==========
if num_logprobs == 0:
ref_token_ids = sampled_token_ids.unsqueeze(-1)
else:
topk_indices = torch.topk(logits, num_logprobs, dim=-1).indices
ref_token_ids = torch.cat((sampled_token_ids.unsqueeze(-1), topk_indices), dim=1)
ref_all_logprobs = torch.log_softmax(logits, dim=-1)
ref_logprobs = torch.gather(ref_all_logprobs, dim=1, index=ref_token_ids)
sampled_logits = torch.gather(logits, 1, sampled_token_ids.unsqueeze(-1))
ref_ranks = (logits > sampled_logits).sum(dim=1).to(torch.int64)
# ========== 4. Verify results ==========
assert torch.equal(triton_output.logprob_token_ids, ref_token_ids), (
"Token IDs (Sampled + TopK) do not match between Triton and PyTorch."
)
assert torch.equal(triton_output.selected_token_ranks, ref_ranks), (
f"Token Ranks do not match.\nTriton: {triton_output.selected_token_ranks}\nPyTorch: {ref_ranks}"
)
assert torch.allclose(triton_output.logprobs, ref_logprobs, rtol=1e-4, atol=1e-4), (
f"Logprobs values differ between Triton and PyTorch.\n"
f"Max diff: {torch.max(torch.abs(triton_output.logprobs - ref_logprobs))}"
)

View File

@@ -0,0 +1,51 @@
import torch
from vllm_ascend._310p.ops.fla.fused_gdn_gating import fused_gdn_gating_pytorch
from vllm_ascend.ops.triton.fused_gdn_gating import fused_gdn_gating_patch
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
def test_fused_gdn_gating_310p_parity_precision():
init_device_properties_triton()
torch.manual_seed(0)
device = "npu"
num_tokens = 37
num_heads = 8
A_log = torch.randn(num_heads, dtype=torch.float16, device=device)
dt_bias = torch.randn(num_heads, dtype=torch.float16, device=device)
a = torch.randn(num_tokens, num_heads, dtype=torch.float16, device=device)
b = torch.randn(num_tokens, num_heads, dtype=torch.float16, device=device)
triton_g, triton_beta = fused_gdn_gating_patch(
A_log=A_log,
a=a,
b=b,
dt_bias=dt_bias,
beta=1.0,
threshold=20.0,
)
ref_g, ref_beta = fused_gdn_gating_pytorch(
A_log=A_log,
a=a,
b=b,
dt_bias=dt_bias,
beta=1.0,
threshold=20.0,
)
torch.testing.assert_close(
triton_g.to(torch.float32).cpu(),
ref_g.to(torch.float32).cpu(),
rtol=1e-2,
atol=1e-2,
equal_nan=True,
)
torch.testing.assert_close(
triton_beta.to(torch.float32).cpu(),
ref_beta.to(torch.float32).cpu(),
rtol=1e-2,
atol=1e-2,
equal_nan=True,
)

View File

@@ -0,0 +1,88 @@
import gc
import pytest
import torch
from einops import rearrange
from vllm.model_executor.layers.mamba.gdn.base import GatedDeltaNetAttention # type: ignore[import-not-found]
from vllm_ascend.ops.triton.fla.fused_qkvzba_split_reshape import fused_qkvzba_split_reshape_cat
def validate_cmp(y_cal, y_ref, dtype, device="npu"):
y_cal = y_cal.to(device)
y_ref = y_ref.to(device)
if dtype == torch.float16 or dtype == torch.bfloat16:
torch.testing.assert_close(y_ref, y_cal, rtol=5e-03, atol=5e-03, equal_nan=True)
elif dtype == torch.float32:
torch.testing.assert_close(y_ref, y_cal, rtol=1e-03, atol=1e-03, equal_nan=True)
elif (
dtype == torch.int32
or dtype == torch.int64
or dtype == torch.int16
or dtype == torch.int8
or dtype == torch.uint32
or dtype == torch.bool
):
assert torch.equal(y_cal, y_ref)
else:
raise ValueError('Invalid parameter "dtype" is found : {}'.format(dtype))
@pytest.mark.parametrize("seq_len", [1, 64, 1024, 2048])
@pytest.mark.parametrize("num_heads_qk", [2, 4, 8])
@pytest.mark.parametrize("num_heads_v", [8])
@pytest.mark.parametrize("head_qk_dim", [256])
@pytest.mark.parametrize("head_v_dim", [128])
@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16])
def test_fused_qkvzba_split_reshape_cat(
seq_len,
num_heads_qk,
num_heads_v,
head_qk_dim,
head_v_dim,
dtype,
):
if num_heads_v % num_heads_qk != 0:
pytest.skip("num_heads_v must be divisible by num_heads_qk")
torch.random.manual_seed(0)
device = "npu"
projected_states_qkvz = torch.randn(
seq_len, 2 * head_qk_dim * num_heads_qk + 2 * head_v_dim * num_heads_v, dtype=dtype, device=device
)
projected_states_ba = torch.randn(seq_len, 2 * num_heads_v, dtype=dtype, device=device)
projected_states_qkvz_copy = projected_states_qkvz.clone()
projected_states_ba_copy = projected_states_ba.clone()
mixed_qkv, z, b, a = fused_qkvzba_split_reshape_cat(
projected_states_qkvz_copy,
projected_states_ba_copy,
num_heads_qk,
num_heads_v,
head_qk_dim,
head_v_dim,
)
gdn = GatedDeltaNetAttention.__new__(GatedDeltaNetAttention)
gdn.num_k_heads = num_heads_qk
gdn.num_v_heads = num_heads_v
gdn.head_k_dim = head_qk_dim
gdn.head_v_dim = head_v_dim
gdn.tp_size = 1
query, key, value, z_ref, b_ref, a_ref = gdn.fix_query_key_value_ordering(
mixed_qkvz=projected_states_qkvz, mixed_ba=projected_states_ba
)
query, key, value = map(lambda x: rearrange(x, "l p d -> l (p d)"), (query, key, value))
mixed_qkv_ref = torch.cat((query, key, value), dim=-1)
validate_cmp(mixed_qkv, mixed_qkv_ref, dtype)
validate_cmp(z, z_ref, dtype)
validate_cmp(b, b_ref, dtype)
validate_cmp(a, a_ref, dtype)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,114 @@
import pytest
import torch
from vllm.model_executor.layers.fla.ops import fused_recurrent_gated_delta_rule
from vllm_ascend._310p.ops.fla.fused_recurrent_gated_delta_rule import fused_recurrent_gated_delta_rule_pytorch
@pytest.mark.skip("Probabilistic failure, need zengtian after fix")
def test_fused_recurrent_gated_delta_rule_310p_parity_precision():
torch.manual_seed(0)
device = "npu"
bsz = 1
total_tokens = 9
num_qk_heads = 2
num_v_heads = 4
kdim = 64
vdim = 48
q = torch.randn(bsz, total_tokens, num_qk_heads, kdim, dtype=torch.float16, device=device)
k = torch.randn(bsz, total_tokens, num_qk_heads, kdim, dtype=torch.float16, device=device)
v = torch.randn(bsz, total_tokens, num_v_heads, vdim, dtype=torch.float16, device=device)
g = torch.randn(bsz, total_tokens, num_v_heads, dtype=torch.float32, device=device)
beta = torch.sigmoid(torch.randn(bsz, total_tokens, num_v_heads, dtype=torch.float32, device=device)).to(
torch.float16
)
initial_state = torch.randn(2, num_v_heads, vdim, kdim, dtype=torch.float16, device=device)
cu_seqlens = torch.tensor([0, 4, 9], dtype=torch.long, device=device)
# For inplace_final_state=True, Ascend triton kernel expects explicit per-token state indices.
# seq0 (len=4) -> state 0, seq1 (len=5) -> state 1.
ssm_state_indices = torch.tensor(
[
[0, 0, 0, 0, 0],
[1, 1, 1, 1, 1],
],
dtype=torch.long,
device=device,
)
triton_out, triton_state = fused_recurrent_gated_delta_rule(
q=q,
k=k,
v=v,
g=g,
beta=beta,
initial_state=initial_state.clone(),
inplace_final_state=True,
cu_seqlens=cu_seqlens,
ssm_state_indices=ssm_state_indices,
use_qk_l2norm_in_kernel=True,
)
ref_out, ref_state = fused_recurrent_gated_delta_rule_pytorch(
q=q,
k=k,
v=v,
g=g,
beta=beta,
initial_state=initial_state.clone(),
inplace_final_state=True,
cu_seqlens=cu_seqlens,
ssm_state_indices=ssm_state_indices,
use_qk_l2norm_in_kernel=True,
)
torch.testing.assert_close(
triton_out.to(torch.float32).cpu(),
ref_out.to(torch.float32).cpu(),
rtol=1e-2,
atol=1e-2,
equal_nan=True,
)
torch.testing.assert_close(
triton_state.to(torch.float32).cpu(),
ref_state.to(torch.float32).cpu(),
rtol=1e-2,
atol=1e-2,
equal_nan=True,
)
def test_fused_recurrent_gated_delta_rule_310_state_layout_matches_vllm():
q = torch.tensor([[[[1.0, 0.0]]]], dtype=torch.float32)
k = torch.tensor([[[[1.0, 0.0]]]], dtype=torch.float32)
v = torch.tensor([[[[10.0, 20.0, 30.0]]]], dtype=torch.float32)
g = torch.zeros(1, 1, 1, dtype=torch.float32)
beta = torch.ones(1, 1, 1, dtype=torch.float32)
initial_state = torch.tensor(
[[[[1.0, 2.0], [4.0, 8.0], [16.0, 32.0]]]],
dtype=torch.float32,
)
out, final_state = fused_recurrent_gated_delta_rule_pytorch(
q=q,
k=k,
v=v,
g=g,
beta=beta,
initial_state=initial_state,
inplace_final_state=False,
cu_seqlens=None,
ssm_state_indices=None,
num_accepted_tokens=None,
use_qk_l2norm_in_kernel=False,
)
expected_out = torch.tensor([[[[10.0, 20.0, 30.0]]]], dtype=torch.float32) / (2.0**0.5)
expected_state = torch.tensor(
[[[[10.0, 2.0], [20.0, 8.0], [30.0, 32.0]]]],
dtype=torch.float32,
)
torch.testing.assert_close(out, expected_out, rtol=1e-5, atol=1e-5)
torch.testing.assert_close(final_state, expected_state, rtol=1e-5, atol=1e-5)

View File

@@ -0,0 +1,369 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# mypy: ignore-errors
"""Precision tests for vllm's fused_recurrent_kda Triton operator on NPU.
Tests the recurrent-mode (decode) kernel against a naive PyTorch recurrent
reference implementation. Both ref and kernel use the same token-by-token
recurrence algorithm, so errors are purely from FP accumulation differences
on NPU triton-ascend.
"""
import pytest
import torch
import torch.nn.functional as F
import torch_npu # noqa: F401
from vllm_ascend.ops.triton.kda.kda import fused_recurrent_kda
DEVICE = "npu"
# Both ref and kernel use the same recurrent algorithm; errors come from
# FP accumulation differences on NPU triton-ascend.
NPU_RMSE_RATIO_O = 0.005
NPU_RMSE_RATIO_HT = 0.005
def reference_l2norm(x: torch.Tensor, eps: float = 1e-6) -> torch.Tensor:
dtype = x.dtype
x = x.to(torch.float32)
return (x * torch.rsqrt(torch.sum(x * x, dim=-1, keepdim=True) + eps)).to(dtype)
def naive_recurrent_kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""Naive recurrent KDA reference (pure PyTorch, runs on any device).
Ported from flash-linear-attention/fla/ops/kda/naive.py.
"""
dtype = v.dtype
B, T, H, K, V = *q.shape, v.shape[-1]
if scale is None:
scale = K**-0.5
q, k, v, g, beta = map(lambda x: x.to(torch.float), [q, k, v, g, beta])
q = q * scale
S = k.new_zeros(B, H, K, V).to(q)
if initial_state is not None:
S += initial_state
o = torch.zeros_like(v)
for i in range(T):
q_i, k_i, v_i, g_i, b_i = q[:, i], k[:, i], v[:, i], g[:, i], beta[:, i]
S = S * g_i[..., None].exp()
S = S + torch.einsum(
"bhk,bhv->bhkv",
b_i[..., None] * k_i,
v_i - (k_i[..., None] * S).sum(-2),
)
o[:, i] = torch.einsum("bhk,bhkv->bhv", q_i, S)
if not output_final_state:
S = None
return o.to(dtype), S
def assert_close(
name: str,
ref: torch.Tensor,
tri: torch.Tensor,
ratio: float,
err_atol: float = 1e-6,
):
"""RMSE-based relative error comparison (same logic as FLA's assert_close)."""
abs_err = (ref.detach() - tri.detach()).flatten().abs().max().item()
rmse_diff = (ref.detach() - tri.detach()).flatten().square().mean().sqrt().item()
rmse_base = ref.detach().flatten().square().mean().sqrt().item()
rel_err = rmse_diff / (rmse_base + 1e-8)
print(f"{name:>8} | max abs err: {abs_err:.6f} | rmse ratio: {rel_err:.6f} | threshold: {ratio}")
if abs_err <= err_atol:
return
assert not torch.isnan(ref).any(), f"{name}: NaN detected in ref"
assert not torch.isnan(tri).any(), f"{name}: NaN detected in tri"
assert rel_err < ratio, f"{name}: max abs err {abs_err:.6f}, rmse ratio {rel_err:.6f} >= {ratio}"
# ---------------------------------------------------------------------------
# Test 1: Non-inplace varlen (clean output / state comparison)
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
("H", "D", "cu_seqlens", "dtype"),
[
pytest.param(
*test,
id="H{}-D{}-cu{}-{}".format(*test),
)
for test in [
# Decode: single token per sequence
(32, 128, [0, 1], torch.float16),
(32, 128, [0, 1, 2, 3, 4], torch.float16),
(32, 128, [0, 1, 2, 3, 4, 5, 6, 7, 8], torch.float16),
# Short sequences (multi-token recurrent)
(32, 128, [0, 16], torch.float16),
(32, 128, [0, 8, 24], torch.float16),
(32, 128, [0, 4, 8, 16], torch.float16),
(32, 128, [0, 64], torch.float16),
# Different head count
(64, 128, [0, 1, 2, 3, 4], torch.float16),
# BFloat16
(32, 128, [0, 1, 2, 3, 4], torch.bfloat16),
(32, 128, [0, 8, 24], torch.bfloat16),
]
],
)
@pytest.mark.skip_global_cleanup
@torch.inference_mode()
def test_fused_recurrent_kda(
H: int,
D: int,
cu_seqlens: list[int],
dtype: torch.dtype,
):
"""Non-inplace varlen mode — easiest to verify output and per-token state."""
T = cu_seqlens[-1]
N = len(cu_seqlens) - 1
B = 1
torch.manual_seed(42)
cu_seqlens_t = torch.LongTensor(cu_seqlens).to(DEVICE)
q = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
k = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
v = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
g = F.logsigmoid(torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE)).to(dtype)
beta = torch.rand(B, T, H, dtype=dtype, device=DEVICE).sigmoid()
# Kernel layout: [T, H, V, K] = [T, H, D, D].
# For varlen without ssm_state_indices, seq i reads from h0[cu_seqlens[i]].
h0 = torch.randn(T, H, D, D, dtype=torch.float32, device=DEVICE)
# --- naive reference per sequence ---
ref_outputs = []
ref_states = []
for i in range(N):
s, e = cu_seqlens[i], cu_seqlens[i + 1]
q_i = reference_l2norm(q[:, s:e].contiguous())
k_i = reference_l2norm(k[:, s:e].contiguous())
# Kernel state [H, V, K] -> naive [H, K, V]
init_state_i = h0[s].transpose(-1, -2).unsqueeze(0)
o_i, ht_i = naive_recurrent_kda(
q_i,
k_i,
v[:, s:e],
g[:, s:e],
beta[:, s:e],
initial_state=init_state_i,
output_final_state=True,
)
ref_outputs.append(o_i)
ref_states.append(ht_i)
ref_o = torch.cat(ref_outputs, dim=1)
# --- Triton kernel ---
tri_o, tri_ht = fused_recurrent_kda(
q=q.clone(),
k=k.clone(),
v=v.clone(),
g=g.clone(),
beta=beta.clone(),
initial_state=h0.clone(),
inplace_final_state=False,
use_qk_l2norm_in_kernel=True,
cu_seqlens=cu_seqlens_t,
)
assert not torch.isnan(tri_o).any(), "Triton output o contains NaN"
assert not torch.isnan(tri_ht).any(), "Triton output ht contains NaN"
assert_close("o", ref_o, tri_o, NPU_RMSE_RATIO_O)
# Compare final state per sequence: tri_ht[eos-1] in kernel layout [H,V,K]
for i in range(N):
e = cu_seqlens[i + 1]
tri_state = tri_ht[e - 1].transpose(-1, -2).unsqueeze(0)
assert_close(f"ht_{i}", ref_states[i], tri_state, NPU_RMSE_RATIO_HT)
# ---------------------------------------------------------------------------
# Test 2: Inplace decode with ssm_state_indices (vllm actual pattern)
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
("H", "D", "N", "dtype"),
[
pytest.param(
*test,
id="H{}-D{}-N{}-{}".format(*test),
)
for test in [
(32, 128, 1, torch.float16),
(32, 128, 4, torch.float16),
(32, 128, 16, torch.float16),
(64, 128, 4, torch.float16),
(32, 128, 4, torch.bfloat16),
]
],
)
@pytest.mark.skip_global_cleanup
@torch.inference_mode()
def test_fused_recurrent_kda_decode_inplace(
H: int,
D: int,
N: int,
dtype: torch.dtype,
):
"""Decode with inplace state update + ssm_state_indices — vllm usage."""
B = 1
T = N # one token per sequence
cu_seqlens = list(range(N + 1)) # [0, 1, 2, ..., N]
torch.manual_seed(42)
cu_seqlens_t = torch.LongTensor(cu_seqlens).to(DEVICE)
q = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
k = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
v = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
g = F.logsigmoid(torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE)).to(dtype)
beta = torch.rand(B, T, H, dtype=dtype, device=DEVICE).sigmoid()
# State buffer: slot 0 is NULL (invalid), slots 1..N are valid
max_slots = N + 1
state_buf = torch.randn(max_slots, H, D, D, dtype=torch.float32, device=DEVICE)
state_buf[0] = 0 # NULL slot
# ssm_state_indices: seq i -> slot (i + 1), all valid (> 0)
ssm_state_indices = torch.arange(1, N + 1, dtype=torch.long, device=DEVICE)
# --- naive reference per sequence ---
ref_outputs = []
ref_states = []
for i in range(N):
slot = i + 1
q_i = reference_l2norm(q[:, i : i + 1].contiguous())
k_i = reference_l2norm(k[:, i : i + 1].contiguous())
init_state_i = state_buf[slot].transpose(-1, -2).unsqueeze(0)
o_i, ht_i = naive_recurrent_kda(
q_i,
k_i,
v[:, i : i + 1],
g[:, i : i + 1],
beta[:, i : i + 1],
initial_state=init_state_i,
output_final_state=True,
)
ref_outputs.append(o_i)
ref_states.append(ht_i)
ref_o = torch.cat(ref_outputs, dim=1)
# --- Triton kernel — inplace updates state_buf ---
state_buf_tri = state_buf.clone()
tri_o, _ = fused_recurrent_kda(
q=q.clone(),
k=k.clone(),
v=v.clone(),
g=g.clone(),
beta=beta.clone(),
initial_state=state_buf_tri,
inplace_final_state=True,
use_qk_l2norm_in_kernel=True,
cu_seqlens=cu_seqlens_t,
ssm_state_indices=ssm_state_indices,
)
assert not torch.isnan(tri_o).any(), "Triton output o contains NaN"
assert_close("o", ref_o, tri_o, NPU_RMSE_RATIO_O)
# Verify inplace state update at each slot
for i in range(N):
slot = i + 1
tri_state = state_buf_tri[slot].transpose(-1, -2).unsqueeze(0)
assert_close(f"ht_{i}", ref_states[i], tri_state, NPU_RMSE_RATIO_HT)
# Verify NULL slot was not modified
assert torch.all(state_buf_tri[0] == 0), "NULL slot (0) should not be modified"
# ---------------------------------------------------------------------------
# Test 3: Float32 — isolate algorithmic error from dtype precision
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
("H", "D", "cu_seqlens"),
[
pytest.param(
*test,
id="H{}-D{}-cu{}".format(*test),
)
for test in [
(32, 128, [0, 1]),
(32, 128, [0, 1, 2, 3, 4]),
(32, 128, [0, 16]),
(32, 128, [0, 8, 24]),
]
],
)
@pytest.mark.skip_global_cleanup
@torch.inference_mode()
def test_fused_recurrent_kda_fp32(
H: int,
D: int,
cu_seqlens: list[int],
):
"""Float32 test to isolate algorithmic error from dtype precision."""
T = cu_seqlens[-1]
N = len(cu_seqlens) - 1
B = 1
torch.manual_seed(42)
cu_seqlens_t = torch.LongTensor(cu_seqlens).to(DEVICE)
q = torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE)
k = torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE)
v = torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE)
g = F.logsigmoid(torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE))
beta = torch.rand(B, T, H, dtype=torch.float32, device=DEVICE).sigmoid()
h0 = torch.randn(T, H, D, D, dtype=torch.float32, device=DEVICE)
ref_outputs = []
ref_states = []
for i in range(N):
s, e = cu_seqlens[i], cu_seqlens[i + 1]
q_i = reference_l2norm(q[:, s:e].contiguous())
k_i = reference_l2norm(k[:, s:e].contiguous())
init_state_i = h0[s].transpose(-1, -2).unsqueeze(0)
o_i, ht_i = naive_recurrent_kda(
q_i,
k_i,
v[:, s:e],
g[:, s:e],
beta[:, s:e],
initial_state=init_state_i,
output_final_state=True,
)
ref_outputs.append(o_i)
ref_states.append(ht_i)
ref_o = torch.cat(ref_outputs, dim=1)
tri_o, tri_ht = fused_recurrent_kda(
q=q.clone(),
k=k.clone(),
v=v.clone(),
g=g.clone(),
beta=beta.clone(),
initial_state=h0.clone(),
inplace_final_state=False,
use_qk_l2norm_in_kernel=True,
cu_seqlens=cu_seqlens_t,
)
assert not torch.isnan(tri_o).any(), "Triton output o contains NaN"
assert not torch.isnan(tri_ht).any(), "Triton output ht contains NaN"
assert_close("o", ref_o, tri_o, NPU_RMSE_RATIO_O)
for i in range(N):
e = cu_seqlens[i + 1]
tri_state = tri_ht[e - 1].transpose(-1, -2).unsqueeze(0)
assert_close(f"ht_{i}", ref_states[i], tri_state, NPU_RMSE_RATIO_HT)

View File

@@ -0,0 +1,58 @@
import gc
import torch
from vllm.model_executor.layers.fla.ops import fused_recurrent_gated_delta_rule
from vllm_ascend.ops.triton.fla.sigmoid_gating import fused_sigmoid_gating_delta_rule_update
from vllm_ascend.ops.triton.fused_gdn_gating import fused_gdn_gating_patch
def test_triton_fusion_ops():
q = torch.randn(1, 1, 4, 128, dtype=torch.bfloat16).npu()
k = torch.randn(1, 1, 4, 128, dtype=torch.bfloat16).npu()
v = torch.randn(1, 1, 8, 128, dtype=torch.bfloat16).npu()
a = torch.tensor([[-2.6094, -0.2617, -0.3848, 2.2656, 3.6250, -0.7383, -1.0938, -0.0505]]).bfloat16().npu()
b = torch.tensor([[0.4277, 0.8906, 1.6875, 2.3750, 4.1562, 0.3809, 1.0625, 3.6719]]).bfloat16().npu()
non_spec_state_indices_tensor = torch.tensor([2]).int().npu()
non_spec_query_start_loc = torch.tensor([0, 1]).int().npu()
a_log = torch.tensor([-2.6875, -3.2031, -3.3438, -2.7812, -3.0625, -4.0312, -5.3750, 5.7188]).bfloat16().npu()
dt_bias = torch.tensor([-4.7812, -5.0938, -5.5000, 9.4375, 7.6250, -4.3750, -3.0938, 0.9688]).bfloat16().npu()
ssm_state1 = torch.ones(1, 8, 128, 128, dtype=torch.bfloat16).npu()
core_attn_out_non_spec_fused = fused_sigmoid_gating_delta_rule_update(
A_log=a_log.contiguous(),
dt_bias=dt_bias.contiguous(),
q=q.contiguous(),
k=k.contiguous(),
v=v.contiguous(),
a=a.contiguous(),
b=b.contiguous(),
initial_state_source=ssm_state1,
initial_state_indices=non_spec_state_indices_tensor,
cu_seqlens=non_spec_query_start_loc,
use_qk_l2norm_in_kernel=True,
softplus_beta=1.0,
softplus_threshold=20.0,
)
ssm_state2 = torch.ones(1, 8, 128, 128, dtype=torch.bfloat16).npu()
g, beta = fused_gdn_gating_patch(a_log, a, b, dt_bias)
g_non_spec = g
beta_non_spec = beta
core_attn_out_non_spec_split, last_recurrent_state = fused_recurrent_gated_delta_rule(
q=q,
k=k,
v=v,
g=g_non_spec,
beta=beta_non_spec,
initial_state=ssm_state2,
inplace_final_state=True,
cu_seqlens=non_spec_query_start_loc,
ssm_state_indices=non_spec_state_indices_tensor,
use_qk_l2norm_in_kernel=True,
)
torch.testing.assert_close(
core_attn_out_non_spec_fused, core_attn_out_non_spec_split, rtol=1e-02, atol=1e-02, equal_nan=True
)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,39 @@
import gc
import pytest
import torch
import torch.nn.functional as F
from vllm_ascend.ops.triton.fla.l2norm import l2norm_fwd
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
@pytest.mark.parametrize(
("B", "T", "H", "D", "dtype"),
[
pytest.param(*test, id="B{}-T{}-H{}-D{}-{}".format(*test))
for test in [
(1, 63, 1, 60, torch.float),
(2, 500, 4, 64, torch.float),
(2, 1000, 2, 100, torch.float),
(3, 1024, 4, 128, torch.float),
]
],
)
def test_l2norm(B: int, T: int, H: int, D: int, dtype: torch.dtype):
torch.manual_seed(42)
init_device_properties_triton()
device = "npu"
rtol, atol = (3e-4, 1e-3) if dtype == torch.float32 else (3e-3, 5e-3)
if dtype == torch.bfloat16:
rtol, atol = 1e-2, 5e-2
x = torch.randn(B, T, H, D, dtype=dtype).to(device).requires_grad_(True)
x = x * 0.5 + 0.3
ref = F.normalize(x, dim=-1, p=2)
tri = l2norm_fwd(x)
assert torch.allclose(tri, ref, rtol=rtol, atol=atol)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,930 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
"""Tests for lightning attention triton kernels.
Covers the following 4 triton kernels:
- ``_fwd_diag_kernel``: diagonal block causal attention
- ``_fwd_kv_parallel``: key-value outer product per block
- ``_fwd_kv_reduce``: prefix-sum reduction of KV across blocks
- ``_fwd_none_diag_kernel``: non-diagonal block attention
All kernels are exercised through the public APIs:
- ``lightning_attention_npu_`` (``_attention.apply``, single d-chunk)
- ``lightning_attention_npu`` (full function with d-dimension chunking)
- ``AscendLightningAttentionKernel.jit_linear_forward_prefix``
The naive reference implementations replicate the exact triton tiling
algorithm (diagonal blocks -> KV parallel -> KV reduce -> non-diagonal)
to ensure an apples-to-apples numerical comparison. Low-precision inputs use
bounded random values and mirror the kernel's output stores to avoid comparing
against an unrealistically precise PyTorch recurrence.
The production BailingMoE path promotes QKV to float32 before calling these
kernels and commonly uses a float32 Mamba cache. Therefore larger accuracy
cases use float32, while bf16/fp16 cases are kept as bounded compatibility
coverage for the raw operator entry points.
"""
import gc
import pytest
import torch
from einops import rearrange
from vllm_ascend.ops.triton.mamba.lightning_attn import (
AscendLightningAttentionKernel,
lightning_attention_npu,
lightning_attention_npu_,
)
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
# ---------------------------------------------------------------------------
# Naive reference implementations (pure PyTorch, no Triton)
# ---------------------------------------------------------------------------
# The triton lightning attention algorithm processes the sequence in BLOCK-sized
# tiles and follows these steps:
# 1. _fwd_diag_kernel : causal attention within each diagonal block
# 2. _fwd_kv_parallel : per-block KV outer product, decayed to block end
# 3. _fwd_kv_reduce : prefix-sum across blocks, updates kv_history
# 4. _fwd_none_diag_kernel : non-diagonal attention using prefix KV
#
# IMPORTANT: The non-diagonal kernel applies decay as exp(-s * t_in_block).
# The references below mirror that tiled kernel convention directly rather
# than using an equivalent-looking recurrence with a different history state.
BLOCK_SIZE = 256
DIAG_BLOCK_SIZE = 32
KV_BLOCK_SIZE = 64
TEST_INPUT_SCALE = 0.02
TEST_DECAY_SCALE = 0.02
SEMANTIC_INPUT_SCALE = 0.05
FLOAT32_TOLERANCE = (1e-2, 1e-2)
LOW_PRECISION_TOLERANCE = (5e-2, 5e-2)
def _round_like_kernel_store(x, dtype):
"""Mirror Triton stores to output dtype, then compare in float32."""
if dtype == torch.float32:
return x
return x.to(dtype).float()
def _randn(shape, dtype, device, scale=TEST_INPUT_SCALE):
return (torch.randn(*shape, dtype=dtype, device=device) * scale).to(dtype)
def _rand_decay(h, device, scale=TEST_DECAY_SCALE):
return torch.rand(h, dtype=torch.float32, device=device) * scale
def _naive_triton_lightning_attention(q, k, v, s, kv_history, block_size=BLOCK_SIZE):
"""Step-by-step replication of the triton tiling algorithm for a single
d-chunk.
This function mirrors the four-kernel pipeline used in
``_attention.apply`` (a.k.a. ``lightning_attention_npu_``):
1. Diagonal blocks (within-block causal attention)
2. Per-block KV outer product (decayed to block end)
3. Prefix-sum reduce across blocks (updates kv_history in place)
4. Non-diagonal blocks (cross-block attention with prefix KV)
Args:
q: [b, h, n, d] queries
k: [b, h, n, d] keys
v: [b, h, n, e] values
s: [1, h, 1, 1] per-head decay rates
kv_history: [b, h, d, e] accumulated KV state from previous steps
Returns:
(o, kv_return) where
o: [b, h, n, e] attention output
kv_return: [b, h, num_blocks + 1, d, e] block KV plus final state
"""
output_dtype = q.dtype
b, h, n, d = q.shape
e_dim = v.shape[-1]
q = q.float()
k = k.float()
v = v.float()
decay_rate = s.float().reshape(1, h, 1, 1)
num_blocks = (n + block_size - 1) // block_size
# ---- Step 1: Diagonal blocks ----
# Mirror _fwd_diag_kernel's tiled matmul path. The recurrence is
# mathematically equivalent, but fp16/bf16 accumulation order differs.
o = torch.zeros(b, h, n, e_dim, dtype=torch.float32, device=q.device)
for block_idx in range(num_blocks):
blk_start = block_idx * block_size
blk_end = min(blk_start + block_size, n)
block_len = blk_end - blk_start
for q_start in range(0, block_len, DIAG_BLOCK_SIZE):
q_end = min(q_start + DIAG_BLOCK_SIZE, block_len)
q_block = q[:, :, blk_start + q_start : blk_start + q_end, :]
q_pos = torch.arange(q_start, q_end, dtype=torch.float32, device=q.device)
q_len = q_end - q_start
qkv = torch.zeros(b, h, q_len, e_dim, dtype=torch.float32, device=q.device)
for kv_start in range(0, q_start + DIAG_BLOCK_SIZE, DIAG_BLOCK_SIZE):
kv_end = min(kv_start + DIAG_BLOCK_SIZE, block_len)
if kv_start >= kv_end:
continue
k_block = k[:, :, blk_start + kv_start : blk_start + kv_end, :]
v_block = v[:, :, blk_start + kv_start : blk_start + kv_end, :]
kv_pos = torch.arange(kv_start, kv_end, dtype=torch.float32, device=q.device)
kv_len = kv_end - kv_start
diff = q_pos[:, None] - kv_pos[None, :]
causal_mask = (diff >= 0).reshape(1, 1, q_len, kv_len)
decay = torch.exp(-decay_rate * diff.clamp_min(0).reshape(1, 1, q_len, kv_len))
decay = decay * causal_mask.to(torch.float32)
qk = torch.matmul(q_block, k_block.transpose(-1, -2)) * decay
qkv = qkv + torch.matmul(qk, v_block)
o[:, :, blk_start + q_start : blk_start + q_end, :] = _round_like_kernel_store(qkv, output_dtype)
# ---- Step 2: Per-block KV outer product ----
# For each block, accumulate K^T @ V with each row decayed to the end
# of that block. The last partial block uses the same left-shifted
# CBLOCK layout as _fwd_kv_parallel.
kv_block = torch.zeros(b, h, num_blocks, d, e_dim, dtype=torch.float32, device=q.device)
for block_idx in range(num_blocks):
blk_start = block_idx * block_size
blk_end = min(blk_start + block_size, n)
block_len = blk_end - blk_start
num_kv_blocks = min((block_len + KV_BLOCK_SIZE - 1) // KV_BLOCK_SIZE, block_size // KV_BLOCK_SIZE)
left_shift = num_kv_blocks * KV_BLOCK_SIZE - block_len
decay_start = (block_size // KV_BLOCK_SIZE - num_kv_blocks) * KV_BLOCK_SIZE
for kv_block_idx in range(num_kv_blocks):
row_offsets = torch.arange(KV_BLOCK_SIZE, device=q.device)
source_pos = kv_block_idx * KV_BLOCK_SIZE - left_shift + row_offsets
left_bound = (1 - kv_block_idx) * left_shift
valid = (row_offsets >= left_bound) & (source_pos >= 0) & (source_pos < block_len)
safe_pos = source_pos.clamp(0, max(block_len - 1, 0))
k_block = k[:, :, blk_start + safe_pos, :] * valid.reshape(1, 1, KV_BLOCK_SIZE, 1)
v_block = v[:, :, blk_start + safe_pos, :] * valid.reshape(1, 1, KV_BLOCK_SIZE, 1)
decay_pos = decay_start + kv_block_idx * KV_BLOCK_SIZE + row_offsets
k_decay = torch.exp(-decay_rate * (block_size - 1 - decay_pos).float().reshape(1, 1, KV_BLOCK_SIZE, 1))
weighted_k = k_block * k_decay
kv_block[:, :, block_idx, :, :] = kv_block[:, :, block_idx, :, :] + torch.matmul(
weighted_k.transpose(-1, -2),
v_block,
)
# ---- Step 3: Prefix-sum reduce across blocks ----
# Replicates _fwd_kv_reduce exactly:
# kv_pre starts as the existing kv_history
# For each block i:
# block_decay = exp(-s * block_size)
# store kv_pre into KV[i]
# kv_pre = block_decay * kv_pre + kv_cur
kv = kv_block.clone()
kv_pre = kv_history.clone().float()
for i in range(num_blocks):
blk_size = min(n - i * block_size, block_size)
block_decay = torch.exp(-decay_rate * blk_size) # [1, h, 1, 1]
kv_cur = kv[:, :, i, :, :].clone()
kv[:, :, i, :, :] = kv_pre
kv_pre = block_decay * kv_pre + kv_cur
kv_history_updated = kv_pre
# ---- Step 4: Non-diagonal blocks ----
# O[t] += Q[t] @ kv[block_idx] * exp(-s * t_local)
# Note: the triton kernel uses t_local (NOT t_local + 1),
# matching _fwd_none_diag_kernel's q_decay = exp(-s * (off_c*CBLOCK + c)).
for block_idx in range(num_blocks):
blk_start = block_idx * block_size
blk_end = min(blk_start + block_size, n)
block_len = blk_end - blk_start
for q_start in range(0, block_len, KV_BLOCK_SIZE):
q_end = min(q_start + KV_BLOCK_SIZE, block_len)
q_len = q_end - q_start
q_block = q[:, :, blk_start + q_start : blk_start + q_end, :]
q_pos = torch.arange(q_start, q_end, dtype=torch.float32, device=q.device)
q_decay = torch.exp(-decay_rate * q_pos.reshape(1, 1, q_len, 1))
nondiag = torch.matmul(q_block, kv[:, :, block_idx, :, :]) * q_decay
out_slice = o[:, :, blk_start + q_start : blk_start + q_end, :] + nondiag
o[:, :, blk_start + q_start : blk_start + q_end, :] = _round_like_kernel_store(out_slice, output_dtype)
return _round_like_kernel_store(o, output_dtype), torch.cat([kv, kv_history_updated.unsqueeze(2)], dim=2)
def _naive_lightning_attention_npu(q, k, v, ed, block_size, kv_history):
"""Naive reference that replicates ``lightning_attention_npu``.
Handles the d-dimension chunking in the same way as the real
implementation so that the comparison is apples-to-apples.
"""
d = q.shape[-1]
e_dim = v.shape[-1]
if ed.dim() == 1:
ed = ed.view(1, -1, 1, 1)
m = 128 if d >= 128 else 64
arr = [m * i for i in range(d // m + 1)]
if arr[-1] != d:
arr.append(d)
if kv_history is None:
kv_history = torch.zeros(
(q.shape[0], q.shape[1], d, e_dim),
dtype=torch.float32,
device=q.device,
)
else:
kv_history = kv_history.clone().contiguous().float()
output = torch.zeros(
q.shape[0],
q.shape[1],
q.shape[2],
e_dim,
dtype=torch.float32,
device=q.device,
)
kv_state = None
for i in range(len(arr) - 1):
s_idx = arr[i]
e_idx = arr[i + 1]
q1 = q[..., s_idx:e_idx]
k1 = k[..., s_idx:e_idx]
kv_history_chunk = kv_history[:, :, s_idx:e_idx, :]
o, kv_state = _naive_triton_lightning_attention(
q1,
k1,
v,
ed,
kv_history_chunk,
block_size=block_size,
)
output = _round_like_kernel_store(output + o, q.dtype)
kv_history[:, :, s_idx:e_idx, :] = kv_state[:, :, -1, :, :]
return output, kv_state
def _naive_jit_linear_forward_prefix(q, k, v, kv_caches, slope_rate, block_size, layer_idx=None):
"""Naive reference for ``AscendLightningAttentionKernel.jit_linear_forward_prefix``."""
slope_rate = slope_rate.to(torch.float32)
should_squeeze = q.dim() == 3
if should_squeeze:
q = q.unsqueeze(0)
k = k.unsqueeze(0)
v = v.unsqueeze(0)
b, h, n, d = q.shape
e_dim = v.shape[-1]
if slope_rate.dim() == 1:
ed = slope_rate.view(1, -1, 1, 1)
else:
ed = slope_rate
kv_history = kv_caches.reshape(1, h, d, e_dim).contiguous().float()
output, kv_state = _naive_lightning_attention_npu(
q,
k,
v,
ed,
block_size,
kv_history,
)
# The triton kernel updates kv_caches in-place with the final KV state
kv_caches_out = kv_state[:, :, -1, :, :].reshape_as(kv_caches)
assert output.shape[0] == 1, "batch size must be 1"
result = rearrange(output.squeeze(0), "h n d -> n (h d)")
return result.float(), kv_caches_out
# ---------------------------------------------------------------------------
# Tolerance helpers
# ---------------------------------------------------------------------------
def _get_tolerances(dtype):
"""Return (rtol, atol) appropriate for the given dtype."""
if dtype == torch.float32:
return FLOAT32_TOLERANCE
elif dtype in (torch.float16, torch.bfloat16):
return LOW_PRECISION_TOLERANCE
return FLOAT32_TOLERANCE
# ---------------------------------------------------------------------------
# Tests for lightning_attention_npu_ (single d-chunk, exercises all 4 kernels)
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
("b", "h", "n", "d", "e", "dtype"),
[
pytest.param(*case, id=f"b{case[0]}-h{case[1]}-n{case[2]}-d{case[3]}-e{case[4]}-{str(case[5]).split('.')[-1]}")
for case in [
# Small seq, basic sanity
(1, 4, 32, 128, 64, torch.bfloat16),
# n < BLOCK (256), not aligned
(1, 4, 100, 128, 128, torch.bfloat16),
# float16 dtype
(1, 4, 256, 128, 128, torch.float16),
# float32 dtype
(1, 4, 128, 128, 128, torch.float32),
# batch > 1
(2, 4, 128, 128, 64, torch.bfloat16),
# d = 64 (smaller head dim)
(1, 4, 128, 64, 64, torch.bfloat16),
# e != d
(1, 4, 128, 64, 128, torch.bfloat16),
]
],
)
def test_lightning_attention_npu_single_chunk(b, h, n, d, e, dtype):
"""Test lightning_attention_npu_ against naive PyTorch reference.
This exercises all 4 triton kernels through the _attention.apply path.
All n values are <= BLOCK (256) to ensure single-block correctness.
"""
torch.manual_seed(42)
init_device_properties_triton()
device = "npu"
rtol, atol = _get_tolerances(dtype)
q = _randn((b, h, n, d), dtype, device)
k = _randn((b, h, n, d), dtype, device)
v = _randn((b, h, n, e), dtype, device)
ed = _rand_decay(h, device).view(1, h, 1, 1)
kv_history = torch.zeros(b, h, d, e, dtype=torch.float32, device=device)
# NOTE: Must clone kv_history before the triton call because
# _fwd_kv_reduce modifies it in-place. The naive reference must
# receive the ORIGINAL (pre-modification) value.
o_triton, kv_triton = lightning_attention_npu_(q, k, v, ed, kv_history.clone())
o_ref, _ = _naive_triton_lightning_attention(q, k, v, ed, kv_history)
torch.testing.assert_close(
o_triton.float().cpu(),
o_ref.cpu(),
rtol=rtol,
atol=atol,
)
assert kv_triton.shape == (b, h, (n + BLOCK_SIZE - 1) // BLOCK_SIZE + 1, d, e)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.parametrize(
("b", "h", "n", "d", "e", "dtype"),
[
pytest.param(*case, id=f"b{case[0]}-h{case[1]}-n{case[2]}-d{case[3]}-e{case[4]}-{str(case[5]).split('.')[-1]}")
for case in [
(1, 4, 256, 128, 128, torch.bfloat16),
(2, 4, 128, 128, 64, torch.bfloat16),
]
],
)
def test_lightning_attention_npu_single_chunk_with_kv_history(b, h, n, d, e, dtype):
"""Test lightning_attention_npu_ with non-zero initial KV history."""
torch.manual_seed(42)
init_device_properties_triton()
device = "npu"
rtol, atol = _get_tolerances(dtype)
q = _randn((b, h, n, d), dtype, device)
k = _randn((b, h, n, d), dtype, device)
v = _randn((b, h, n, e), dtype, device)
ed = _rand_decay(h, device).view(1, h, 1, 1)
kv_history = _randn((b, h, d, e), torch.float32, device)
o_triton, kv_triton = lightning_attention_npu_(q, k, v, ed, kv_history.clone())
o_ref, kv_ref = _naive_triton_lightning_attention(q, k, v, ed, kv_history)
torch.testing.assert_close(
o_triton.float().cpu(),
o_ref.cpu(),
rtol=rtol,
atol=atol,
)
torch.testing.assert_close(
kv_triton[:, :, -1, :, :].cpu(),
kv_ref[:, :, -1, :, :].cpu(),
rtol=rtol,
atol=atol,
)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
# ---------------------------------------------------------------------------
# Tests for multi-block sequences (n > BLOCK)
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
("b", "h", "n", "d", "e", "dtype"),
[
pytest.param(*case, id=f"b{case[0]}-h{case[1]}-n{case[2]}-d{case[3]}-e{case[4]}-{str(case[5]).split('.')[-1]}")
for case in [
# n > BLOCK, not aligned
(1, 4, 300, 128, 128, torch.bfloat16),
# Larger n, production path promotes q/k/v to float32.
(1, 4, 768, 128, 128, torch.float32),
]
],
)
def test_lightning_attention_npu_multi_block(b, h, n, d, e, dtype):
"""Test lightning_attention_npu_ with multi-block sequences (n > BLOCK).
This exercises the _fwd_kv_parallel, _fwd_kv_reduce, and
_fwd_none_diag_kernel in the multi-block path.
Uses small decay rates so inter-block numerical differences
remain within tolerance.
"""
torch.manual_seed(42)
init_device_properties_triton()
device = "npu"
rtol, atol = _get_tolerances(dtype)
q = _randn((b, h, n, d), dtype, device)
k = _randn((b, h, n, d), dtype, device)
v = _randn((b, h, n, e), dtype, device)
# Use small decay rates to keep errors within tolerance
ed = _rand_decay(h, device, scale=0.01)
ed = ed.view(1, h, 1, 1)
kv_history = torch.zeros(b, h, d, e, dtype=torch.float32, device=device)
o_triton, kv_triton = lightning_attention_npu_(q, k, v, ed, kv_history.clone())
o_ref, kv_ref = _naive_triton_lightning_attention(q, k, v, ed, kv_history)
torch.testing.assert_close(
o_triton.float().cpu(),
o_ref.cpu(),
rtol=rtol,
atol=atol,
)
torch.testing.assert_close(
kv_triton[:, :, -1, :, :].cpu(),
kv_ref[:, :, -1, :, :].cpu(),
rtol=rtol,
atol=atol,
)
# Also verify the output is well-formed
assert not torch.isnan(o_triton).any(), "Output contains NaN values"
assert not torch.isinf(o_triton).any(), "Output contains Inf values"
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
# ---------------------------------------------------------------------------
# Tests for lightning_attention_npu (full function with d-dimension chunking)
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
("b", "h", "n", "d", "e", "dtype"),
[
pytest.param(*case, id=f"b{case[0]}-h{case[1]}-n{case[2]}-d{case[3]}-e{case[4]}-{str(case[5]).split('.')[-1]}")
for case in [
# d=128 -> single chunk (m=128)
(1, 4, 128, 128, 128, torch.bfloat16),
# d=64 -> single chunk (m=64)
(1, 4, 128, 64, 64, torch.bfloat16),
# n not aligned to BLOCK, single-block
(1, 4, 100, 128, 128, torch.bfloat16),
# n == BLOCK, production path promotes q/k/v to float32.
(1, 4, 256, 128, 128, torch.float32),
]
],
)
def test_lightning_attention_npu(b, h, n, d, e, dtype):
"""Test lightning_attention_npu (with d-dimension chunking) against naive reference."""
torch.manual_seed(42)
init_device_properties_triton()
device = "npu"
rtol, atol = _get_tolerances(dtype)
q = _randn((b, h, n, d), dtype, device)
k = _randn((b, h, n, d), dtype, device)
v = _randn((b, h, n, e), dtype, device)
ed = _rand_decay(h, device)
# Triton output
o_triton, kv_triton = lightning_attention_npu(q, k, v, ed, block_size=256, kv_history=None)
# Naive reference output
o_ref, kv_ref = _naive_lightning_attention_npu(q, k, v, ed, block_size=256, kv_history=None)
torch.testing.assert_close(
o_triton.float().cpu(),
o_ref.cpu(),
rtol=rtol,
atol=atol,
)
torch.testing.assert_close(
kv_triton[:, :, -1, :, :].cpu(),
kv_ref[:, :, -1, :, :].cpu(),
rtol=rtol,
atol=atol,
)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.parametrize(
("b", "h", "n", "d", "e", "dtype"),
[
pytest.param(*case, id=f"b{case[0]}-h{case[1]}-n{case[2]}-d{case[3]}-e{case[4]}-{str(case[5]).split('.')[-1]}")
for case in [
# Production path promotes q/k/v to float32 while cache is float32.
(1, 4, 128, 128, 128, torch.float32),
(1, 4, 256, 64, 128, torch.bfloat16),
]
],
)
def test_lightning_attention_npu_with_kv_history(b, h, n, d, e, dtype):
"""Test lightning_attention_npu with pre-existing KV history."""
torch.manual_seed(42)
init_device_properties_triton()
device = "npu"
rtol, atol = _get_tolerances(dtype)
q = _randn((b, h, n, d), dtype, device)
k = _randn((b, h, n, d), dtype, device)
v = _randn((b, h, n, e), dtype, device)
ed = _rand_decay(h, device)
kv_history = _randn((b, h, d, e), torch.float32, device)
o_triton, kv_triton = lightning_attention_npu(q, k, v, ed, block_size=256, kv_history=kv_history.clone())
o_ref, kv_ref = _naive_lightning_attention_npu(q, k, v, ed, block_size=256, kv_history=kv_history)
torch.testing.assert_close(
o_triton.float().cpu(),
o_ref.cpu(),
rtol=rtol,
atol=atol,
)
torch.testing.assert_close(
kv_triton[:, :, -1, :, :].cpu(),
kv_ref[:, :, -1, :, :].cpu(),
rtol=rtol,
atol=atol,
)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
# ---------------------------------------------------------------------------
# Tests for AscendLightningAttentionKernel.jit_linear_forward_prefix
# ---------------------------------------------------------------------------
@pytest.mark.parametrize(
("h", "n", "d", "e", "dtype"),
[
pytest.param(*case, id=f"h{case[0]}-n{case[1]}-d{case[2]}-e{case[3]}-{str(case[4]).split('.')[-1]}")
for case in [
(4, 128, 128, 128, torch.bfloat16),
(8, 256, 128, 128, torch.float32),
(4, 100, 128, 128, torch.bfloat16),
(4, 256, 64, 64, torch.bfloat16),
(4, 128, 128, 128, torch.float16),
]
],
)
def test_ascend_lightning_attention_kernel_prefix(h, n, d, e, dtype):
"""Test AscendLightningAttentionKernel.jit_linear_forward_prefix."""
torch.manual_seed(42)
init_device_properties_triton()
device = "npu"
rtol, atol = _get_tolerances(dtype)
# jit_linear_forward_prefix receives one sequence in [h, n, d] layout.
q = _randn((h, n, d), dtype, device)
k = _randn((h, n, d), dtype, device)
v = _randn((h, n, e), dtype, device)
slope_rate = _rand_decay(h, device)
kv_caches = torch.zeros(h, d, e, dtype=torch.float32, device=device)
# Triton output (clone all shared tensors)
kv_caches_triton = kv_caches.clone()
out_triton = AscendLightningAttentionKernel.jit_linear_forward_prefix(
q.clone(), k.clone(), v.clone(), kv_caches_triton, slope_rate.clone(), block_size=256
)
# Naive reference output
out_ref, kv_caches_ref = _naive_jit_linear_forward_prefix(q, k, v, kv_caches, slope_rate, block_size=256)
torch.testing.assert_close(
out_triton.float().cpu(),
out_ref.cpu(),
rtol=rtol,
atol=atol,
)
torch.testing.assert_close(
kv_caches_triton.cpu(),
kv_caches_ref.cpu(),
rtol=rtol,
atol=atol,
)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.parametrize(
("h", "n", "d", "e", "dtype"),
[
pytest.param(*case, id=f"h{case[0]}-n{case[1]}-d{case[2]}-e{case[3]}-{str(case[4]).split('.')[-1]}")
for case in [
# Production path promotes q/k/v to float32 while cache is float32.
(4, 128, 128, 128, torch.float32),
]
],
)
def test_ascend_lightning_attention_kernel_prefix_with_history(h, n, d, e, dtype):
"""Test AscendLightningAttentionKernel.jit_linear_forward_prefix with
non-zero initial kv_caches."""
torch.manual_seed(42)
init_device_properties_triton()
device = "npu"
rtol, atol = _get_tolerances(dtype)
q = _randn((h, n, d), dtype, device)
k = _randn((h, n, d), dtype, device)
v = _randn((h, n, e), dtype, device)
slope_rate = _rand_decay(h, device)
kv_caches = _randn((h, d, e), torch.float32, device)
kv_caches_triton = kv_caches.clone()
out_triton = AscendLightningAttentionKernel.jit_linear_forward_prefix(
q.clone(), k.clone(), v.clone(), kv_caches_triton, slope_rate.clone(), block_size=256
)
out_ref, kv_caches_ref = _naive_jit_linear_forward_prefix(q, k, v, kv_caches, slope_rate, block_size=256)
torch.testing.assert_close(
out_triton.float().cpu(),
out_ref.cpu(),
rtol=rtol,
atol=atol,
)
torch.testing.assert_close(
kv_caches_triton.cpu(),
kv_caches_ref.cpu(),
rtol=rtol,
atol=atol,
)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.parametrize(
("h", "n", "d", "e"),
[
pytest.param(*case, id=f"h{case[0]}-n{case[1]}-d{case[2]}-e{case[3]}")
for case in [
# Multi-block sequence via prefix kernel
(4, 512, 128, 128),
]
],
)
def test_ascend_lightning_attention_kernel_prefix_multi_block(h, n, d, e):
"""Smoke test the prefix wrapper with multi-block sequences."""
torch.manual_seed(42)
init_device_properties_triton()
device = "npu"
dtype = torch.bfloat16
q = _randn((h, n, d), dtype, device)
k = _randn((h, n, d), dtype, device)
v = _randn((h, n, e), dtype, device)
# Small decay rates for multi-block tolerance
slope_rate = _rand_decay(h, device, scale=0.01)
kv_caches = torch.zeros(h, d, e, dtype=torch.float32, device=device)
kv_caches_triton = kv_caches.clone()
out_triton = AscendLightningAttentionKernel.jit_linear_forward_prefix(
q.clone(), k.clone(), v.clone(), kv_caches_triton, slope_rate.clone(), block_size=256
)
assert out_triton.shape == (n, h * e)
assert not torch.isnan(out_triton).any(), "Output contains NaN values"
assert not torch.isinf(out_triton).any(), "Output contains Inf values"
assert kv_caches_triton.shape == (h, d, e)
assert not torch.isnan(kv_caches_triton).any(), "KV cache contains NaN values"
assert not torch.isinf(kv_caches_triton).any(), "KV cache contains Inf values"
assert not torch.allclose(kv_caches_triton, kv_caches), "KV cache should be updated"
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
# ---------------------------------------------------------------------------
# Output shape and basic sanity tests
# ---------------------------------------------------------------------------
def test_lightning_attention_npu_output_shapes():
"""Verify output tensor shapes match expected shapes."""
torch.manual_seed(42)
init_device_properties_triton()
device = "npu"
b, h, n, d, e = 1, 4, 256, 128, 128
q = _randn((b, h, n, d), torch.bfloat16, device)
k = _randn((b, h, n, d), torch.bfloat16, device)
v = _randn((b, h, n, e), torch.bfloat16, device)
ed = _rand_decay(h, device)
o, kv = lightning_attention_npu(q, k, v, ed, block_size=256, kv_history=None)
assert o.shape == (b, h, n, e), f"Expected output shape {(b, h, n, e)}, got {o.shape}"
assert kv.dim() == 5, f"Expected kv_return to be 5D, got {kv.dim()}D"
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
def test_ascend_kernel_prefix_output_shape():
"""Verify AscendLightningAttentionKernel.jit_linear_forward_prefix output shape."""
torch.manual_seed(42)
init_device_properties_triton()
device = "npu"
h, n, d, e = 4, 256, 128, 128
q = _randn((h, n, d), torch.bfloat16, device)
k = _randn((h, n, d), torch.bfloat16, device)
v = _randn((h, n, e), torch.bfloat16, device)
slope_rate = _rand_decay(h, device)
kv_caches = torch.zeros(h, d, e, dtype=torch.float32, device=device)
out = AscendLightningAttentionKernel.jit_linear_forward_prefix(q, k, v, kv_caches, slope_rate, block_size=256)
expected_shape = (n, h * e)
assert out.shape == expected_shape, f"Expected output shape {expected_shape}, got {out.shape}"
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
def test_lightning_attention_output_no_nan():
"""Verify the output does not contain NaN or Inf values."""
torch.manual_seed(42)
init_device_properties_triton()
device = "npu"
b, h, n, d, e = 1, 4, 256, 128, 128
q = _randn((b, h, n, d), torch.bfloat16, device)
k = _randn((b, h, n, d), torch.bfloat16, device)
v = _randn((b, h, n, e), torch.bfloat16, device)
ed = _rand_decay(h, device).view(1, h, 1, 1)
o, _ = lightning_attention_npu_(q, k, v, ed, torch.zeros(b, h, d, e, dtype=torch.float32, device=device))
assert not torch.isnan(o).any(), "Output contains NaN values"
assert not torch.isinf(o).any(), "Output contains Inf values"
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
def test_lightning_attention_causal_property():
"""Verify causal property: output at position t should not depend on
keys/values at positions > t.
If we modify V at positions after t, the output at position t should
remain unchanged. This tests the diagonal kernel's causal mask.
Uses float32 for numerical precision.
"""
torch.manual_seed(42)
init_device_properties_triton()
device = "npu"
b, h, n, d, e = 1, 4, 128, 64, 64
q = _randn((b, h, n, d), torch.float32, device, scale=SEMANTIC_INPUT_SCALE)
k = _randn((b, h, n, d), torch.float32, device, scale=SEMANTIC_INPUT_SCALE)
v = _randn((b, h, n, e), torch.float32, device, scale=SEMANTIC_INPUT_SCALE)
ed = _rand_decay(h, device).view(1, h, 1, 1)
kv_history = torch.zeros(b, h, d, e, dtype=torch.float32, device=device)
# Output with original V
o_orig, _ = lightning_attention_npu_(q, k, v, ed, kv_history.clone())
# Scramble V at positions after t=50
v_scrambled = v.clone()
v_scrambled[:, :, 50:, :] = _randn(v_scrambled[:, :, 50:, :].shape, torch.float32, device)
# Output with scrambled V (should be identical at positions 0..49)
o_scrambled, _ = lightning_attention_npu_(q, k, v_scrambled, ed, kv_history.clone())
# Output at positions 0..49 should be unchanged
torch.testing.assert_close(
o_orig[:, :, :50, :].cpu(),
o_scrambled[:, :, :50, :].cpu(),
rtol=1e-5,
atol=1e-5,
)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
def test_lightning_attention_decay_effect():
"""Verify that different decay rates produce different outputs.
With a larger decay rate, the output is more dominated by recent tokens.
This test verifies the decay mechanism is active, using the naive
reference as ground truth since both should produce valid outputs.
"""
torch.manual_seed(42)
init_device_properties_triton()
device = "npu"
b, h, n, d, e = 1, 4, 128, 64, 64
q = _randn((b, h, n, d), torch.float32, device, scale=SEMANTIC_INPUT_SCALE)
k = _randn((b, h, n, d), torch.float32, device, scale=SEMANTIC_INPUT_SCALE)
v = _randn((b, h, n, e), torch.float32, device, scale=SEMANTIC_INPUT_SCALE)
kv_history = torch.zeros(b, h, d, e, dtype=torch.float32, device=device)
# Small decay rate (slow decay, long memory)
ed_small = torch.full((1, h, 1, 1), 0.01, dtype=torch.float32, device=device)
o_small, _ = lightning_attention_npu_(q, k, v, ed_small, kv_history.clone())
# Large decay rate (fast decay, short memory)
ed_large = torch.full((1, h, 1, 1), 1.0, dtype=torch.float32, device=device)
o_large, _ = lightning_attention_npu_(q, k, v, ed_large, kv_history.clone())
# Verify both produce valid (non-NaN, non-Inf) outputs
assert not torch.isnan(o_small).any(), "Small decay output contains NaN"
assert not torch.isinf(o_small).any(), "Small decay output contains Inf"
assert not torch.isnan(o_large).any(), "Large decay output contains NaN"
assert not torch.isinf(o_large).any(), "Large decay output contains Inf"
# The outputs should differ because different decay rates produce
# different attention distributions.
assert not torch.allclose(o_small, o_large, rtol=1e-3, atol=1e-3), (
"Different decay rates should produce different outputs"
)
# Verify both outputs match naive reference
o_ref_small, _ = _naive_triton_lightning_attention(q, k, v, ed_small, kv_history)
o_ref_large, _ = _naive_triton_lightning_attention(q, k, v, ed_large, kv_history)
torch.testing.assert_close(
o_small.cpu(),
o_ref_small.cpu(),
rtol=1e-3,
atol=1e-3,
)
torch.testing.assert_close(
o_large.cpu(),
o_ref_large.cpu(),
rtol=1e-3,
atol=1e-3,
)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,64 @@
import pytest
import torch
from vllm.triton_utils import triton
from vllm_ascend.worker.v2.sample.logprob import _topk_log_softmax_kernel
@pytest.mark.parametrize(
"batch_size,vocab_size,num_logprobs",
[
(48, 102400, 50),
(96, 102400, 1),
(24, 151936, 8),
],
)
def test_topk_log_softmax_kernel(batch_size, vocab_size, num_logprobs):
"""Test _topk_log_softmax_kernel for computing log probabilities
Args:
batch_size: Number of sequences in the batch
vocab_size: Size of the vocabulary
num_logprobs: Number of tokens to compute log probabilities for
"""
# ========== Setup test data ==========
torch.manual_seed(42)
# Generate random logits
logits = torch.randn(batch_size, vocab_size, device="npu", dtype=torch.float32)
# Generate token_ids for which to compute logprobs
token_ids = torch.randint(0, vocab_size, (batch_size, num_logprobs), device="npu", dtype=torch.int64)
# ========== Execute test ==========
# Prepare output tensor
triton_output = torch.empty(batch_size, num_logprobs, dtype=torch.float32, device="npu")
# Invoke Triton kernel
_topk_log_softmax_kernel[(batch_size,)](
triton_output,
logits,
logits.stride(0),
token_ids,
num_logprobs,
vocab_size,
BLOCK_SIZE=1024,
PADDED_TOPK=max(triton.next_power_of_2(num_logprobs), 2),
)
torch.npu.synchronize()
# Compute reference values using PyTorch
torch_logprobs = torch.log_softmax(logits, dim=-1)
# Extract logprobs for each batch and token_id
ref_output = torch.zeros_like(triton_output)
for i in range(batch_size):
for j in range(num_logprobs):
token_id = token_ids[i, j]
ref_output[i, j] = torch_logprobs[i, token_id]
# ========== Verify results ==========
assert torch.allclose(triton_output, ref_output, rtol=1e-3, atol=1e-3), (
f"Triton output differs from PyTorch reference.\n"
f"Max diff: {torch.max(torch.abs(triton_output - ref_output))}\n"
f"Mean diff: {torch.mean(torch.abs(triton_output - ref_output))}"
)

View File

@@ -0,0 +1,91 @@
import pytest
import torch
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
from vllm_ascend.worker.v2.sample.min_p import apply_min_p
def torch_min_p_torch(
logits: torch.Tensor,
expanded_idx_mapping: torch.Tensor,
min_p: torch.Tensor,
):
num_tokens, _ = logits.shape
out = logits.clone()
for token_idx in range(num_tokens):
req_state_idx = expanded_idx_mapping[token_idx].item()
min_p_val = min_p[req_state_idx].item()
if min_p_val == 0.0:
continue
token_logits = out[token_idx]
max_val = token_logits.max()
threshold = max_val + torch.log(torch.tensor(min_p_val, device=logits.device))
token_logits = torch.where(token_logits < threshold, -torch.inf, token_logits)
out[token_idx] = token_logits
return out
@pytest.mark.parametrize(
"num_reqs,vocab_size",
[
(48, 102400),
(96, 102400),
(24, 151936),
(1, 32000),
],
)
def test_apply_min_p_kernel(num_reqs, vocab_size):
"""Test apply_min_p for computing Min-P sampling mask
Args:
num_reqs: Number of sequences in the batch
vocab_size: Size of the vocabulary
"""
init_device_properties_triton()
# ========== Setup test data ==========
torch.manual_seed(42)
# Generate random logits (using float32 as specified in your kernel)
device = "npu"
original_logits = torch.randn(num_reqs, vocab_size, device=device, dtype=torch.float32)
triton_logits = original_logits.clone()
ref_logits = original_logits.clone()
expanded_idx_mapping = torch.arange(num_reqs - 1, -1, -1, device=device, dtype=torch.int32)
# Generate random min_p values (valid range is typically (0, 1.0])
min_p = torch.empty(num_reqs, device=device, dtype=torch.float32).uniform_(0.01, 0.5)
# ========== Execute test ==========
# 1. Invoke your Triton kernel wrapper
apply_min_p(triton_logits, expanded_idx_mapping, min_p)
torch.npu.synchronize()
# 2. Compute reference values using PyTorch
ref_logits = torch_min_p_torch(
ref_logits,
expanded_idx_mapping,
min_p,
)
# ========== Verify results ==========
triton_inf_mask = torch.isinf(triton_logits)
ref_inf_mask = torch.isinf(ref_logits)
assert torch.equal(triton_inf_mask, ref_inf_mask), (
"Masked positions (where logits == -inf) do not match between Triton and PyTorch."
)
valid_triton_logits = triton_logits[~triton_inf_mask]
valid_ref_logits = ref_logits[~ref_inf_mask]
assert torch.allclose(valid_triton_logits, valid_ref_logits, rtol=1e-4, atol=1e-4), (
f"Logits values differ between Triton and PyTorch reference.\n"
f"Max diff: {torch.max(torch.abs(valid_triton_logits - valid_ref_logits))}"
)

View File

@@ -0,0 +1,166 @@
import gc
import pytest
import torch
from vllm.model_executor.layers.rotary_embedding.mrope import triton_mrope
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
MROPE_SECTION = [[32, 32, 32]]
DTYPES = [torch.bfloat16, torch.float16]
HEAD_SIZES = [128]
ROTARY_DIMS = [128]
NUM_Q_HEADS = [64]
NUM_K_HEADS = [1]
NUM_TOKENS = [1, 4, 8, 16]
SEEDS = [0]
DEVICES = [f"npu:{0}"]
DEFAULT_ATOL = 1e-3
DEFAULT_RTOL = 1e-3
def pytorch_forward_native(q, k, cos, sin, mrope_section, head_size, rotary_dim, mrope_interleaved):
"""PyTorch-native implementation equivalent to forward()."""
num_tokens = q.shape[0]
n_q_head = q.shape[1] // head_size
n_kv_head = k.shape[1] // head_size
q_reshaped = q.view(num_tokens, n_q_head, head_size)
k_reshaped = k.view(num_tokens, n_kv_head, head_size)
cos_reshaped = cos.permute(1, 2, 0)
sin_reshaped = sin.permute(1, 2, 0)
half_rd = rotary_dim // 2
for token_idx in range(num_tokens):
token_cos = cos_reshaped[token_idx]
token_sin = sin_reshaped[token_idx]
cos_row = torch.zeros(head_size // 2, device=q.device, dtype=q.dtype)
sin_row = torch.zeros(head_size // 2, device=q.device, dtype=q.dtype)
if mrope_interleaved:
cos_offsets = torch.arange(0, head_size // 2, device=q.device)
h_mask = ((cos_offsets % 3) == 1) & (cos_offsets <= 3 * mrope_section[1])
w_mask = ((cos_offsets % 3) == 2) & (cos_offsets <= 3 * mrope_section[2])
t_mask = ~(h_mask | w_mask)
cos_row[t_mask] = token_cos[t_mask, 0]
cos_row[h_mask] = token_cos[h_mask, 1]
cos_row[w_mask] = token_cos[w_mask, 2]
sin_row[t_mask] = token_sin[t_mask, 0]
sin_row[h_mask] = token_sin[h_mask, 1]
sin_row[w_mask] = token_sin[w_mask, 2]
else:
t_end = mrope_section[0]
h_end = t_end + mrope_section[1]
if t_end > 0:
cos_row[:t_end] = token_cos[:t_end, 0]
sin_row[:t_end] = token_sin[:t_end, 0]
if mrope_section[1] > 0:
cos_row[t_end:h_end] = token_cos[t_end:h_end, 1]
sin_row[t_end:h_end] = token_sin[t_end:h_end, 1]
if mrope_section[2] > 0:
w_start = h_end
cos_row[w_start:half_rd] = token_cos[w_start:half_rd, 2]
sin_row[w_start:half_rd] = token_sin[w_start:half_rd, 2]
q_token = q_reshaped[token_idx]
k_token = k_reshaped[token_idx]
q1 = q_token[:, :half_rd]
q2 = q_token[:, half_rd:]
k1 = k_token[:, :half_rd]
k2 = k_token[:, half_rd:]
cos_half = cos_row.unsqueeze(0)
sin_half = sin_row.unsqueeze(0)
new_q1 = q1 * cos_half - q2 * sin_half
new_q2 = q2 * cos_half + q1 * sin_half
new_k1 = k1 * cos_half - k2 * sin_half
new_k2 = k2 * cos_half + k1 * sin_half
q_reshaped[token_idx] = torch.cat([new_q1, new_q2], dim=1)
k_reshaped[token_idx] = torch.cat([new_k1, new_k2], dim=1)
q_result = q_reshaped.view(num_tokens, -1)
k_result = k_reshaped.view(num_tokens, -1)
return q_result, k_result
def create_test_data(num_tokens, n_q_head, n_kv_head, rotary_dim, head_size, device, dtype):
q = torch.randn(num_tokens, n_q_head * head_size, dtype=dtype, device=device)
k = torch.randn(num_tokens, n_kv_head * head_size, dtype=dtype, device=device)
sin = torch.randn(3, num_tokens, rotary_dim // 2, dtype=dtype, device=device)
cos = torch.randn(3, num_tokens, rotary_dim // 2, dtype=dtype, device=device)
norm = torch.sqrt(cos**2 + sin**2)
cos = cos / norm
sin = sin / norm
return q, k, cos, sin
@pytest.mark.parametrize("mrope_section", MROPE_SECTION)
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
@pytest.mark.parametrize("num_q_heads", NUM_Q_HEADS)
@pytest.mark.parametrize("num_k_heads", NUM_K_HEADS)
@pytest.mark.parametrize("head_size", HEAD_SIZES)
@pytest.mark.parametrize("rotary_dim", ROTARY_DIMS)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("seed", SEEDS)
@pytest.mark.parametrize("device", DEVICES)
@torch.inference_mode()
def test_mrotary_embedding_triton_kernel(
mrope_section: list[int],
num_tokens: int,
num_q_heads: int,
num_k_heads: int,
head_size: int,
rotary_dim: int,
dtype: torch.dtype,
seed: int,
device: str,
) -> None:
torch.manual_seed(seed)
torch.set_default_device(device)
init_device_properties_triton()
if rotary_dim == -1:
rotary_dim = head_size
q_trt, k_trt, cos, sin = create_test_data(
num_tokens=num_tokens,
n_q_head=num_q_heads,
n_kv_head=num_k_heads,
head_size=head_size,
rotary_dim=rotary_dim,
device=device,
dtype=dtype,
)
q_gold, k_gold = q_trt.clone(), k_trt.clone()
q_trt, k_trt = triton_mrope(q_trt, k_trt, cos, sin, mrope_section, head_size, rotary_dim, True)
q_gold, k_gold = pytorch_forward_native(q_gold, k_gold, cos, sin, mrope_section, head_size, rotary_dim, True)
atol = DEFAULT_ATOL
rtol = DEFAULT_RTOL
if dtype == torch.bfloat16:
atol = 1e-02
rtol = 1e-02
# Compare the results.
torch.testing.assert_close(q_trt.view(q_gold.size()), q_gold, atol=atol, rtol=rtol)
torch.testing.assert_close(k_trt.view(k_gold.size()), k_gold, atol=atol, rtol=rtol)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,33 @@
import pytest
import torch
from vllm_ascend.ops.triton.muls_add import muls_add_triton
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
@pytest.mark.parametrize(
("shape", "dtype", "scale"),
[
((1, 2048), torch.float16, 1.25),
((4000, 2048), torch.float16, 0.75),
((4, 2048), torch.bfloat16, 1.0),
],
)
@torch.inference_mode()
def test_muls_add_triton_correctness(shape, dtype, scale):
"""compare the correctness of muls_add_triton with the PyTorch baseline implementation."""
init_device_properties_triton()
device = "npu"
torch.manual_seed(0)
x = torch.randn(*shape, dtype=dtype, device=device)
y = torch.randn(*shape, dtype=dtype, device=device)
out_triton = muls_add_triton(x, y, scale)
out_ref = x * scale + y
rtol, atol = 1e-3, 1e-3
assert out_triton.shape == out_ref.shape
assert out_triton.dtype == out_ref.dtype
assert torch.allclose(out_triton, out_ref, rtol=rtol, atol=atol)

View File

@@ -0,0 +1,246 @@
import gc
import pytest
import torch
from vllm_ascend.worker.v2.sample.penalties import apply_penalties
NUM_TOKENS = [1, 4]
VOCAB_SIZE = [1000]
NUM_STATUS = [1, 4]
NUM_SPECULATIVE_TOKENS = [0, 1, 3]
DTYPES = [torch.bfloat16, torch.float16]
SEEDS = [42]
DEVICES = [f"npu:{0}"]
DEFAULT_ATOL = 1e-3
DEFAULT_RTOL = 1e-3
def pytorch_apply_penalties(
logits: torch.Tensor,
idx_mapping: torch.Tensor,
token_ids: torch.Tensor,
expanded_local_pos: torch.Tensor,
repetition_penalty: torch.Tensor,
frequency_penalty: torch.Tensor,
presence_penalty: torch.Tensor,
prompt_bin_mask: torch.Tensor,
output_bin_counts: torch.Tensor,
) -> torch.Tensor:
"""
Pytorch equivalent implementation
"""
num_tokens, vocab_size = logits.shape
device = logits.device
dtype = logits.dtype
logits_float = logits.float()
num_status = prompt_bin_mask.shape[0]
num_packed = prompt_bin_mask.shape[1]
prompt_masks_unpacked = torch.zeros(num_status, vocab_size, dtype=torch.bool, device=device)
for state_idx in range(num_status):
for packed_idx in range(num_packed):
packed_val = prompt_bin_mask[state_idx, packed_idx].item()
if packed_val == 0:
continue
start_idx = packed_idx * 32
end_idx = min(start_idx + 32, vocab_size)
for bit_pos in range(end_idx - start_idx):
if (packed_val >> bit_pos) & 1:
prompt_masks_unpacked[state_idx, start_idx + bit_pos] = True
start_idx_in_batch = torch.arange(num_tokens, device=device) - expanded_local_pos
for token_idx in range(num_tokens):
req_state_idx = idx_mapping[token_idx].item()
rep_penalty = repetition_penalty[req_state_idx].item()
freq_penalty = frequency_penalty[req_state_idx].item()
pres_penalty = presence_penalty[req_state_idx].item()
use_rep_penalty = rep_penalty != 1.0
use_freq_penalty = freq_penalty != 0.0
use_pres_penalty = pres_penalty != 0.0
use_penalty = use_rep_penalty or use_freq_penalty or use_pres_penalty
if not use_penalty:
continue
current_prompt_mask = prompt_masks_unpacked[req_state_idx]
base_counts = output_bin_counts[req_state_idx].clone()
# Compute cumulative draft counts
pos = expanded_local_pos[token_idx].item()
draft_counts = torch.zeros(vocab_size, device=device, dtype=torch.int32)
for prev_pos in range(pos):
prev_token_idx = start_idx_in_batch[token_idx] + prev_pos + 1
if 0 <= prev_token_idx < num_tokens:
prev_token = token_ids[prev_token_idx].item()
if 0 <= prev_token < vocab_size:
draft_counts[prev_token] += 1
# Total counts = base output counts + cumulative draft counts
total_counts = base_counts + draft_counts
output_bin_mask = total_counts > 0
if use_rep_penalty:
need_scale = current_prompt_mask | output_bin_mask
scale = torch.where(need_scale, rep_penalty, 1.0)
pos_mask = logits_float[token_idx] > 0
scale_factor = torch.where(pos_mask, 1.0 / scale, scale)
logits_float[token_idx] *= scale_factor
if use_freq_penalty:
logits_float[token_idx] -= freq_penalty * total_counts.float()
if use_pres_penalty:
logits_float[token_idx] -= pres_penalty * output_bin_mask.float()
return logits_float.to(dtype)
def create_test_data(
num_tokens: int = 8,
vocab_size: int = 51200,
num_status: int = 16,
num_speculative_tokens: int = 3,
device: str = "npu",
dtype: torch.dtype = torch.bfloat16,
seed: int = 42,
):
"""Create test data for penalties"""
torch.manual_seed(seed)
logits = torch.randn(num_tokens, vocab_size, device=device, dtype=dtype)
repetition_penalty = torch.ones(num_status, device=device, dtype=torch.float32)
for i in range(num_status):
if torch.rand(1) > 0.3:
repetition_penalty[i] = torch.rand(1, device=device).item() * 0.8 + 0.6
frequency_penalty = torch.zeros(num_status, device=device, dtype=torch.float32)
for i in range(num_status):
if torch.rand(1) > 0.5:
frequency_penalty[i] = torch.rand(1, device=device).item() * 0.2
presence_penalty = torch.zeros(num_status, device=device, dtype=torch.float32)
for i in range(num_status):
if torch.rand(1) > 0.5:
presence_penalty[i] = torch.rand(1, device=device).item() * 0.2
idx_mapping = torch.randint(0, num_status, (num_tokens,), device=device, dtype=torch.int32)
# Create token_ids for speculative decoding
token_ids = torch.randint(0, vocab_size, (num_tokens,), device=device, dtype=torch.int32)
# Create expanded_local_pos (position within speculative decoding window)
expanded_local_pos = torch.zeros(num_tokens, device=device, dtype=torch.int32)
for i in range(num_tokens):
expanded_local_pos[i] = torch.randint(0, num_speculative_tokens + 1, (1,)).item()
num_packed = (vocab_size + 31) // 32
prompt_bin_mask = torch.zeros(num_status, num_packed, device=device, dtype=torch.int32)
for state_idx in range(num_status):
num_tokens_in_prompt = max(1, vocab_size // 20)
prompt_tokens = torch.randperm(vocab_size, device=device)[:num_tokens_in_prompt]
for token_id in prompt_tokens:
packed_idx = token_id // 32
bit_pos = token_id % 32
prompt_bin_mask[state_idx, packed_idx] |= 1 << bit_pos
output_bin_counts = torch.zeros(num_status, vocab_size, device=device, dtype=torch.int32)
for state_idx in range(num_status):
num_output_tokens = max(1, vocab_size // 20)
output_tokens = torch.randint(0, vocab_size, (num_output_tokens,), device=device)
counts = torch.randint(1, 10, (num_output_tokens,), device=device)
for token, count in zip(output_tokens, counts):
output_bin_counts[state_idx, token] = count
return (
logits,
idx_mapping,
token_ids,
expanded_local_pos,
repetition_penalty,
frequency_penalty,
presence_penalty,
prompt_bin_mask,
output_bin_counts,
)
class TestApplyPenalties:
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
@pytest.mark.parametrize("vocab_size", VOCAB_SIZE)
@pytest.mark.parametrize("num_status", NUM_STATUS)
@pytest.mark.parametrize("num_speculative_tokens", NUM_SPECULATIVE_TOKENS)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("seed", SEEDS)
@pytest.mark.parametrize("device", DEVICES)
@torch.inference_mode()
def test_apply_penalties(self, num_tokens, vocab_size, num_status, num_speculative_tokens, dtype, seed, device):
(
logits_triton,
idx_mapping,
token_ids,
expanded_local_pos,
repetition_penalty,
frequency_penalty,
presence_penalty,
prompt_bin_mask,
output_bin_counts,
) = create_test_data(
num_tokens=num_tokens,
vocab_size=vocab_size,
num_status=num_status,
num_speculative_tokens=num_speculative_tokens,
device=device,
dtype=dtype,
seed=seed,
)
logits_pytorch = logits_triton.clone()
apply_penalties(
logits_triton,
idx_mapping,
token_ids,
expanded_local_pos,
repetition_penalty,
frequency_penalty,
presence_penalty,
prompt_bin_mask,
output_bin_counts,
)
logits_pytorch_result = pytorch_apply_penalties(
logits_pytorch,
idx_mapping,
token_ids,
expanded_local_pos,
repetition_penalty,
frequency_penalty,
presence_penalty,
prompt_bin_mask,
output_bin_counts,
)
atol = DEFAULT_ATOL
rtol = DEFAULT_RTOL
if dtype == torch.bfloat16:
atol = 1e-02
rtol = 1e-02
assert torch.allclose(logits_triton, logits_pytorch_result, atol=atol, rtol=rtol)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,105 @@
from typing import Any
import pytest
import torch
from vllm.v1.worker.gpu.input_batch import post_update as post_update_gpu
from vllm_ascend.worker.v2.input_batch import post_update as post_update_npu
def generate_test_data(
num_reqs: int, max_num_reqs: int, vocab_size: int, num_speculative_steps: int, device: str
) -> dict[str, Any]:
"""
Generate random test data.
Return a dictionary containing all input tensors and the additional field 'expected_query_lens' for validation.
"""
if num_reqs > max_num_reqs:
raise ValueError("num_reqs cannot be larger than max_num_reqs")
idx_mapping = torch.arange(num_reqs, dtype=torch.int32, device=device)
num_computed_tokens = torch.randint(0, 100, (max_num_reqs,), dtype=torch.int32, device=device)
last_sampled_tokens = torch.randint(0, vocab_size, (max_num_reqs,), dtype=torch.int32, device=device)
output_bin_counts = torch.randint(0, 10, (max_num_reqs, vocab_size), dtype=torch.int32, device=device)
sampled_tokens = torch.randint(
0, vocab_size, (num_reqs, num_speculative_steps + 1), dtype=torch.int32, device=device
)
num_sampled = torch.randint(1, num_speculative_steps + 2, (num_reqs,), dtype=torch.int32, device=device)
num_rejected = torch.randint(0, num_speculative_steps + 1, (num_reqs,), dtype=torch.int32, device=device)
num_rejected = torch.min(num_rejected, num_sampled - 1)
query_lengths = torch.randint(1, 20, (num_reqs,), dtype=torch.int32, device=device)
query_start_loc = torch.cat(
[torch.tensor([0], dtype=torch.int32, device=device), torch.cumsum(query_lengths, dim=0)]
)
total_len = torch.randint(50, 200, (max_num_reqs,), dtype=torch.int32, device=device)
max_model_len = 3000 # 或者可以从total_len的最大值获取
all_token_ids = torch.randint(0, vocab_size, (max_num_reqs, max_model_len), dtype=torch.int32, device=device)
return {
"idx_mapping": idx_mapping,
"num_computed_tokens": num_computed_tokens,
"last_sampled_tokens": last_sampled_tokens,
"output_bin_counts": output_bin_counts,
"sampled_tokens": sampled_tokens,
"num_sampled": num_sampled,
"num_rejected": num_rejected,
"query_start_loc": query_start_loc,
"all_token_ids": all_token_ids,
"total_len": total_len,
}
@pytest.mark.parametrize(
"num_reqs,max_num_reqs,vocab_size,num_speculative_steps",
[
(36, 36, 200, 2),
(48, 48, 32000, 5),
(128, 128, 32000, 5),
],
)
def test_post_update(num_reqs: int, max_num_reqs: int, vocab_size: int, num_speculative_steps: int):
"""Test _topk_log_softmax_kernel for computing log probabilities
Args:
batch_size: Number of sequences in the batch
vocab_size: Size of the vocabulary
num_logprobs: Number of tokens to compute log probabilities for
"""
torch.manual_seed(42)
post_update_params = [
"idx_mapping",
"num_computed_tokens",
"last_sampled_tokens",
"output_bin_counts",
"sampled_tokens",
"num_sampled",
"num_rejected",
"query_start_loc",
"all_token_ids",
"total_len",
]
data = generate_test_data(num_reqs, max_num_reqs, vocab_size, num_speculative_steps, device="npu")
kernel_inputs_gpu = {k: data[k].clone() for k in post_update_params}
kernel_inputs_npu = {k: data[k].clone() for k in post_update_params}
# Invoke Triton kernel
post_update_gpu(**kernel_inputs_gpu)
torch.npu.synchronize()
post_update_npu(**kernel_inputs_npu)
torch.npu.synchronize()
# ========== Verify results ==========
assert torch.allclose(
kernel_inputs_gpu["output_bin_counts"], kernel_inputs_npu["output_bin_counts"], rtol=1e-3, atol=1e-3
), (
f"Triton output differs from PyTorch reference.\n"
f"Max diff: "
f"{torch.max(torch.abs(kernel_inputs_gpu['output_bin_counts'] - kernel_inputs_npu['output_bin_counts']))}\n"
f"Mean diff: "
f"{torch.mean(torch.abs(kernel_inputs_gpu['output_bin_counts'] - kernel_inputs_npu['output_bin_counts']))}"
)

View File

@@ -0,0 +1,318 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from collections.abc import Callable
from dataclasses import dataclass
from typing import Any
from unittest.mock import MagicMock
import numpy as np
import pytest
import torch
from vllm.model_executor.layers.mamba.mamba_utils import (
get_conv_copy_spec,
get_temporal_copy_spec,
)
from vllm.v1.core.sched.output import CachedRequestData, SchedulerOutput
from vllm.v1.kv_cache_interface import KVCacheConfig, KVCacheGroupSpec, MambaSpec
from vllm.v1.worker.mamba_utils import (
MambaCopyBuffers,
MambaSpecDecodeGPUContext,
collect_mamba_copy_meta,
do_mamba_copy_block,
)
import vllm_ascend.patch.worker.patch_mamba_utils # noqa: F401
MambaStateCopyFunc = Callable[..., Any]
_COPY_FUNCS: tuple[MambaStateCopyFunc, ...] = (
get_conv_copy_spec,
get_temporal_copy_spec,
)
def postprocess_mamba(
scheduler_output: SchedulerOutput,
kv_cache_config: KVCacheConfig,
input_batch: Any,
requests: dict[str, Any],
forward_context: dict[str, Any],
mamba_state_copy_funcs: tuple[MambaStateCopyFunc, ...],
copy_bufs: MambaCopyBuffers,
):
assert input_batch.mamba_state_idx_cpu is not None
num_scheduled_tokens_dict = scheduler_output.num_scheduled_tokens
scheduled_spec_decode_tokens_dict = scheduler_output.scheduled_spec_decode_tokens
num_accepted_tokens_cpu = input_batch.num_accepted_tokens_cpu
mamba_state_idx_cpu = input_batch.mamba_state_idx_cpu
mamba_group_ids = copy_bufs.mamba_group_ids
mamba_spec = copy_bufs.mamba_spec
copy_bufs.offset = 0
for i, req_id in enumerate(input_batch.req_ids):
req_state = requests[req_id]
num_computed_tokens = req_state.num_computed_tokens
num_draft_tokens = len(scheduled_spec_decode_tokens_dict.get(req_id, []))
num_scheduled_tokens = num_scheduled_tokens_dict[req_id]
num_accepted_tokens = num_accepted_tokens_cpu[i]
num_tokens_running_state = num_computed_tokens + num_scheduled_tokens - num_draft_tokens
new_num_computed_tokens = num_tokens_running_state + num_accepted_tokens - 1
aligned_new_computed_tokens = new_num_computed_tokens // mamba_spec.block_size * mamba_spec.block_size
if aligned_new_computed_tokens >= num_tokens_running_state:
accept_token_bias = aligned_new_computed_tokens - num_tokens_running_state
src_block_idx = mamba_state_idx_cpu[i]
dest_block_idx = aligned_new_computed_tokens // mamba_spec.block_size - 1
collect_mamba_copy_meta(
copy_bufs,
kv_cache_config,
mamba_state_copy_funcs,
mamba_group_ids,
src_block_idx,
dest_block_idx,
accept_token_bias,
req_state,
forward_context,
)
if src_block_idx == dest_block_idx:
num_accepted_tokens_cpu[i] = 1
do_mamba_copy_block(copy_bufs)
@dataclass
class _TestConfig:
block_size: int = 16
num_blocks: int = 32
num_layers: int = 2
max_num_reqs: int = 8
conv_width: int = 4
conv_inner_dim: int = 64
temporal_state_dim: int = 128
dtype: torch.dtype = torch.float16
class _MockCpuGpuBuffer:
def __init__(self, size: int, dtype: torch.dtype, device: torch.device):
self.cpu = torch.zeros(size, dtype=dtype, device="cpu")
self.gpu = torch.zeros(size, dtype=dtype, device=device)
self.np = self.cpu.numpy()
def copy_to_gpu(self, n: int | None = None) -> torch.Tensor:
if n is None:
return self.gpu.copy_(self.cpu, non_blocking=True)
return self.gpu[:n].copy_(self.cpu[:n], non_blocking=True)
def _make_scheduler_output(
num_scheduled_tokens: dict[str, int],
scheduled_spec_decode_tokens: dict[str, list] | None = None,
) -> SchedulerOutput:
cached = CachedRequestData.make_empty()
return SchedulerOutput(
scheduled_new_reqs=[],
scheduled_cached_reqs=cached,
num_scheduled_tokens=num_scheduled_tokens,
total_num_scheduled_tokens=sum(num_scheduled_tokens.values()),
scheduled_spec_decode_tokens=scheduled_spec_decode_tokens or {},
scheduled_encoder_inputs={},
num_common_prefix_blocks=[],
finished_req_ids=set(),
free_encoder_mm_hashes=[],
preempted_req_ids=set(),
)
def _make_mock_attention(conv_state: torch.Tensor, temporal_state: torch.Tensor) -> MagicMock:
attention = MagicMock()
attention.kv_cache = [conv_state, temporal_state]
return attention
def _make_states(
cfg: _TestConfig, layer_names: list[str], device: torch.device
) -> tuple[
list[torch.Tensor],
list[torch.Tensor],
list[torch.Tensor],
list[torch.Tensor],
dict[str, MagicMock],
dict[str, MagicMock],
]:
conv_py = [
torch.randn(cfg.num_blocks, cfg.conv_width, cfg.conv_inner_dim, dtype=cfg.dtype, device=device)
for _ in layer_names
]
temporal_py = [
torch.randn(cfg.num_blocks, cfg.temporal_state_dim, dtype=cfg.dtype, device=device) for _ in layer_names
]
conv_gpu = [s.clone() for s in conv_py]
temporal_gpu = [s.clone() for s in temporal_py]
fwd_py = {name: _make_mock_attention(c, t) for name, c, t in zip(layer_names, conv_py, temporal_py)}
fwd_gpu = {name: _make_mock_attention(c, t) for name, c, t in zip(layer_names, conv_gpu, temporal_gpu)}
return conv_py, temporal_py, conv_gpu, temporal_gpu, fwd_py, fwd_gpu
def _make_kv_cache_config(cfg: _TestConfig, layer_names: list[str]) -> KVCacheConfig:
mamba_spec = MambaSpec(
block_size=cfg.block_size,
shapes=((cfg.conv_width, cfg.conv_inner_dim), (cfg.temporal_state_dim,)),
dtypes=(cfg.dtype, cfg.dtype),
mamba_cache_mode="all",
)
return KVCacheConfig(
num_blocks=cfg.num_blocks,
kv_cache_tensors=[],
kv_cache_groups=[KVCacheGroupSpec(layer_names=layer_names, kv_cache_spec=mamba_spec)],
)
def _make_input_batch(req_ids: list[str], num_accepted_tokens: list[int], mamba_state_idx: list[int]) -> MagicMock:
batch = MagicMock()
batch.req_ids = req_ids
batch.req_id_to_index = {rid: i for i, rid in enumerate(req_ids)}
batch.num_accepted_tokens_cpu = np.array(num_accepted_tokens, dtype=np.int32)
batch.mamba_state_idx_cpu = np.array(mamba_state_idx, dtype=np.int32)
return batch
def _make_requests(
req_ids: list[str],
num_computed_tokens: list[int],
block_ids_per_req: list[list[int]],
) -> dict[str, MagicMock]:
requests: dict[str, MagicMock] = {}
for i, req_id in enumerate(req_ids):
req = MagicMock()
req.num_computed_tokens = num_computed_tokens[i]
req.block_ids = {0: block_ids_per_req[i]}
requests[req_id] = req
return requests
def _make_copy_bufs(cfg: _TestConfig, kv_cache_config: KVCacheConfig, device: torch.device) -> MambaCopyBuffers:
return MambaCopyBuffers.create(
max_num_reqs=cfg.max_num_reqs,
kv_cache_config=kv_cache_config,
copy_funcs=_COPY_FUNCS,
make_buffer=lambda n, dtype: _MockCpuGpuBuffer(n, dtype, device),
)
def _make_gpu_ctx(cfg: _TestConfig, kv_cache_config: KVCacheConfig, device: torch.device) -> MambaSpecDecodeGPUContext:
return MambaSpecDecodeGPUContext.create(
max_num_reqs=cfg.max_num_reqs,
kv_cache_config=kv_cache_config,
num_state_types=2,
device=device,
make_buffer=lambda n, dtype: _MockCpuGpuBuffer(n, dtype, device),
)
def _run_gpu_postprocess(
gpu_ctx: MambaSpecDecodeGPUContext,
*,
kv_cache_config: KVCacheConfig,
forward_context: dict[str, Any],
copy_funcs: tuple,
block_table: torch.Tensor,
req_ids: list[str],
num_accepted_tokens: list[int],
mamba_state_idx: list[int],
num_scheduled_tokens: dict[str, int],
num_computed_tokens: list[int],
num_draft_tokens: dict[str, int],
device: torch.device,
) -> None:
def t(values):
return torch.tensor(values, dtype=torch.int32, device=device)
gpu_ctx.initialize_from_forward_context(kv_cache_config, forward_context, copy_funcs, [block_table])
gpu_ctx.run_fused_postprocess(
num_reqs=len(req_ids),
num_accepted_tokens_gpu=t(num_accepted_tokens),
mamba_state_idx_gpu=t(mamba_state_idx),
num_scheduled_tokens_gpu=t([num_scheduled_tokens[r] for r in req_ids]),
num_computed_tokens_gpu=t(num_computed_tokens),
num_draft_tokens_gpu=t([num_draft_tokens.get(r, 0) for r in req_ids]),
)
torch.accelerator.synchronize()
@pytest.mark.skipif(not torch.npu.is_available(), reason="NPU required")
def test_matches_python_postprocess_mamba():
cfg = _TestConfig()
device = torch.device("npu:0")
torch.manual_seed(42)
req_ids = ["req_0", "req_1", "req_2", "req_3"]
num_computed_tokens = [60, 30, 45, 10]
num_scheduled_tokens = {"req_0": 5, "req_1": 3, "req_2": 8, "req_3": 6}
num_draft_tokens = {"req_0": 2, "req_1": 0, "req_2": 3, "req_3": 0}
num_accepted_tokens = [3, 2, 4, 2]
mamba_state_idx = [3, 1, 2, 0]
block_ids_per_req = [
list(range(8)),
list(range(8, 16)),
list(range(16, 24)),
list(range(24, 32)),
]
layer_names = [f"layer_{i}" for i in range(cfg.num_layers)]
kv_cache_config = _make_kv_cache_config(cfg, layer_names)
(
conv_states_py,
temporal_states_py,
conv_states_gpu,
temporal_states_gpu,
forward_context_py,
forward_context_gpu,
) = _make_states(cfg, layer_names, device)
scheduler_output = _make_scheduler_output(
num_scheduled_tokens,
{k: [None] * v for k, v in num_draft_tokens.items() if v > 0},
)
input_batch_py = _make_input_batch(req_ids, num_accepted_tokens.copy(), mamba_state_idx.copy())
requests = _make_requests(req_ids, num_computed_tokens, block_ids_per_req)
copy_bufs = _make_copy_bufs(cfg, kv_cache_config, device)
postprocess_mamba(
scheduler_output,
kv_cache_config,
input_batch_py,
requests,
forward_context_py,
_COPY_FUNCS,
copy_bufs,
)
torch.accelerator.synchronize()
gpu_ctx = _make_gpu_ctx(cfg, kv_cache_config, device)
block_table_gpu = torch.zeros(len(req_ids), 8, dtype=torch.int32, device=device)
for i, block_ids in enumerate(block_ids_per_req):
block_table_gpu[i, : len(block_ids)] = torch.tensor(block_ids, dtype=torch.int32)
_run_gpu_postprocess(
gpu_ctx,
kv_cache_config=kv_cache_config,
forward_context=forward_context_gpu,
copy_funcs=_COPY_FUNCS,
block_table=block_table_gpu,
req_ids=req_ids,
num_accepted_tokens=num_accepted_tokens,
mamba_state_idx=mamba_state_idx,
num_scheduled_tokens=num_scheduled_tokens,
num_computed_tokens=num_computed_tokens,
num_draft_tokens=num_draft_tokens,
device=device,
)
for i in range(cfg.num_layers):
torch.testing.assert_close(conv_states_gpu[i], conv_states_py[i])
torch.testing.assert_close(temporal_states_gpu[i], temporal_states_py[i])
expected_accepted = torch.tensor(
input_batch_py.num_accepted_tokens_cpu[: len(req_ids)],
dtype=torch.int32,
device=device,
)
torch.testing.assert_close(gpu_ctx.num_accepted_tokens_out[: len(req_ids)], expected_accepted)

View File

@@ -0,0 +1,76 @@
import gc
import pytest
import torch
from vllm.triton_utils import triton
from vllm_ascend.ops.triton.spec_decode.utils import prepare_inputs_padded_kernel
from vllm_ascend.ops.triton.triton_utils import get_vectorcore_num
from vllm_ascend.spec_decode.llm_base_proposer import _PREPARE_INPUTS_BLOCK_SIZE as BLOCK_SIZE
def prepare_inputs_padded_ref(
cu_num_draft_tokens,
valid_sampled_tokens_count,
query_start_loc,
):
num_draft_tokens = torch.cat(
[
cu_num_draft_tokens[0:1],
cu_num_draft_tokens[1:] - cu_num_draft_tokens[:-1],
]
)
num_rejected_tokens = torch.where(
num_draft_tokens > 0,
num_draft_tokens + 1 - valid_sampled_tokens_count,
torch.zeros_like(num_draft_tokens),
)
token_indices_to_sample = query_start_loc[1:] - 1 - num_rejected_tokens
return token_indices_to_sample.to(torch.int32)
@pytest.mark.parametrize("num_reqs", [1, 7, 32, 128, 2048])
def test_prepare_inputs_padded(num_reqs):
device = "npu"
torch.manual_seed(0)
draft_lens = torch.randint(1, 6, (num_reqs,), device=device, dtype=torch.int32)
cu_num_draft_tokens = torch.cumsum(draft_lens, dim=0).to(torch.int32)
valid_sampled_tokens_count = torch.zeros_like(draft_lens)
for i in range(num_reqs):
valid_sampled_tokens_count[i] = torch.randint(0, draft_lens[i] + 2, (1,)).item()
seq_lens = draft_lens + 1
query_start_loc = torch.zeros(num_reqs + 1, device=device, dtype=torch.int32)
query_start_loc[1:] = torch.cumsum(seq_lens, dim=0)
# Run PyTorch reference
out_ref = prepare_inputs_padded_ref(cu_num_draft_tokens, valid_sampled_tokens_count, query_start_loc)
# Run Triton kernel
out_tri = torch.empty(num_reqs, dtype=torch.int32, device=device)
num_rejected_tokens = torch.empty(num_reqs, dtype=torch.int32, device=device)
num_blocks_needed = triton.cdiv(num_reqs, BLOCK_SIZE)
num_vector_core = get_vectorcore_num()
grid_size = min(num_blocks_needed, num_vector_core)
grid = (grid_size,)
prepare_inputs_padded_kernel[grid](
cu_num_draft_tokens,
valid_sampled_tokens_count,
query_start_loc,
out_tri,
num_rejected_tokens,
num_reqs,
BLOCK_SIZE=BLOCK_SIZE,
)
torch.testing.assert_close(out_tri, out_ref)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,210 @@
import gc
import pytest
import torch
from vllm.v1.sample.rejection_sampler import rejection_random_sample_kernel as original_rejection_random_sample_kernel
from vllm_ascend.ops.triton.reject_sample import (
cal_grid_and_block_size,
rejection_random_sample_block_verify_kernel,
rejection_random_sample_kernel,
)
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
from vllm_ascend.sample.rejection_sampler import rejection_random_sample_block_verify_pytorch
@pytest.fixture(scope="function", autouse=True)
def setup_device_properties():
init_device_properties_triton()
yield
@pytest.mark.skip("Probabilistic failure, need zengtian after fix")
@pytest.mark.parametrize("max_spec_len", [1, 2, 3])
@pytest.mark.parametrize("vocab_size", [1024])
@pytest.mark.parametrize("batch_size", [1, 256, 512, 1024])
@torch.inference_mode()
def test_rejection_random_sample(max_spec_len, vocab_size, batch_size):
device = "npu"
torch.manual_seed(0)
draft_probs = torch.rand(batch_size * max_spec_len, vocab_size, dtype=torch.float32, device=device)
target_probs = torch.rand(batch_size * max_spec_len, vocab_size, dtype=torch.float32, device=device)
bonus_token_ids = torch.randint(low=0, high=vocab_size, size=(batch_size, 1), dtype=torch.int64, device=device)
draft_token_ids = torch.randint(
low=0, high=vocab_size, size=(batch_size * max_spec_len,), dtype=torch.int64, device=device
)
output_token_ids = torch.empty((batch_size, max_spec_len + 1), dtype=torch.int64, device=device)
original_output_token_ids = output_token_ids.clone()
num_tokens = draft_token_ids.shape[0]
uniform_probs = torch.rand((num_tokens,), dtype=torch.float32, device=device)
num_draft_tokens = [max_spec_len] * batch_size
num_draft_tokens = torch.tensor(num_draft_tokens, dtype=torch.int32, device=device)
cu_num_draft_tokens = torch.cumsum(num_draft_tokens, dim=0, dtype=torch.int32)
is_greedy_ptr = torch.full((batch_size,), False, dtype=torch.bool, device=device)
recovered_ids = torch.zeros_like(draft_token_ids, dtype=torch.int64, device=device)
grid, block_size = cal_grid_and_block_size(batch_size)
synthetic_conditional_rates = None
original_rejection_random_sample_kernel[(batch_size,)](
original_output_token_ids,
cu_num_draft_tokens,
draft_token_ids,
draft_probs,
target_probs,
bonus_token_ids,
recovered_ids,
uniform_probs,
is_greedy_ptr,
max_spec_len,
vocab_size,
synthetic_conditional_rates,
NO_DRAFT_PROBS=draft_probs is None,
SYNTHETIC_MODE=False,
)
rejection_random_sample_kernel[(grid,)](
output_token_ids,
cu_num_draft_tokens,
draft_token_ids,
draft_probs,
target_probs,
bonus_token_ids,
recovered_ids,
uniform_probs,
is_greedy_ptr,
max_spec_len,
vocab_size,
batch_size,
NO_DRAFT_PROBS=draft_probs is None,
BLOCK_SIZE=block_size,
)
torch.npu.synchronize()
assert torch.equal(original_output_token_ids, output_token_ids)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
DEVICE = "npu"
BATCH_SIZE = 7
MAX_SPEC_LEN = 3
VOCAB_SIZE = 5
CU_NUM_DRAFT_TOKENS = torch.tensor([2, 2, 5, 8, 11, 14, 15], dtype=torch.int32, device=DEVICE)
DRAFT_TOKEN_IDS = torch.tensor([0, 1, 0, 1, 2, 0, 1, 2, 0, 1, 2, 0, 1, 2, 0], dtype=torch.int64, device=DEVICE)
NUM_TOKENS = DRAFT_TOKEN_IDS.shape[0]
DRAFT_PROBS = None
TARGET_PROBS = torch.tensor(
[
[0.4, 0.3, 0.1, 0.1, 0.1], # 0
[0.1, 0.9, 0.0, 0.0, 0.0], # 1
[0.2, 0.1, 0.2, 0.4, 0.1], # 0
[0.1, 0.4, 0.1, 0.1, 0.3], # 0
[0.2, 0.1, 0.4, 0.1, 0.2], # 0
[0.4, 0.2, 0.1, 0.2, 0.1], # 0
[0.1, 0.6, 0.1, 0.1, 0.1], # 1
[0.2, 0.2, 0.2, 0.3, 0.1], # 0
[0.4, 0.2, 0.1, 0.2, 0.1], # 0
[0.1, 0.6, 0.1, 0.1, 0.1], # 1
[0.2, 0.2, 0.2, 0.3, 0.1], # 0
[0.4, 0.4, 0.1, 0.0, 0.1], # 1
[0.4, 0.3, 0.1, 0.1, 0.1], # 0
[0.4, 0.0, 0.5, 0.0, 0.1], # 1
[0.4, 0.1, 0.3, 0.1, 0.1], # 1
],
dtype=torch.float32,
device=DEVICE,
)
UNIFORM_PROBS = torch.tensor(
[
0.9,
0.0,
0.9,
0.7,
0.8,
0.5,
0.45,
1.0,
0.5,
0.45,
1.0,
0.39,
0.4,
0.1,
0.3,
],
dtype=torch.float32,
device=DEVICE,
)
BONUS_TOKEN_IDS = torch.full((BATCH_SIZE,), MAX_SPEC_LEN + 1, dtype=torch.int64, device=DEVICE)
RECOVERED_TOKEN_IDS = torch.full((NUM_TOKENS,), MAX_SPEC_LEN, dtype=torch.int64, device=DEVICE)
IS_GREEDY = torch.zeros(BATCH_SIZE, dtype=torch.bool, device=DEVICE)
IS_GREEDY[4] = True
@pytest.mark.skip(reason="Probabilistic failure, #TODO: tracking issue 9852")
@pytest.mark.parametrize("cu_num_draft_tokens", [CU_NUM_DRAFT_TOKENS])
@pytest.mark.parametrize("draft_token_ids", [DRAFT_TOKEN_IDS])
@pytest.mark.parametrize("draft_probs", [DRAFT_PROBS])
@pytest.mark.parametrize("target_probs", [TARGET_PROBS])
@pytest.mark.parametrize("bonus_token_ids", [BONUS_TOKEN_IDS])
@pytest.mark.parametrize("recovered_token_ids", [RECOVERED_TOKEN_IDS])
@pytest.mark.parametrize("uniform_probs", [UNIFORM_PROBS])
@pytest.mark.parametrize("is_greedy", [IS_GREEDY])
@pytest.mark.parametrize("batch_size", [BATCH_SIZE])
@pytest.mark.parametrize("max_spec_len", [MAX_SPEC_LEN])
@pytest.mark.parametrize("vocab_size", [VOCAB_SIZE])
@torch.inference_mode()
def test_rejection_sampler_block_verify_triton_kernel(
cu_num_draft_tokens, # [batch_size]
draft_token_ids, # [num_tokens]
draft_probs, # [num_tokens, vocab_size] or None
target_probs, # [num_tokens, vocab_size]
bonus_token_ids, # [batch_size]
recovered_token_ids, # [num_tokens]
uniform_probs, # [num_tokens]
is_greedy, # [batch_size]
batch_size, # int
max_spec_len, # int
vocab_size, # int
) -> None:
grid, block_size = cal_grid_and_block_size(batch_size)
output_token_ids_ref = torch.full((batch_size, max_spec_len + 1), -1, dtype=torch.int64, device=DEVICE)
output_token_ids_triton = output_token_ids_ref.clone()
rejection_random_sample_block_verify_pytorch(
output_token_ids=output_token_ids_ref,
cu_num_draft_tokens=cu_num_draft_tokens,
draft_token_ids=draft_token_ids,
draft_probs=draft_probs,
target_probs=target_probs,
bonus_token_ids=bonus_token_ids,
recovered_token_ids=recovered_token_ids,
uniform_probs=uniform_probs,
is_greedy=is_greedy,
max_spec_len=max_spec_len,
vocab_size=vocab_size,
IS_NGRAM=draft_probs is None,
)
rejection_random_sample_block_verify_kernel[(grid,)](
output_token_ids_ptr=output_token_ids_triton,
cu_num_draft_tokens_ptr=cu_num_draft_tokens,
draft_token_ids_ptr=draft_token_ids,
draft_probs_ptr=draft_probs,
target_probs_ptr=target_probs,
bonus_token_ids_ptr=bonus_token_ids,
recovered_token_ids_ptr=recovered_token_ids,
uniform_probs_ptr=uniform_probs,
is_greedy_ptr=is_greedy,
max_spec_len=max_spec_len,
vocab_size=vocab_size,
vec_len=batch_size,
NO_DRAFT_PROBS=draft_probs is None,
BLOCK_SIZE=block_size,
SUB_BLOCK=32,
)
torch.npu.synchronize()
assert torch.equal(output_token_ids_ref, output_token_ids_triton)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,224 @@
import gc
import pytest
import torch
from vllm_ascend.ops.triton.rope import rope_forward_triton, rope_forward_triton_siso
IS_NEOX_STYLE = [True, False]
DTYPES = [torch.bfloat16, torch.float16]
MAX_POSITION_EMBEDDINGS = [262144]
# parameters for test_rotary_embedding_triton_kernel only
# (head_size, rotary_dim)
HEAD_ROTARY_DIMS = [
(64, 32),
(128, 128),
]
# (num_q_heads, num_k_heads)
NUM_QK_HEADS = [
(64, 1),
(96, 8),
]
# parameters for test_rotary_embedding_triton_kernel_siso only
SISO_HEAD_SIZES = [64, 128]
SISO_ROTARY_DIMS = [32, 64]
SISO_NUM_HEADS = [64]
NUM_TOKENS = [1, 4, 8, 16, 1024]
SEEDS = [0]
DEVICES = [f"npu:{0}"]
DEFAULT_ATOL = 1e-3
DEFAULT_RTOL = 1e-3
def rotate_neox(x: torch.Tensor) -> torch.Tensor:
x1 = x[..., : x.shape[-1] // 2]
x2 = x[..., x.shape[-1] // 2 :]
return torch.cat((-x2, x1), dim=-1)
def rotate_gptj(x: torch.Tensor) -> torch.Tensor:
x1 = x[..., ::2]
x2 = x[..., 1::2]
x = torch.stack((-x2, x1), dim=-1)
return x.flatten(-2)
def _rope_pytorch_native(query, key, cos, sin, rope_dim, is_neox_style) -> tuple[torch.Tensor, torch.Tensor | None]:
"""PyTorch-native implementation equivalent to forward()."""
assert key is not None
orig_dtype = query.dtype
query_rot = query[..., :rope_dim].to(torch.float32)
key_rot = key[..., :rope_dim].to(torch.float32)
head_size = query.shape[-1]
if rope_dim < head_size:
query_pass = query[..., rope_dim:]
key_pass = key[..., rope_dim:]
if is_neox_style:
cos = cos.repeat(1, 2).unsqueeze(-2).to(torch.float32)
sin = sin.repeat(1, 2).unsqueeze(-2).to(torch.float32)
else:
cos = cos.repeat_interleave(2, dim=-1).unsqueeze(-2).to(torch.float32)
sin = sin.repeat_interleave(2, dim=-1).unsqueeze(-2).to(torch.float32)
rotate_fn = rotate_neox if is_neox_style else rotate_gptj
query_rot = query_rot * cos + rotate_fn(query_rot) * sin
key_rot = key_rot * cos + rotate_fn(key_rot) * sin
if rope_dim < head_size:
query = torch.cat((query_rot.to(orig_dtype), query_pass), dim=-1)
key = torch.cat((key_rot.to(orig_dtype), key_pass), dim=-1)
else:
query = query_rot.to(orig_dtype)
key = key_rot.to(orig_dtype)
return query, key
def _rope_siso_pytorch_native(query, cos, sin, rope_dim, is_neox_style) -> tuple[torch.Tensor, torch.Tensor | None]:
"""PyTorch-native implementation equivalent to forward()."""
assert query is not None
orig_dtype = query.dtype
query_rot = query[..., :rope_dim].to(torch.float32)
head_size = query.shape[-1]
if rope_dim < head_size:
query_pass = query[..., rope_dim:]
if is_neox_style:
cos = cos.repeat(1, 2).unsqueeze(-2).to(torch.float32)
sin = sin.repeat(1, 2).unsqueeze(-2).to(torch.float32)
else:
cos = cos.repeat_interleave(2, dim=-1).unsqueeze(-2).to(torch.float32)
sin = sin.repeat_interleave(2, dim=-1).unsqueeze(-2).to(torch.float32)
rotate_fn = rotate_neox if is_neox_style else rotate_gptj
query_rot = query_rot * cos + rotate_fn(query_rot) * sin
if rope_dim < head_size:
query = torch.cat((query_rot.to(orig_dtype), query_pass), dim=-1)
else:
query = query_rot.to(orig_dtype)
return query
@pytest.mark.parametrize("is_neox_style", IS_NEOX_STYLE)
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
@pytest.mark.parametrize("num_q_heads,num_k_heads", NUM_QK_HEADS)
@pytest.mark.parametrize("head_size,rotary_dim", HEAD_ROTARY_DIMS)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("seed", SEEDS)
@pytest.mark.parametrize("device", DEVICES)
@torch.inference_mode()
def test_rotary_embedding_triton_kernel(
is_neox_style: bool,
num_tokens: int,
num_q_heads: int,
num_k_heads: int,
head_size: int,
rotary_dim: int,
dtype: torch.dtype,
seed: int,
device: str,
) -> None:
torch.manual_seed(seed)
torch.set_default_device(device)
sin = torch.randn(num_tokens, rotary_dim // 2, dtype=dtype, device=device)
cos = torch.randn(num_tokens, rotary_dim // 2, dtype=dtype, device=device)
q_trt = torch.randn(num_tokens, num_q_heads, head_size, dtype=dtype, device=device)
k_trt = torch.randn(num_tokens, num_k_heads, head_size, dtype=dtype, device=device)
q_gold = torch.randn(num_tokens, num_q_heads, head_size, dtype=dtype, device=device)
k_gold = torch.randn(num_tokens, num_k_heads, head_size, dtype=dtype, device=device)
q_trt.copy_(q_gold)
k_trt.copy_(k_gold)
q_trt, k_trt = rope_forward_triton(q_trt, k_trt, cos, sin, rope_dim=rotary_dim, is_neox_style=is_neox_style)
q_gold, k_gold = _rope_pytorch_native(q_gold, k_gold, cos, sin, rope_dim=rotary_dim, is_neox_style=is_neox_style)
# Compare the results.
torch.testing.assert_close(q_trt.view(q_gold.size()), q_gold, atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
torch.testing.assert_close(k_trt.view(k_gold.size()), k_gold, atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.parametrize("max_position_embeddings", MAX_POSITION_EMBEDDINGS)
@pytest.mark.parametrize("is_neox_style", IS_NEOX_STYLE)
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
@pytest.mark.parametrize("num_q_heads,num_k_heads", NUM_QK_HEADS)
@pytest.mark.parametrize("head_size,rotary_dim", HEAD_ROTARY_DIMS)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("seed", SEEDS)
@pytest.mark.parametrize("device", DEVICES)
@torch.inference_mode()
def test_rotary_embedding_triton_kernel_with_cos_sin_cache(
max_position_embeddings: int,
is_neox_style: bool,
num_tokens: int,
num_q_heads: int,
num_k_heads: int,
head_size: int,
rotary_dim: int,
dtype: torch.dtype,
seed: int,
device: str,
) -> None:
torch.manual_seed(seed)
torch.set_default_device(device)
cos_sin_cache = torch.randn(max_position_embeddings, rotary_dim, dtype=dtype, device=device)
positions = torch.randint(low=0, high=max_position_embeddings, size=(num_tokens,), dtype=torch.int64, device=device)
q_trt = torch.randn(num_tokens, num_q_heads, head_size, dtype=dtype, device=device)
k_trt = torch.randn(num_tokens, num_k_heads, head_size, dtype=dtype, device=device)
q_gold = torch.randn(num_tokens, num_q_heads, head_size, dtype=dtype, device=device)
k_gold = torch.randn(num_tokens, num_k_heads, head_size, dtype=dtype, device=device)
q_trt.copy_(q_gold)
k_trt.copy_(k_gold)
q_trt, k_trt = rope_forward_triton(
q_trt, k_trt, cos_sin_cache=cos_sin_cache, positions=positions, rope_dim=rotary_dim, is_neox_style=is_neox_style
)
cos, sin = cos_sin_cache.index_select(0, positions).chunk(2, dim=-1)
q_gold, k_gold = _rope_pytorch_native(q_gold, k_gold, cos, sin, rope_dim=rotary_dim, is_neox_style=is_neox_style)
# Compare the results.
torch.testing.assert_close(q_trt.view(q_gold.size()), q_gold, atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
torch.testing.assert_close(k_trt.view(k_gold.size()), k_gold, atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.parametrize("is_neox_style", IS_NEOX_STYLE)
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
@pytest.mark.parametrize("num_q_heads", SISO_NUM_HEADS)
@pytest.mark.parametrize("head_size", SISO_HEAD_SIZES)
@pytest.mark.parametrize("rotary_dim", SISO_ROTARY_DIMS)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("seed", SEEDS)
@pytest.mark.parametrize("device", DEVICES)
@torch.inference_mode()
def test_rotary_embedding_triton_kernel_siso(
is_neox_style: bool,
num_tokens: int,
num_q_heads: int,
head_size: int,
rotary_dim: int,
dtype: torch.dtype,
seed: int,
device: str,
) -> None:
torch.manual_seed(seed)
torch.set_default_device(device)
if rotary_dim == -1:
rotary_dim = head_size
sin = torch.randn(num_tokens, rotary_dim // 2, dtype=dtype, device=device)
cos = torch.randn(num_tokens, rotary_dim // 2, dtype=dtype, device=device)
q_trt = torch.randn(num_tokens, num_q_heads, head_size, dtype=dtype, device=device)
q_gold = torch.randn(num_tokens, num_q_heads, head_size, dtype=dtype, device=device)
q_trt.copy_(q_gold)
q_trt = rope_forward_triton_siso(q_trt, cos, sin, rope_dim=rotary_dim, is_neox_style=is_neox_style)
q_gold = _rope_siso_pytorch_native(q_gold, cos, sin, rope_dim=rotary_dim, is_neox_style=is_neox_style)
# Compare the results.
torch.testing.assert_close(q_trt.view(q_gold.size()), q_gold, atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,324 @@
import gc
import pytest
import torch
NUM_TOKENS = [1, 4096]
NUM_QKV_HEADS = [(2, 1), (16, 2)]
HEAD_SIZES = [128, 256]
EPS = [1e-6]
MROPE_SECTION = [[11, 11, 10], [24, 20, 20]]
IS_INTERLEAVED = [True, False]
HAS_GATE = [True, False]
DTYPES = [torch.bfloat16, torch.float16]
DEVICES = [f"npu:{0}"]
DEFAULT_ATOL = 1e-2
DEFAULT_RTOL = 1e-2
def apply_interleaved_rope(x: torch.Tensor, mrope_section: list[int]) -> torch.Tensor:
"""Apply interleaved MRoPE to 3D rotary embeddings.
Reorganizes frequency layout from chunked [TTT...HHH...WWW] to
interleaved [THTHWHTHW...TT], preserving frequency continuity.
"""
x_t = x[0].clone()
x_t[..., 1 : mrope_section[1] * 3 : 3] = x[1, ..., 1 : mrope_section[1] * 3 : 3]
x_t[..., 2 : mrope_section[2] * 3 : 3] = x[2, ..., 2 : mrope_section[2] * 3 : 3]
return x_t
def rms_norm(
x: torch.Tensor,
norm_weight: torch.Tensor,
eps,
norm_bias=None,
):
x = x.cpu()
norm_weight = norm_weight.cpu()
x = x.to(torch.float32)
norm_weight = norm_weight.to(torch.float32).cpu()
reciprocal_std = 1 / torch.sqrt(torch.mean(x**2, axis=-1, keepdims=True) + eps)
out = x * reciprocal_std * norm_weight
if norm_bias is not None:
norm_bias = norm_bias.cpu().to(torch.float32)
out = out + norm_bias
return out
def naive_split_qkv_rmsnorm_mrope(
qkv: torch.Tensor,
q_weight: torch.Tensor,
q_bias: torch.Tensor,
k_weight: torch.Tensor,
k_bias: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
num_q_heads: int,
num_kv_heads: int,
head_size: int,
eps: float,
mrope_section: list[int],
rope_dim: int,
):
q_size = num_q_heads * head_size
kv_size = num_kv_heads * head_size
# split
qkv = qkv.cpu()
q, k, v = qkv.split([q_size, kv_size, kv_size], dim=-1)
# norm
q = rms_norm(q.reshape(-1, head_size), q_weight, eps, norm_bias=q_bias)
k = rms_norm(k.reshape(-1, head_size), k_weight, eps, norm_bias=k_bias)
# mrope
rotary_dim = rope_dim
num_tokens = qkv.shape[0]
n_q_head = num_q_heads
n_kv_head = num_kv_heads
q_reshaped = q.view(num_tokens, n_q_head, head_size)
k_reshaped = k.view(num_tokens, n_kv_head, head_size)
cos_reshaped = cos.permute(1, 2, 0)
sin_reshaped = sin.permute(1, 2, 0)
half_rd = rotary_dim // 2
for token_idx in range(num_tokens):
token_cos = cos_reshaped[token_idx]
token_sin = sin_reshaped[token_idx]
cos_row = torch.zeros(half_rd, device=q.device, dtype=q.dtype)
sin_row = torch.zeros(half_rd, device=q.device, dtype=q.dtype)
t_end = mrope_section[0]
h_end = t_end + mrope_section[1]
if t_end > 0:
cos_row[:t_end] = token_cos[:t_end, 0]
sin_row[:t_end] = token_sin[:t_end, 0]
if mrope_section[1] > 0:
cos_row[t_end:h_end] = token_cos[t_end:h_end, 1]
sin_row[t_end:h_end] = token_sin[t_end:h_end, 1]
if mrope_section[2] > 0:
w_start = h_end
cos_row[w_start:half_rd] = token_cos[w_start:half_rd, 2]
sin_row[w_start:half_rd] = token_sin[w_start:half_rd, 2]
q_token = q_reshaped[token_idx]
k_token = k_reshaped[token_idx]
q1 = q_token[:, :half_rd]
q2 = q_token[:, half_rd:rotary_dim]
k1 = k_token[:, :half_rd]
k2 = k_token[:, half_rd:rotary_dim]
cos_half = cos_row.unsqueeze(0)
sin_half = sin_row.unsqueeze(0)
new_q1 = q1 * cos_half - q2 * sin_half
new_q2 = q2 * cos_half + q1 * sin_half
new_k1 = k1 * cos_half - k2 * sin_half
new_k2 = k2 * cos_half + k1 * sin_half
q_reshaped[token_idx, :, :rotary_dim] = torch.cat([new_q1, new_q2], dim=1)
k_reshaped[token_idx, :, :rotary_dim] = torch.cat([new_k1, new_k2], dim=1)
q_result = q_reshaped.view(num_tokens, -1)
k_result = k_reshaped.view(num_tokens, -1)
q = q_result.to(qkv.dtype)
k = k_result.to(qkv.dtype)
v = v.to(qkv.dtype)
return q, k, v
def naive_split_qkv_rmsnorm_mrope_interleaved(
qkv: torch.Tensor,
q_weight: torch.Tensor,
q_bias: torch.Tensor,
k_weight: torch.Tensor,
k_bias: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
num_q_heads: int,
num_kv_heads: int,
head_size: int,
eps: float,
mrope_section: list[int],
rope_dim: int,
):
q_size = num_q_heads * head_size
kv_size = num_kv_heads * head_size
# split
qkv = qkv.cpu()
q, k, v = qkv.split([q_size, kv_size, kv_size], dim=-1)
# norm
q = rms_norm(q.reshape(-1, head_size), q_weight, eps, norm_bias=q_bias)
k = rms_norm(k.reshape(-1, head_size), k_weight, eps, norm_bias=k_bias)
# mrope
rotary_dim = rope_dim
num_tokens = qkv.shape[0]
n_q_head = num_q_heads
n_kv_head = num_kv_heads
q_reshaped = q.view(num_tokens, n_q_head, head_size)
k_reshaped = k.view(num_tokens, n_kv_head, head_size)
cos_reshaped = apply_interleaved_rope(cos, mrope_section)
sin_reshaped = apply_interleaved_rope(sin, mrope_section)
half_rd = rotary_dim // 2
for token_idx in range(num_tokens):
cos_row = cos_reshaped[token_idx]
sin_row = sin_reshaped[token_idx]
q_token = q_reshaped[token_idx]
k_token = k_reshaped[token_idx]
q1 = q_token[:, :half_rd]
q2 = q_token[:, half_rd:rotary_dim]
k1 = k_token[:, :half_rd]
k2 = k_token[:, half_rd:rotary_dim]
cos_half = cos_row.unsqueeze(0)
sin_half = sin_row.unsqueeze(0)
new_q1 = q1 * cos_half - q2 * sin_half
new_q2 = q2 * cos_half + q1 * sin_half
new_k1 = k1 * cos_half - k2 * sin_half
new_k2 = k2 * cos_half + k1 * sin_half
q_reshaped[token_idx, :, :rotary_dim] = torch.cat([new_q1, new_q2], dim=1)
k_reshaped[token_idx, :, :rotary_dim] = torch.cat([new_k1, new_k2], dim=1)
q_result = q_reshaped.view(num_tokens, -1)
k_result = k_reshaped.view(num_tokens, -1)
q = q_result.to(qkv.dtype)
k = k_result.to(qkv.dtype)
v = v.to(qkv.dtype)
return q, k, v
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
@pytest.mark.parametrize("num_q_heads, num_kv_heads", NUM_QKV_HEADS)
@pytest.mark.parametrize("head_size", HEAD_SIZES)
@pytest.mark.parametrize("eps", EPS)
@pytest.mark.parametrize("mrope_section", MROPE_SECTION)
@pytest.mark.parametrize("is_interleaved", IS_INTERLEAVED)
@pytest.mark.parametrize("has_gate", HAS_GATE)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("device", DEVICES)
@torch.inference_mode()
def test_split_qkv_rmsnorm_mrope(
num_tokens: int,
num_q_heads: int,
num_kv_heads: int,
head_size: int,
mrope_section: list[int],
eps: float,
dtype: torch.dtype,
device: str,
is_interleaved: bool,
has_gate: bool,
):
torch.set_default_device(device)
rope_dim = 2 * sum(mrope_section)
q_size = num_q_heads * head_size
kv_size = num_kv_heads * head_size
# input tensor
if has_gate:
qkv = torch.randn(num_tokens, 2 * q_size + kv_size * 2, dtype=dtype, device=device)
else:
qkv = torch.randn(num_tokens, q_size + kv_size * 2, dtype=dtype, device=device)
q_weight = torch.randn(head_size, dtype=dtype, device=device)
k_weight = torch.randn(head_size, dtype=dtype, device=device)
q_bias = None
k_bias = None
cos_sin = torch.randn(3, num_tokens, rope_dim, dtype=dtype, device=device)
cos, sin = cos_sin.chunk(2, dim=-1)
cos = cos.contiguous()
sin = sin.contiguous()
if has_gate:
q_gate_data = qkv[:, : q_size * 2].view(-1, num_q_heads, head_size * 2)
q_data, golden_gate = torch.chunk(q_gate_data, 2, dim=-1)
golden_gate = golden_gate.reshape(-1, q_size)
q_data = q_data.reshape(-1, q_size)
k_data = qkv[:, 2 * q_size : 2 * q_size + kv_size]
v_data = qkv[:, 2 * q_size + kv_size :]
qkv_for_ref = torch.cat([q_data, k_data, v_data], dim=-1)
else:
qkv_for_ref = qkv
if is_interleaved:
golden_q, golden_k, golden_v = naive_split_qkv_rmsnorm_mrope_interleaved(
qkv_for_ref.cpu(),
q_weight.cpu(),
q_bias,
k_weight.cpu(),
k_bias,
cos.cpu(),
sin.cpu(),
num_q_heads,
num_kv_heads,
head_size,
eps,
mrope_section,
rope_dim,
)
else:
golden_q, golden_k, golden_v = naive_split_qkv_rmsnorm_mrope(
qkv_for_ref.cpu(),
q_weight.cpu(),
q_bias,
k_weight.cpu(),
k_bias,
cos.cpu(),
sin.cpu(),
num_q_heads,
num_kv_heads,
head_size,
eps,
mrope_section,
rope_dim,
)
real_q, real_k, real_v, real_gate = torch.ops.vllm.triton_split_qkv_rmsnorm_mrope(
qkv=qkv,
q_weight=q_weight,
k_weight=k_weight,
cos_sin=cos_sin,
num_q_heads=num_q_heads,
num_kv_heads=num_kv_heads,
head_size=head_size,
eps=eps,
mrope_section=mrope_section,
is_interleaved=is_interleaved,
rope_dim=rope_dim,
has_gate=has_gate,
)
torch.testing.assert_close(real_q.cpu(), golden_q.cpu(), atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
torch.testing.assert_close(real_k.cpu(), golden_k.cpu(), atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
torch.testing.assert_close(real_v.cpu(), golden_v.cpu(), atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
if has_gate:
torch.testing.assert_close(real_gate.cpu(), golden_gate.cpu(), atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,202 @@
import gc
import numpy as np
import pytest
import torch
from vllm_ascend.device.device_op import DeviceOperator
MAX_POSITION_EMBEDDINGS = [262144]
NUM_TOKENS = [1, 16, 1024, 10240]
NUM_QKV_HEADS = [(12, 1), (64, 4)]
HEAD_SIZES = [128]
ROPE_DIMS = [64, 128]
EPS = [1e-6]
DTYPES = [torch.bfloat16]
SEEDS = [0]
DEVICES = [f"npu:{0}"]
DEFAULT_ATOL = 5e-2
DEFAULT_RTOL = 5e-3
def custom_rope(q, k, sin, cos):
rotary_dim = sin.shape[-1]
sin = sin.to(torch.float32)
cos = cos.to(torch.float32)
q_rot = q[..., :rotary_dim]
k_rot = k[..., :rotary_dim]
q_pass = q[..., rotary_dim:]
k_pass = k[..., rotary_dim:]
x1 = q_rot[..., : rotary_dim // 2]
x2 = q_rot[..., rotary_dim // 2 :]
cat_x = torch.cat([-x2, x1], axis=-1)
mul1 = cat_x * sin
mul2 = q_rot * cos
q_rot = mul1 + mul2
res1 = torch.cat([q_rot, q_pass], dim=-1)
x1 = k_rot[..., : rotary_dim // 2]
x2 = k_rot[..., rotary_dim // 2 :]
cat_x = torch.cat([-x2, x1], axis=-1)
mul1 = cat_x * sin
mul2 = k_rot * cos
k_rot = mul1 + mul2
res2 = torch.cat([k_rot, k_pass], dim=-1)
return res1, res2
def rms_norm(
input,
norm_weight,
eps,
norm_bias=None,
):
input = input.to(torch.float32)
norm_weight = norm_weight.to(torch.float32)
reciprocal_std = 1 / torch.sqrt(torch.mean(input**2, axis=-1, keepdims=True) + eps)
out = input * reciprocal_std * norm_weight
if norm_bias is not None:
norm_bias = norm_bias.to(torch.float32)
out = out + norm_bias
return out
@pytest.mark.parametrize("max_position_embeddings", MAX_POSITION_EMBEDDINGS)
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
@pytest.mark.parametrize("num_q_heads, num_kv_heads", NUM_QKV_HEADS)
@pytest.mark.parametrize("head_size", HEAD_SIZES)
@pytest.mark.parametrize("eps", EPS)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("seed", SEEDS)
@pytest.mark.parametrize("device", DEVICES)
@pytest.mark.parametrize("rope_dim", ROPE_DIMS)
@torch.inference_mode()
def test_split_qkv_rmsnorm_rope(
max_position_embeddings, num_tokens, num_q_heads, num_kv_heads, head_size, eps, dtype, seed, device, rope_dim
):
torch.manual_seed(seed)
torch.set_default_device(device)
q_hidden_size = num_q_heads * head_size
kv_hidden_size = num_kv_heads * head_size
qkv = torch.randn(num_tokens, q_hidden_size + kv_hidden_size * 2, dtype=dtype, device=device)
q_weight = torch.randn(head_size, dtype=dtype, device=device)
k_weight = torch.randn(head_size, dtype=dtype, device=device)
cos_sin_cache = torch.from_numpy(np.random.uniform(0, 1, [max_position_embeddings, rope_dim])).to(dtype).npu()
positions = torch.randint(low=0, high=max_position_embeddings, size=(num_tokens,), dtype=torch.int64, device=device)
# fused kernel
q, k, v = DeviceOperator.split_qkv_rmsnorm_rope(
input=qkv,
q_weight=q_weight,
k_weight=k_weight,
q_hidden_size=q_hidden_size,
kv_hidden_size=kv_hidden_size,
head_dim=head_size,
eps=eps,
q_bias=None,
k_bias=None,
cos_sin_cache=cos_sin_cache,
positions=positions,
)
cos, sin = cos_sin_cache.index_select(0, positions).view(num_tokens, 2, -1).repeat(1, 1, 2).chunk(2, dim=-2)
cos = cos.unsqueeze(1)
sin = sin.unsqueeze(1)
# split
_q, _k, v_gold = qkv.cpu().split([q_hidden_size, kv_hidden_size, kv_hidden_size], dim=-1)
# norm
_q = rms_norm(_q.reshape(-1, head_size), q_weight.cpu(), eps)
_k = rms_norm(_k.reshape(-1, head_size), k_weight.cpu(), eps)
_q = _q.reshape(num_tokens, 1, -1, head_size)
_k = _k.reshape(num_tokens, 1, -1, head_size)
# rope
q_gold, k_gold = custom_rope(_q, _k, sin.cpu(), cos.cpu())
q_gold = q_gold.reshape(num_tokens, -1)
k_gold = k_gold.reshape(num_tokens, -1)
# Compare the results.
torch.testing.assert_close(q.to(torch.float32).cpu(), q_gold, atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
torch.testing.assert_close(k.to(torch.float32).cpu(), k_gold, atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
torch.testing.assert_close(
v.to(torch.float32).cpu(), v_gold.to(torch.float32), atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL
)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.parametrize("max_position_embeddings", MAX_POSITION_EMBEDDINGS)
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
@pytest.mark.parametrize("num_q_heads, num_kv_heads", NUM_QKV_HEADS)
@pytest.mark.parametrize("head_size", HEAD_SIZES)
@pytest.mark.parametrize("eps", EPS)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("seed", SEEDS)
@pytest.mark.parametrize("device", DEVICES)
@pytest.mark.parametrize("rope_dim", ROPE_DIMS)
@torch.inference_mode()
def test_split_qkv_rmsnorm_rope_with_bias(
max_position_embeddings, num_tokens, num_q_heads, num_kv_heads, head_size, eps, dtype, seed, device, rope_dim
):
torch.manual_seed(seed)
torch.set_default_device(device)
q_hidden_size = num_q_heads * head_size
kv_hidden_size = num_kv_heads * head_size
qkv = torch.randn(num_tokens, q_hidden_size + kv_hidden_size * 2, dtype=dtype, device=device)
q_weight = torch.randn(head_size, dtype=dtype, device=device)
k_weight = torch.randn(head_size, dtype=dtype, device=device)
q_bias = torch.randn(head_size, dtype=dtype, device=device)
k_bias = torch.randn(head_size, dtype=dtype, device=device)
cos_sin_cache = torch.from_numpy(np.random.uniform(0, 1, [max_position_embeddings, rope_dim])).to(dtype).npu()
positions = torch.randint(low=0, high=max_position_embeddings, size=(num_tokens,), dtype=torch.int64, device=device)
# fused kernel
q, k, v = DeviceOperator.split_qkv_rmsnorm_rope(
input=qkv,
q_weight=q_weight,
k_weight=k_weight,
q_hidden_size=q_hidden_size,
kv_hidden_size=kv_hidden_size,
head_dim=head_size,
eps=eps,
q_bias=q_bias,
k_bias=k_bias,
cos_sin_cache=cos_sin_cache,
positions=positions,
)
cos, sin = cos_sin_cache.index_select(0, positions).view(num_tokens, 2, -1).repeat(1, 1, 2).chunk(2, dim=-2)
cos = cos.unsqueeze(1)
sin = sin.unsqueeze(1)
# split
_q, _k, v_gold = qkv.cpu().split([q_hidden_size, kv_hidden_size, kv_hidden_size], dim=-1)
# norm
_q = rms_norm(_q.reshape(-1, head_size), q_weight.cpu(), eps, norm_bias=q_bias.cpu())
_k = rms_norm(_k.reshape(-1, head_size), k_weight.cpu(), eps, norm_bias=k_bias.cpu())
_q = _q.reshape(num_tokens, 1, -1, head_size)
_k = _k.reshape(num_tokens, 1, -1, head_size)
# rope
q_gold, k_gold = custom_rope(_q, _k, sin.cpu(), cos.cpu())
q_gold = q_gold.reshape(num_tokens, -1)
k_gold = k_gold.reshape(num_tokens, -1)
# Compare the results.
torch.testing.assert_close(q.to(torch.float32).cpu(), q_gold, atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
torch.testing.assert_close(k.to(torch.float32).cpu(), k_gold, atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
torch.testing.assert_close(
v.to(torch.float32).cpu(), v_gold.to(torch.float32), atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL
)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,167 @@
import gc
import numpy as np
import pytest
import torch
NUM_TOKENS = [1, 8, 32]
NUM_QKV_HEADS = [(6, 1), (8, 2)]
HEAD_DIMS = [128]
ROTARY_DIMS = [64, 128]
TP_WORLDS = [1]
EPS = [1e-6]
DTYPES = [torch.bfloat16]
SEEDS = [0]
DEVICES = [f"npu:{0}"]
DEFAULT_ATOL = 5e-2
DEFAULT_RTOL = 5e-3
def _build_rope(num_tokens, rotary_dim, dtype, device):
cos = torch.from_numpy(np.random.uniform(0, 1, [num_tokens, rotary_dim // 2])).to(dtype).to(device)
sin = torch.from_numpy(np.random.uniform(0, 1, [num_tokens, rotary_dim // 2])).to(dtype).to(device)
return cos.contiguous(), sin.contiguous()
def _apply_rope_neox(q, k, cos, sin, rotary_dim):
half = rotary_dim // 2
cos = cos.to(torch.float32).unsqueeze(1)
sin = sin.to(torch.float32).unsqueeze(1)
q_f32 = q.to(torch.float32)
k_f32 = k.to(torch.float32)
q1 = q_f32[..., :half]
q2 = q_f32[..., half:rotary_dim]
q_rot = torch.cat([q1 * cos - q2 * sin, q2 * cos + q1 * sin], dim=-1)
q_out = torch.cat([q_rot, q_f32[..., rotary_dim:]], dim=-1).to(q.dtype)
k1 = k_f32[..., :half]
k2 = k_f32[..., half:rotary_dim]
k_rot = torch.cat([k1 * cos - k2 * sin, k2 * cos + k1 * sin], dim=-1)
k_out = torch.cat([k_rot, k_f32[..., rotary_dim:]], dim=-1).to(k.dtype)
return q_out.contiguous(), k_out.contiguous()
def _fused_impl(
qkv,
q_weight,
k_weight,
q_hidden_size,
kv_hidden_size,
head_dim,
rotary_dim,
eps,
tp_world,
cos,
sin,
):
return torch.ops.vllm.split_qkv_tp_rmsnorm_rope(
input=qkv,
q_weight=q_weight,
k_weight=k_weight,
q_hidden_size=q_hidden_size,
kv_hidden_size=kv_hidden_size,
head_dim=head_dim,
rotary_dim=rotary_dim,
eps=eps,
tp_world=tp_world,
cos=cos,
sin=sin,
)
def _reference_impl(
qkv,
q_weight,
k_weight,
q_hidden_size,
kv_hidden_size,
head_dim,
rotary_dim,
eps,
tp_world,
cos,
sin,
):
q, k, v = qkv.split([q_hidden_size, kv_hidden_size, kv_hidden_size], dim=-1)
orig_dtype = q.dtype
q_f32 = q.to(torch.float32)
k_f32 = k.to(torch.float32)
q_var = q_f32.pow(2).mean(dim=-1, keepdim=True)
k_var = k_f32.pow(2).mean(dim=-1, keepdim=True)
q_out = (q_f32 * torch.rsqrt(q_var + eps) * q_weight.to(torch.float32)).to(orig_dtype)
k_out = (k_f32 * torch.rsqrt(k_var + eps) * k_weight.to(torch.float32)).to(orig_dtype)
q_3d = q_out.view(q.shape[0], -1, head_dim).contiguous()
k_3d = k_out.view(k.shape[0], -1, head_dim).contiguous()
q_3d, k_3d = _apply_rope_neox(q_3d, k_3d, cos.contiguous(), sin.contiguous(), rotary_dim)
return (
q_3d.view(q.shape[0], q_hidden_size).contiguous(),
k_3d.view(k.shape[0], kv_hidden_size).contiguous(),
v.contiguous(),
)
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
@pytest.mark.parametrize("num_q_heads, num_kv_heads", NUM_QKV_HEADS)
@pytest.mark.parametrize("head_dim", HEAD_DIMS)
@pytest.mark.parametrize("rotary_dim", ROTARY_DIMS)
@pytest.mark.parametrize("tp_world", TP_WORLDS)
@pytest.mark.parametrize("eps", EPS)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("seed", SEEDS)
@pytest.mark.parametrize("device", DEVICES)
@torch.inference_mode()
def test_split_qkv_tp_rmsnorm_rope(
num_tokens, num_q_heads, num_kv_heads, head_dim, rotary_dim, tp_world, eps, dtype, seed, device
):
torch.manual_seed(seed)
np.random.seed(seed)
torch.set_default_device(device)
q_hidden_size = num_q_heads * head_dim
kv_hidden_size = num_kv_heads * head_dim
qkv = torch.randn(num_tokens, q_hidden_size + kv_hidden_size * 2, dtype=dtype, device=device)
q_weight = torch.randn(q_hidden_size, dtype=torch.float32, device=device) * 0.1 + 1.0
k_weight = torch.randn(kv_hidden_size, dtype=torch.float32, device=device) * 0.1 + 1.0
cos, sin = _build_rope(num_tokens, rotary_dim, dtype, device)
q_fused, k_fused, v_fused = _fused_impl(
qkv=qkv.clone(),
q_weight=q_weight.clone(),
k_weight=k_weight.clone(),
q_hidden_size=q_hidden_size,
kv_hidden_size=kv_hidden_size,
head_dim=head_dim,
rotary_dim=rotary_dim,
eps=eps,
tp_world=tp_world,
cos=cos,
sin=sin,
)
q_ref, k_ref, v_ref = _reference_impl(
qkv=qkv.clone(),
q_weight=q_weight.clone(),
k_weight=k_weight.clone(),
q_hidden_size=q_hidden_size,
kv_hidden_size=kv_hidden_size,
head_dim=head_dim,
rotary_dim=rotary_dim,
eps=eps,
tp_world=tp_world,
cos=cos,
sin=sin,
)
torch.testing.assert_close(q_fused.to(torch.float32), q_ref.to(torch.float32), atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
torch.testing.assert_close(k_fused.to(torch.float32), k_ref.to(torch.float32), atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
torch.testing.assert_close(v_fused.to(torch.float32), v_ref.to(torch.float32), atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,81 @@
import random
import pytest
import torch
from vllm_ascend.worker.v2.sample.gumbel import apply_temperature
# Common vocab sizes from mainstream models
VOCAB_SIZES = [
32000, # LLaMA / LLaMA2 / Mistral
50257, # GPT-2
65024, # ChatGLM
128256, # LLaMA3
151936, # Qwen2
]
def torch_apply_temperature(
logits: torch.Tensor,
expanded_idx_mapping: torch.Tensor,
temperature: torch.Tensor,
) -> None:
"""Pure PyTorch reference implementation of temperature scaling.
Args:
logits: Tensor of shape (num_tokens, vocab_size) containing the logits.
expanded_idx_mapping: Tensor containing the mapping from token index
to request index of tensor temperature.
temperature: Tensor containing the temperature value for each request.
"""
for token_idx in range(logits.shape[0]):
req_state_idx = expanded_idx_mapping[token_idx].item()
temp = temperature[req_state_idx].item()
if temp == 0.0 or temp == 1.0:
continue
logits[token_idx] = logits[token_idx].float() / temp
@pytest.mark.parametrize(
"num_tokens, vocab_size",
[(random.randint(1, 64), vocab_size) for vocab_size in VOCAB_SIZES],
)
def test_temperature_kernel(num_tokens, vocab_size):
"""
Test the Triton _temperature_kernel against a pure PyTorch reference.
The kernel divides logits by per-request temperature values. Tokens
whose temperature is 0.0 or 1.0 are skipped (logits unchanged).
Args:
num_tokens: Number of tokens (rows) in the logits tensor.
vocab_size: Vocabulary size (columns) in the logits tensor.
"""
torch.manual_seed(42)
# Build input tensors
logits_triton = torch.randn((num_tokens, vocab_size), dtype=torch.float32).npu()
logits_ref = logits_triton.clone()
num_requests = num_tokens
expanded_idx_mapping = torch.arange(num_tokens, dtype=torch.int32).npu()
# Include edge cases: 0.0 and 1.0 should leave logits unchanged
temperature = torch.rand(num_requests, dtype=torch.float32).npu()
temperature = temperature * 1.8 + 0.2 # range [0.2, 2.0]
if num_requests >= 3:
temperature[0] = 0.0
temperature[1] = 1.0
# ========== Run Triton kernel ==========
apply_temperature(logits_triton, expanded_idx_mapping, temperature)
# ========== Run PyTorch reference ==========
torch_apply_temperature(logits_ref, expanded_idx_mapping, temperature)
# ========== Verify results ==========
assert torch.allclose(logits_triton, logits_ref, atol=1e-4, rtol=1e-5), (
f"Triton temperature kernel output differs from torch reference.\n"
f"Max diff: {torch.max(torch.abs(logits_triton - logits_ref))}\n"
f"Mean diff: {torch.mean(torch.abs(logits_triton - logits_ref).float())}"
)