@@ -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()
|
||||
@@ -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")
|
||||
@@ -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)
|
||||
@@ -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))}"
|
||||
)
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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!")
|
||||
@@ -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))}"
|
||||
)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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))}"
|
||||
)
|
||||
@@ -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))}"
|
||||
)
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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']))}"
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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())}"
|
||||
)
|
||||
Reference in New Issue
Block a user