init v0.23.0

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

View File

@@ -0,0 +1,84 @@
# ==========================================
# Shared Configurations
# ==========================================
_envs: &envs
OMP_NUM_THREADS: "10"
OMP_PROC_BIND: "false"
HCCL_BUFFSIZE: "1024"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
SERVER_PORT: "DEFAULT_PORT"
VLLM_ENGINE_READY_TIMEOUT_S: "3000"
_server_cmd: &server_cmd
- "--quantization"
- "ascend"
- "--data-parallel-size"
- "2"
- "--tensor-parallel-size"
- "8"
- "--enable-expert-parallel"
- "--port"
- "$SERVER_PORT"
- "--seed"
- "1024"
- "--max-model-len"
- "36864"
- "--max-num-batched-tokens"
- "4096"
- "--max-num-seqs"
- "16"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.9"
- "--speculative-config"
- '{"num_speculative_tokens": 1, "method": "mtp"}'
- "--additional-config"
- '{"enable_weight_nz_layout": true}'
_benchmarks_acc: &benchmarks_acc
acc:
case_type: accuracy
dataset_path: vllm-ascend/gsm8k-lite
request_conf: vllm_api_general_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_chat_prompt
max_out_len: 32768
batch_size: 32
baseline: 95
threshold: 10
_benchmarks_perf: &benchmarks_perf
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs400
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 400
max_out_len: 1500
batch_size: 1000
baseline: 1
threshold: 0.97
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "DeepSeek-R1-0528-W8A8-single"
model: "vllm-ascend/DeepSeek-R1-0528-W8A8"
envs:
<<: *envs
server_cmd: *server_cmd
server_cmd_extra:
- "--enforce-eager"
benchmarks:
- name: "DeepSeek-R1-0528-W8A8-aclgraph"
model: "vllm-ascend/DeepSeek-R1-0528-W8A8"
envs:
<<: *envs
server_cmd: *server_cmd
benchmarks:
<<: *benchmarks_acc
<<: *benchmarks_perf

View File

@@ -0,0 +1,81 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "DeepSeek-V3.2-W8A8-DCP-replicated-indexer"
model: "vllm-ascend/DeepSeek-V3.2-W8A8"
envs:
VLLM_ASCEND_ENABLE_NZ: "1"
HCCL_OP_EXPANSION_MODE: "AIV"
OMP_PROC_BIND: "false"
OMP_NUM_THREADS: "20"
HCCL_BUFFSIZE: "768"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
VLLM_SERVER_DEV_MODE: "1"
VLLM_WORKER_MULTIPROC_METHOD: "spawn"
ASCEND_LAUNCH_BLOCKING: "0"
ASCEND_ENABLE_USE_FABRIC_MEM: "1"
VLLM_ASCEND_ENABLE_FLASHCOMM1: "1"
VLLM_ASCEND_ENABLE_FUSED_MC2: "0"
VLLM_ASCEND_ENABLE_MLAPO: "0"
PYTHONHASHSEED: "0"
ASCEND_A3_ENABLE: "1"
VLLM_ENGINE_READY_TIMEOUT_S: "10000"
VLLM_RPC_TIMEOUT: "3600000"
VLLM_EXECUTE_MODEL_TIMEOUT_SECONDS: "30000"
TASK_QUEUE_ENABLE: "1"
CPU_AFFINITY_CONF: "1"
SERVER_PORT: "DEFAULT_PORT"
server_cmd:
- "--port"
- "$SERVER_PORT"
- "--seed"
- "1024"
- "--max-model-len"
- "8192"
- "--max-num-batched-tokens"
- "1024"
- "--max-num-seqs"
- "32"
- "--trust-remote-code"
- "--quantization"
- "ascend"
- "--data-parallel-size"
- "1"
- "--pipeline-parallel-size"
- "1"
- "--tensor-parallel-size"
- "16"
- "--prefill-context-parallel-size"
- "1"
- "--decode-context-parallel-size"
- "16"
- "--cp-kv-cache-interleave-size"
- "1"
- "--block-size"
- "128"
- "--enable-expert-parallel"
- "--gpu-memory-utilization"
- "0.95"
- "--api-server-count"
- "1"
- "--safetensors-load-strategy"
- "prefetch"
- "--compilation-config"
- '{"cudagraph_mode": "FULL_DECODE_ONLY", "cudagraph_capture_sizes":[4, 16, 64, 128]}'
- "--additional-config"
- '{"enable_dsa_cp": true, "ascend_compilation_config":{"enable_npugraph_ex": true, "enable_static_kernel": false}, "multistream_overlap_shared_expert": true, "enable_mc2_hierarchy_comm": false, "enable_sparse_sfa_c8": true, "enable_sparse_li_c8": true, "enable_cpu_binding": true, "recompute_scheduler_enable": false}'
- "--speculative-config"
- '{"num_speculative_tokens": 3, "method": "deepseek_mtp"}'
test_content: []
benchmarks:
acc_gsm8k:
case_type: accuracy
dataset_path: vllm-ascend/gsm8k-lite
request_conf: vllm_api_general_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_chat_prompt
max_out_len: 8192
batch_size: 32
baseline: 95
threshold: 5

View File

@@ -0,0 +1,80 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "DeepSeek-V3.2-W8A8-TP8-DP2"
model: "vllm-ascend/DeepSeek-V3.2-W8A8"
envs:
OMP_PROC_BIND: "false"
OMP_NUM_THREADS: "1"
HCCL_OP_EXPANSION_MODE: "AIV"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
VLLM_USE_V1: "1"
HCCL_BUFFSIZE: "256"
ASCEND_AGGREGATE_ENABLE: "1"
ASCEND_TRANSPORT_PRINT: "1"
ACL_OP_INIT_MODE: "1"
ASCEND_A3_ENABLE: "1"
VLLM_NIXL_ABORT_REQUEST_TIMEOUT: "300000"
TASK_QUEUE_ENABLE: "1"
VLLM_ASCEND_ENABLE_MLAPO: "1"
VLLM_ASCEND_ENABLE_FLASHCOMM1: "1"
SERVER_PORT: "DEFAULT_PORT"
VLLM_ENGINE_READY_TIMEOUT_S: "3000"
server_cmd:
- "--tensor-parallel-size"
- "8"
- "--data-parallel-size"
- "2"
- "--port"
- "$SERVER_PORT"
- "--seed"
- "1024"
- "--max-model-len"
- "67000"
- "--max-num-batched-tokens"
- "4096"
- "--max-num-seqs"
- "8"
- "--trust-remote-code"
- "--quantization"
- "ascend"
- "--async-scheduling"
- "--no-enable-prefix-caching"
- "--enable-expert-parallel"
- "--gpu-memory-utilization"
- "0.95"
- "--compilation-config"
- '{"cudagraph_capture_sizes":[1,2,4,8,16,24,32,40,48], "cudagraph_mode":"FULL_DECODE_ONLY"}'
- "--speculative-config"
- '{"num_speculative_tokens": 3, "method":"deepseek_mtp"}'
- "--reasoning-parser"
- "deepseek_v3"
- "--tokenizer_mode"
- "deepseek_v32"
benchmarks:
acc_aime2025:
case_type: accuracy
dataset_path: vllm-ascend/aime2025
request_conf: vllm_api_general_chat
dataset_conf: aime2025/aime2025_gen_0_shot_chat_prompt
max_out_len: 32768
batch_size: 32
baseline: 86.67
temperature: 1.0
top_p: 0.95
thinking: true
threshold: 10
perf_2:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs400
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 16
max_out_len: 1500
batch_size: 4
request_rate: 0
baseline: 1
threshold: 0.97

View File

@@ -0,0 +1,80 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "DeepSeek-V4-Flash-W8A8-A3"
model: "Eco-Tech/DeepSeek-V4-Flash-w8a8-mtp"
special_dependencies:
transformers: "5.9.0"
envs:
OMP_PROC_BIND: "false"
OMP_NUM_THREADS: "1"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
HCCL_BUFFSIZE: "1024"
VLLM_ASCEND_ENABLE_FUSED_MC2: "0"
VLLM_ASCEND_ENABLE_FLASHCOMM1: "1"
ASCEND_LAUNCH_BLOCKING: "0"
SERVER_PORT: "DEFAULT_PORT"
VLLM_ENGINE_READY_TIMEOUT_S: "3000"
server_cmd:
- "--enable-prefix-caching"
- "--max-model-len"
- "1048576"
- "--max-num-batched-tokens"
- "10240"
- "--gpu-memory-utilization"
- "0.9"
- "--max-num-seqs"
- "64"
- "--data-parallel-size"
- "4"
- "--tensor-parallel-size"
- "4"
- "--enable-expert-parallel"
- "--tokenizer-mode"
- "deepseek_v4"
- "--tool-call-parser"
- "deepseek_v4"
- "--enable-auto-tool-choice"
- "--reasoning-parser"
- "deepseek_v4"
- "--safetensors-load-strategy"
- "prefetch"
- "--quantization"
- "ascend"
- "--api-server-count"
- "1"
- "--speculative-config"
- '{"num_speculative_tokens": 1,"method": "mtp","enforce_eager": true}'
- "--port"
- "$SERVER_PORT"
- "--block-size"
- "128"
- "--compilation-config"
- '{"cudagraph_mode": "FULL_DECODE_ONLY"}'
- "--async-scheduling"
- "--additional-config"
- '{"ascend_compilation_config":{"enable_npugraph_ex":true,"enable_static_kernel":false},"enable_cpu_binding":"true","enable_shared_expert_dp":true,"multistream_overlap_shared_expert":true}'
benchmarks:
acc-gpqa:
case_type: accuracy
dataset_path: vllm-ascend/gpqa
request_conf: vllm_api_general_chat
dataset_conf: gpqa/gpqa_gen_0_shot_cot_chat_prompt
max_out_len: 65536
batch_size: 32
baseline: 86.36
threshold: 5
thinking: true
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K_prefix0_in32768_bs1000_deepseek
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 64
max_out_len: 1024
batch_size: 16
request_rate: 0
baseline: 1
threshold: 0.97

View File

@@ -0,0 +1,66 @@
# ==========================================
# Shared Configurations
# ==========================================
_envs: &envs
HCCL_BUFFSIZE: "512"
SERVER_PORT: "DEFAULT_PORT"
OMP_PROC_BIND: "false"
OMP_NUM_THREADS: "1"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
HCCL_OP_EXPANSION_MODE: "AIV"
VLLM_ASCEND_BALANCE_SCHEDULING: "1"
VLLM_ASCEND_ENABLE_TOPK_OPTIMIZE: "1"
VLLM_ASCEND_ENABLE_FLASHCOMM1: "1"
VLLM_ASCEND_ENABLE_FUSED_MC2: "1"
_server_cmd: &server_cmd
- "--enable-expert-parallel"
- "--tensor-parallel-size"
- "8"
- "--data-parallel-size"
- "2"
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "8192"
- "--max-num-batched-tokens"
- "8192"
- "--max-num-seqs"
- "16"
- "--quantization"
- "ascend"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.9"
- "--speculative-config"
- '{"num_speculative_tokens": 3, "method":"mtp"}'
- "--additional-config"
- '{"enable_shared_expert_dp": true, "ascend_fusion_config": {"fusion_ops_gmmswigluquant": false}}'
_benchmarks: &benchmarks
acc:
case_type: accuracy
dataset_path: vllm-ascend/gsm8k-lite
request_conf: vllm_api_general_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_chat_prompt
max_out_len: 4096
batch_size: 8
baseline: 95
threshold: 10
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "GLM-4.7-TP8-DP2-decodegraph"
model: "Eco-Tech/GLM-4.7-W8A8-floatmtp"
envs:
<<: *envs
server_cmd: *server_cmd
server_cmd_extra:
- "--compilation-config"
- '{"cudagraph_capture_sizes": [1,2,4,8,16,32,64,128,256,512], "cudagraph_mode": "FULL_DECODE_ONLY"}'
benchmarks:
<<: *benchmarks

View File

@@ -0,0 +1,82 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "GLM-5.1-W8A8-PrefillMC2"
model: "Eco-Tech/GLM-5.1-w8a8" #need update
envs:
OMP_PROC_BIND: "false"
OMP_NUM_THREADS: "1"
HCCL_OP_EXPANSION_MODE: "AIV"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
VLLM_USE_V1: "1"
HCCL_BUFFSIZE: "1800"
ASCEND_AGGREGATE_ENABLE: "1"
ASCEND_TRANSPORT_PRINT: "1"
ACL_OP_INIT_MODE: "1"
ASCEND_A3_ENABLE: "1"
VLLM_NIXL_ABORT_REQUEST_TIMEOUT: "300000"
TASK_QUEUE_ENABLE: "1"
VLLM_ASCEND_ENABLE_MLAPO: "1"
VLLM_ASCEND_ENABLE_FLASHCOMM1: "1"
SERVER_PORT: "DEFAULT_PORT"
server_cmd:
- "--tensor-parallel-size"
- "16"
- "--data-parallel-size"
- "1"
- "--port"
- "$SERVER_PORT"
- "--seed"
- "1024"
- "--max-model-len"
- "10240"
- "--max-num-batched-tokens"
- "4096"
- "--max-num-seqs"
- "32"
- "--trust-remote-code"
- "--quantization"
- "ascend"
- "--async-scheduling"
- "--no-enable-prefix-caching"
- "--enable-expert-parallel"
- "--gpu-memory-utilization"
- "0.94"
- "--compilation-config"
- '{"cudagraph_mode":"FULL_DECODE_ONLY"}'
- "--speculative-config"
- '{"num_speculative_tokens": 3, "method":"deepseek_mtp", "enforce_eager": true}'
- "--additional_config"
- '{"enable_prefill_mc2": true}'
- "--reasoning-parser"
- "glm45"
- "--tool-call-parser"
- "glm47"
benchmarks:
acc_gsm8k:
case_type: accuracy
dataset_path: vllm-ascend/gsm8k-lite
request_conf: vllm_api_general_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_chat_prompt
max_out_len: 8192
batch_size: 32
baseline: 96.88
temperature: 1.0
top_p: 0.95
thinking: true
threshold: 5
perf_2:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs400
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 64
max_out_len: 1500
batch_size: 32
request_rate: 0
baseline: 1
threshold: 0.97

View File

@@ -0,0 +1,58 @@
# ==========================================
# Shared Configurations
# ==========================================
_envs: &envs
VLLM_USE_MODELSCOPE: "true"
HCCL_OP_EXPANSION_MODE: "AIV"
HCCL_BUFFSIZE: "1024"
OMP_PROC_BIND: "false"
OMP_NUM_THREADS: "1"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
SERVER_PORT: "DEFAULT_PORT"
_server_cmd: &server_cmd
- "--tensor-parallel-size"
- "16"
- "--enable-expert-parallel"
- "--enable-ep-weight-filter"
- "--tool-call-parser"
- "hy_v3"
- "--reasoning-parser"
- "hy_v3"
- "--enable-auto-tool-choice"
- "--max-model-len"
- "32768"
- "--max-num-seqs"
- "8"
- "--port"
- "$SERVER_PORT"
- "--speculative-config"
- '{"method": "mtp", "num_speculative_tokens": 1}'
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.9"
_benchmarks: &benchmarks
acc_gsm8k:
case_type: accuracy
dataset_path: vllm-ascend/gsm8k-lite
request_conf: vllm_api_general_chat
dataset_conf: gsm8k/gsm8k_gen_4_shot_cot_chat_prompt
max_out_len: 4096
batch_size: 8
baseline: 93.07
threshold: 10
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Hy3-preview-TP16-EP-MTP"
model: "Tencent-Hunyuan/Hy3-preview"
envs:
<<: *envs
server_cmd: *server_cmd
benchmarks:
<<: *benchmarks

View File

@@ -0,0 +1,52 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Kimi-K2-Thinking-TP16-Case"
model: "moonshotai/Kimi-K2-Thinking"
envs:
HCCL_BUFFSIZE: "1024"
TASK_QUEUE_ENABLE: "1"
OMP_PROC_BIND: "false"
HCCL_OP_EXPANSION_MODE: "AIV"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
SERVER_PORT: "DEFAULT_PORT"
server_cmd:
- "--tensor-parallel-size"
- "16"
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "8192"
- "--max-num-batched-tokens"
- "8192"
- "--max-num-seqs"
- "12"
- "--gpu-memory-utilization"
- "0.9"
- "--trust-remote-code"
- "--enable-expert-parallel"
- "--no-enable-prefix-caching"
benchmarks:
acc:
case_type: accuracy
dataset_path: vllm-ascend/gsm8k-lite
request_conf: vllm_api_general_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_chat_prompt
max_out_len: 4096
batch_size: 32
baseline: 95
threshold: 10
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs400
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 512
max_out_len: 256
batch_size: 64
trust_remote_code: true
request_rate: 11.2
baseline: 1
threshold: 0.97

View File

@@ -0,0 +1,91 @@
# ==========================================
# Shared Configurations
# ==========================================
_envs: &envs
HCCL_BUFFSIZE: "512"
SERVER_PORT: "DEFAULT_PORT"
HCCL_OP_EXPANSION_MODE: "AIV"
OMP_PROC_BIND: "false"
OMP_NUM_THREADS: "1"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
TASK_QUEUE_ENABLE: "1"
VLLM_ASCEND_ENABLE_MLAPO: "1"
VLLM_ASCEND_ENABLE_FLASHCOMM1: "1"
VLLM_ASCEND_ENABLE_NZ: "1"
_server_cmd: &server_cmd
- "--enable-expert-parallel"
- "--enable-prefix-caching"
- "--enable-chunked-prefill"
- "--allowed-local-media-path"
- "/"
- "--tensor-parallel-size"
- "4"
- "--data-parallel-size"
- "4"
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "133120"
- "--max-num-batched-tokens"
- "8192"
- "--max-num-seqs"
- "16"
- "--quantization"
- "ascend"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.9"
- "--seed"
- "42"
- "--compilation-config"
- '{"cudagraph_capture_sizes":[4,8,12,16,32], "cudagraph_mode":"FULL_DECODE_ONLY"}'
- "--speculative-config"
- '{"method":"eagle3", "model":"lightseekorg/kimi-k2.5-eagle3", "num_speculative_tokens":3}'
- "--additional-config"
- '{"enable_shared_expert_dp":true}'
- "--mm-processor-cache-gb"
- "0"
- "--mm-encoder-tp-mode"
- "data"
_benchmarks: &benchmarks
acc_aime2025:
case_type: accuracy
dataset_path: vllm-ascend/aime2025
request_conf: vllm_api_general_chat
dataset_conf: aime2025/aime2025_gen_0_shot_chat_prompt
max_out_len: 65536
temperature: 0.0
top_p: 1
top_k: -1
repetition_penalty: 1.0
batch_size: 32
baseline: 95
threshold: 10
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K_prefix90_in131072_bs1000_kimi
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 8
max_out_len: 1024
batch_size: 2
trust_remote_code: true
request_rate: 0
baseline: 1
threshold: 0.97
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Kimi-K2.5-W4A8-Case"
model: "Eco-Tech/Kimi-K2.5-W4A8"
envs:
<<: *envs
server_cmd: *server_cmd
benchmarks:
<<: *benchmarks

View File

@@ -0,0 +1,70 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Kimi-K2.6-W4A8-in3.5k-out1.5k-TPOT50-0-128-32"
model: "Eco-Tech/Kimi-K2.6-w4a8"
envs:
HCCL_OP_EXPANSION_MODE: "AIV"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
OMP_PROC_BIND: "false"
OMP_NUM_THREADS: "1"
TASK_QUEUE_ENABLE: "1"
HCCL_BUFFSIZE: "800"
VLLM_ASCEND_ENABLE_FLASHCOMM1: "1"
VLLM_ASCEND_ENABLE_MLAPO: "1"
VLLM_ASCEND_BALANCE_SCHEDULING: "1"
VLLM_ASCEND_ENABLE_FUSED_MC2: "1"
DYNAMIC_EPLB: "true"
SERVER_PORT: "DEFAULT_PORT"
server_cmd:
- "--quantization"
- "ascend"
- "--port"
- "$SERVER_PORT"
- "--allowed-local-media-path"
- "/"
- "--trust-remote-code"
- "--safetensors-load-strategy"
- 'prefetch'
- "--tensor-parallel-size"
- "8"
- "--data-parallel-size"
- "2"
- "--no-enable-prefix-caching"
- "--enable-expert-parallel"
- "--max-num-seqs"
- "24"
- "--max-model-len"
- "6144"
- "--max-num-batched-tokens"
- "4096"
- "--gpu-memory-utilization"
- "0.85"
- "--seed"
- "42"
- "--async-scheduling"
- "--compilation-config"
- '{"cudagraph_mode":"FULL_DECODE_ONLY"}'
- "--additional-config"
- '{"eplb_config": {"dynamic_eplb": true},"ascend_compilation_config": {"enable_static_kernel": false}}'
- "--profiler-config"
- '{"profiler": "torch", "torch_profiler_dir": "./vllm_profile", "torch_profiler_with_stack": true}'
- "--mm-processor-cache-gb"
- "0"
- "--mm-encoder-tp-mode"
- "data"
- "--speculative-config"
- '{"method": "dflash","model": "z-lab/Kimi-K2.5-DFlash", "num_speculative_tokens": 7}'
benchmarks:
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs4096-prefix0-kimi
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 128
max_out_len: 1500
batch_size: 32
request_rate: 0
baseline: 1433.4454
threshold: 0.97

View File

@@ -0,0 +1,91 @@
# ==========================================
# Shared Configurations
# ==========================================
_envs: &envs
OMP_NUM_THREADS: "100"
OMP_PROC_BIND: "false"
HCCL_BUFFSIZE: "1024"
VLLM_RPC_TIMEOUT: "3600000"
VLLM_EXECUTE_MODEL_TIMEOUT_SECONDS: "3600000"
SERVER_PORT: "DEFAULT_PORT"
VLLM_ENGINE_READY_TIMEOUT_S: "7200"
_server_cmd: &server_cmd
- "--quantization"
- "ascend"
- "--seed"
- "1024"
- "--no-enable-prefix-caching"
- "--data-parallel-size"
- "2"
- "--tensor-parallel-size"
- "8"
- "--enable-expert-parallel"
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "40960"
- "--max-num-seqs"
- "14"
- "--trust-remote-code"
_benchmarks_gsm8k: &benchmarks_gsm8k
acc:
case_type: accuracy
dataset_path: vllm-ascend/gsm8k-lite
request_conf: vllm_api_general_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_chat_prompt
max_out_len: 32768
batch_size: 32
baseline: 95
threshold: 10
_benchmarks_aime: &benchmarks_aime
acc:
case_type: accuracy
dataset_path: vllm-ascend/aime2024
request_conf: vllm_api_general_chat
dataset_conf: aime2024/aime2024_gen_0_shot_chat_prompt
max_out_len: 32768
batch_size: 32
baseline: 86.67
threshold: 10
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "MTPX-DeepSeek-R1-0528-W8A8-mtp2"
model: "vllm-ascend/DeepSeek-R1-0528-W8A8"
envs:
<<: *envs
server_cmd: *server_cmd
server_cmd_extra:
- "--max-num-batched-tokens"
- "4096"
- "--speculative-config"
- '{"num_speculative_tokens": 2, "method": "mtp"}'
- "--gpu-memory-utilization"
- "0.92"
benchmarks:
<<: *benchmarks_gsm8k
- name: "MTPX-DeepSeek-R1-0528-W8A8-mtp3"
model: "vllm-ascend/DeepSeek-R1-0528-W8A8"
envs:
<<: *envs
HCCL_OP_EXPANSION_MODE: "AIV"
server_cmd: *server_cmd
server_cmd_extra:
- "--max-num-batched-tokens"
- "2048"
- "--speculative-config"
- '{"num_speculative_tokens": 3, "method": "mtp"}'
- "--gpu-memory-utilization"
- "0.9"
- "--compilation-config"
- '{"cudagraph_capture_sizes": [56], "cudagraph_mode": "FULL_DECODE_ONLY"}'
benchmarks:
<<: *benchmarks_aime

View File

@@ -0,0 +1,88 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "MiniMax-M2.5-w8a8"
model: "Eco-Tech/MiniMax-M2.5-w8a8-QuaRot"
envs:
HCCL_BUFFSIZE: "512"
HCCL_OP_EXPANSION_MODE: "AIV"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
VLLM_ASCEND_ENABLE_FLASHCOMM1: "1"
OMP_NUM_THREADS: "1"
TASK_QUEUE_ENABLE: "1"
VLLM-ASCEND_BALANCE_SCHEDULING: "1"
HCCL_INTRA_PCIE_ENABLE: "1"
HCCL_INTRA_ROCE_ENABLE: "0"
OMP_PROC_BIND: "false"
VLLM_TORCH_PROFILER_WITH_STACK: "0"
VLLM_TORCH_PROFILER_DIR: "./profile"
VLLM_USE_MODELSCOPE: "true"
SERVER_PORT: "DEFAULT_PORT"
server_cmd:
- "--tensor-parallel-size"
- "8"
- "--port"
- "$SERVER_PORT"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.9"
- "--quantization"
- "ascend"
- "--additional-config"
- '{"enable_cpu_binding":true}'
- "--model-loader-extra-config"
- '{"enable_multithread_load":true,"num_threads":16}'
- "--speculative_config"
- '{"method":"eagle3","model":"vllm-ascend/MiniMax-M2.5-eagle-model-0318","num_speculative_tokens":3}'
- "--enable-expert-parallel"
- "--enable-chunked-prefill"
- "--enable-prefix-caching"
- "--max-num-seqs"
- "100"
- "--max-model-len"
- "196608"
- "--seed"
- "1024"
- "--max-num-batched-tokens"
- "6144"
- "--enable-auto-tool-choice"
- "--tool-call-parser"
- "minimax_m2"
- "--reasoning-parser"
- "minimax_m2_append_think"
- "--enable-force-include-usage"
- "--profiler-config"
- '{"profiler":"torch","torch_profiler_dir":"./profile","torch_profiler_with_stack":false}'
- "--compilation-config"
- '{"cudagraph_mode":"FULL_DECODE_ONLY","cudagraph_capture_sizes":[4,16,40,80,160,256,400]}'
benchmarks:
acc:
case_type: accuracy
dataset_path: vllm-ascend/gpqa
request_conf: vllm_api_general_chat
dataset_conf: gsm8k/gpqa_gen_0_shot_str
max_out_len: 131072
batch_size: 64
baseline: 83
threshold: 5
bos_token_id: 200019
do_sample: true
eos_token_id: 200020
temperature: 1.0
top_p: 0.95
top_k: 40
transformers_version: 4.46.1
ignore_eos: false
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs2800
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 360
max_out_len: 1500
batch_size: 120
request_rate: 0
baseline: 2042
threshold: 0.97

View File

@@ -0,0 +1,73 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "MiniMax-M2.5-w8a8"
model: "Eco-Tech/MiniMax-M2.5-w8a8-QuaRot"
envs:
HCCL_BUFFSIZE: "512"
HCCL_OP_EXPANSION_MODE: "AIV"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
OMP_NUM_THREADS: "1"
TASK_QUEUE_ENABLE: "1"
VLLM-ASCEND_ENABLE_NZ: "1"
VLLM_ASCEND_ENABLE_FLASHCOMM1: "1"
VLLM-ASCEND_BALANCE_SCHEDULING: "1"
VLLM_USE_MODELSCOPE: "true"
SERVER_PORT: "DEFAULT_PORT"
server_cmd:
- "--tensor-parallel-size"
- "8"
- "--data-parallel-size"
- "1"
- "--port"
- "$SERVER_PORT"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.85"
- "--quantization"
- "ascend"
- "--no-enable-prefix-caching"
- "--additional-config"
- '{"enable_cpu_binding":true}'
- "--speculative_config"
- '{"method":"eagle3","model":"vllm-ascend/MiniMax-M2.5-eagle-model-0318","num_speculative_tokens":3}'
- "--enable-expert-parallel"
- "--max-num-seqs"
- "128"
- "--max-num-batched-tokens"
- "16384"
- "--max-model-len"
- "196608"
- "--compilation-config"
- '{"cudagraph_mode":"FULL_DECODE_ONLY"}'
benchmarks:
acc_aime2025:
case_type: accuracy
dataset_path: vllm-ascend/aime2025
request_conf: vllm_api_general_chat
dataset_conf: aime2025/aime2025_gen_0_shot_chat_prompt
max_out_len: 131072
batch_size: 32
baseline: 90
threshold: 10
bos_token_id: 200019
do_sample: true
eos_token_id: 200020
temperature: 1.0
top_p: 0.95
top_k: 40
transformers_version: 4.46.1
ignore_eos: false
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs2800
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 512
max_out_len: 1500
batch_size: 128
request_rate: 0
baseline: 1116
threshold: 0.97

View File

@@ -0,0 +1,85 @@
# ==========================================
# Shared Configurations
# ==========================================
_envs: &envs
SERVER_PORT: "DEFAULT_PORT"
HCCL_OP_EXPANSION_MODE: "AIV"
HCCL_BUFFSIZE: "1200"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
OMP_NUM_THREADS: "1"
LD_PRELOAD: "/usr/lib/aarch64-linux-gnu/libjemalloc.so.2:$LD_PRELOAD"
TASK_QUEUE_ENABLE: "1"
_server_cmd: &server_cmd
- "--port"
- "$SERVER_PORT"
- "--host"
- "0.0.0.0"
- "--tensor-parallel-size"
- "8"
- "--data-parallel-size"
- "2"
- "--enable-expert-parallel"
- "--async-scheduling"
- "--max-num-seqs"
- "128"
- "--safetensors-load-strategy"
- 'prefetch'
- "--max-num-batched-tokens"
- "16384"
- "--trust-remote-code"
- "--quantization"
- "ascend"
- "--enable-auto-tool-choice"
- "--tool-call-parser"
- "minimax_m2"
- "--speculative-config"
- '{"method":"eagle3","model":"Eco-Tech/MiniMax-M2.7-eagle-model-short","num_speculative_tokens":3}'
- "--compilation-config"
- '{"cudagraph_mode": "FULL_DECODE_ONLY"}'
- "--additional-config"
- '{"enable_cpu_binding": true, "enable_npugraph_ex": true, "enable_static_kernel": true,"enable_fused_mc2":true,"weight_nz_mode":true,"enable_flashcomm1":true}'
_benchmarks_3500: &benchmarks_3500
perf_50:
case_type: performance
dataset_path: vllm-ascend/GSM8K_in3500_bs4000_minimax
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 760
max_out_len: 1500
batch_size: 190
request_rate: 0
baseline: 4573.02
threshold: 0.97
perf_20:
case_type: performance
dataset_path: vllm-ascend/GSM8K_in3500_bs4000_minimax
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 192
max_out_len: 1500
batch_size: 48
request_rate: 0
baseline: 2229.147
threshold: 0.97
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "MiniMax-M2.7-3500"
model: "vllm-ascend/MiniMax-M2.7-w8a8-QuaRot"
envs:
<<: *envs
server_cmd: *server_cmd
server_cmd_extra:
- "--max-model-len"
- "70000"
- "--gpu-memory-utilization"
- "0.8"
- "--no-enable-prefix-caching"
benchmarks:
<<: *benchmarks_3500

View File

@@ -0,0 +1,78 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "prefix-cache-deepseek-r1-0528-w8a8"
model: "vllm-ascend/DeepSeek-R1-0528-W8A8"
envs:
OMP_NUM_THREADS: "10"
OMP_PROC_BIND: "false"
HCCL_BUFFSIZE: "1024"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
SERVER_PORT: "DEFAULT_PORT"
VLLM_ENGINE_READY_TIMEOUT_S: "7200"
server_cmd:
- "--quantization"
- "ascend"
- "--data-parallel-size"
- "2"
- "--tensor-parallel-size"
- "8"
- "--enable-expert-parallel"
- "--port"
- "$SERVER_PORT"
- "--seed"
- "1024"
- "--max-model-len"
- "5200"
- "--max-num-batched-tokens"
- "4096"
- "--max-num-seqs"
- "16"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.9"
- "--additional-config"
- '{"enable_weight_nz_layout": true}'
- "--speculative-config"
- '{"num_speculative_tokens": 1, "method": "mtp"}'
test_content:
- "benchmark_comparisons"
benchmark_comparisons_args:
- metric: "TTFT"
baseline: "prefix0"
target: "prefix75"
ratio: 0.5
operator: "<"
benchmarks:
warm_up:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in1024-bs210
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 210
max_out_len: 1
batch_size: 1000
baseline: 0
threshold: 0.97
prefix0:
case_type: performance
dataset_path: vllm-ascend/prefix0-in3500-bs210
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 210
max_out_len: 1
batch_size: 18
baseline: 1
threshold: 0.97
prefix75:
case_type: performance
dataset_path: vllm-ascend/prefix75-in3500-bs210
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 210
max_out_len: 1
batch_size: 18
baseline: 1
threshold: 0.97

View File

@@ -0,0 +1,70 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "prefix-cache-qwen3-32b-w8a8"
model: "vllm-ascend/Qwen3-32B-W8A8"
envs:
TASK_QUEUE_ENABLE: "1"
HCCL_OP_EXPANSION_MODE: "AIV"
SERVER_PORT: "DEFAULT_PORT"
server_cmd:
- "--quantization"
- "ascend"
- "--reasoning-parser"
- "qwen3"
- "--tensor-parallel-size"
- "4"
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "8192"
- "--max-num-batched-tokens"
- "8192"
- "--max-num-seqs"
- "256"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.9"
- "--additional-config"
- '{"enable_weight_nz_layout": true}'
test_content:
- "benchmark_comparisons"
benchmark_comparisons_args:
- metric: "TTFT"
baseline: "prefix0"
target: "prefix75"
ratio: 0.4
operator: "<"
benchmarks:
warm_up:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in1024-bs210
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 210
max_out_len: 1
batch_size: 1000
baseline: 0
threshold: 0.97
prefix0:
case_type: performance
dataset_path: vllm-ascend/prefix0-in3500-bs210
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 210
max_out_len: 1
batch_size: 48
baseline: 1
threshold: 0.97
prefix75:
case_type: performance
dataset_path: vllm-ascend/prefix75-in3500-bs210
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 210
max_out_len: 1
batch_size: 48
baseline: 1
threshold: 0.97

View File

@@ -0,0 +1,84 @@
# ==========================================
# Shared Configurations
# ==========================================
_envs: &envs
OMP_NUM_THREADS: "10"
OMP_PROC_BIND: "false"
HCCL_BUFFSIZE: "1024"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
VLLM_ASCEND_ENABLE_FLASHCOMM1: "1"
SERVER_PORT: "DEFAULT_PORT"
_server_cmd: &server_cmd
- "--quantization"
- "ascend"
- "--data-parallel-size"
- "4"
- "--tensor-parallel-size"
- "4"
- "--enable-expert-parallel"
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "40960"
- "--max-num-batched-tokens"
- "8192"
- "--max-num-seqs"
- "12"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.9"
_benchmarks: &benchmarks
acc:
case_type: accuracy
dataset_path: vllm-ascend/gsm8k-lite
request_conf: vllm_api_general_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_chat_prompt
max_out_len: 32768
batch_size: 32
top_k: 20
baseline: 95
threshold: 10
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Qwen3-235B-A22B-W8A8-full_graph"
model: "vllm-ascend/Qwen3-235B-A22B-W8A8"
envs:
<<: *envs
server_cmd: *server_cmd
server_cmd_extra:
- "--compilation-config"
- '{"cudagraph_mode": "FULL_DECODE_ONLY"}'
benchmarks:
<<: *benchmarks
- name: "Qwen3-235B-A22B-W8A8-piecewise"
model: "vllm-ascend/Qwen3-235B-A22B-W8A8"
envs:
<<: *envs
server_cmd: *server_cmd
server_cmd_extra:
- "--compilation-config"
- '{"cudagraph_mode": "PIECEWISE"}'
benchmarks:
<<: *benchmarks
- name: "Qwen3-235B-A22B-W8A8-EPLB"
model: "vllm-ascend/Qwen3-235B-A22B-W8A8"
envs:
<<: *envs
DYNAMIC_EPLB: "true"
server_cmd: *server_cmd
server_cmd_extra:
- "--additional-config"
- '{"eplb_config": {"dynamic_eplb": "true", "expert_heat_collection_interval": 600, "algorithm_execution_interval": 50, "num_redundant_experts": 16, "eplb_policy_type": 2}}'
- "--compilation-config"
- '{"cudagraph_mode": "FULL_DECODE_ONLY"}'
benchmarks:
<<: *benchmarks

View File

@@ -0,0 +1,41 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Qwen3-30B-A3B-W4A8-llm-compressor"
model: "vllm-ascend/Qwen3-30B-A3B-Instruct-2507-quantized.w4a8"
envs:
OMP_PROC_BIND: "false"
OMP_NUM_THREADS: "1"
HCCL_BUFFSIZE: "1024"
HCCL_OP_EXPANSION_MODE: "AIV"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
SERVER_PORT: "DEFAULT_PORT"
server_cmd:
- "--no-enable-prefix-caching"
- "--tensor-parallel-size"
- "2"
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "40960"
- "--max-num-batched-tokens"
- "16384"
- "--max-num-seqs"
- "128"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.8"
- "--compilation-config"
- '{"cudagraph_mode": "FULL_DECODE_ONLY"}'
benchmarks:
acc:
case_type: accuracy
dataset_path: vllm-ascend/gsm8k-lite
request_conf: vllm_api_general_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_chat_prompt
max_out_len: 32768
batch_size: 32
baseline: 95
threshold: 10

View File

@@ -0,0 +1,43 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Qwen3-30B-A3B-W8A8-TP1"
model: "vllm-ascend/Qwen3-30B-A3B-W8A8"
envs:
OMP_PROC_BIND: "false"
OMP_NUM_THREADS: "10"
HCCL_BUFFSIZE: "1024"
HCCL_OP_EXPANSION_MODE: "AIV"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
SERVER_PORT: "DEFAULT_PORT"
server_cmd:
- "--no-enable-prefix-caching"
- "--tensor-parallel-size"
- "1"
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "5600"
- "--max-num-batched-tokens"
- "16384"
- "--max-num-seqs"
- "100"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.9"
- "--compilation-config"
- '{"cudagraph_mode": "FULL_DECODE_ONLY"}'
benchmarks:
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs400
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 180
max_out_len: 1500
batch_size: 45
request_rate: 0
baseline: 1
threshold: 0.97

View File

@@ -0,0 +1,40 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Qwen3-30B-QuaRot"
model: "vllm-ascend/Qwen3-30B-A3B-W8A8-QuaRot"
envs:
VLLM_WORKER_MULTIPROC_METHOD: "spawn"
SERVER_PORT: "DEFAULT_PORT"
HCCL_BUFFSIZE: "768"
server_cmd:
- "--enforce-eager"
- "--no-enable-prefix-caching"
- "--enable-expert-parallel"
- "--tensor-parallel-size"
- "2"
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "8192"
- "--trust-remote-code"
- "--distributed-executor-backend"
- "mp"
- "--gpu-memory-utilization"
- "0.9"
- "--speculative-config"
- '{"method": "eagle3", "model": "AngelSlim/Qwen3-a3B_eagle3", "num_speculative_tokens": 3}'
benchmarks:
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs400
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 80
max_out_len: 1500
batch_size: 20
request_rate: 0
baseline: 1
threshold: 0.97

View File

@@ -0,0 +1,78 @@
# ==========================================
# Shared Configurations
# ==========================================
_envs: &envs
TASK_QUEUE_ENABLE: "1"
HCCL_OP_EXPANSION_MODE: "AIV"
VLLM_ASCEND_ENABLE_FLASHCOMM: "1"
SERVER_PORT: "DEFAULT_PORT"
_server_cmd: &server_cmd
- "--quantization"
- "ascend"
- "--no-enable-prefix-caching"
- "--tensor-parallel-size"
- "4"
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "40960"
- "--max-num-batched-tokens"
- "40960"
- "--block-size"
- "128"
- "--trust-remote-code"
- "--reasoning-parser"
- "qwen3"
- "--gpu-memory-utilization"
- "0.9"
- "--additional-config"
- '{"weight_prefetch_config":{"enabled":true}}'
_benchmarks: &benchmarks
acc:
case_type: accuracy
dataset_path: vllm-ascend/aime2024
request_conf: vllm_api_general_chat
dataset_conf: aime2024/aime2024_gen_0_shot_chat_prompt
max_out_len: 32768
batch_size: 32
baseline: 83.33
threshold: 10
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs400
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 288
max_out_len: 1500
batch_size: 72
baseline: 1
threshold: 0.97
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Qwen3-32B-W8A8-aclgraph-a2"
model: "vllm-ascend/Qwen3-32B-W8A8"
envs:
<<: *envs
server_cmd: *server_cmd
server_cmd_extra:
- "--compilation-config"
- '{"cudagraph_mode":"FULL_DECODE_ONLY","cudagraph_capture_sizes":[1,12,16,20,24,32,48,60,64,68,72,76,80]}'
benchmarks:
<<: *benchmarks
- name: "Qwen3-32B-W8A8-single-a2"
model: "vllm-ascend/Qwen3-32B-W8A8"
envs:
<<: *envs
server_cmd: *server_cmd
server_cmd_extra:
- "--enforce-eager"
benchmarks:

View File

@@ -0,0 +1,78 @@
# ==========================================
# Shared Configurations
# ==========================================
_envs: &envs
TASK_QUEUE_ENABLE: "1"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
HCCL_OP_EXPANSION_MODE: "AIV"
VLLM_ASCEND_ENABLE_FLASHCOMM: "1"
SERVER_PORT: "DEFAULT_PORT"
_server_cmd: &server_cmd
- "--quantization"
- "ascend"
- "--no-enable-prefix-caching"
- "--tensor-parallel-size"
- "4"
- "--max-num-seqs"
- "80"
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "40960"
- "--max-num-batched-tokens"
- "40960"
- "--block-size"
- "128"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.9"
- "--additional-config"
- '{"weight_prefetch_config":{"enabled":true}}'
_benchmarks: &benchmarks
acc:
case_type: accuracy
dataset_path: vllm-ascend/gsm8k-lite
request_conf: vllm_api_general_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_noncot_chat_prompt
max_out_len: 10240
batch_size: 32
baseline: 96
threshold: 10
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs400
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 304
max_out_len: 1500
batch_size: 76
baseline: 1
threshold: 0.97
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Qwen3-32B-W8A8-aclgraph-a3"
model: "vllm-ascend/Qwen3-32B-W8A8"
envs:
<<: *envs
server_cmd: *server_cmd
server_cmd_extra:
- "--compilation-config"
- '{"cudagraph_mode":"FULL_DECODE_ONLY","cudagraph_capture_sizes":[1,12,16,20,24,32,48,60,64,68,72,76,80]}'
benchmarks:
<<: *benchmarks
- name: "Qwen3-32B-W8A8-single-a3"
model: "vllm-ascend/Qwen3-32B-W8A8"
envs:
<<: *envs
server_cmd: *server_cmd
server_cmd_extra:
- "--enforce-eager"
benchmarks:

View File

@@ -0,0 +1,38 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Qwen3-32B-QuaRot"
model: "vllm-ascend/Qwen3-32B-W8A8-QuaRot"
envs:
VLLM_WORKER_MULTIPROC_METHOD: "spawn"
SERVER_PORT: "DEFAULT_PORT"
server_cmd:
- "--enforce-eager"
- "--no-enable-prefix-caching"
- "--tensor-parallel-size"
- "2"
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "8192"
- "--trust-remote-code"
- "--distributed-executor-backend"
- "mp"
- "--gpu-memory-utilization"
- "0.9"
- "--speculative-config"
- '{"method": "eagle3", "model": "RedHatAI/Qwen3-32B-speculator.eagle3", "num_speculative_tokens": 3}'
benchmarks:
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs400
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 80
max_out_len: 1500
batch_size: 20
request_rate: 0
baseline: 1
threshold: 0.97

View File

@@ -0,0 +1,70 @@
# ==========================================
# Shared Configurations
# ==========================================
_envs: &envs
OMP_NUM_THREADS: "1"
OMP_PROC_BIND: "false"
TASK_QUEUE_ENABLE: "1"
HCCL_OP_EXPANSION_MODE: "AIV"
HCCL_BUFFSIZE: "1536"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
VLLM_ASCEND_ENABLE_FLASHCOMM1: "1"
VLLM_ASCEND_ENABLE_FUSED_MC2: "1"
VLLM_ASCEND_ENABLE_NZ: "2"
VLLM_ASCEND_BALANCE_SCHEDULING: "1"
SERVER_PORT: "DEFAULT_PORT"
_server_cmd: &server_cmd
- "--quantization"
- "ascend"
- "--no-enable-prefix-caching"
- "--mm-processor-cache-gb"
- "0"
- "--tensor-parallel-size"
- "4"
- "--data-parallel-size"
- "4"
- "--enable-expert-parallel"
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "32768"
- "--max-num-batched-tokens"
- "16384"
- "--max-num-seqs"
- "32"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.92"
_benchmarks: &benchmarks
acc:
case_type: accuracy
dataset_path: vllm-ascend/textvqa-lite
request_conf: vllm_api_stream_chat
dataset_conf: textvqa/textvqa_gen_base64
max_out_len: 2048
batch_size: 128
baseline: 83
temperature: 0
top_k: -1
top_p: 1
repetition_penalty: 1
threshold: 5
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Qwen3-VL-235B-A22B-Instruct-W8A8"
model: "Eco-Tech/Qwen3-VL-235B-A22B-Instruct-w8a8-QuaRot"
envs:
<<: *envs
server_cmd: *server_cmd
server_cmd_extra:
- "--compilation_config"
- '{"cudagraph_mode": "FULL_DECODE_ONLY", "cudagraph_capture_sizes": [1,2,4,8,16,24,32]}'
benchmarks:
<<: *benchmarks

View File

@@ -0,0 +1,67 @@
# ==========================================
# Shared Configurations
# ==========================================
_envs: &envs
OMP_NUM_THREADS: "1"
OMP_PROC_BIND: "false"
TASK_QUEUE_ENABLE: "1"
HCCL_OP_EXPANSION_MODE: "AIV"
VLLM_ASCEND_ENABLE_FLASHCOMM1: "1"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
VLLM_ASCEND_ENABLE_PREFETCH_MLP: "1"
SERVER_PORT: "DEFAULT_PORT"
_server_cmd: &server_cmd
- "--quantization"
- "ascend"
- "--no-enable-prefix-caching"
- "--mm-processor-cache-gb"
- "0"
- "--tensor-parallel-size"
- "2"
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "20000"
- "--max-num-batched-tokens"
- "8192"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.9"
_benchmarks: &benchmarks
acc:
case_type: accuracy
dataset_path: vllm-ascend/textvqa-lite
request_conf: vllm_api_stream_chat
dataset_conf: textvqa/textvqa_gen_base64
max_out_len: 2048
batch_size: 128
baseline: 80
temperature: 0
top_k: -1
top_p: 1
repetition_penalty: 1
threshold: 5
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Qwen3-VL-32B-Instruct-W8A8"
model: "Eco-Tech/Qwen3-VL-32B-Instruct-w8a8-QuaRot"
envs:
<<: *envs
server_cmd: *server_cmd
server_cmd_extra:
- "--compilation_config"
- '{"cudagraph_mode": "FULL_DECODE_ONLY", "cudagraph_capture_sizes": [1,12,16,20,24,32,48,64,68,72,76,80,128]}'
benchmarks:
<<: *benchmarks

View File

@@ -0,0 +1,79 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
_envs: &envs
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
HCCL_OP_EXPANSION_MODE: "AIV"
HCCL_BUFFSIZE: "1024"
OMP_NUM_THREADS: "1"
TASK_QUEUE_ENABLE: "1"
VLLM_ASCEND_ENABLE_FUSED_MC2: "1"
VLLM_ASCEND_ENABLE_FLASHCOMM1: "1"
SERVER_PORT: "DEFAULT_PORT"
_server_cmd: &server_cmd
- "--port"
- "$SERVER_PORT"
- "--max-model-len"
- "131072"
- "--max-num-batched-tokens"
- "16384"
- "--data-parallel-size"
- "1"
- "--tensor-parallel-size"
- "4"
- "--enable-expert-parallel"
- "--max-num-seqs"
- "128"
- "--gpu-memory-utilization"
- "0.9"
- "--compilation-config"
- '{"cudagraph_capture_sizes":[1,4,8,12,16,24,32,48,56,64,72,84,96,108,112,128,160,172,196,200,212,232,256,160,172,196,200,212,232,256,272,288,312,328,344,360,384,400,416,432,448,480,512], "cudagraph_mode":"FULL_DECODE_ONLY"}'
- "--speculative_config"
- '{"method": "qwen3_5_mtp", "num_speculative_tokens": 3, "enforce_eager": true}'
- "--trust-remote-code"
- "--allowed-local-media-path"
- "/"
- "--quantization"
- "ascend"
- "--mm-processor-cache-gb"
- "0"
- "--additional-config"
- '{"enable_cpu_binding": true, "enable_shared_expert_dp": true}'
test_cases:
- name: "Qwen3.5-122B-A10B-W8A8-A3"
model: "Eco-Tech/Qwen3.5-122B-A10B-w8a8-mtp"
envs:
<<: *envs
server_cmd: *server_cmd
benchmarks:
acc_aime2025:
case_type: accuracy
dataset_path: vllm-ascend/aime2025
request_conf: vllm_api_general_chat
dataset_conf: aime2025/aime2025_gen_0_shot_chat_prompt
max_out_len: 65536
batch_size: 32
baseline: 90
threshold: 10
thinking: true
temperature: 1.0
top_p: 0.95
top_k: 20
min_p: 0.0
presence_penalty: 1.5
repetition_penalty: 1.0
ignore_eos: false
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs8000-qwen3
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 320
max_out_len: 1500
batch_size: 80
request_rate: 0
baseline: 1
threshold: 0.97

View File

@@ -0,0 +1,57 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Qwen3.5-27B-w8a8"
model: "Eco-Tech/Qwen3.5-27B-w8a8-mtp"
envs:
VLLM_USE_MODELSCOPE: "true"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
HCCL_BUFFSIZE: "1024"
HCCL_OP_EXPANSION_MODE: "AIV"
OMP_NUM_THREADS: "1"
TASK_QUEUE_ENABLE: "1"
VLLM_ASCEND_ENABLE_PREFETCH_MLP: "1"
VLLM_ASCEND_ENABLE_DENSE_OPTIMIZE: "1"
VLLM_ASCEND_ENABLE_NZ: "1"
VLLM_ASCEND_ENABLE_FUSED_MC2: "1"
SERVER_PORT: "DEFAULT_PORT"
server_cmd:
- "--tensor-parallel-size"
- "2"
- "--port"
- "$SERVER_PORT"
- "--max-num-seqs"
- "128"
- "--quantization"
- "ascend"
- "--max-model-len"
- "262144"
- "--max-num-batched-tokens"
- "8192"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.95"
- "--additional-config"
- '{"enable_cpu_binding":true, "enable_weight_nz_layout":true}'
- "--speculative_config"
- '{"method": "qwen3_5_mtp", "num_speculative_tokens": 3,"enforce_eager": true}'
- "--compilation-config"
- '{"cudagraph_mode":"FULL_DECODE_ONLY", "cudagraph_capture_sizes":[4,8,12,16,20,24,28,32,36,40,44,48,52,56,60,64,68,72,76,80,84,88,92,96,100,104,108,112,116,120,124,128,132,136,140,144]}'
- "--mm-processor-cache-gb"
- "0"
- "--mm_processor_cache_type"
- "shm"
benchmarks:
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs8000-qwen3
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 128
max_out_len: 1500
batch_size: 32
request_rate: 0
baseline: 604
threshold: 0.97

View File

@@ -0,0 +1,73 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Qwen3.5-27B-w8a8"
model: "Eco-Tech/Qwen3.5-27B-w8a8-mtp"
envs:
VLLM_USE_MODELSCOPE: "true"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
HCCL_BUFFSIZE: "1024"
HCCL_OP_EXPANSION_MODE: "AIV"
OMP_NUM_THREADS: "1"
TASK_QUEUE_ENABLE: "1"
SERVER_PORT: "DEFAULT_PORT"
server_cmd:
- "--tensor-parallel-size"
- "2"
- "--data-parallel-size"
- "1"
- "--port"
- "$SERVER_PORT"
- "--max-num-seqs"
- "128"
- "--quantization"
- "ascend"
- "--max-model-len"
- "196608"
- "--max-num-batched-tokens"
- "16384"
- "--trust-remote-code"
- "--async-scheduling"
- "--allowed-local-media-path"
- "/"
- "--gpu-memory-utilization"
- "0.9"
- "--additional-config"
- '{"enable_cpu_binding":true}'
- "--speculative_config"
- '{"method": "qwen3_5_mtp", "num_speculative_tokens": 3,"enforce_eager": true}'
- "--compilation-config"
- '{"cudagraph_capture_sizes":[1,4,8,12,16,24,32,48,56,64,72,84,96,108,112,128,160,172,196,200,212,232,256,272,288,312,328,344,360,384,400,416,432,448,480,512], "cudagraph_mode":"FULL_DECODE_ONLY"}'
- "--mm-processor-cache-gb"
- "0"
benchmarks:
acc_aime2025:
case_type: accuracy
dataset_path: vllm-ascend/aime2025
request_conf: vllm_api_general_chat
dataset_conf: aime2025/aime2025_gen_0_shot_chat_prompt
max_out_len: 65536
batch_size: 32
baseline: 90
threshold: 10
ignore_eos: false
thinking: true
temperature: 1.0
top_p: 0.95
top_k: 20
min_p: 0.0
presence_penalty: 1.5
repetition_penalty: 1.0
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs8000-qwen3
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 140
max_out_len: 1500
batch_size: 35
request_rate: 0
baseline: 610.22
threshold: 0.97

View File

@@ -0,0 +1,76 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Qwen3.5-397B-A17B-w8a8-mtp"
model: "Eco-Tech/Qwen3.5-397B-A17B-w8a8-mtp"
envs:
VLLM_USE_MODELSCOPE: "true"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
HCCL_OP_EXPANSION_MODE: "AIV"
HCCL_BUFFSIZE: "1024"
OMP_NUM_THREADS: "1"
TASK_QUEUE_ENABLE: "1"
SERVER_PORT: "DEFAULT_PORT"
VLLM_ASCEND_ENABLE_FUSED_MC2: "1"
VLLM_ENGINE_READY_TIMEOUT_S: "3000"
VLLM_RPC_TIMEOUT: "6000"
server_cmd:
- "--tensor-parallel-size"
- "16"
- "--data-parallel-size"
- "1"
- "--enable-expert-parallel"
- "--port"
- "$SERVER_PORT"
- "--max-num-seqs"
- "128"
- "--quantization"
- "ascend"
- "--max-model-len"
- "133120"
- "--max-num-batched-tokens"
- "16384"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.9"
- "--additional-config"
- '{"enable_cpu_binding":true}'
- "--speculative_config"
- '{"method": "qwen3_5_mtp", "num_speculative_tokens": 5, "enforce_eager": true}'
- "--compilation-config"
- '{"cudagraph_capture_sizes":[1,6,12,18,24,30,36,42,48,54,72,78,84,90,96,102,108,144,192], "cudagraph_mode":"FULL_DECODE_ONLY"}'
- "--allowed-local-media-path"
- "/"
- "--mm-processor-cache-gb"
- "0"
- "--safetensors-load-strategy"
- "lazy"
benchmarks:
acc_aime2025:
case_type: accuracy
dataset_path: vllm-ascend/aime2025
request_conf: vllm_api_general_chat
dataset_conf: aime2025/aime2025_gen_0_shot_chat_prompt
max_out_len: 32768
batch_size: 32
baseline: 90
threshold: 10
temperature: 0.6
top_p: 0.95
top_k: 20
min_p: 0.0
presence_penalty: 0.0
repetition_penalty: 1.0
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in131072-bs100-qwen3
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 8
max_out_len: 1024
batch_size: 2
request_rate: 0
baseline: 56.357
threshold: 0.97

View File

@@ -0,0 +1,74 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Qwen3.5-397B-A17B-w4a8-mtp"
model: "Eco-Tech/Qwen3.5-397B-A17B-w4a8-mtp"
envs:
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
HCCL_OP_EXPANSION_MODE: "AIV"
HCCL_BUFFSIZE: "1024"
OMP_NUM_THREADS: "1"
TASK_QUEUE_ENABLE: "1"
SERVER_PORT: "DEFAULT_PORT"
server_cmd:
- "--tensor-parallel-size"
- "8"
- "--data-parallel-size"
- "1"
- "--enable-expert-parallel"
- "--port"
- "$SERVER_PORT"
- "--max-num-seqs"
- "128"
- "--quantization"
- "ascend"
- "--max-model-len"
- "133120"
- "--max-num-batched-tokens"
- "16384"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.9"
- "--no-enable-prefix-caching"
- "--additional-config"
- '{"enable_cpu_binding":true,"multistream_overlap_shared_expert":true}'
- "--speculative_config"
- '{"method": "qwen3_5_mtp", "num_speculative_tokens": 3, "enforce_eager": true}'
- "--compilation-config"
- '{"cudagraph_capture_sizes":[1,4,8,12,16,20,24,28,32,36,40,44,48,52,56,60,64,68,72,76,80,84,88,92,96,100,108,112,128,160,172,196,200,212,232,256,260,288,320,360,400], "cudagraph_mode":"FULL_DECODE_ONLY"}'
- "--allowed-local-media-path"
- "/"
- "--mm-processor-cache-gb"
- "0"
benchmarks:
acc_GPQA:
case_type: accuracy
dataset_path: vllm-ascend/gpqa
request_conf: vllm_api_general_chat
dataset_conf: gpqa/gpqa_gen_0_shot_cot_chat_prompt
max_out_len: 32768
batch_size: 32
baseline: 88.38
threshold: 5
# acc_aime2025:
# case_type: accuracy
# dataset_path: vllm-ascend/aime2025
# request_conf: vllm_api_general_chat
# dataset_conf: aime2025/aime2025_gen_0_shot_chat_prompt
# max_out_len: 72348
# batch_size: 32
# baseline: 93.33
# threshold: 5
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs8000-qwen3
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 256
max_out_len: 1500
batch_size: 64
request_rate: 0
baseline: 1040
threshold: 0.97

View File

@@ -0,0 +1,62 @@
# ==========================================
# ACTUAL TEST CASES
# ==========================================
test_cases:
- name: "Qwen3.5-397B-A17B-w8a8-mtp-longseq"
model: "Eco-Tech/Qwen3.5-397B-A17B-w8a8-mtp"
envs:
VLLM_USE_MODELSCOPE: "true"
PYTORCH_NPU_ALLOC_CONF: "expandable_segments:True"
HCCL_OP_EXPANSION_MODE: "AIV"
HCCL_BUFFSIZE: "1024"
OMP_NUM_THREADS: "1"
TASK_QUEUE_ENABLE: "1"
SERVER_PORT: "DEFAULT_PORT"
VLLM_ASCEND_ENABLE_FUSED_MC2: "1"
VLLM_ENGINE_READY_TIMEOUT_S: "3000"
VLLM_RPC_TIMEOUT: "6000"
server_cmd:
- "--tensor-parallel-size"
- "8"
- "--data-parallel-size"
- "1"
- "--prefill-context-parallel-size"
- "2"
- "--decode-context-parallel-size"
- "4"
- "--enable-expert-parallel"
- "--port"
- "$SERVER_PORT"
- "--max-num-seqs"
- "32"
- "--quantization"
- "ascend"
- "--max-model-len"
- "133120"
- "--max-num-batched-tokens"
- "16384"
- "--trust-remote-code"
- "--gpu-memory-utilization"
- "0.9"
- "--speculative_config"
- '{"method": "qwen3_5_mtp", "num_speculative_tokens": 3, "enforce_eager": true}'
- "--compilation-config"
- '{"cudagraph_capture_sizes":[1,4,8,16,32,64], "cudagraph_mode":"FULL_DECODE_ONLY"}'
- "--allowed-local-media-path"
- "/"
- "--mm-processor-cache-gb"
- "0"
- "--safetensors-load-strategy"
- "lazy"
benchmarks:
acc_GPQA:
case_type: accuracy
dataset_path: vllm-ascend/gpqa
request_conf: vllm_api_general_chat
dataset_conf: gpqa/gpqa_gen_0_shot_cot_chat_prompt
num_prompts: 50
max_out_len: 32768
batch_size: 32
baseline: 84
threshold: 8

View File

@@ -0,0 +1,312 @@
# vLLM-Ascend Single-Node E2E Test Developer Guide
This document is intended to help developers understand the architecture of the single-node E2E (End-to-End) testing framework in `vllm-ascend`, how to run test scripts, and how to add custom testing functionality by writing YAML configuration files and extending the code.
## 1. Test Architecture Overview
To achieve high readability, extensibility, and decoupling of configuration from code, the single-node E2E test adopts a **"YAML-driven + Dispatcher"** architectural structure.
It consists of the following core components:
* **Configuration Parser (`single_node_config.py`)**: Responsible for reading `models/configs/*.yaml` files and parsing them into a strongly-typed `@dataclass` (`SingleNodeConfig`) via `SingleNodeConfigLoader`, while handling regex replacement for environment variables.
* **Service Manager Framework (`test_single_node.py` and `conftest.py`)**: Based on the `service_mode` (`openai` or `epd`), it utilizes context managers to safely start/stop server processes.
* **Test Function Dispatcher (`TEST_HANDLERS` Registry)**: Specific test logic is encapsulated into independent functions and registered in the global `TEST_HANDLERS` dictionary.
* **Performance Benchmarking (`_run_benchmarks`)**: Calls `aisbench` for performance and TTFT testing based on the `benchmarks` parameters in the YAML.
### 1.1 Key Files and Responsibilities
* `tests/e2e/nightly/single_node/models/scripts/single_node_config.py`
* Defines `SingleNodeConfig` and `SingleNodeConfigLoader`
* Loads YAML from `tests/e2e/nightly/single_node/models/configs/<CONFIG_YAML_PATH>`
* Auto-assigns ports when `envs` contains `DEFAULT_PORT` / missing values
* Expands `$VAR` / `${VAR}` placeholders inside commands via `_expand_values`
* `tests/e2e/nightly/single_node/models/scripts/test_single_node.py`
* Declares `configs = SingleNodeConfigLoader.from_yaml_cases()` (loaded at import time)
* `pytest.mark.parametrize("config", configs, ids=[config.name for config in configs])` runs one test per YAML case
* Controls server lifecycle via context managers
* Dispatches `test_content` to functions registered in `TEST_HANDLERS`
* Runs `aisbench` and optional benchmark assertions
### 1.2 End-to-End Flow (High Level)
```txt
pytest starts
|
v
import tests/e2e/nightly/single_node/models/scripts/test_single_node.py
|
v
configs = SingleNodeConfigLoader.from_yaml_cases()
|
v
pytest parametrize("config", configs) # one config == one test case
|
v
test_single_node(config)
|
+-----------------------------------------------+
| Start service (depends on service_mode) |
| |
| openai: start one vLLM OpenAI-compatible |
| service process |
| epd: start (encode service + decode/PD |
| service) + start proxy process |
+-----------------------------------------------+
|
v
Run test phases (test_content)
|
v
Optional benchmarks (if benchmarks is configured)
|
v
Shutdown all started processes
Notes:
- One YAML file may contain multiple test_cases; pytest will run them one by one.
- The framework is "YAML-driven": changes are typically done by editing YAML rather than editing Python code.
```
### 1.3 Function Call Relationships (Dispatcher)
`test_content` is a list of “phases”. Each phase maps to one handler function.
```txt
For each test_case:
test_content (list of phases)
|
v
[Dispatcher]
|
+--> phase "completion" -> send completion request(s)
|
+--> phase "chat_completion" -> send chat completion request(s)
|
+--> phase "image" -> send multimodal image request(s)
|
\--> (extendable) add your own phase by registering a new handler
After phases:
if benchmarks is configured -> run aisbench
Notes:
- The dispatcher only controls "what to run"; service lifecycle is controlled by the service manager.
- Phases are intentionally small & composable so you can reuse them across YAML cases.
```
## 2. Running and Debugging Steps
### 2.1 Dependencies
Ensure you are in an NPU environment and have installed `pytest`, `pyyaml`, `openai`, and `aisbench`.
### 2.2 Local Execution
The framework uses the `CONFIG_YAML_PATH` environment variable to specify the configuration file.
```bash
# Switch to the project root directory
cd /vllm-workspace/vllm-ascend
# Run a specific yaml test
export CONFIG_YAML_PATH="Qwen3-32B.yaml"
pytest -sv tests/e2e/nightly/single_node/models/scripts/test_single_node.py
```
### 2.3 Tips for Debugging
* Only run a subset of cases: `pytest -sv ... -k <keyword>` (matches case names in the report output)
* Stop on first failure: `pytest -sv ... -x`
* Keep server logs visible: use `-s` (already included in `-sv`) and increase log verbosity via standard Python logging configuration if needed.
## 3. How to Write YAML Configuration Files
### 3.1 File Location and Selection Rules
* YAML files live under: `tests/e2e/nightly/single_node/models/configs/`
* Selected by env var: `CONFIG_YAML_PATH=<YourConfig>.yaml`
* If not set, the loader uses `SingleNodeConfigLoader.DEFAULT_CONFIG_NAME`
### 3.2 Field Descriptions
| Field Name | Type | Required | Default Value | Description |
| :--------------- | :--------- | :------- | :-------------- | :------------------------------------------------------------------ |
| `test_cases` | list | **Yes** | - | List of test case objects |
| `name` | string | **Yes** | - | Human-readable case ID shown in pytest output and logs |
| `model` | string | **Yes** | - | Model name or local path |
| `service_mode` | string | No | `openai` | Service mode: `openai` or `epd` (disaggregated) |
| `envs` | map | **Yes** | `{}` | Environment variables for the server process |
| `server_cmd` | list | Cond. | `[]` | vLLM startup arguments (Required for non-EPD) |
| `server_cmd_extra` | list | No | `[]` | Extra vLLM startup arguments appended after `server_cmd` |
| `prompts` | list | No | built-in default | Prompts for completion/chat tests |
| `api_keyword_args` | map | No | built-in default | OpenAI API keyword args (e.g., `max_tokens`, sampling params) |
| `test_content` | list | No | `["completion"]` | Test phases: `completion`, `chat_completion`, `image` etc. |
| `benchmarks` | map | No | `{}` | Configuration for `aisbench` performance verification |
| `epd_server_cmds`| list[list] | Cond. | `[]` | (EPD Only) Command arrays for starting dual Encode/Decode processes |
| `epd_proxy_args` | list | Cond. | `[]` | (EPD Only) Startup arguments for the EPD routing gateway |
**Notes / Behaviors**
* `name` is mandatory and must be a non-empty string.
* It is used directly as pytest case id (e.g., `test_single_node[DeepSeek-R1-0528-W8A8-single]`).
* It is also printed in `[single-node][START]` marker for log navigation.
* `envs` (ports): the config object recognizes these keys: `SERVER_PORT`, `ENCODE_PORT`, `PD_PORT`, `PROXY_PORT`.
* If a port key is missing or set to `DEFAULT_PORT`, it will be automatically filled with an available open port.
* `$SERVER_PORT` / `${SERVER_PORT}` placeholders in commands will be expanded using `envs`.
* `server_cmd` vs `server_cmd_extra`:
* YAML can define `server_cmd_extra` to append additional args after `server_cmd`.
* The loader merges them into a single `server_cmd` list.
* Extra fields:
* Any non-standard fields in a case are stored in `config.extra_config`.
* This is how extension configs are passed through without changing the dataclass.
### 3.3 YAML Examples
#### Single-Case (similar to DeepSeek-R1-W8A8-HBM)
```yaml
test_cases:
- name: "<your-case-name>"
model: "<model-repo-or-local-path>"
# Optional: The default values are as follows
prompts:
- "San Francisco is a"
api_keyword_args:
max_tokens: 10
envs:
SERVER_PORT: "DEFAULT_PORT"
# Add only what you need.
server_cmd:
- "--port"
- "$SERVER_PORT"
# plus your vLLM serve args...
# Optional: omit -> defaults to ["completion"]
test_content:
- "chat_completion"
# Optional: leave empty if you don't run aisbench
benchmarks:
```
#### Multi-Case + Shared Anchors
```yaml
_envs: &envs
SERVER_PORT: "DEFAULT_PORT"
# shared envs...
_server_cmd: &server_cmd
- "--port"
- "$SERVER_PORT"
# shared vLLM serve args...
_benchmarks: &benchmarks
perf:
case_type: performance
dataset_path: vllm-ascend/GSM8K-in3500-bs400
request_conf: vllm_api_stream_chat
dataset_conf: gsm8k/gsm8k_gen_0_shot_cot_str_perf
num_prompts: 400
max_out_len: 1500
batch_size: 1000
baseline: 1
threshold: 0.97
test_cases:
- name: "case-a"
model: "<model>"
envs:
<<: *envs
DYNAMIC_EPLB: "true"
# private envs...
server_cmd: *server_cmd
server_cmd_extra:
- "--enforce-eager"
benchmarks:
- name: "case-b"
model: "<model>"
envs:
<<: *envs
server_cmd: *server_cmd
benchmarks:
<<: *benchmarks
```
#### EPD / Disaggregated Case
```yaml
test_cases:
- name: "<your-epd-case>"
model: "<model>"
service_mode: "epd"
envs:
ENCODE_PORT: "DEFAULT_PORT"
PD_PORT: "DEFAULT_PORT"
PROXY_PORT: "DEFAULT_PORT"
epd_server_cmds:
- ["--port", "$ENCODE_PORT", "--model", "<encode-model>"]
- ["--port", "$PD_PORT", "--model", "<decode-model>"]
epd_proxy_args:
- "--host"
- "127.0.0.1"
- "--port"
- "$PROXY_PORT"
- "--encode-servers-urls"
- "http://localhost:$ENCODE_PORT"
- "--decode-servers-urls"
- "http://localhost:$PD_PORT"
- "--prefill-servers-urls"
- "disable"
test_content:
- "chat_completion"
```
## 4. How to Add Custom Tests (Extension)
### Step 1: Write your test logic in `test_single_node.py`
```python
async def run_video_test(config: SingleNodeConfig, server: 'RemoteOpenAIServer | DisaggEpdProxy') -> None:
client = server.get_async_client()
# Your custom logic here...
```
### Step 2: Register your function in `TEST_HANDLERS`
```python
TEST_HANDLERS = {
"completion": run_completion_test,
"video": run_video_test, # Registered!
}
```
### Step 3: Enable in YAML
```yaml
test_content:
- "completion"
- "video"
```
## 5. Checklist (Before Submitting a New YAML)
* `test_cases` exists and is a list
* Each case contains required fields for its `service_mode`
* Common required: `name`, `model`, `envs`
* `openai`: `server_cmd`
* `epd`: `epd_server_cmds`, `epd_proxy_args`
* Port envs are set to `DEFAULT_PORT` (or to explicit free ports)
* If using `benchmarks`, ensure each benchmark case includes required aisbench fields (e.g., `case_type`, `dataset_path`, `request_conf`, `dataset_conf`, `max_out_len`, `batch_size`)

View File

@@ -0,0 +1,16 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# 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.
# This file is a part of the vllm-ascend project.
#

View File

@@ -0,0 +1,188 @@
import logging
import os
from dataclasses import dataclass, field
from typing import Any
import regex as re
import yaml
from vllm.utils.network_utils import get_open_port
CONFIG_BASE_PATH = os.getenv("CONFIG_BASE_PATH") or "tests/e2e/nightly/single_node/models/configs"
logger = logging.getLogger(__name__)
# Default prompts and API args fallback
PROMPTS = [
"San Francisco is a",
]
API_KEYWORD_ARGS = {
"max_tokens": 10,
}
@dataclass
class SingleNodeConfig:
name: str
model: str
envs: dict[str, Any] = field(default_factory=dict)
special_dependencies: dict[str, Any] = field(default_factory=dict)
prompts: list[str] = field(default_factory=lambda: PROMPTS)
api_keyword_args: dict[str, Any] = field(default_factory=lambda: API_KEYWORD_ARGS)
benchmarks: dict[str, Any] = field(default_factory=dict)
server_cmd: list[str] = field(default_factory=list)
test_content: list[str] = field(default_factory=lambda: ["completion"])
service_mode: str = "openai"
epd_server_cmds: list[list[str]] = field(default_factory=list)
epd_proxy_args: list[str] = field(default_factory=list)
extra_config: dict[str, Any] = field(default_factory=dict)
def __post_init__(self) -> None:
port_keys = ["SERVER_PORT", "ENCODE_PORT", "PD_PORT", "PROXY_PORT"]
for env_key in port_keys:
if self.envs.get(env_key) in ["DEFAULT_PORT", None]:
self.envs[env_key] = str(get_open_port())
if self.prompts is None:
self.prompts = PROMPTS
if self.api_keyword_args is None:
self.api_keyword_args = API_KEYWORD_ARGS
if self.benchmarks is None:
self.benchmarks = {}
if self.special_dependencies is None:
self.special_dependencies = {}
if self.test_content is None:
self.test_content = []
self.server_cmd = self._expand_values(self.server_cmd or [], self.envs)
self.epd_server_cmds = [self._expand_values(cmd, self.envs) for cmd in self.epd_server_cmds]
self.epd_proxy_args = self._expand_values(self.epd_proxy_args or [], self.envs)
for key, value in self.extra_config.items():
setattr(self, key, value)
@staticmethod
def _expand_values(values: list[str], envs: dict[str, Any]) -> list[str]:
"""Interpolate $VAR/${VAR} placeholders with provided env values."""
pattern = re.compile(r"\$(\w+)|\$\{(\w+)\}")
def repl(m: re.Match[str]) -> str:
key = m.group(1) or m.group(2)
return str(envs.get(key, m.group(0)))
return [pattern.sub(repl, str(arg)) for arg in values]
def _get_required_port(self, key: str) -> int:
value = self.envs.get(key)
if value is None:
raise ValueError(f"Missing required port env: {key}")
return int(value)
@property
def server_port(self) -> int:
return self._get_required_port("SERVER_PORT")
@property
def encode_port(self) -> int:
return self._get_required_port("ENCODE_PORT")
@property
def pd_port(self) -> int:
return self._get_required_port("PD_PORT")
@property
def proxy_port(self) -> int:
return self._get_required_port("PROXY_PORT")
class SingleNodeConfigLoader:
"""Load SingleNodeConfig from yaml file."""
DEFAULT_CONFIG_NAME = "Kimi-K2-Thinking.yaml"
STANDARD_CASE_FIELDS = {
"name",
"model",
"envs",
"special_dependencies",
"prompts",
"api_keyword_args",
"benchmarks",
"service_mode",
"server_cmd",
"server_cmd_extra",
"test_content",
"epd_server_cmds",
"epd_proxy_args",
}
@classmethod
def from_yaml_cases(cls, yaml_path: str | None = None) -> list[SingleNodeConfig]:
config = cls._load_yaml(yaml_path)
if "test_cases" not in config:
raise KeyError("test_cases field is required in config yaml")
cases = config.get("test_cases")
if not isinstance(cases, list):
raise TypeError("test_cases must be a list")
cls._validate_para(cases)
return cls._parse_test_cases(cases)
@classmethod
def _load_yaml(cls, yaml_path: str | None) -> dict[str, Any]:
if not yaml_path:
yaml_path = os.getenv("CONFIG_YAML_PATH", cls.DEFAULT_CONFIG_NAME)
full_path = os.path.join(CONFIG_BASE_PATH, yaml_path)
logger.info("Loading config yaml: %s", full_path)
with open(full_path) as f:
return yaml.safe_load(f)
@staticmethod
def _validate_para(cases: list[dict[str, Any]]) -> None:
if not cases:
raise ValueError("test_cases is empty")
for case in cases:
mode = case.get("service_mode", "openai")
required = ["name", "model", "envs"]
if mode == "epd":
required.extend(["epd_server_cmds", "epd_proxy_args"])
else:
required.append("server_cmd")
missing = [k for k in required if k not in case]
if missing:
raise KeyError(f"Missing required config fields: {missing}")
if not isinstance(case["name"], str) or not case["name"].strip():
raise ValueError("test case field 'name' must be a non-empty string")
@classmethod
def _parse_test_cases(cls, cases: list[dict[str, Any]]) -> list[SingleNodeConfig]:
result: list[SingleNodeConfig] = []
for case in cases:
server_cmd = case.get("server_cmd", [])
server_cmd_extra = case.get("server_cmd_extra", [])
full_cmd = list(server_cmd) + list(server_cmd_extra)
extra_case_fields = {key: value for key, value in case.items() if key not in cls.STANDARD_CASE_FIELDS}
# Safe parsing mapping
result.append(
SingleNodeConfig(
name=case["name"],
model=case["model"],
envs=case.get("envs", {}),
special_dependencies=case.get("special_dependencies", {}),
server_cmd=full_cmd,
epd_server_cmds=case.get("epd_server_cmds", []),
epd_proxy_args=case.get("epd_proxy_args", []),
benchmarks=case.get("benchmarks", {}),
prompts=case.get("prompts", PROMPTS),
api_keyword_args=case.get("api_keyword_args", API_KEYWORD_ARGS),
test_content=case.get("test_content", ["completion"]),
service_mode=case.get("service_mode", "openai"),
extra_config=extra_case_fields,
)
)
return result

View File

@@ -0,0 +1,443 @@
import asyncio
import json
import logging
import os
import shlex
import subprocess
import sys
from typing import Any
import openai
import psutil
import pytest
import vllm
from tests.e2e.conftest import DisaggEpdProxy, RemoteEPDServer, RemoteOpenAIServer
from tests.e2e.nightly.single_node.models.scripts.single_node_config import (
SingleNodeConfig,
SingleNodeConfigLoader,
)
from tools.aisbench import run_aisbench_cases
logger = logging.getLogger(__name__)
configs = SingleNodeConfigLoader.from_yaml_cases()
async def run_completion_test(config: SingleNodeConfig, server: "RemoteOpenAIServer | DisaggEpdProxy") -> None:
client = server.get_async_client()
batch = await client.completions.create(
model=config.model,
prompt=config.prompts,
**config.api_keyword_args,
)
choices: list[openai.types.CompletionChoice] = batch.choices
assert choices[0].text, "empty response"
print(choices)
async def run_image_test(config: SingleNodeConfig, server: "RemoteOpenAIServer | DisaggEpdProxy") -> None:
from tools.send_mm_request import send_image_request
send_image_request(config.model, server)
async def run_chat_completion_test(config: SingleNodeConfig, server: "RemoteOpenAIServer | DisaggEpdProxy") -> None:
from tools.send_request import send_v1_chat_completions
send_v1_chat_completions(
config.prompts[0],
model=config.model,
server=server,
request_args=config.api_keyword_args,
)
def run_benchmark_comparisons(config: SingleNodeConfig, results: Any) -> None:
"""General assertion engine for aisbench outcomes mapped directly from YAML."""
comparisons = config.extra_config.get("benchmark_comparisons_args", [])
if not comparisons:
return
# Valid task keys defined in benchmarks mapping
valid_keys = [k for k, v in config.benchmarks.items() if v]
metrics_cache = {}
for comp in comparisons:
metric = comp.get("metric", "TTFT")
baseline_key = comp.get("baseline")
target_key = comp.get("target")
ratio = comp.get("ratio", 1.0)
op = comp.get("operator", "<")
if not baseline_key or not target_key:
logger.warning("Invalid comparison config: missing baseline or target. %s", comp)
continue
if metric not in metrics_cache:
if metric == "TTFT":
from tools.aisbench import get_TTFT
# map TTFT outputs directly to their corresponding benchmark test case names
metrics_cache[metric] = dict(zip(valid_keys, get_TTFT(results)))
else:
logger.warning("Unsupported metric for comparison: %s", metric)
continue
metric_dict = metrics_cache[metric]
baseline_val = metric_dict.get(baseline_key)
target_val = metric_dict.get(target_key)
if baseline_val is None or target_val is None:
logger.warning("Missing data to compare %s and %s in metrics: %s", baseline_key, target_key, metric_dict)
continue
expected_threshold = baseline_val * ratio
eval_str = f"metric {metric}: {target_key}({target_val}) {op} {baseline_key}({baseline_val}) * {ratio}"
if op == "<":
assert target_val < expected_threshold, f"Assertion Failed: {eval_str} [threshold: {expected_threshold}]"
elif op == ">":
assert target_val > expected_threshold, f"Assertion Failed: {eval_str} [threshold: {expected_threshold}]"
elif op == "<=":
assert target_val <= expected_threshold, f"Assertion Failed: {eval_str} [threshold: {expected_threshold}]"
elif op == ">=":
assert target_val >= expected_threshold, f"Assertion Failed: {eval_str} [threshold: {expected_threshold}]"
else:
logger.warning("Unsupported comparison operator: %s", op)
continue
print(f"✅ Comparison passed: {eval_str} [threshold: {expected_threshold}]")
async def run_check_rank0_process_count(
config: SingleNodeConfig, server: "RemoteOpenAIServer | DisaggEpdProxy"
) -> None:
proc = await asyncio.create_subprocess_exec(
"npu-smi",
"info",
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
)
stdout_bytes, stderr_bytes = await proc.communicate()
if proc.returncode == 0:
logger.info("npu-smi info:\n%s", stdout_bytes.decode(errors="ignore"))
else:
logger.warning("npu-smi info failed: %s", stderr_bytes.decode(errors="ignore"))
vllm_serve_procs = [
p
for p in psutil.process_iter(attrs=["pid", "cmdline"], ad_value=None)
if p.info["cmdline"]
and any("vllm" in arg for arg in p.info["cmdline"])
and any("serve" in arg for arg in p.info["cmdline"])
]
count = len(vllm_serve_procs)
assert count == 1, (
f"rank0 process count check failed: expected exactly 1 vllm serve process on rank0, found {count}"
)
# Extend this dictionary to add new test capabilities
TEST_HANDLERS = {
"completion": run_completion_test,
"image": run_image_test,
"chat_completion": run_chat_completion_test,
"check_rank0_process_count": run_check_rank0_process_count,
}
async def _dispatch_tests(config: SingleNodeConfig, server: "RemoteOpenAIServer | DisaggEpdProxy") -> None:
"""Dispatches requested tests defined in yaml."""
for test_name in config.test_content:
if test_name == "benchmark_comparisons":
continue
handler = TEST_HANDLERS.get(test_name)
if handler:
await handler(config, server)
else:
logger.warning("No handler registered for test content type: %s", test_name)
def _extract_server_cmd_value(server_cmd: list[str], flag: str) -> str | None:
"""Return the value following `flag` in a server_cmd list, or None."""
try:
idx = server_cmd.index(flag)
return server_cmd[idx + 1]
except (ValueError, IndexError):
return None
def _extract_hardware(runner: str) -> str:
"""Derive hardware label (e.g. 'A2', 'A3') from runner name."""
runner_lower = runner.lower()
for label in ("a3", "a2"):
if label in runner_lower:
return label.upper()
return runner
_PORT_ENV_KEYS = {"SERVER_PORT", "ENCODE_PORT", "PD_PORT", "PROXY_PORT"}
_FEATURE_ENVS: dict[str, str] = {
"VLLM_ASCEND_ENABLE_FLASHCOMM": "flashcomm",
"VLLM_ASCEND_ENABLE_FLASHCOMM1": "flashcomm1",
"VLLM_ASCEND_ENABLE_TOPK_OPTIMIZE": "topk_optimize",
"VLLM_ASCEND_ENABLE_MATMUL_ALLREDUCE": "matmul_allreduce",
"VLLM_ASCEND_ENABLE_MLAPO": "mlapo",
"VLLM_ASCEND_ENABLE_FUSED_MC2": "fused_mc2",
}
_PERF_METRIC_RENAME: dict[str, str] = {
"Benchmark Duration": "Benchmark_Duration(BD)",
"Prefill Token Throughput": "Prefill_Token_Throughput(PTT)",
"Input Token Throughput": "Input_Token_Throughput(ITT)",
"Output Token Throughput": "Output_Token_Throughput(OTT)",
"Total Token Throughput": "Total_Token_Throughput(TTT)",
}
def _extract_dtype(config: SingleNodeConfig) -> str:
"""Determine weight dtype: w8a8 if model name contains 'w8a8' and --quantization ascend is set, else bf16."""
has_w8a8 = "w8a8" in config.model.lower()
has_quant_ascend = _extract_server_cmd_value(config.server_cmd, "--quantization") == "ascend"
return "w8a8" if (has_w8a8 and has_quant_ascend) else "bf16"
def _parse_json_flag(cmd_list: list[str], flag: str) -> dict[str, Any]:
"""Extract and JSON-parse the value following `flag` in a command list."""
val = _extract_server_cmd_value(cmd_list, flag)
if not val:
return {}
try:
return json.loads(val)
except (json.JSONDecodeError, ValueError):
return {}
def _extract_features(server_cmd: list[str] | str, envs: dict[str, Any]) -> list[str]:
"""Extract enabled feature names from server_cmd and environment variables."""
if isinstance(server_cmd, str):
try:
cmd_list = shlex.split(server_cmd)
except ValueError:
cmd_list = server_cmd.split()
else:
cmd_list = list(server_cmd)
features: list[str] = []
# Features from --additional-config JSON
additional = _parse_json_flag(cmd_list, "--additional-config")
if additional.get("enable_weight_nz_layout"):
features.append("weight_nz_layout")
wp = additional.get("weight_prefetch_config") or {}
if isinstance(wp, dict) and wp.get("enabled"):
features.append("weight_prefetch")
tc = additional.get("torchair_graph_config") or {}
if isinstance(tc, dict) and tc.get("enabled"):
features.append("torchair_graph")
asc = additional.get("ascend_scheduler_config") or {}
if isinstance(asc, dict) and asc.get("enabled"):
features.append("ascend_scheduler")
# Features from --compilation-config JSON
compilation = _parse_json_flag(cmd_list, "--compilation-config")
if compilation.get("cudagraph_mode"):
features.append("aclgraph")
# Features from --speculative-config JSON
speculative = _parse_json_flag(cmd_list, "--speculative-config")
if speculative:
features.append(speculative.get("method", "speculative"))
# Features from direct flags
if "--enable-expert-parallel" in cmd_list:
features.append("expert_parallel")
# Features from environment variables
for env_key, feature_name in _FEATURE_ENVS.items():
val = str(envs.get(env_key, "0"))
if val not in ("0", "", "false", "False"):
features.append(feature_name)
if int(envs.get("VLLM_ASCEND_FLASHCOMM2_PARALLEL_SIZE", 0)) > 0:
features.append("flashcomm2")
return features
def _build_serve_cmd(config: SingleNodeConfig) -> dict[str, str]:
"""Build serve_cmd dict with mix key for single-node deployments."""
args = " ".join(config.server_cmd)
return {"mix": f"vllm serve {config.model} {args}".strip()}
def _filter_environment(envs: dict[str, Any]) -> dict[str, Any]:
"""Return env vars with internal port keys removed."""
return {k: v for k, v in envs.items() if k not in _PORT_ENV_KEYS}
def _task_passed(case_config: dict[str, Any], result: Any) -> bool:
"""Return True if a single benchmark result meets its baseline/threshold."""
if result == "":
return False
case_type = case_config.get("case_type")
baseline = case_config.get("baseline")
threshold = case_config.get("threshold")
if baseline is None or threshold is None:
return True
if case_type == "accuracy" and isinstance(result, (int, float)):
return abs(float(result) - float(baseline)) <= float(threshold)
if case_type == "performance" and isinstance(result, list) and len(result) == 2:
_, result_json = result
throughput_str = result_json.get("Output Token Throughput", {}).get("total", "")
try:
throughput_val = float(throughput_str.replace("token/s", "").strip())
return throughput_val >= float(threshold) * float(baseline)
except (ValueError, AttributeError):
return False
return True
def _build_task_entry(case_key: str, case_config: dict[str, Any], result: Any) -> dict[str, Any]:
"""Build a single task dict in the required format."""
dataset_path = case_config.get("dataset_path", "")
dataset_conf = case_config.get("dataset_conf", "")
if dataset_path:
task_name = dataset_path.split("/", 1)[-1]
elif dataset_conf:
task_name = dataset_conf.split("/")[0]
else:
task_name = case_key
case_type = case_config.get("case_type", "unknown")
metrics: dict[str, float] = {}
if result == "":
# benchmark run failed — no metrics available
pass
elif case_type == "accuracy" and isinstance(result, (int, float)):
metrics["accuracy"] = round(float(result), 4)
elif case_type == "performance" and isinstance(result, list) and len(result) == 2:
_, result_json = result
for metric_name, metric_data in result_json.items():
if not isinstance(metric_data, dict):
continue
total_str = metric_data.get("total", "")
try:
value = float(total_str.replace("token/s", "").replace("ms", "").replace("s", "").strip())
metrics[_PERF_METRIC_RENAME.get(metric_name, metric_name)] = round(value, 4)
except (ValueError, AttributeError):
pass
test_input_keys = ("num_prompts", "max_out_len", "batch_size", "request_rate")
test_input = {k: case_config[k] for k in test_input_keys if k in case_config}
target: dict[str, Any] = {}
if case_config.get("baseline") is not None:
target["baseline"] = case_config["baseline"]
if case_config.get("threshold") is not None:
target["threshold"] = case_config["threshold"]
entry: dict[str, Any] = {"name": task_name, "metrics": metrics, "test_input": test_input}
if target:
entry["target"] = target
entry["pass_fail"] = "pass" if _task_passed(case_config, result) else "fail"
return entry
def _all_passed(case_configs: list[dict[str, Any]], results: list[Any]) -> bool:
"""Return True only when every benchmark result meets its baseline/threshold."""
return all(_task_passed(cfg, res) for cfg, res in zip(case_configs, results))
def _save_benchmark_results_json(config: SingleNodeConfig, benchmark_keys: list[str], results: list[Any]) -> None:
"""Serialize acc & perf benchmark results to a JSON file under benchmark_results/."""
runner = os.environ.get("VLLM_CI_RUNNER", "")
case_configs = [config.benchmarks[k] for k in benchmark_keys]
tasks = [
_build_task_entry(key, case_cfg, result) for key, case_cfg, result in zip(benchmark_keys, case_configs, results)
]
passed = _all_passed(case_configs, results)
output: dict[str, Any] = {
"model_name": config.model,
"hardware": _extract_hardware(runner),
"dtype": _extract_dtype(config),
"feature": _extract_features(config.server_cmd, config.envs),
"vllm_version": vllm.__version__,
"vllm_ascend_version": os.environ.get("VLLM_ASCEND_VERSION", ""),
"tasks": tasks,
"serve_cmd": _build_serve_cmd(config),
"environment": _filter_environment(config.envs),
"pass_fail": "pass" if passed else "fail",
}
os.makedirs("benchmark_results", exist_ok=True)
job_name = os.environ.get("BENCHMARK_JOB_NAME") or config.name
safe_name = job_name.replace("/", "_").replace(" ", "_")
output_path = os.path.join("benchmark_results", f"{safe_name}.json")
with open(output_path, "w", encoding="utf-8") as f:
json.dump(output, f, indent=2, ensure_ascii=False)
logger.info("Benchmark results saved to %s", output_path)
print(f"Benchmark results saved to {output_path}")
def _run_benchmarks(config: SingleNodeConfig, port: int) -> None:
"""Run Aisbench benchmarks and process benchmark-dependent custom assertions."""
benchmark_keys = [k for k, v in config.benchmarks.items() if v]
aisbench_cases = [config.benchmarks[k] for k in benchmark_keys]
if not aisbench_cases:
return
result = run_aisbench_cases(
model=config.model,
port=port,
aisbench_cases=aisbench_cases,
)
_save_benchmark_results_json(config, benchmark_keys, result)
if "benchmark_comparisons" in config.test_content:
run_benchmark_comparisons(config, result)
@pytest.mark.asyncio
@pytest.mark.parametrize("config", configs, ids=[config.name for config in configs])
async def test_single_node(config: SingleNodeConfig) -> None:
# TODO: remove this part after the transformers version upgraded
if config.special_dependencies:
for k, v in config.special_dependencies.items():
command = [
sys.executable,
"-m",
"pip",
"install",
f"{k}=={v}",
]
subprocess.call(command)
if config.service_mode == "epd":
with (
RemoteEPDServer(vllm_serve_args=config.epd_server_cmds, env_dict=config.envs) as _,
DisaggEpdProxy(proxy_args=config.epd_proxy_args, env_dict=config.envs) as proxy,
):
await _dispatch_tests(config, proxy)
_run_benchmarks(config, proxy.port)
return
# Standard OpenAI service mode
with RemoteOpenAIServer(
model=config.model,
vllm_serve_args=config.server_cmd,
server_port=config.server_port,
env_dict=config.envs,
auto_port=False,
) as server:
await _dispatch_tests(config, server)
_run_benchmarks(config, config.server_port)

View File

@@ -0,0 +1,61 @@
import time
from datetime import datetime
import pytest
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
from vllm_ascend.utils import enable_custom_op
init_device_properties_triton()
enable_custom_op()
DURATION_THRESHOLD = 120
SLOW_COUNT_LIMIT = 5
_per_file_slow_cases = {}
_current_file = None
def pytest_runtest_setup(item):
item.start_time = time.time()
def pytest_runtest_teardown(item, nextitem):
global _current_file
file_path = item.fspath
if not hasattr(item, "start_time"):
return
duration = time.time() - item.start_time
if file_path not in _per_file_slow_cases:
_per_file_slow_cases[file_path] = 0
if duration > DURATION_THRESHOLD:
_per_file_slow_cases[file_path] += 1
cnt = _per_file_slow_cases[file_path]
print(f" Detected that the test case took too long, ({cnt}/{SLOW_COUNT_LIMIT}):{duration:.2f}s")
if cnt >= SLOW_COUNT_LIMIT:
print(f"\n The number of timeout test cases {file_path} ≥{SLOW_COUNT_LIMIT}\n")
_current_file = file_path
def pytest_runtest_call(item):
if _current_file == item.fspath:
print(f"CASE SKIP:{item.nodeid}")
pytest.skip("The use case takes too long.")
@pytest.hookimpl(tryfirst=True, hookwrapper=True)
def pytest_runtest_makereport(item, call):
"""Hook to add timestamp to test reports"""
start_time = datetime.now().strftime("[%H:%M:%S]")
outcome = yield
report = outcome.get_result()
if report.when == "call":
print(f"{start_time}")

View File

@@ -0,0 +1,135 @@
import gc
import os
import npugraph_ex as nge
import numpy as np
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
import torch_npu
from vllm_ascend.utils import enable_custom_op
config = nge.CompilerConfig()
npu_backend = nge.get_npu_backend(compiler_config=config)
torch_npu.npu.config.allow_internal_format = True
enable_custom_op()
global_rank_id = 0
def golden_op_matmul_allreduce_add_rmsnorm(a, b, residual, gamma, epsilon):
c_ret = torch.nn.functional.linear(a, b)
dist.all_reduce(c_ret)
rmsnorm_ret, _, add_ret = torch_npu.npu_add_rms_norm(c_ret, residual, gamma, epsilon)
return rmsnorm_ret, add_ret
def worker(rank, ep_world_size, batch_size, m, k, n):
global global_rank_id
global_rank_id = rank
rank = rank
os.environ["MASTER_ADDR"] = "127.0.0.1"
os.environ["MASTER_PORT"] = "29500"
dist.init_process_group(backend="hccl", rank=rank, world_size=ep_world_size)
ep_ranks_list = list(np.arange(0, ep_world_size))
ep_group = dist.new_group(backend="hccl", ranks=ep_ranks_list)
torch_npu.npu.set_device(rank)
ep_hcomm_info = ep_group._get_backend(torch.device("npu")).get_hccl_comm_name(rank)
torch_npu.npu.synchronize(rank)
class Module(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
def forward(self, x1, x2, residual, gamma, ep_hcomm_info, epsilon, is_trans_b, is_allgather_add_out):
out1, add_out1 = torch.ops._C_ascend.matmul_allreduce_add_rmsnorm(
x1,
x2,
residual,
gamma,
ep_hcomm_info,
ep_world_size,
global_rank_id,
epsilon,
is_trans_b,
is_allgather_add_out,
)
return out1, add_out1
DTYPE = torch.bfloat16
USE_ONES = False
torch.manual_seed(42)
if USE_ONES:
x1 = torch.ones([m, k], dtype=DTYPE).npu(rank)
x2 = torch.ones([n, k], dtype=DTYPE).npu(rank)
else:
x1 = torch.normal(0, 0.1, [m, k], dtype=DTYPE).npu(rank)
x2 = torch.normal(0, 0.1, [n, k], dtype=DTYPE).npu(rank)
if USE_ONES:
residual = torch.full([m, n], 2048, dtype=DTYPE).npu(rank)
else:
residual = torch.full([m, n], 0, dtype=DTYPE).npu(rank)
gamma = torch.full([n], 1, dtype=DTYPE).npu(rank)
epsilon = 1e-5
is_trans_b = True
is_allgather_add_out = True
warnup_cnt = 5
repeat_cnt = 10
def run_golden_case(loop_cnt):
for _ in range(loop_cnt):
golden_out, golden_add_out = golden_op_matmul_allreduce_add_rmsnorm(x1, x2, residual, gamma, epsilon)
torch_npu.npu.synchronize(rank)
return golden_out, golden_add_out
run_golden_case(warnup_cnt)
golden_out, golden_add_out = run_golden_case(repeat_cnt)
golden_out = golden_out.detach().cpu()
golden_add_out = golden_add_out.detach().cpu()
mod = Module().npu()
opt_model = torch.compile(mod, backend=npu_backend)
def run_custom_case(loop_cnt):
for _ in range(loop_cnt):
out, add_out = opt_model(x1, x2, residual, gamma, ep_hcomm_info, epsilon, is_trans_b, is_allgather_add_out)
torch_npu.npu.synchronize(rank)
return out, add_out
# warn up
run_custom_case(warnup_cnt)
out, add_out = run_custom_case(repeat_cnt)
out = out.detach().cpu()
add_out = add_out.detach().cpu()
dist.destroy_process_group()
torch.testing.assert_close(golden_out, out, atol=0.1, rtol=0.005)
torch.testing.assert_close(golden_add_out, add_out, atol=0.1, rtol=0.005)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@torch.inference_mode()
def test_matmul_allreduce_add_rmsnorm_kernel():
ep_world_size = 4
batch_size = 1
m = 10000
k = 1024
n = 5120
args = (ep_world_size, batch_size, m, k, n)
mp.spawn(worker, args=args, nprocs=ep_world_size, join=True)

View File

@@ -0,0 +1,231 @@
import random
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
import torch_npu
from torch.distributed.distributed_c10d import _get_default_group
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
class TestDispatchFFNCombine:
def __init__(self, rank, world_size, port):
self.rank = rank
self.world_size = world_size
self.master_ip = "127.0.0.1"
self.port = port
def get_hcomm(self, comm_group):
hcomm_info = None
if torch.__version__ > "2.0.1":
hcomm_info = comm_group._get_backend(torch.device("npu")).get_hccl_comm_name(self.rank)
else:
hcomm_info = comm_group.get_hccl_comm_name(self.rank)
return hcomm_info
def setup_ep_tp(
self,
rank,
tp_size,
ep_size,
backend_type,
ep_ranks_list=None,
tp_ranks_list=None,
):
for i in range(tp_size):
if ep_ranks_list:
ep_ranks = ep_ranks_list[i]
else:
ep_ranks = [x + ep_size * i for x in range(ep_size)]
ep_group = dist.new_group(backend=backend_type, ranks=ep_ranks)
if rank in ep_ranks:
ep_group_tmp = ep_group
for i in range(ep_size):
if tp_ranks_list:
tp_ranks = tp_ranks_list[i]
else:
tp_ranks = [x * ep_size + i for x in range(tp_size)]
tp_group = dist.new_group(backend=backend_type, ranks=tp_ranks)
if rank in tp_ranks:
tp_group_tmp = tp_group
return ep_group_tmp, tp_group_tmp
def generate_hcom(self):
torch_npu.npu.set_device(self.rank)
dist.init_process_group(
backend="hccl",
rank=self.rank,
world_size=self.world_size,
init_method=f"tcp://127.0.0.1:{self.port}",
)
ep_size = 0
tp_size = self.world_size
hcomm_info_dist = {
"default_pg_info": None,
"ep_hcomm_info": None,
"group_ep": None,
"tp_hcomm_info": None,
"group_tp": None,
}
if ep_size and tp_size:
group_ep, group_tp = self.setup_ep_tp(self.rank, tp_size, ep_size, "hccl", None, None)
hcomm_info_dist["ep_hcomm_info"] = self.get_hcomm(group_ep)
hcomm_info_dist["tp_hcomm_info"] = self.get_hcomm(group_tp)
hcomm_info_dist["group_ep"] = group_ep
hcomm_info_dist["group_tp"] = group_tp
else:
if dist.is_available():
default_pg = _get_default_group()
hcomm_info_dist["default_pg_info"] = self.get_hcomm(default_pg)
hcomm_info = hcomm_info_dist["default_pg_info"]
self.hcomm_info = hcomm_info
def run_tensor_list(self) -> bool:
torch_npu.npu.set_device(self.rank)
m = 64
k = 1024
n = 1024
topk = 8
e = 8
k2 = n // 2
n2 = k
torch_npu.npu.config.allow_internal_format = True
x = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
weight1 = self.generate_random_tensor((e, k, n), dtype=torch.int8).npu()
weight1 = torch_npu.npu_format_cast(weight1, 29)
weight2 = self.generate_random_tensor((e, k2, n2), dtype=torch.int8).npu()
weight2 = torch_npu.npu_format_cast(weight2, 29)
expert_idx = torch.randint(0, self.world_size * e, (m, topk), dtype=torch.int32).npu()
scale1 = torch.randint(0, 1, (e, n), dtype=torch.int64).npu()
scale2 = torch.randint(0, 1, (e, n2), dtype=torch.int64).npu()
probs = torch.randn(size=(m, topk), dtype=torch.float32).npu()
xactmask = torch.randint(0, 2, (m,), torch.bool).npu()
weight1_nz_npu = []
weight2_nz_npu = []
scale1_npu = []
scale2_npu = []
for i in range(e):
weight1_nz_npu.append(torch_npu.npu_format_cast(weight1[i].npu(), 29))
scale1_npu.append(scale1[i].npu())
weight2_nz_npu.append(torch_npu.npu_format_cast(weight2[i].npu(), 29))
scale2_npu.append(scale2[i].npu())
out = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
expert_token_nums = self.generate_random_tensor((1, e), dtype=torch.int32).npu()
torch.ops._C_ascend.dispatch_ffn_combine(
x=x,
weight1=weight1_nz_npu,
weight2=weight2_nz_npu,
expert_idx=expert_idx,
bias1=torch.tensor([]),
bias2=torch.tensor([]),
scale1=scale1_npu,
scale2=scale2_npu,
probs=probs,
group=self.hcomm_info,
max_output_size=512,
x_active_mask=xactmask,
out=out,
expert_token_nums=expert_token_nums,
)
return True
def run_normal(self) -> bool:
torch_npu.npu.set_device(self.rank)
m = 64
k = 1024
n = 1024
topk = 8
e = 8
k2 = n // 2
n2 = k
torch_npu.npu.config.allow_internal_format = True
x = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
weight1 = self.generate_random_tensor((e, k, n), dtype=torch.int8).npu()
weight1 = torch_npu.npu_format_cast(weight1, 29)
weight2 = self.generate_random_tensor((e, k2, n2), dtype=torch.int8).npu()
weight2 = torch_npu.npu_format_cast(weight2, 29)
expert_idx = torch.randint(0, self.world_size * e, (m, topk), dtype=torch.int32).npu()
scale1 = torch.randint(0, 1, (e, n), dtype=torch.int64).npu()
scale2 = torch.randint(0, 1, (e, n2), dtype=torch.int64).npu()
probs = torch.randn(size=(m, topk), dtype=torch.float32).npu()
weight1_nz_npu = []
weight2_nz_npu = []
scale1_npu = []
scale2_npu = []
weight1_nz_npu.append(torch_npu.npu_format_cast(weight1.npu(), 29))
scale1_npu.append(scale1.npu())
weight2_nz_npu.append(torch_npu.npu_format_cast(weight2.npu(), 29))
scale2_npu.append(scale2.npu())
out = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
expert_token_nums = self.generate_random_tensor((1, e), dtype=torch.int32).npu()
torch.ops._C_ascend.dispatch_ffn_combine(
x=x,
weight1=weight1_nz_npu,
weight2=weight2_nz_npu,
expert_idx=expert_idx,
bias1=torch.tensor([]),
bias2=torch.tensor([]),
scale1=scale1_npu,
scale2=scale2_npu,
probs=probs,
group=self.hcomm_info,
max_output_size=512,
out=out,
expert_token_nums=expert_token_nums,
)
return True
def generate_random_tensor(self, size, dtype):
if dtype in [torch.float16, torch.bfloat16, torch.float32]:
return torch.randn(size=size, dtype=dtype)
elif dtype is torch.int8:
return torch.randint(-16, 16, size=size, dtype=dtype)
elif dtype is torch.int32:
return torch.randint(-1024, 1024, size=size, dtype=dtype)
else:
raise ValueError(f"Invalid dtype: {dtype}")
def worker(rank: int, world_size: int, port: int, q: mp.SimpleQueue):
op = TestDispatchFFNCombine(rank, world_size, port)
op.generate_hcom()
out1 = op.run_tensor_list()
q.put(out1)
out2 = op.run_normal()
q.put(out2)
@torch.inference_mode()
def test_dispatch_ffn_combine_kernel():
world_size = 2
mp.set_start_method("fork", force=True)
q = mp.SimpleQueue()
p_list = []
port = 29501 + random.randint(0, 10000)
for rank in range(world_size):
p = mp.Process(target=worker, args=(rank, world_size, port, q))
p.start()
p_list.append(p)
results = [q.get() for _ in range(world_size)]
for p in p_list:
p.join()
assert all(results)

View File

@@ -0,0 +1,229 @@
import random
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
import torch_npu
from torch.distributed.distributed_c10d import _get_default_group
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
class TestDispatchFFNCombine:
def __init__(self, rank, world_size, port):
self.rank = rank
self.world_size = world_size
self.master_ip = "127.0.0.1"
self.port = port
def get_hcomm(self, comm_group):
hcomm_info = None
if torch.__version__ > "2.0.1":
hcomm_info = comm_group._get_backend(torch.device("npu")).get_hccl_comm_name(self.rank)
else:
hcomm_info = comm_group.get_hccl_comm_name(self.rank)
return hcomm_info
def setup_ep_tp(
self,
rank,
tp_size,
ep_size,
backend_type,
ep_ranks_list=None,
tp_ranks_list=None,
):
for i in range(tp_size):
if ep_ranks_list:
ep_ranks = ep_ranks_list[i]
else:
ep_ranks = [x + ep_size * i for x in range(ep_size)]
ep_group = dist.new_group(backend=backend_type, ranks=ep_ranks)
if rank in ep_ranks:
ep_group_tmp = ep_group
for i in range(ep_size):
if tp_ranks_list:
tp_ranks = tp_ranks_list[i]
else:
tp_ranks = [x * ep_size + i for x in range(tp_size)]
tp_group = dist.new_group(backend=backend_type, ranks=tp_ranks)
if rank in tp_ranks:
tp_group_tmp = tp_group
return ep_group_tmp, tp_group_tmp
def generate_hcom(self):
torch_npu.npu.set_device(self.rank)
dist.init_process_group(
backend="hccl",
rank=self.rank,
world_size=self.world_size,
init_method=f"tcp://127.0.0.1:{self.port}",
)
ep_size = 0
tp_size = self.world_size
hcomm_info_dist = {
"default_pg_info": None,
"ep_hcomm_info": None,
"group_ep": None,
"tp_hcomm_info": None,
"group_tp": None,
}
if ep_size and tp_size:
group_ep, group_tp = self.setup_ep_tp(self.rank, tp_size, ep_size, "hccl", None, None)
hcomm_info_dist["ep_hcomm_info"] = self.get_hcomm(group_ep)
hcomm_info_dist["tp_hcomm_info"] = self.get_hcomm(group_tp)
hcomm_info_dist["group_ep"] = group_ep
hcomm_info_dist["group_tp"] = group_tp
else:
if dist.is_available():
default_pg = _get_default_group()
hcomm_info_dist["default_pg_info"] = self.get_hcomm(default_pg)
hcomm_info = hcomm_info_dist["default_pg_info"]
self.hcomm_info = hcomm_info
def run_tensor_list(self) -> bool:
torch_npu.npu.set_device(self.rank)
m = 64
k = 1024
n = 1024
topk = 8
e = 8
k2 = n // 2
n2 = k
torch_npu.npu.config.allow_internal_format = True
x = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
weight1 = self.generate_random_tensor((e, k, n), dtype=torch.bfloat16).npu()
weight1 = torch_npu.npu_format_cast(weight1, 29)
weight2 = self.generate_random_tensor((e, k2, n2), dtype=torch.bfloat16).npu()
weight2 = torch_npu.npu_format_cast(weight2, 29)
expert_idx = torch.randint(0, self.world_size * e, (m, topk), dtype=torch.int32).npu()
scale1 = torch.randint(0, 1, (e, n), dtype=torch.int64).npu()
scale2 = torch.randint(0, 1, (e, n2), dtype=torch.int64).npu()
probs = torch.randn(size=(m, topk), dtype=torch.float32).npu()
weight1_nz_npu = []
weight2_nz_npu = []
scale1_npu = []
scale2_npu = []
for i in range(e):
weight1_nz_npu.append(torch_npu.npu_format_cast(weight1[i].npu(), 29))
scale1_npu.append(scale1[i].npu())
weight2_nz_npu.append(torch_npu.npu_format_cast(weight2[i].npu(), 29))
scale2_npu.append(scale2[i].npu())
out = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
expert_token_nums = self.generate_random_tensor((1, e), dtype=torch.int32).npu()
torch.ops._C_ascend.dispatch_ffn_combine(
x=x,
weight1=weight1_nz_npu,
weight2=weight2_nz_npu,
expert_idx=expert_idx,
scale1=scale1_npu,
scale2=scale2_npu,
bias1=torch.tensor([]),
bias2=torch.tensor([]),
probs=probs,
group=self.hcomm_info,
max_output_size=512,
out=out,
expert_token_nums=expert_token_nums,
)
return True
def run_normal(self) -> bool:
torch_npu.npu.set_device(self.rank)
m = 64
k = 1024
n = 1024
topk = 8
e = 8
k2 = n // 2
n2 = k
torch_npu.npu.config.allow_internal_format = True
x = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
weight1 = self.generate_random_tensor((e, k, n), dtype=torch.bfloat16).npu()
weight1 = torch_npu.npu_format_cast(weight1, 29)
weight2 = self.generate_random_tensor((e, k2, n2), dtype=torch.bfloat16).npu()
weight2 = torch_npu.npu_format_cast(weight2, 29)
expert_idx = torch.randint(0, self.world_size * e, (m, topk), dtype=torch.int32).npu()
scale1 = torch.randint(0, 1, (e, n), dtype=torch.int64).npu()
scale2 = torch.randint(0, 1, (e, n2), dtype=torch.int64).npu()
probs = torch.randn(size=(m, topk), dtype=torch.float32).npu()
weight1_nz_npu = []
weight2_nz_npu = []
scale1_npu = []
scale2_npu = []
weight1_nz_npu.append(torch_npu.npu_format_cast(weight1.npu(), 29))
scale1_npu.append(scale1.npu())
weight2_nz_npu.append(torch_npu.npu_format_cast(weight2.npu(), 29))
scale2_npu.append(scale2.npu())
out = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
expert_token_nums = self.generate_random_tensor((1, e), dtype=torch.int32).npu()
torch.ops._C_ascend.dispatch_ffn_combine(
x=x,
weight1=weight1_nz_npu,
weight2=weight2_nz_npu,
expert_idx=expert_idx,
scale1=scale1_npu,
scale2=scale2_npu,
bias1=torch.tensor([]),
bias2=torch.tensor([]),
probs=probs,
group=self.hcomm_info,
max_output_size=512,
out=out,
expert_token_nums=expert_token_nums,
)
return True
def generate_random_tensor(self, size, dtype):
if dtype in [torch.float16, torch.bfloat16, torch.float32]:
return torch.randn(size=size, dtype=dtype)
elif dtype is torch.int8:
return torch.randint(-16, 16, size=size, dtype=dtype)
elif dtype is torch.int32:
return torch.randint(-1024, 1024, size=size, dtype=dtype)
else:
raise ValueError(f"Invalid dtype: {dtype}")
def worker(rank: int, world_size: int, port: int, q: mp.SimpleQueue):
op = TestDispatchFFNCombine(rank, world_size, port)
op.generate_hcom()
out1 = op.run_tensor_list()
q.put(out1)
out2 = op.run_normal()
q.put(out2)
@torch.inference_mode()
def test_dispatch_ffn_combine_kernel():
world_size = 2
mp.set_start_method("fork", force=True)
q = mp.SimpleQueue()
p_list = []
port = 29501 + random.randint(0, 10000)
for rank in range(world_size):
p = mp.Process(target=worker, args=(rank, world_size, port, q))
p.start()
p_list.append(p)
results = [q.get() for _ in range(world_size)]
for p in p_list:
p.join()
assert all(results)

View File

@@ -0,0 +1,335 @@
import random
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
import torch_npu
from torch.distributed.distributed_c10d import _get_default_group
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
DEVICE_OFFSET = 0
def int32_to_8x_int4_float(tensor_int32):
"""
Unpack each int32 value in the tensor into 8 signed int4 values and convert them to float32.
Logic:
1. Extract the lower 4 bits -> 0th int4
2. Shift right by 4 bits, extract the lower 4 bits -> 1st int4
...
3. Shift right by 28 bits, extract the lower 4 bits -> 7th int4
For signed int4 (Two's complement):
Binary 0000 ~ 0111 (0~7) -> float 0.0 ~ 7.0
Binary 1000 ~ 1111 (8~15) -> float -8.0 ~ -1.0
"""
# Ensure the dtype is int32 (for robustness, even if the input is already int32)
if tensor_int32.dtype != torch.int32:
tensor_int32 = tensor_int32.to(torch.int32)
original_shape = tensor_int32.shape
# 1. Create shift amounts [0, 4, 8, 12, 16, 20, 24, 28]
# Reshape to (1, 1, ..., 8) for broadcasting
shifts = torch.arange(0, 32, 4, device=tensor_int32.device).view(*([1] * len(original_shape)), -1)
# 2. Expand dimension and shift right
# unsqueeze(-1) adds a dimension -> [..., 1]
# After shifting -> [..., 8]
shifted = tensor_int32.unsqueeze(-1) >> shifts
# 3. Apply mask to keep only the lower 4 bits (0xF = 1111 binary)
# The value range here is 0 ~ 15 (unsigned view)
unpacked_unsigned = shifted & 0xF
# 4. Convert to signed int4 (-8 ~ 7)
# If value >= 8, the highest bit is 1, representing a negative number.
# In two's complement, 4-bit values 8~15 correspond to -8~-1.
# Algorithm: val = val - 16 (if val >= 8)
unpacked_signed = unpacked_unsigned.to(torch.int32) # Ensure calculation precision
mask = unpacked_signed >= 8
unpacked_signed[mask] -= 16
# 5. Convert to float32
result_float = unpacked_signed.to(torch.float32)
result_flat = result_float.flatten(start_dim=-2)
return result_flat
class TestDispatchFFNCombine:
def __init__(self, rank, world_size, port):
self.rank = rank
self.world_size = world_size
self.master_ip = "127.0.0.1"
self.port = port
def get_hcomm(self, comm_group):
hcomm_info = None
if torch.__version__ > "2.0.1":
hcomm_info = comm_group._get_backend(torch.device("npu")).get_hccl_comm_name(self.rank)
else:
hcomm_info = comm_group.get_hccl_comm_name(self.rank)
return hcomm_info
def setup_ep_tp(
self,
rank,
tp_size,
ep_size,
backend_type,
ep_ranks_list=None,
tp_ranks_list=None,
):
for i in range(tp_size):
if ep_ranks_list:
ep_ranks = ep_ranks_list[i]
else:
ep_ranks = [x + ep_size * i for x in range(ep_size)]
ep_group = dist.new_group(backend=backend_type, ranks=ep_ranks)
if rank in ep_ranks:
ep_group_tmp = ep_group
for i in range(ep_size):
if tp_ranks_list:
tp_ranks = tp_ranks_list[i]
else:
tp_ranks = [x * ep_size + i for x in range(tp_size)]
tp_group = dist.new_group(backend=backend_type, ranks=tp_ranks)
if rank in tp_ranks:
tp_group_tmp = tp_group
return ep_group_tmp, tp_group_tmp
def generate_hcom(self):
torch_npu.npu.set_device(DEVICE_OFFSET + self.rank)
dist.init_process_group(
backend="hccl",
rank=self.rank,
world_size=self.world_size,
init_method=f"tcp://127.0.0.1:{self.port}",
)
ep_size = 0
tp_size = self.world_size
hcomm_info_dist = {
"default_pg_info": None,
"ep_hcomm_info": None,
"group_ep": None,
"tp_hcomm_info": None,
"group_tp": None,
}
if ep_size and tp_size:
group_ep, group_tp = self.setup_ep_tp(self.rank, tp_size, ep_size, "hccl", None, None)
hcomm_info_dist["ep_hcomm_info"] = self.get_hcomm(group_ep)
hcomm_info_dist["tp_hcomm_info"] = self.get_hcomm(group_tp)
hcomm_info_dist["group_ep"] = group_ep
hcomm_info_dist["group_tp"] = group_tp
else:
if dist.is_available():
default_pg = _get_default_group()
hcomm_info_dist["default_pg_info"] = self.get_hcomm(default_pg)
hcomm_info = hcomm_info_dist["default_pg_info"]
self.hcomm_info = hcomm_info
def run_tensor_list(self) -> bool:
torch_npu.npu.set_device(DEVICE_OFFSET + self.rank)
m = 64
k = 1024
n = 1024
topk = 8
e = 8
k2 = n // 2
n2 = k
active_num = m // 8
torch_npu.npu.config.allow_internal_format = True
x = self.generate_random_tensor((m, k), dtype=torch.bfloat16)
weight1 = self.generate_random_tensor((e, k, n // 8), dtype=torch.int32).npu()
weight1 = torch_npu.npu_format_cast(weight1, 29)
weight2 = self.generate_random_tensor((e, k2, n2 // 8), dtype=torch.int32).npu()
weight2 = torch_npu.npu_format_cast(weight2, 29)
bias1 = int32_to_8x_int4_float(weight1.cpu())
bias1_npu = bias1.sum(dim=-1).npu() # shape: [e, n]
bias2 = int32_to_8x_int4_float(weight2.cpu())
bias2_npu = bias2.sum(dim=-1).npu() # shape: [e, n2]
expert_idx = torch.randint(0, self.world_size * e, (m, topk), dtype=torch.int32)
scale1 = torch.randint(0, 1, (e, n), dtype=torch.int64)
scale2 = torch.randint(0, 1, (e, n2), dtype=torch.int64)
probs = torch.randn(size=(m, topk), dtype=torch.float32)
x_active_mask = torch.cat(
[
torch.ones(active_num, dtype=torch.bool),
torch.zeros(m - active_num, dtype=torch.bool),
]
)
x[active_num:, :] = 0
expert_idx[active_num:, :] = torch.arange(topk, dtype=torch.int32)
x = x.npu()
expert_idx = expert_idx.npu()
scale1 = scale1.npu()
scale2 = scale2.npu()
probs = probs.npu()
x_active_mask = x_active_mask.npu()
weight1_nz_npu = []
weight2_nz_npu = []
scale1_npu = []
scale2_npu = []
bias1_list = []
bias2_list = []
for i in range(e):
weight1_nz_npu.append(torch_npu.npu_format_cast(weight1[i].npu(), 29))
scale1_npu.append(scale1[i].npu())
bias1_list.append(bias1_npu[i])
weight2_nz_npu.append(torch_npu.npu_format_cast(weight2[i].npu(), 29))
scale2_npu.append(scale2[i].npu())
bias2_list.append(bias2_npu[i])
out = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
expert_token_nums = self.generate_random_tensor((1, e), dtype=torch.int32).npu()
torch.ops._C_ascend.dispatch_ffn_combine(
x=x,
weight1=weight1_nz_npu,
weight2=weight2_nz_npu,
expert_idx=expert_idx,
scale1=scale1_npu,
scale2=scale2_npu,
bias1=bias1_list,
bias2=bias2_list,
probs=probs,
group=self.hcomm_info,
max_output_size=512,
out=out,
expert_token_nums=expert_token_nums,
x_active_mask=x_active_mask,
)
return True
def run_normal(self) -> bool:
torch_npu.npu.set_device(DEVICE_OFFSET + self.rank)
m = 64
k = 1024
n = 1024
topk = 8
e = 8
k2 = n // 2
n2 = k
active_num = m // 2
torch_npu.npu.config.allow_internal_format = True
x = self.generate_random_tensor((m, k), dtype=torch.bfloat16)
weight1 = self.generate_random_tensor((e, k, n // 8), dtype=torch.int32).npu()
weight1 = torch_npu.npu_format_cast(weight1, 29)
weight2 = self.generate_random_tensor((e, k2, n2 // 8), dtype=torch.int32).npu()
weight2 = torch_npu.npu_format_cast(weight2, 29)
bias1 = int32_to_8x_int4_float(weight1.cpu())
bias1_npu = bias1.sum(dim=-1).npu() # shape: [e, n]
bias2 = int32_to_8x_int4_float(weight2.cpu())
bias2_npu = bias2.sum(dim=-1).npu() # shape: [e, n2]
expert_idx = torch.randint(0, self.world_size * e, (m, topk), dtype=torch.int32)
scale1 = torch.randint(0, 1, (e, n), dtype=torch.int64)
scale2 = torch.randint(0, 1, (e, n2), dtype=torch.int64)
probs = torch.randn(size=(m, topk), dtype=torch.float32)
x_active_mask = torch.cat(
[
torch.ones(active_num, dtype=torch.bool),
torch.zeros(m - active_num, dtype=torch.bool),
]
)
x[active_num:, :] = 0
expert_idx[active_num:, :] = torch.arange(topk, dtype=torch.int32)
x = x.npu()
expert_idx = expert_idx.npu()
scale1 = scale1.npu()
scale2 = scale2.npu()
probs = probs.npu()
x_active_mask = x_active_mask.npu()
weight1_nz_npu = []
weight2_nz_npu = []
scale1_npu = []
scale2_npu = []
bias1_list = []
bias2_list = []
weight1_nz_npu.append(torch_npu.npu_format_cast(weight1.npu(), 29))
scale1_npu.append(scale1.npu())
bias1_list.append(bias1_npu)
weight2_nz_npu.append(torch_npu.npu_format_cast(weight2.npu(), 29))
scale2_npu.append(scale2.npu())
bias2_list.append(bias2_npu)
out = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
expert_token_nums = self.generate_random_tensor((1, e), dtype=torch.int32).npu()
torch.ops._C_ascend.dispatch_ffn_combine(
x=x,
weight1=weight1_nz_npu,
weight2=weight2_nz_npu,
expert_idx=expert_idx,
scale1=scale1_npu,
scale2=scale2_npu,
bias1=bias1_list,
bias2=bias2_list,
probs=probs,
group=self.hcomm_info,
max_output_size=512,
out=out,
expert_token_nums=expert_token_nums,
x_active_mask=x_active_mask,
)
return True
def generate_random_tensor(self, size, dtype):
if dtype in [torch.float16, torch.bfloat16, torch.float32]:
return torch.randn(size=size, dtype=dtype)
elif dtype is torch.int8:
return torch.randint(-16, 16, size=size, dtype=dtype)
elif dtype is torch.int32:
return torch.randint(-127, 127, size=size, dtype=dtype)
else:
raise ValueError(f"Invalid dtype: {dtype}")
def worker(rank: int, world_size: int, port: int, q: mp.SimpleQueue):
op = TestDispatchFFNCombine(rank, world_size, port)
op.generate_hcom()
out1 = op.run_tensor_list()
q.put(out1)
out2 = op.run_normal()
q.put(out2)
@torch.inference_mode()
def test_dispatch_ffn_combine_kernel():
world_size = 2
mp.set_start_method("fork", force=True)
q = mp.SimpleQueue()
p_list = []
port = 29501 + random.randint(0, 10000)
for rank in range(world_size):
p = mp.Process(target=worker, args=(rank, world_size, port, q))
p.start()
p_list.append(p)
results = [q.get() for _ in range(world_size)]
for p in p_list:
p.join()
assert all(results)

View File

@@ -0,0 +1,576 @@
import gc
import os
import sys
import time
from pathlib import Path
import npugraph_ex as nge
import numpy as np
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
import torch_npu
from vllm_ascend.utils import enable_custom_op
torch.manual_seed(42)
torch_npu.npu.config.allow_internal_format = True
enable_custom_op()
LOG_NAME = "dispatch_gmm_combine_decode_test_logs"
BASE_KWARGS = {
"batch_size": 64,
"token_hidden_size": 7168,
"moe_intermediate_size": 2048,
"ep_world_size": 16,
"moe_expert_num": 64,
"shared_expert_rank_num": 0,
"top_k": 8,
"test_bfloat16": True,
"enable_dynamic_bs": False,
"test_graph": False,
"with_mc2_mask": False,
"dynamic_eplb": False,
"w8a8_dynamic": True,
"is_nz": True,
}
def redirect_output(log_file_path):
log_path = Path(LOG_NAME) / log_file_path
log_path.parent.mkdir(parents=True, exist_ok=True)
f = open(LOG_NAME + "/" + log_file_path, "w") # noqa: SIM115
os.dup2(f.fileno(), sys.stdout.fileno())
os.dup2(f.fileno(), sys.stderr.fileno())
return f
def permute_weight(w: torch.Tensor, tile_n):
*dims, n = w.shape
order = list(range(len(dims))) + [-2, -3, -1]
return w.reshape(*dims, 2, n // tile_n, tile_n // 2).permute(order).reshape(*dims, n).contiguous()
def output_to_file(rank_id):
return rank_id > 0
class DecodeMoeOps(torch.nn.Module):
def __init__(
self,
gmm1_weight,
gmm1_weight_scale,
gmm2_weight,
gmm2_weight_scale,
ep_hcomm_info,
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
global_rank_id,
shared_expert_rank_num=0,
dynamic_eplb=False,
w8a8_dynamic=True,
is_nz=True,
):
super().__init__()
if w8a8_dynamic:
assert gmm1_weight_scale is not None and gmm2_weight_scale is not None, (
"gmm1_weight_scale and gmm2_weight_scale must be provided for w8a8_dynamic"
)
else:
assert gmm1_weight_scale is None and gmm2_weight_scale is None, (
"gmm1_weight_scale and gmm2_weight_scale must be None for w8a8_dynamic"
)
self.ep_hcomm_info = ep_hcomm_info
self.batch_size = batch_size
self.token_hidden_size = token_hidden_size
self.moe_intermediate_size = moe_intermediate_size
self.ep_world_size = ep_world_size
self.moe_expert_num = moe_expert_num
self.global_rank_id = global_rank_id
self.shared_expert_rank_num = shared_expert_rank_num
is_shared_expert = global_rank_id < shared_expert_rank_num
moe_expert_num_per_rank = moe_expert_num // (ep_world_size - shared_expert_rank_num)
self.local_expert_num = 1 if is_shared_expert else moe_expert_num_per_rank
self.ep_recv_count_size = self.local_expert_num * ep_world_size
self.dynamic_eplb = dynamic_eplb
self.w8a8_dynamic = w8a8_dynamic
self.is_nz = is_nz
self.gmm1_weight = torch.empty([self.local_expert_num, self.token_hidden_size, self.moe_intermediate_size * 2])
self.gmm2_weight = torch.empty([self.local_expert_num, self.moe_intermediate_size, self.token_hidden_size])
if self.w8a8_dynamic:
self.gmm1_weight_scale = torch.empty([self.local_expert_num, self.moe_intermediate_size * 2])
self.gmm2_weight_scale = torch.empty([self.local_expert_num, self.token_hidden_size])
else:
self.gmm1_weight_scale = None
self.gmm2_weight_scale = None
self.gmm1_weight_scale_fp32 = None
self.gmm2_weight_scale_fp32 = None
self._process_weights_after_loading(gmm1_weight, gmm1_weight_scale, gmm2_weight, gmm2_weight_scale)
def _process_weights_after_loading(self, gmm1_weight, gmm1_weight_scale, gmm2_weight, gmm2_weight_scale):
if self.w8a8_dynamic:
gmm1_weight = torch_npu.npu_format_cast(gmm1_weight, torch_npu.Format.FRACTAL_NZ)
gmm2_weight = torch_npu.npu_format_cast(gmm2_weight, torch_npu.Format.FRACTAL_NZ)
self.gmm1_weight = torch.nn.Parameter(gmm1_weight, requires_grad=False)
self.gmm2_weight = torch.nn.Parameter(gmm2_weight, requires_grad=False)
if self.w8a8_dynamic:
self.gmm1_weight_scale = torch.nn.Parameter(gmm1_weight_scale, requires_grad=False)
self.gmm2_weight_scale = torch.nn.Parameter(gmm2_weight_scale, requires_grad=False)
self.gmm1_weight_scale_fp32 = torch.nn.Parameter(gmm1_weight_scale.float(), requires_grad=False)
self.gmm2_weight_scale_fp32 = torch.nn.Parameter(gmm2_weight_scale.float(), requires_grad=False)
def _apply_ops(self, x, expert_ids, smooth_scales, expert_scales, x_active_mask):
raise NotImplementedError("To be implemented in subclass")
def forward(self, x, expert_ids, smooth_scales, expert_scales, x_active_mask):
return self._apply_ops(x, expert_ids, smooth_scales, expert_scales, x_active_mask)
class SmallOps(DecodeMoeOps):
def __init__(
self,
gmm1_weight,
gmm1_weight_scale,
gmm2_weight,
gmm2_weight_scale,
ep_hcomm_info,
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
global_rank_id,
shared_expert_rank_num=0,
dynamic_eplb=False,
w8a8_dynamic=True,
is_nz=True,
):
super().__init__(
gmm1_weight,
gmm1_weight_scale,
gmm2_weight,
gmm2_weight_scale,
ep_hcomm_info,
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
global_rank_id,
shared_expert_rank_num,
dynamic_eplb,
w8a8_dynamic,
is_nz,
)
self.tp_hcomm_info = ""
def _apply_ops(self, x, expert_ids, smooth_scales, expert_scales, x_active_mask):
outputs = torch_npu.npu_moe_distribute_dispatch_v2(
x=x,
expert_ids=expert_ids,
expert_scales=expert_scales,
x_active_mask=x_active_mask,
group_ep=self.ep_hcomm_info,
ep_world_size=self.ep_world_size,
ep_rank_id=self.global_rank_id,
moe_expert_num=self.moe_expert_num,
group_tp=self.tp_hcomm_info,
tp_world_size=1,
tp_rank_id=0,
expert_shard_type=0,
shared_expert_num=1,
shared_expert_rank_num=self.shared_expert_rank_num,
quant_mode=2 if self.w8a8_dynamic else 0,
global_bs=self.batch_size * self.ep_world_size,
expert_token_nums_type=1, # 0代表前缀和,1代表各自数量
)
(
expand_x,
dynamic_scales,
assist_info_for_combine,
expert_token_nums,
ep_send_counts,
tp_send_counts,
expand_scales,
) = outputs
output_dtype = x.dtype
y1_int32 = torch_npu.npu_grouped_matmul(
x=[expand_x],
weight=[self.gmm1_weight],
split_item=3,
group_list_type=1, # 默认为0,代表前缀和形式
group_type=0, # 0代表m轴分组
group_list=expert_token_nums,
output_dtype=torch.int32 if self.w8a8_dynamic else output_dtype,
)[0]
y1_scale = None
if self.w8a8_dynamic:
y1, y1_scale = torch_npu.npu_dequant_swiglu_quant(
x=y1_int32,
weight_scale=self.gmm1_weight_scale.to(torch.float32),
activation_scale=dynamic_scales,
bias=None,
quant_scale=None,
quant_offset=None,
group_index=expert_token_nums,
activate_left=True,
quant_mode=1,
)
else:
y1 = torch_npu.npu_swiglu(y1_int32)
y2 = torch_npu.npu_grouped_matmul(
x=[y1],
weight=[self.gmm2_weight],
scale=[self.gmm2_weight_scale] if self.w8a8_dynamic else None,
per_token_scale=[y1_scale] if self.w8a8_dynamic else None,
split_item=2,
group_list_type=1,
group_type=0,
group_list=expert_token_nums,
output_dtype=output_dtype,
)[0]
combine_output = torch_npu.npu_moe_distribute_combine_v2(
expand_x=y2,
expert_ids=expert_ids,
assist_info_for_combine=assist_info_for_combine,
ep_send_counts=ep_send_counts,
expert_scales=expert_scales,
x_active_mask=x_active_mask,
group_ep=self.ep_hcomm_info,
ep_world_size=self.ep_world_size,
ep_rank_id=self.global_rank_id,
moe_expert_num=self.moe_expert_num,
tp_send_counts=tp_send_counts,
expand_scales=expand_scales,
group_tp=self.tp_hcomm_info,
tp_world_size=1,
tp_rank_id=0,
expert_shard_type=0,
shared_expert_num=1,
shared_expert_rank_num=self.shared_expert_rank_num,
global_bs=self.batch_size * self.ep_world_size,
)
return (combine_output, expert_token_nums)
class FusionOp(DecodeMoeOps):
def __init__(
self,
gmm1_weight,
gmm1_weight_scale,
gmm2_weight,
gmm2_weight_scale,
ep_hcomm_info,
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
global_rank_id,
shared_expert_rank_num=0,
dynamic_eplb=False,
w8a8_dynamic=True,
is_nz=True,
):
super().__init__(
gmm1_weight,
gmm1_weight_scale,
gmm2_weight,
gmm2_weight_scale,
ep_hcomm_info,
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
global_rank_id,
shared_expert_rank_num,
dynamic_eplb,
w8a8_dynamic,
is_nz,
)
def _apply_ops(self, x, expert_ids, smooth_scales, expert_scales, x_active_mask):
smooth_scales = torch.zeros(128 * 1024 * 1024).npu()
output = torch.ops._C_ascend.dispatch_gmm_combine_decode(
x=x,
expert_ids=expert_ids,
gmm1_permuted_weight=self.gmm1_weight,
gmm1_permuted_weight_scale=self.gmm1_weight_scale_fp32,
gmm2_weight=self.gmm2_weight,
gmm2_weight_scale=self.gmm2_weight_scale_fp32,
expert_scales=expert_scales,
expert_smooth_scales=smooth_scales,
x_active_mask=x_active_mask,
group_ep=self.ep_hcomm_info,
ep_rank_size=self.ep_world_size,
ep_rank_id=self.global_rank_id,
moe_expert_num=self.moe_expert_num,
shared_expert_num=1,
shared_expert_rank_num=self.shared_expert_rank_num,
quant_mode=0,
global_bs=self.batch_size * self.ep_world_size,
)
return output
def _process_weights_after_loading(self, gmm1_weight, gmm1_weight_scale, gmm2_weight, gmm2_weight_scale):
if self.is_nz:
gmm1_weight = torch_npu.npu_format_cast(gmm1_weight, torch_npu.Format.FRACTAL_NZ)
gmm2_weight = torch_npu.npu_format_cast(gmm2_weight, torch_npu.Format.FRACTAL_NZ)
if self.dynamic_eplb:
self.gmm1_weight = [weight.clone() for weight in gmm1_weight.unbind(dim=0)]
self.gmm2_weight = [weight.clone() for weight in gmm2_weight.unbind(dim=0)]
if self.w8a8_dynamic:
self.gmm1_weight_scale_fp32 = [weight.clone() for weight in gmm1_weight_scale.unbind(dim=0)]
self.gmm2_weight_scale_fp32 = [weight.clone() for weight in gmm2_weight_scale.unbind(dim=0)]
else:
self.gmm1_weight_scale_fp32 = [torch.ones(1).npu().to(gmm1_weight.dtype)]
self.gmm2_weight_scale_fp32 = [torch.ones(1).npu().to(gmm2_weight.dtype)]
else:
self.gmm1_weight = [gmm1_weight.clone()]
self.gmm2_weight = [gmm2_weight.clone()]
if self.w8a8_dynamic:
self.gmm1_weight_scale_fp32 = [gmm1_weight_scale.clone()]
self.gmm2_weight_scale_fp32 = [gmm2_weight_scale.clone()]
else:
self.gmm1_weight_scale_fp32 = [torch.ones(1).npu().to(gmm1_weight.dtype)]
self.gmm2_weight_scale_fp32 = [torch.ones(1).npu().to(gmm2_weight.dtype)]
def generate_datas(
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
global_rank_id,
shared_expert_rank_num=0,
top_k=8,
test_bfloat16=True,
enable_dynamic_bs=False,
with_mc2_mask=False,
w8a8_dynamic=True,
):
is_shared_expert = global_rank_id < shared_expert_rank_num
moe_expert_num_per_rank = moe_expert_num // (ep_world_size - shared_expert_rank_num)
actual_bs = int(
torch.randint(2 if with_mc2_mask else 1, batch_size, [1]).item() if enable_dynamic_bs else batch_size
)
local_expert_num = 1 if is_shared_expert else moe_expert_num_per_rank
gmm1_input_dim = token_hidden_size
gmm1_output_dim = moe_intermediate_size * 2
gmm2_input_dim = moe_intermediate_size
gmm2_output_dim = token_hidden_size
x = torch.rand([actual_bs, token_hidden_size]) * 0.5 - 0.5
expert_ids = (
torch.arange(global_rank_id * batch_size * top_k, global_rank_id * batch_size * top_k + actual_bs * top_k)
.to(torch.int32)
.view(actual_bs, top_k)
)
expert_ids = expert_ids % moe_expert_num
gmm1_weight_scale = None
gmm2_weight_scale = None
if w8a8_dynamic:
if is_shared_expert:
gmm1_weight = torch.ones([local_expert_num, gmm1_input_dim, gmm1_output_dim]).to(torch.int8) * 4
gmm2_weight = torch.ones([local_expert_num, gmm2_input_dim, gmm2_output_dim]).to(torch.int8) * 4
gmm1_weight[:, :, ::2] = gmm1_weight[:, :, ::2] * -1
gmm2_weight[:, :, ::2] = gmm2_weight[:, :, ::2] * -1
gmm1_weight_scale = torch.ones([local_expert_num, gmm1_output_dim]) * 0.0015
gmm2_weight_scale = torch.ones([local_expert_num, gmm2_output_dim]) * 0.0015
else:
gmm1_weight = torch.randint(-16, 16, [local_expert_num, gmm1_input_dim, gmm1_output_dim]).to(torch.int8)
gmm2_weight = torch.randint(-16, 16, [local_expert_num, gmm2_input_dim, gmm2_output_dim]).to(torch.int8)
gmm1_weight_scale = torch.rand([local_expert_num, gmm1_output_dim]) * 0.003 + 0.0015
gmm2_weight_scale = torch.rand([local_expert_num, gmm2_output_dim]) * 0.003 + 0.0015
else:
if is_shared_expert:
gmm1_weight = (
torch.ones([local_expert_num, gmm1_input_dim, gmm1_output_dim]).to(
torch.bfloat16 if test_bfloat16 else torch.float16
)
* 0.5
)
gmm2_weight = (
torch.ones([local_expert_num, gmm2_input_dim, gmm2_output_dim]).to(
torch.bfloat16 if test_bfloat16 else torch.float16
)
* 0.5
)
else:
gmm1_weight = (
torch.rand([local_expert_num, gmm1_input_dim, gmm1_output_dim]).to(
torch.bfloat16 if test_bfloat16 else torch.float16
)
* 0.25
)
gmm2_weight = (
torch.rand([local_expert_num, gmm2_input_dim, gmm2_output_dim]).to(
torch.bfloat16 if test_bfloat16 else torch.float16
)
* 0.25
)
gmm1_weight[:, ::2, :] = gmm1_weight[:, ::2, :] * -1
gmm2_weight[:, ::2, :] = gmm2_weight[:, ::2, :] * -1
expert_scales = torch.rand(actual_bs, top_k)
if test_bfloat16:
x = x.bfloat16()
if w8a8_dynamic:
assert gmm1_weight_scale is not None and gmm2_weight_scale is not None, (
"gmm1_weight_scale and gmm2_weight_scale must be provided for w8a8_dynamic"
)
gmm1_weight_scale = gmm1_weight_scale.bfloat16()
gmm2_weight_scale = gmm2_weight_scale.bfloat16()
else:
x = x.half()
smooth_sales = None
x_active_mask = None
valid_token_num = actual_bs
if with_mc2_mask:
valid_token_num = int(torch.randint(1, actual_bs, [1]).item())
x_active_mask = torch.cat((torch.ones(valid_token_num), torch.zeros(actual_bs - valid_token_num))).bool()
return (
(x, expert_ids, smooth_sales, expert_scales, x_active_mask),
(gmm1_weight, gmm1_weight_scale, gmm2_weight, gmm2_weight_scale),
actual_bs,
valid_token_num,
)
def run_once(
local_rank_id,
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
shared_expert_rank_num=0,
top_k=8,
test_bfloat16=True,
enable_dynamic_bs=False,
test_graph=False,
with_mc2_mask=False,
dynamic_eplb=False,
w8a8_dynamic=True,
is_nz=True,
):
log_file = redirect_output(f"local_rank_{local_rank_id}.log") if output_to_file(local_rank_id) else None
global_rank_id = local_rank_id # 单机
device_id = local_rank_id % 16
torch_npu.npu.set_device(device_id)
# 初始化分布式环境
os.environ["MASTER_ADDR"] = "127.0.0.1"
os.environ["MASTER_PORT"] = "29500" # 端口号随意
dist.init_process_group(backend="hccl", rank=local_rank_id, world_size=ep_world_size)
ep_ranks_list = list(np.arange(0, ep_world_size))
ep_group = dist.new_group(backend="hccl", ranks=ep_ranks_list)
ep_group_small = dist.new_group(backend="hccl", ranks=ep_ranks_list)
ep_hcomm_info_fused = ep_group._get_backend(torch.device("npu")).get_hccl_comm_name(local_rank_id)
ep_hcomm_info_small = ep_group_small._get_backend(torch.device("npu")).get_hccl_comm_name(local_rank_id)
torch_npu.npu.synchronize(device_id)
parameter = (
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
global_rank_id,
shared_expert_rank_num,
)
input_datas, weight_datas, actual_bs, valid_token_num = generate_datas(
*parameter, top_k, test_bfloat16, enable_dynamic_bs, with_mc2_mask, w8a8_dynamic
)
input_datas = [data.npu() if data is not None else None for data in input_datas]
weight_datas = [data.npu() if data is not None else None for data in weight_datas]
small_ops = SmallOps(*weight_datas, ep_hcomm_info_small, *parameter, dynamic_eplb, w8a8_dynamic, is_nz).npu() # type: ignore
fused_ops = FusionOp(*weight_datas, ep_hcomm_info_fused, *parameter, dynamic_eplb, w8a8_dynamic, is_nz).npu() # type: ignore
if test_graph:
config = nge.CompilerConfig()
npu_backend = nge.get_npu_backend(compiler_config=config)
fused_ops = torch.compile(fused_ops, backend=npu_backend)
# test performance
start_time = time.perf_counter()
for _ in range(100):
small_op_token_output, small_op_count_output = small_ops(*input_datas)
torch_npu.npu.synchronize(device_id)
end_time = time.perf_counter()
elapsed_time = end_time - start_time
elapsed_time_us = elapsed_time * 1000000
print(f"rank-{global_rank_id} small {elapsed_time_us} us")
start_time = time.perf_counter()
for _ in range(100):
fused_op_token_output, fused_op_count_output = fused_ops(*input_datas)
torch_npu.npu.synchronize(device_id)
end_time = time.perf_counter()
elapsed_time = end_time - start_time
elapsed_time_us = elapsed_time * 1000000
print(f"rank-{global_rank_id} fused {elapsed_time_us} us")
small_op_token_output, small_op_count_output = small_ops(*input_datas)
torch_npu.npu.synchronize(device_id)
print(f"rank-{global_rank_id} Small op End")
fused_op_token_output, fused_op_count_output = fused_ops(*input_datas)
torch_npu.npu.synchronize(device_id)
print(f"rank-{global_rank_id} Fused op End")
dist.destroy_process_group()
if log_file is not None:
log_file.close()
try:
torch.testing.assert_close(
small_op_token_output[0:valid_token_num].cpu(),
fused_op_token_output[0:valid_token_num].cpu(),
atol=2.0,
rtol=0.02,
)
torch.testing.assert_close(small_op_count_output.cpu(), fused_op_count_output.cpu())
except Exception as e:
print(f"rank-{global_rank_id} Assert close Failed: {e}")
else:
print(f"rank-{global_rank_id} Assert close Pass")
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@torch.inference_mode()
def test_dispatch_gmm_combine_decode_base():
custom_kwargs = BASE_KWARGS.copy()
custom_kwargs["batch_size"] = 32
custom_kwargs["ep_world_size"] = 8
custom_kwargs["moe_expert_num"] = 32
custom_kwargs["w8a8_dynamic"] = False
custom_kwargs["is_nz"] = True
ep_world_size = custom_kwargs["ep_world_size"]
custom_args = tuple(custom_kwargs.values())
print(f"{custom_kwargs=}")
mp.spawn(run_once, args=custom_args, nprocs=ep_world_size, join=True)
print(f"{custom_kwargs=}")
@torch.inference_mode()
def test_dispatch_gmm_combine_decode_with_mc2_mask():
custom_kwargs = BASE_KWARGS.copy()
custom_kwargs["with_mc2_mask"] = True
ep_world_size = custom_kwargs["ep_world_size"]
custom_args = tuple(custom_kwargs.values())
mp.spawn(run_once, args=custom_args, nprocs=ep_world_size, join=True)
@torch.inference_mode()
def test_dispatch_gmm_combine_decode_dynamic_eplb():
custom_kwargs = BASE_KWARGS.copy()
custom_kwargs["dynamic_eplb"] = True
ep_world_size = custom_kwargs["ep_world_size"]
custom_args = tuple(custom_kwargs.values())
mp.spawn(run_once, args=custom_args, nprocs=ep_world_size, join=True)
if __name__ == "__main__":
test_dispatch_gmm_combine_decode_base()

View File

@@ -0,0 +1,139 @@
import gc
import random
import numpy as np
import pytest
import torch
seed = 45
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
def npu_add_rms_norm_bias_golden(input_x1, input_x2, input_gamma, input_beta, kernelType, epsilon=0.000001):
ori_x_shape = input_x1.shape
ori_gamma_shape = input_gamma.shape
xlength = len(ori_x_shape)
gammaLength = len(ori_gamma_shape)
torchType32 = torch.float32
rstdShape = []
rstdSize = 1
for i in range(xlength):
if i < (xlength - gammaLength):
rstdShape.append(ori_x_shape[i])
rstdSize = rstdSize * ori_x_shape[i]
else:
rstdShape.append(1)
n = xlength - gammaLength
gammaSize = np.multiply.reduce(np.array(ori_gamma_shape))
input_gamma = input_gamma.reshape(gammaSize)
input_beta = input_beta.reshape(gammaSize)
x1_shape = ori_x_shape[0:n] + input_gamma.shape
input_x1 = input_x1.reshape(x1_shape)
input_x2 = input_x2.reshape(x1_shape)
if kernelType == 1:
oriType = torch.float16
xOut = input_x1.to(oriType) + input_x2.to(oriType)
elif kernelType == 2:
oriType = torch.bfloat16
x_fp32 = input_x1.to(torchType32) + input_x2.to(torchType32)
xOut = x_fp32.to(oriType)
else:
oriType = torch.float32
xOut = input_x1.to(torchType32) + input_x2.to(torchType32)
x_fp32 = xOut.to(torchType32)
avgFactor = 1 / gammaSize
x_2 = torch.pow(x_fp32, 2)
x_2_mean = x_2 * avgFactor
tmp_sum = torch.sum(x_2_mean, axis=-1, keepdims=True)
tmp_add_eps = tmp_sum + epsilon
std = torch.sqrt(tmp_add_eps)
rstd = 1 / std
result_mid = x_fp32 * rstd
if kernelType == 1:
result_mid_ori = result_mid.to(oriType)
y_array = result_mid_ori * input_gamma.to(oriType)
y_array = y_array + input_beta.to(oriType)
elif kernelType == 2:
result_mid_ori = result_mid.to(oriType)
y_array = result_mid_ori.to(torchType32) * input_gamma.to(torchType32)
y_array = y_array + input_beta.to(torchType32)
else:
y_array = result_mid.to(torchType32) * input_gamma.to(torchType32)
y_array = y_array + input_beta.to(torchType32)
rstdOut = rstd.reshape(rstdShape).to(torchType32)
yOut = y_array.reshape(ori_x_shape).to(oriType)
xOut = x_fp32.reshape(ori_x_shape).to(oriType)
return yOut, rstdOut, xOut
@pytest.mark.parametrize(
"row",
[1, 16, 64, 77, 128, 255, 1000],
)
@pytest.mark.parametrize(
"col",
[
8,
16,
128,
3000,
7168,
15000,
],
)
@pytest.mark.parametrize(
"dtype, atol, rtol, kernelType",
[
(torch.float16, 0.0010986328125, 0.0010986328125, 1),
(torch.bfloat16, 0.0079345703125, 0.0079345703125, 2),
(torch.float32, 0.000244140625, 0.000244140625, 3),
],
)
def test_quant_fpx_linear(row: int, col: int, dtype, atol, rtol, kernelType):
shape_x = [row, col]
shape_gamma = [col]
dataType = dtype
input_x1 = np.random.uniform(1, 10, size=tuple(shape_x)).astype(np.float32)
input_x1_tensor = torch.tensor(input_x1).type(dataType)
input_x2 = np.random.uniform(1, 10, size=tuple(shape_x)).astype(np.float32)
input_x2_tensor = torch.tensor(input_x2).type(dataType)
input_gamma = np.random.uniform(1, 10, size=tuple(shape_gamma)).astype(np.float32)
input_gamma_tensor = torch.tensor(input_gamma).type(dataType)
input_beta = np.random.uniform(1, 10, size=tuple(shape_gamma)).astype(np.float32)
grad_bias = torch.tensor(input_beta).type(dataType)
y, rstd, x = torch.ops._C_ascend.npu_add_rms_norm_bias(
input_x1_tensor.npu(), input_x2_tensor.npu(), input_gamma_tensor.npu(), grad_bias.npu(), 1e-6
)
y = y.cpu()
rstd = rstd.cpu()
x = x.cpu()
y1, rstd1, x1 = npu_add_rms_norm_bias_golden(
input_x1_tensor, input_x2_tensor, input_gamma_tensor, grad_bias, kernelType, epsilon=0.000001
)
a = y1 > 1
a1 = y1 <= 1
b = rstd1 > 1
b1 = rstd1 <= 1
c = x1 > 1
c1 = x1 <= 1
torch.testing.assert_close(y * a, y1 * a, atol=atol, rtol=100)
torch.testing.assert_close(y * a1, y1 * a1, rtol=rtol, atol=100)
torch.testing.assert_close(rstd * b, rstd1 * b, atol=atol, rtol=100)
torch.testing.assert_close(rstd * b1, rstd1 * b1, rtol=rtol, atol=100)
torch.testing.assert_close(x * c, x1 * c, atol=atol, rtol=100)
torch.testing.assert_close(x * c1, x1 * c1, rtol=rtol, atol=100)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,140 @@
import gc
import random
import unittest
import torch
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
torch.set_printoptions(threshold=float("inf"))
class TestMatrixMultiplication(unittest.TestCase):
def compute_golden(self, a, b, res1, m, n):
"""Compute reference result (golden)"""
torch.bmm(a.transpose(0, 1), b, out=res1.view(-1, m, n).transpose(0, 1))
def assert_tensors_almost_equal(self, actual, expected, dtype):
"""Check if two tensors are approximately equal (considering floating point errors)"""
self.assertEqual(actual.shape, expected.shape, "Shape mismatch")
# Check for NaN
self.assertFalse(torch.isnan(actual).any(), "Actual result contains NaN")
self.assertFalse(torch.isnan(expected).any(), "Expected result contains NaN")
# Check for Inf
self.assertFalse(torch.isinf(actual).any(), "Actual result contains Inf")
self.assertFalse(torch.isinf(expected).any(), "Expected result contains Inf")
# Set different tolerances based on data type
if dtype == torch.float16:
rtol, atol = 1e-5, 1e-5
else: # bfloat16
rtol, atol = 1.5e-5, 1.5e-5
# Compare values
diff = torch.abs(actual - expected)
max_diff = diff.max().item()
max_expected = torch.abs(expected).max().item()
# Check relative and absolute errors
if max_expected > 0:
relative_diff = max_diff / max_expected
self.assertLessEqual(
relative_diff,
rtol,
f"Relative error too large: {relative_diff} > {rtol}. Max difference: {max_diff}",
)
self.assertLessEqual(max_diff, atol, f"Absolute error too large: {max_diff} > {atol}")
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
def test_boundary_conditions(self):
"""Test boundary conditions"""
test_cases = [
# (b, m, k, n)
(1, 1, 1, 1), # Minimum size
(1, 10, 1, 1), # b=1
(10, 1, 1, 10), # m=1
(5, 5, 1, 5), # k=1
(2, 2, 2, 1), # n=1
(100, 1, 1, 100), # Flat case
(1, 100, 100, 1), # Flat case
(2, 3, 4, 5), # Random small size
(10, 20, 30, 40), # Medium size
(36, 128, 512, 128), # target case
(8, 160, 512, 128),
]
dtypes = [torch.float16, torch.bfloat16]
for dtype in dtypes:
for b, m, k, n in test_cases:
with self.subTest(dtype=dtype, shape=f"({b}, {m}, {k}, {n})"):
a = torch.randn(b, m, k, dtype=dtype, device="npu")
b_tensor = torch.randn(m, k, n, dtype=dtype, device="npu")
res1 = torch.empty((b, m * n), dtype=dtype, device="npu")
res2 = torch.empty((b, m, n), dtype=dtype, device="npu")
self.compute_golden(a, b_tensor, res1, m, n)
torch.ops._C_ascend.batch_matmul_transpose(a, b_tensor, res2)
self.assert_tensors_almost_equal(res1.view(-1, m, n), res2, dtype)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
def test_random_shapes(self):
"""Test randomly generated shapes"""
num_tests = 1
dtypes = [torch.float16, torch.bfloat16]
for dtype in dtypes:
for _ in range(num_tests):
# Generate reasonable random sizes
b = random.randint(1, 500)
m = random.randint(1, 500)
k = random.randint(1, 500)
n = random.randint(1, 500)
with self.subTest(dtype=dtype, shape=f"Random ({b}, {m}, {k}, {n})"):
a = torch.randn(b, m, k, dtype=dtype, device="npu")
b_tensor = torch.randn(m, k, n, dtype=dtype, device="npu")
res1 = torch.empty((b, m * n), dtype=dtype, device="npu")
res2 = torch.empty((b, m, n), dtype=dtype, device="npu")
self.compute_golden(a, b_tensor, res1, m, n)
torch.ops._C_ascend.batch_matmul_transpose(a, b_tensor, res2)
self.assert_tensors_almost_equal(res1.view(-1, m, n), res2, dtype)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
def test_zero_values(self):
"""Test zero input values"""
dtypes = [torch.float16, torch.bfloat16]
b, m, k, n = 5, 4, 3, 2
for dtype in dtypes:
with self.subTest(dtype=dtype):
a = torch.zeros(b, m, k, dtype=dtype, device="npu")
b_tensor = torch.zeros(m, k, n, dtype=dtype, device="npu")
res1 = torch.empty((b, m * n), dtype=dtype, device="npu")
res2 = torch.empty((b, m, n), dtype=dtype, device="npu")
self.compute_golden(a, b_tensor, res1, m, n)
torch.ops._C_ascend.batch_matmul_transpose(a, b_tensor, res2)
self.assert_tensors_almost_equal(res1.view(-1, m, n), res2, dtype)
self.assertTrue(torch.all(res2 == 0))
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
if __name__ == "__main__":
unittest.main(verbosity=2)

View File

@@ -0,0 +1,42 @@
import gc
import torch
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
DEFAULT_ATOL = 1e-3
DEFAULT_RTOL = 1e-3
def bgmv_expand_cpu_impl(
x: torch.Tensor, w: torch.Tensor, indices: torch.Tensor, y: torch.tensor, slice_offset: int, slice_size: int
) -> torch.Tensor:
W = w[indices, :, :].transpose(-1, -2).to(torch.float32)
z = torch.bmm(x.unsqueeze(1).to(torch.float32), W).squeeze()
y[:, slice_offset : slice_offset + slice_size] += z
return y
@torch.inference_mode()
def test_bgmv_expand():
B = 1
x = torch.randn([B, 16], dtype=torch.float)
w = torch.randn([64, 128, 16], dtype=torch.float16)
indices = torch.zeros([B], dtype=torch.int64)
y = torch.randn([B, 128 * 3], dtype=torch.float16)
x_npu = x.npu()
w_npu = w.npu()
indices_npu = indices.npu()
y_npu = y.npu()
y_out = bgmv_expand_cpu_impl(x, w, indices, y, 0, 128)
y_out_npu = torch.ops._C_ascend.bgmv_expand(x_npu, w_npu, indices_npu, y_npu, 0, 128)
# Compare the results.
torch.testing.assert_close(y_out_npu.cpu(), y_out.cpu(), atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,42 @@
import gc
import torch
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
DEFAULT_ATOL = 1e-3
DEFAULT_RTOL = 1e-3
def bgmv_shrink_cpu_impl(
x: torch.Tensor, w: torch.Tensor, indices: torch.Tensor, y: torch.tensor, scaling: float
) -> torch.Tensor:
W = w[indices, :, :].transpose(-1, -2).to(torch.float32)
z = torch.bmm(x.unsqueeze(1).to(torch.float32), W).squeeze()
y[:, :] += z * scaling
return y
@torch.inference_mode()
def test_bgmv_shrink():
B = 1
x = torch.randn([B, 128], dtype=torch.float16)
w = torch.randn([64, 16, 128], dtype=torch.float16)
indices = torch.zeros([B], dtype=torch.int64)
y = torch.zeros([B, 16])
x_npu = x.npu()
w_npu = w.npu()
indices_npu = indices.npu()
y_npu = y.npu()
y = bgmv_shrink_cpu_impl(x, w, indices, y, 0.5)
torch.ops._C_ascend.bgmv_shrink(x_npu, w_npu, indices_npu, y_npu, 0.5)
# Compare the results.
torch.testing.assert_close(y_npu.cpu(), y.cpu(), atol=DEFAULT_ATOL, rtol=DEFAULT_RTOL)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,150 @@
import pytest
import torch
import torch_npu
from vllm.v1.attention.backends.utils import PAD_SLOT_ID
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.utils import enable_custom_op
from vllm_ascend.utils import is_310p as is_310p_hw
torch_npu.npu.set_compile_mode(jit_compile=False)
def validate_cmp(y_cal, y_ref, device="npu"):
y_cal = y_cal.to(device)
y_ref = y_ref.to(device)
torch.testing.assert_close(y_ref, y_cal, rtol=3e-03, atol=1e-02, equal_nan=True)
@pytest.mark.skipif(not is_310p_hw(), reason="Tested separately on a 310P machine.")
@pytest.mark.parametrize("has_initial_state", [False, True])
@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_310_fn(
dim, width, extra_state_len, seq_len, has_bias, silu_activation, 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=torch.float16).transpose(0, 1)
weight = torch.randn(dim, width, device=device, dtype=torch.float16)
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=torch.float16).transpose(-1, -2)
conv_states_ref = (
torch.randn((num_seq, state_len, dim), device=device, dtype=torch.float16)
.transpose(-1, -2)
.copy_(conv_states)
)
else:
conv_states = torch.zeros((num_seq, state_len, dim), device=device, dtype=torch.float16).transpose(-1, -2)
conv_states_ref = torch.zeros((num_seq, state_len, dim), device=device, dtype=torch.float16).transpose(-1, -2)
if has_bias:
bias = torch.randn(dim, device=device, dtype=torch.float16)
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,
)
x_origin = x.transpose(-1, -2)
weight_origin = weight.transpose(-1, -2)
conv_states_origin = conv_states.transpose(-1, -2)
activation_mode = 1 if activation else 0
out = torch.ops._C_ascend.npu_causal_conv1d_310(
x_origin,
weight_origin,
bias=bias,
conv_states=conv_states_origin,
query_start_loc=query_start_loc,
cache_indices=cache_indices,
initial_state_mode=has_initial_state_tensor,
num_accepted_tokens=None,
activation_mode=activation_mode,
pad_slot_id=PAD_SLOT_ID,
run_mode=0,
).transpose(-1, -2)
validate_cmp(out, out_ref)
validate_cmp(conv_states, conv_states_ref)
@pytest.mark.skipif(not is_310p_hw(), reason="Tested separately on a 310P machine.")
@pytest.mark.parametrize("itype", [torch.float16])
@pytest.mark.parametrize("silu_activation", [True])
@pytest.mark.parametrize("has_bias", [False, True])
@pytest.mark.parametrize("seqlen", [1, 3])
@pytest.mark.parametrize("width", [4])
@pytest.mark.parametrize("dim", [2048, 4096, 8192])
@pytest.mark.parametrize("batch_size", [4, 8, 16, 32, 64])
def test_causal_conv1d_310_update(batch_size, dim, width, seqlen, has_bias, silu_activation, itype):
device = "npu"
# 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(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
# conv_states will be (cache_lines, dim, state_len)
# with contiguous along dim-axis
conv_states = torch.randn(total_entries, width, dim, device=device, dtype=itype).transpose(-1, -2)
conv_state_for_padding_test = conv_states.detach().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_states[conv_state_indices, :].detach().clone()
activation = None if not silu_activation else "silu"
activation_mode = 1 if activation else 0
conv_states_origin = conv_states.transpose(-1, -2)
out = torch.ops._C_ascend.npu_causal_conv1d_310(
x.transpose(-1, -2),
weight.transpose(-1, -2),
bias=bias,
conv_states=conv_states_origin,
query_start_loc=None,
cache_indices=conv_state_indices,
initial_state_mode=None,
num_accepted_tokens=None,
activation_mode=activation_mode,
pad_slot_id=PAD_SLOT_ID,
run_mode=1,
).transpose(-1, -2)
out_ref = causal_conv1d_update_ref(
x_ref[:batch_size].transpose(-1, -2), conv_state_ref, weight, bias, activation=activation
).transpose(-1, -2)
validate_cmp(out[:batch_size], out_ref)
validate_cmp(conv_states_origin[conv_state_indices, :], conv_state_ref.transpose(-1, -2))
validate_cmp(
conv_states_origin[unused_states_bool], conv_state_for_padding_test[unused_states_bool].transpose(-1, -2)
)

View File

@@ -0,0 +1,337 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from typing import Any
import pytest
import torch
import torch_npu # noqa: F401
from vllm_ascend.utils import bootstrap_custom_op_env
bootstrap_custom_op_env(include_vendor_lib=True)
import vllm_ascend.vllm_ascend_C # type: ignore[import-untyped] # noqa: E402,F401
KV_BLOCK_SIZE = 128
SLOT_MAPPING_FLAT = 1
SLOT_MAPPING_BLOCK_OFFSET = 2
ROPE_DIM = 64
ROPE_ROWS = 2048
@dataclass(frozen=True)
class CompressorMetadataCase:
name: str
compress_ratio: int
query_start_loc: tuple[int, ...]
start_pos: tuple[int, ...]
block_table: tuple[tuple[int, ...], ...]
expected_slot_mapping: tuple[tuple[int, int], ...] | tuple[int, ...] | None = None
slot_mapping_format: int = SLOT_MAPPING_BLOCK_OFFSET
num_rows: int | None = None
def _invalid_block_offset_rows(count: int) -> tuple[tuple[int, int], ...]:
return tuple((-1, KV_BLOCK_SIZE - 1) for _ in range(count))
def _case_with_format(
name: str,
slot_mapping_format: int,
**kwargs,
) -> CompressorMetadataCase:
suffix = "a5_flat" if slot_mapping_format == SLOT_MAPPING_FLAT else "a2a3_block_offset"
return CompressorMetadataCase(
name=f"{name}_{suffix}",
slot_mapping_format=slot_mapping_format,
**kwargs,
)
DSV4_FLASH_MIXED_CASES: list[dict[str, Any]] = [
{
"name": "dsv4_flash_c4_four_request_prefill_mixed_lengths_padded",
"compress_ratio": 4,
"query_start_loc": (0, 131, 151, 171, 178),
"start_pos": (0, 125, 508, 640),
"block_table": (
(10, 11, 12, 13, 14, 15),
(20, 21, 22, 23, 24, 25),
(30, 31, 32, 33, 34, 35),
(40, 41, 42, 43, 44, 45),
),
"num_rows": 48,
},
{
"name": "dsv4_flash_c128_four_request_prefill_mixed_lengths_padded",
"compress_ratio": 128,
"query_start_loc": (0, 131, 391, 521, 531),
"start_pos": (0, 127, 254, 512),
"block_table": (
(50, 51, 52, 53),
(60, 61, 62, 63),
(70, 71, 72, 73),
(80, 81, 82, 83),
),
"num_rows": 10,
},
{
"name": "dsv4_flash_c4_ten_request_decode_mixed_boundaries_padded",
"compress_ratio": 4,
"query_start_loc": (0, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20),
"start_pos": (126, 127, 128, 129, 130, 131, 132, 133, 134, 135),
"block_table": (
(90, 91, 92, 93),
(100, 101, 102, 103),
(110, 111, 112, 113),
(120, 121, 122, 123),
(130, 131, 132, 133),
(140, 141, 142, 143),
(150, 151, 152, 153),
(160, 161, 162, 163),
(170, 171, 172, 173),
(180, 181, 182, 183),
),
"num_rows": 12,
},
{
"name": "dsv4_flash_c128_ten_request_decode_mixed_boundaries_padded",
"compress_ratio": 128,
"query_start_loc": (0, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20),
"start_pos": (126, 127, 128, 254, 255, 256, 382, 383, 384, 510),
"block_table": (
(190, 191, 192, 193),
(200, 201, 202, 203),
(210, 211, 212, 213),
(220, 221, 222, 223),
(230, 231, 232, 233),
(240, 241, 242, 243),
(250, 251, 252, 253),
(260, 261, 262, 263),
(270, 271, 272, 273),
(280, 281, 282, 283),
),
"num_rows": 10,
},
]
DSV4_FLASH_CASES = [
CompressorMetadataCase(
name="dsv4_flash_c4_single_request_prefill_127",
compress_ratio=4,
query_start_loc=(0, 131),
start_pos=(0,),
block_table=((1, 0, 0, 0, 0, 0),),
expected_slot_mapping=tuple((1, offset) for offset in range(32)) + _invalid_block_offset_rows(1),
),
CompressorMetadataCase(
name="dsv4_flash_c4_two_request_prefill_padded_127",
compress_ratio=4,
query_start_loc=(0, 131, 262),
start_pos=(0, 0),
block_table=((1, 0, 0, 0, 0, 0), (3, 0, 0, 0, 0, 0)),
expected_slot_mapping=(
tuple((1, offset) for offset in range(32))
+ tuple((3, offset) for offset in range(32))
+ _invalid_block_offset_rows(3)
),
),
CompressorMetadataCase(
name="dsv4_flash_c4_two_request_decode_127",
compress_ratio=4,
query_start_loc=(0, 2, 4, 6, 8),
start_pos=(133, 0, 0, 0),
block_table=(
(1, 0, 0, 0, 0, 0),
(0, 0, 0, 0, 0, 0),
(0, 0, 0, 0, 0, 0),
(0, 0, 0, 0, 0, 0),
),
expected_slot_mapping=_invalid_block_offset_rows(2),
),
CompressorMetadataCase(
name="dsv4_flash_c4_decode_one_valid_one_invalid_127",
compress_ratio=4,
query_start_loc=(0, 2, 4),
start_pos=(130, 133),
block_table=((5, 0, 0, 0, 0, 0), (6, 0, 0, 0, 0, 0)),
expected_slot_mapping=((5, 32),) + _invalid_block_offset_rows(2),
),
CompressorMetadataCase(
name="dsv4_flash_c4_single_request_prefill_flat_127",
compress_ratio=4,
query_start_loc=(0, 131),
start_pos=(0,),
block_table=((1, 0, 0, 0, 0, 0),),
expected_slot_mapping=tuple(128 + offset for offset in range(32)) + (-1,),
slot_mapping_format=1,
),
CompressorMetadataCase(
name="dsv4_flash_c128_single_request_prefill_127",
compress_ratio=128,
query_start_loc=(0, 131),
start_pos=(0,),
block_table=((2,),),
expected_slot_mapping=((2, 0),) + _invalid_block_offset_rows(1),
),
CompressorMetadataCase(
name="dsv4_flash_c128_two_request_prefill_padded_127",
compress_ratio=128,
query_start_loc=(0, 131, 262),
start_pos=(0, 0),
block_table=((2,), (4,)),
expected_slot_mapping=((2, 0), (4, 0)) + _invalid_block_offset_rows(2),
),
CompressorMetadataCase(
name="dsv4_flash_c128_two_request_decode_127",
compress_ratio=128,
query_start_loc=(0, 2, 4, 6, 8),
start_pos=(133, 0, 0, 0),
block_table=((2,), (0,), (0,), (0,)),
expected_slot_mapping=_invalid_block_offset_rows(2),
),
CompressorMetadataCase(
name="dsv4_flash_c128_decode_one_valid_one_invalid_127",
compress_ratio=128,
query_start_loc=(0, 2, 4),
start_pos=(254, 257),
block_table=((7,), (8,)),
expected_slot_mapping=((7, 1),) + _invalid_block_offset_rows(1),
),
*(
_case_with_format(case["name"], slot_mapping_format, **{k: v for k, v in case.items() if k != "name"})
for case in DSV4_FLASH_MIXED_CASES
for slot_mapping_format in (SLOT_MAPPING_BLOCK_OFFSET, SLOT_MAPPING_FLAT)
),
]
def _make_rope() -> tuple[torch.Tensor, torch.Tensor]:
values = torch.arange(ROPE_ROWS * ROPE_DIM, dtype=torch.float32).reshape(ROPE_ROWS, ROPE_DIM)
return values.to(torch.bfloat16), (values * 0.25).to(torch.bfloat16)
def _reference_outputs(
case: CompressorMetadataCase,
rope_cos: torch.Tensor,
rope_sin: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
query_start_loc = torch.tensor(case.query_start_loc, dtype=torch.int32)
start_pos = torch.tensor(case.start_pos, dtype=torch.int32)
block_table = torch.tensor(case.block_table, dtype=torch.int32)
if case.expected_slot_mapping is None:
assert case.num_rows is not None
num_rows = case.num_rows
else:
num_rows = len(case.expected_slot_mapping)
prefix = [0]
total_rows = 0
for req_idx, start in enumerate(start_pos.tolist()):
seq_len = int(query_start_loc[req_idx + 1] - query_start_loc[req_idx])
compressed_rows = 0
if start >= 0 and seq_len > 0:
compressed_rows = ((start + seq_len) // case.compress_ratio) - (start // case.compress_ratio)
total_rows += compressed_rows
prefix.append(total_rows)
valid_rows = min(total_rows, num_rows)
ref_cos = torch.empty((num_rows, 1, 1, rope_cos.shape[1]), dtype=rope_cos.dtype)
ref_sin = torch.empty_like(ref_cos)
if case.slot_mapping_format == SLOT_MAPPING_BLOCK_OFFSET:
ref_slot = torch.empty((num_rows, 2), dtype=torch.int32)
else:
ref_slot = torch.empty((num_rows,), dtype=torch.int32)
req_idx = 0
for row in range(num_rows):
valid = row < valid_rows
if valid:
while req_idx < len(start_pos) and prefix[req_idx + 1] <= row:
req_idx += 1
compressed_pos = int(start_pos[req_idx]) // case.compress_ratio + row - prefix[req_idx]
block_id_offset = compressed_pos // KV_BLOCK_SIZE
rope_pos = compressed_pos * case.compress_ratio
if block_id_offset >= block_table.shape[1] or rope_pos >= rope_cos.shape[0]:
valid = False
else:
block_id = int(block_table[req_idx, block_id_offset])
valid = block_id >= 0
if valid:
ref_cos[row, 0, 0].copy_(rope_cos[rope_pos])
ref_sin[row, 0, 0].copy_(rope_sin[rope_pos])
slot_offset = compressed_pos % KV_BLOCK_SIZE
if case.slot_mapping_format == SLOT_MAPPING_BLOCK_OFFSET:
ref_slot[row, 0] = block_id
ref_slot[row, 1] = slot_offset
else:
ref_slot[row] = block_id * KV_BLOCK_SIZE + slot_offset
else:
ref_cos[row, 0, 0].fill_(1)
ref_sin[row, 0, 0].fill_(0)
if case.slot_mapping_format == SLOT_MAPPING_BLOCK_OFFSET:
ref_slot[row, 0] = -1
ref_slot[row, 1] = KV_BLOCK_SIZE - 1
else:
ref_slot[row] = -1
return ref_cos, ref_sin, ref_slot
@pytest.mark.parametrize("case", DSV4_FLASH_CASES, ids=[case.name for case in DSV4_FLASH_CASES])
def test_compressor_metadata_dsv4_flash_real_metadata(case: CompressorMetadataCase):
torch.npu.set_device(0)
rope_cos, rope_sin = _make_rope()
expected_cos, expected_sin, reference_slot = _reference_outputs(case, rope_cos, rope_sin)
if case.expected_slot_mapping is None:
expected_slot = reference_slot
else:
expected_slot = torch.tensor(case.expected_slot_mapping, dtype=torch.int32)
assert torch.equal(reference_slot, expected_slot)
num_rows = expected_slot.shape[0]
rope_cos_npu = rope_cos.npu()
rope_sin_npu = rope_sin.npu()
query_start_loc_npu = torch.tensor(case.query_start_loc, dtype=torch.int32, device="npu")
start_pos_npu = torch.tensor(case.start_pos, dtype=torch.int32, device="npu")
block_table_npu = torch.tensor(case.block_table, dtype=torch.int32, device="npu")
actual_cos, actual_sin, actual_slot = torch.ops._C_ascend.compressor_metadata(
rope_cos_npu,
rope_sin_npu,
query_start_loc_npu,
start_pos_npu,
block_table_npu,
KV_BLOCK_SIZE,
case.slot_mapping_format,
case.compress_ratio,
num_rows,
len(case.start_pos),
)
out_cos = torch.empty_like(actual_cos)
out_sin = torch.empty_like(actual_sin)
out_slot = torch.empty_like(actual_slot)
out_cos, out_sin, out_slot = torch.ops._C_ascend.compressor_metadata_out(
rope_cos_npu,
rope_sin_npu,
query_start_loc_npu,
start_pos_npu,
block_table_npu,
KV_BLOCK_SIZE,
case.slot_mapping_format,
case.compress_ratio,
len(case.start_pos),
out_cos,
out_sin,
out_slot,
)
torch.npu.synchronize()
assert torch.equal(actual_slot.cpu(), expected_slot)
assert torch.equal(actual_cos.cpu(), expected_cos)
assert torch.equal(actual_sin.cpu(), expected_sin)
assert torch.equal(out_slot.cpu(), expected_slot)
assert torch.equal(out_cos.cpu(), expected_cos)
assert torch.equal(out_sin.cpu(), expected_sin)

View File

@@ -0,0 +1,494 @@
"""E2E accuracy test for CopyAndExpandEagleInputs custom operator.
Tests the Ascend C kernel against a CPU golden reference implementation
with parametrized test cases covering various configurations.
"""
import numpy as np
import pytest
import torch
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
SEED = 42
# ---------------------------------------------------------------------------
# Golden reference (CPU, pure Python/NumPy)
# ---------------------------------------------------------------------------
def golden_copy_and_expand(
target_token_ids: np.ndarray,
target_positions: np.ndarray,
next_token_ids: np.ndarray,
query_start_loc: np.ndarray,
query_end_loc: np.ndarray,
padding_token_id: int,
parallel_drafting_token_id: int,
num_padding_slots: int,
shift_input_ids: bool,
):
"""CPU golden reference for CopyAndExpandEagleInputs.
Returns:
(out_input_ids, out_positions, out_is_rejected_token_mask,
out_is_masked_token_mask, out_new_token_indices,
out_hidden_state_mapping)
"""
num_reqs = len(next_token_ids)
# Compute total_draft_tokens
total_draft_tokens = 0
for r in range(num_reqs):
qs = query_start_loc[r]
nqs = query_start_loc[r + 1]
qe = query_end_loc[r]
num_rejected = max(nqs - qe - 1, 0)
if shift_input_ids:
num_valid = max(qe - qs, 0)
else:
num_valid = max(qe - qs + 1, 0)
total_draft_tokens += num_valid + num_padding_slots + num_rejected
out_ids = np.zeros(total_draft_tokens, dtype=np.int32)
out_pos = np.zeros(total_draft_tokens, dtype=np.int32)
out_rej = np.zeros(total_draft_tokens, dtype=np.int8)
out_msk = np.zeros(total_draft_tokens, dtype=np.int8)
out_nti = np.zeros(num_reqs * num_padding_slots, dtype=np.int32)
total_input_tokens = len(target_token_ids)
out_hsm = np.zeros(total_input_tokens, dtype=np.int32)
for r in range(num_reqs):
qs = query_start_loc[r]
nqs = query_start_loc[r + 1]
qe = query_end_loc[r]
num_rejected = max(nqs - qe - 1, 0)
if shift_input_ids:
num_valid = max(qe - qs, 0)
output_start = qs + r * (num_padding_slots - 1)
else:
num_valid = max(qe - qs + 1, 0)
output_start = qs + r * num_padding_slots
start_pos = target_positions[qs]
next_token_id = next_token_ids[r]
# Valid region
if shift_input_ids:
read_start = qs + 1
read_count = min(num_valid, total_input_tokens - read_start)
if read_count < 0:
read_count = 0
for j in range(num_valid):
idx = min(j, read_count - 1) if read_count > 0 else 0
out_ids[output_start + j] = target_token_ids[read_start + idx] if read_count > 0 else 0
out_pos[output_start + j] = start_pos + j
out_rej[output_start + j] = 0
out_msk[output_start + j] = 0
else:
num_input = nqs - qs
for j in range(num_valid):
idx = min(j, num_input - 1)
out_ids[output_start + j] = target_token_ids[qs + idx]
out_pos[output_start + j] = start_pos + j
out_rej[output_start + j] = 0
out_msk[output_start + j] = 0
# Bonus token
out_ids[output_start + num_valid] = next_token_id
out_pos[output_start + num_valid] = start_pos + num_valid
out_rej[output_start + num_valid] = 0
out_msk[output_start + num_valid] = 0
# Parallel draft tokens
for k in range(1, num_padding_slots):
j = num_valid + k
out_ids[output_start + j] = parallel_drafting_token_id
out_pos[output_start + j] = start_pos + j
out_rej[output_start + j] = 0
out_msk[output_start + j] = 1
# Rejected tokens
for k in range(num_rejected):
j = num_valid + num_padding_slots + k
out_ids[output_start + j] = padding_token_id
out_pos[output_start + j] = 0
out_rej[output_start + j] = 1
out_msk[output_start + j] = 0
# New token indices
for k in range(num_padding_slots):
out_nti[r * num_padding_slots + k] = output_start + num_valid + k
# Hidden state mapping (shift_input_ids=true only)
if shift_input_ids:
num_input = nqs - qs
for j in range(num_input):
out_hsm[qs + j] = output_start + j
return out_ids, out_pos, out_rej, out_msk, out_nti, out_hsm
# ---------------------------------------------------------------------------
# NPU operator wrapper
# ---------------------------------------------------------------------------
def npu_op_exec(
target_token_ids,
target_positions,
next_token_ids,
query_start_loc,
query_end_loc,
padding_token_id,
parallel_drafting_token_id,
num_padding_slots,
shift_input_ids,
total_draft_tokens,
):
"""Execute the custom Ascend NPU operator."""
result = torch.ops._C_ascend.npu_copy_and_expand_eagle_inputs(
target_token_ids.to(torch.int32).npu(),
target_positions.to(torch.int32).npu(),
next_token_ids.to(torch.int32).npu(),
query_start_loc.to(torch.int32).npu(),
query_end_loc.to(torch.int32).npu(),
padding_token_id,
parallel_drafting_token_id,
num_padding_slots,
shift_input_ids,
total_draft_tokens,
)
return tuple(t.cpu() for t in result)
# ---------------------------------------------------------------------------
# Test case generator
# ---------------------------------------------------------------------------
def generate_test_case(
rng,
num_reqs,
num_padding_slots,
shift_input_ids,
min_tokens_per_req=2,
max_tokens_per_req=64,
max_rejected_per_req=5,
):
"""Generate a random test case.
Returns dict with all input arrays and expected parameters.
"""
padding_token_id = 0
parallel_drafting_token_id = 100
# Generate per-request token counts
tokens_per_req = rng.integers(min_tokens_per_req, max_tokens_per_req + 1, size=num_reqs)
rejected_per_req = rng.integers(0, max_rejected_per_req + 1, size=num_reqs)
# Build query_start_loc (cumulative)
query_start_loc = np.zeros(num_reqs + 1, dtype=np.int32)
for i in range(num_reqs):
query_start_loc[i + 1] = query_start_loc[i] + tokens_per_req[i] + rejected_per_req[i]
total_input_tokens = int(query_start_loc[num_reqs])
# Build query_end_loc: queryEnd = queryStart + numAccepted - 1
# where numAccepted = tokens_per_req[i]
# For shift=false: numValid = queryEnd - queryStart + 1 = tokens_per_req[i]
# For shift=true: numValid = queryEnd - queryStart = tokens_per_req[i] - 1
query_end_loc = np.zeros(num_reqs, dtype=np.int32)
for i in range(num_reqs):
if shift_input_ids:
query_end_loc[i] = query_start_loc[i] + tokens_per_req[i]
else:
query_end_loc[i] = query_start_loc[i] + tokens_per_req[i] - 1
# Generate input tokens and positions
target_token_ids = rng.integers(1, 50000, size=total_input_tokens, dtype=np.int32)
target_positions = np.zeros(total_input_tokens, dtype=np.int32)
for i in range(num_reqs):
qs = query_start_loc[i]
nqs = query_start_loc[i + 1]
for j in range(nqs - qs):
target_positions[qs + j] = j
next_token_ids = rng.integers(1, 50000, size=num_reqs, dtype=np.int32)
# Compute total_draft_tokens
total_draft_tokens = 0
for r in range(num_reqs):
qs = query_start_loc[r]
nqs = query_start_loc[r + 1]
qe = query_end_loc[r]
num_rejected = max(nqs - qe - 1, 0)
if shift_input_ids:
num_valid = max(qe - qs, 0)
else:
num_valid = max(qe - qs + 1, 0)
total_draft_tokens += num_valid + num_padding_slots + num_rejected
return {
"target_token_ids": target_token_ids,
"target_positions": target_positions,
"next_token_ids": next_token_ids,
"query_start_loc": query_start_loc,
"query_end_loc": query_end_loc,
"padding_token_id": padding_token_id,
"parallel_drafting_token_id": parallel_drafting_token_id,
"num_padding_slots": num_padding_slots,
"shift_input_ids": shift_input_ids,
"total_draft_tokens": total_draft_tokens,
}
# ---------------------------------------------------------------------------
# Parametrized tests
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("num_reqs", [1, 2, 4, 8, 16])
@pytest.mark.parametrize("num_padding_slots", [1, 2, 3, 5])
@pytest.mark.parametrize("shift_input_ids", [False, True])
@pytest.mark.parametrize("seed_offset", [0, 1])
def test_copy_and_expand_eagle_inputs(num_reqs, num_padding_slots, shift_input_ids, seed_offset):
"""Test CopyAndExpandEagleInputs with parametrized configurations."""
rng = np.random.default_rng(SEED + seed_offset)
case = generate_test_case(rng, num_reqs, num_padding_slots, shift_input_ids)
# Golden reference
g_ids, g_pos, g_rej, g_msk, g_nti, g_hsm = golden_copy_and_expand(
case["target_token_ids"],
case["target_positions"],
case["next_token_ids"],
case["query_start_loc"],
case["query_end_loc"],
case["padding_token_id"],
case["parallel_drafting_token_id"],
case["num_padding_slots"],
case["shift_input_ids"],
)
# NPU execution
n_ids, n_pos, n_rej, n_msk, n_nti, n_hsm = npu_op_exec(
torch.from_numpy(case["target_token_ids"]),
torch.from_numpy(case["target_positions"]),
torch.from_numpy(case["next_token_ids"]),
torch.from_numpy(case["query_start_loc"]),
torch.from_numpy(case["query_end_loc"]),
case["padding_token_id"],
case["parallel_drafting_token_id"],
case["num_padding_slots"],
case["shift_input_ids"],
case["total_draft_tokens"],
)
# Convert golden to tensors
g_ids_t = torch.from_numpy(g_ids)
g_pos_t = torch.from_numpy(g_pos)
g_rej_t = torch.from_numpy(g_rej)
g_msk_t = torch.from_numpy(g_msk)
g_nti_t = torch.from_numpy(g_nti)
g_hsm_t = torch.from_numpy(g_hsm)
# Compare outputs
torch.testing.assert_close(n_ids, g_ids_t, atol=0, rtol=0, msg="out_input_ids mismatch")
torch.testing.assert_close(n_pos, g_pos_t, atol=0, rtol=0, msg="out_positions mismatch")
torch.testing.assert_close(n_rej, g_rej_t, atol=0, rtol=0, msg="out_is_rejected_token_mask mismatch")
torch.testing.assert_close(n_msk, g_msk_t, atol=0, rtol=0, msg="out_is_masked_token_mask mismatch")
torch.testing.assert_close(n_nti, g_nti_t, atol=0, rtol=0, msg="out_new_token_indices mismatch")
if shift_input_ids:
torch.testing.assert_close(n_hsm, g_hsm_t, atol=0, rtol=0, msg="out_hidden_state_mapping mismatch")
@pytest.mark.parametrize("num_reqs", [1])
@pytest.mark.parametrize("num_padding_slots", [1])
@pytest.mark.parametrize("shift_input_ids", [False, True])
def test_minimal_case(num_reqs, num_padding_slots, shift_input_ids):
"""Test with minimal input (1 request, 1 padding slot)."""
rng = np.random.default_rng(SEED + 100)
case = generate_test_case(
rng,
num_reqs,
num_padding_slots,
shift_input_ids,
min_tokens_per_req=2,
max_tokens_per_req=3,
max_rejected_per_req=1,
)
g_ids, g_pos, g_rej, g_msk, g_nti, g_hsm = golden_copy_and_expand(
case["target_token_ids"],
case["target_positions"],
case["next_token_ids"],
case["query_start_loc"],
case["query_end_loc"],
case["padding_token_id"],
case["parallel_drafting_token_id"],
case["num_padding_slots"],
case["shift_input_ids"],
)
n_ids, n_pos, n_rej, n_msk, n_nti, n_hsm = npu_op_exec(
torch.from_numpy(case["target_token_ids"]),
torch.from_numpy(case["target_positions"]),
torch.from_numpy(case["next_token_ids"]),
torch.from_numpy(case["query_start_loc"]),
torch.from_numpy(case["query_end_loc"]),
case["padding_token_id"],
case["parallel_drafting_token_id"],
case["num_padding_slots"],
case["shift_input_ids"],
case["total_draft_tokens"],
)
torch.testing.assert_close(n_ids, torch.from_numpy(g_ids), atol=0, rtol=0)
torch.testing.assert_close(n_pos, torch.from_numpy(g_pos), atol=0, rtol=0)
torch.testing.assert_close(n_rej, torch.from_numpy(g_rej), atol=0, rtol=0)
torch.testing.assert_close(n_msk, torch.from_numpy(g_msk), atol=0, rtol=0)
torch.testing.assert_close(n_nti, torch.from_numpy(g_nti), atol=0, rtol=0)
@pytest.mark.parametrize("num_reqs", [3, 7, 13])
def test_large_tokens_per_request(num_reqs):
"""Test with larger token counts per request."""
rng = np.random.default_rng(SEED + 200)
case = generate_test_case(
rng,
num_reqs,
num_padding_slots=3,
shift_input_ids=False,
min_tokens_per_req=100,
max_tokens_per_req=512,
max_rejected_per_req=10,
)
g_ids, g_pos, g_rej, g_msk, g_nti, g_hsm = golden_copy_and_expand(
case["target_token_ids"],
case["target_positions"],
case["next_token_ids"],
case["query_start_loc"],
case["query_end_loc"],
case["padding_token_id"],
case["parallel_drafting_token_id"],
case["num_padding_slots"],
case["shift_input_ids"],
)
n_ids, n_pos, n_rej, n_msk, n_nti, n_hsm = npu_op_exec(
torch.from_numpy(case["target_token_ids"]),
torch.from_numpy(case["target_positions"]),
torch.from_numpy(case["next_token_ids"]),
torch.from_numpy(case["query_start_loc"]),
torch.from_numpy(case["query_end_loc"]),
case["padding_token_id"],
case["parallel_drafting_token_id"],
case["num_padding_slots"],
case["shift_input_ids"],
case["total_draft_tokens"],
)
torch.testing.assert_close(n_ids, torch.from_numpy(g_ids), atol=0, rtol=0)
torch.testing.assert_close(n_pos, torch.from_numpy(g_pos), atol=0, rtol=0)
torch.testing.assert_close(n_rej, torch.from_numpy(g_rej), atol=0, rtol=0)
torch.testing.assert_close(n_msk, torch.from_numpy(g_msk), atol=0, rtol=0)
torch.testing.assert_close(n_nti, torch.from_numpy(g_nti), atol=0, rtol=0)
@pytest.mark.parametrize("num_reqs", [3, 7, 13])
def test_large_tokens_shift_true(num_reqs):
"""Test with larger token counts and shift_input_ids=True."""
rng = np.random.default_rng(SEED + 300)
case = generate_test_case(
rng,
num_reqs,
num_padding_slots=4,
shift_input_ids=True,
min_tokens_per_req=50,
max_tokens_per_req=256,
max_rejected_per_req=8,
)
g_ids, g_pos, g_rej, g_msk, g_nti, g_hsm = golden_copy_and_expand(
case["target_token_ids"],
case["target_positions"],
case["next_token_ids"],
case["query_start_loc"],
case["query_end_loc"],
case["padding_token_id"],
case["parallel_drafting_token_id"],
case["num_padding_slots"],
case["shift_input_ids"],
)
n_ids, n_pos, n_rej, n_msk, n_nti, n_hsm = npu_op_exec(
torch.from_numpy(case["target_token_ids"]),
torch.from_numpy(case["target_positions"]),
torch.from_numpy(case["next_token_ids"]),
torch.from_numpy(case["query_start_loc"]),
torch.from_numpy(case["query_end_loc"]),
case["padding_token_id"],
case["parallel_drafting_token_id"],
case["num_padding_slots"],
case["shift_input_ids"],
case["total_draft_tokens"],
)
torch.testing.assert_close(n_ids, torch.from_numpy(g_ids), atol=0, rtol=0)
torch.testing.assert_close(n_pos, torch.from_numpy(g_pos), atol=0, rtol=0)
torch.testing.assert_close(n_rej, torch.from_numpy(g_rej), atol=0, rtol=0)
torch.testing.assert_close(n_msk, torch.from_numpy(g_msk), atol=0, rtol=0)
torch.testing.assert_close(n_nti, torch.from_numpy(g_nti), atol=0, rtol=0)
torch.testing.assert_close(n_hsm, torch.from_numpy(g_hsm), atol=0, rtol=0)
@pytest.mark.parametrize("num_reqs", [1, 4, 8])
def test_no_rejected_tokens(num_reqs):
"""Test cases with zero rejected tokens."""
rng = np.random.default_rng(SEED + 400)
case = generate_test_case(
rng,
num_reqs,
num_padding_slots=2,
shift_input_ids=False,
min_tokens_per_req=5,
max_tokens_per_req=20,
max_rejected_per_req=0,
)
g_ids, g_pos, g_rej, g_msk, g_nti, g_hsm = golden_copy_and_expand(
case["target_token_ids"],
case["target_positions"],
case["next_token_ids"],
case["query_start_loc"],
case["query_end_loc"],
case["padding_token_id"],
case["parallel_drafting_token_id"],
case["num_padding_slots"],
case["shift_input_ids"],
)
n_ids, n_pos, n_rej, n_msk, n_nti, n_hsm = npu_op_exec(
torch.from_numpy(case["target_token_ids"]),
torch.from_numpy(case["target_positions"]),
torch.from_numpy(case["next_token_ids"]),
torch.from_numpy(case["query_start_loc"]),
torch.from_numpy(case["query_end_loc"]),
case["padding_token_id"],
case["parallel_drafting_token_id"],
case["num_padding_slots"],
case["shift_input_ids"],
case["total_draft_tokens"],
)
torch.testing.assert_close(n_ids, torch.from_numpy(g_ids), atol=0, rtol=0)
torch.testing.assert_close(n_pos, torch.from_numpy(g_pos), atol=0, rtol=0)
torch.testing.assert_close(n_rej, torch.from_numpy(g_rej), atol=0, rtol=0)
torch.testing.assert_close(n_msk, torch.from_numpy(g_msk), atol=0, rtol=0)
torch.testing.assert_close(n_nti, torch.from_numpy(g_nti), atol=0, rtol=0)

View File

@@ -0,0 +1,105 @@
import gc
import pytest
import torch
import torch.nn.functional as F
import torch_npu
from vllm_ascend.utils import enable_custom_op
# enable internal format
torch_npu.npu.config.allow_internal_format = True
# enable vllm-ascend custom ops
enable_custom_op()
def _shared_dequant_swiglu_quant(
hidden_states: torch.Tensor,
weight_scale: torch.Tensor,
activation_scale: torch.Tensor,
swiglu_limit: int | float,
output_dtype: torch.dtype,
) -> tuple[torch.Tensor, torch.Tensor]:
if hidden_states.shape[0] == 0:
output_shape = hidden_states.shape[:-1] + (hidden_states.shape[-1] // 2,)
return (
hidden_states.new_empty(output_shape, dtype=torch.int8),
torch.empty(hidden_states.shape[:-1], dtype=torch.float32, device=hidden_states.device),
)
weight_scale = weight_scale.to(torch.float32).reshape((1,) * (hidden_states.dim() - 1) + (-1,))
activation_scale = activation_scale.to(torch.float32).reshape(hidden_states.shape[:-1] + (1,))
gate_up = hidden_states.to(torch.float32) * weight_scale * activation_scale
half = gate_up.shape[-1] // 2
limit = float(swiglu_limit)
gate = gate_up[..., :half]
up = gate_up[..., half:]
# Skip clamp when limit == 0 (treated as "no clamp")
if limit > 0.0:
gate = torch.clamp(gate, max=limit)
up = torch.clamp(up, min=-limit, max=limit)
swiglu = F.silu(gate) * up
if swiglu.dtype not in (torch.float16, torch.bfloat16):
swiglu = swiglu.to(output_dtype if output_dtype in (torch.float16, torch.bfloat16) else torch.bfloat16)
return torch_npu.npu_dynamic_quant(swiglu)
_REPRO_CASES = [
([4608, 2048], 0.0, "large_2048_aligned"),
([2, 192], 0.0, "small_192_misaligned"),
([4, 192], 0.0, "small_192_misaligned_4rows"),
([8, 384], 0.0, "small_384_aligned"),
([1, 256], 0.0, "single_row_256_aligned"),
]
@torch.inference_mode()
@pytest.mark.parametrize("x_shape,clamp_limit,desc", _REPRO_CASES, ids=[c[2] for c in _REPRO_CASES])
def test_npu_dequant_swiglu_quant_with_limit(x_shape, clamp_limit, desc):
# Use values with non-trivial abs() so reduce-max is meaningful
x = torch.randint(-100, 100, x_shape, dtype=torch.int32)
weight_scale = torch.randn(x_shape[1], dtype=torch.float32) * 0.1
activate_scale = torch.randn((x_shape[0], 1), dtype=torch.float32) * 0.5
x = x.npu()
weight_scale = weight_scale.npu()
activate_scale = activate_scale.npu()
# 1. Golden reference (pure PyTorch + npu_dynamic_quant)
output_golden, output_scale_golden = _shared_dequant_swiglu_quant(
x,
weight_scale,
activate_scale,
clamp_limit,
torch.bfloat16,
)
# 2. Fused op (NPUGraph, same as production code)
graph = torch.npu.NPUGraph()
with torch.npu.graph(graph, capture_error_mode="thread_local", auto_dispatch_capture=True):
output, output_scale = torch.ops._C_ascend.npu_dequant_swiglu_quant(
x=x,
weight_scale=weight_scale,
activation_scale=activate_scale,
bias=None,
quant_scale=None,
quant_offset=None,
group_index=None,
activate_left=True,
quant_mode=1,
swiglu_mode=1,
clamp_limit=clamp_limit,
glu_alpha=1.0,
glu_bias=0.0,
)
graph.replay()
# int8 quantization output: atol=1 covers the rounding error (max_abs=1 in all cases)
torch.testing.assert_close(output.cpu(), output_golden.cpu(), atol=1, rtol=0.1)
# Dynamic quant scale: relax tolerance to cover both aligned and non-64-aligned outDimy
torch.testing.assert_close(output_scale.cpu(), output_scale_golden.cpu(), atol=1e-4, rtol=5e-3)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,342 @@
# SPDX-License-Identifier: Apache-2.0
"""E2E correctness test for AscendC fused_gdn_gating kernel.
Validates torch.ops._C_ascend.npu_fused_gdn_gating against a CPU golden
reference across num_heads / batch / dtype combinations.
Prerequisite: the AscendC kernel must be compiled and installed via
bash csrc/build_aclnn.sh <ROOT_DIR> <SOC_VERSION>
Run:
pytest tests/e2e/nightly/single_node/ops/singlecard_ops/test_fused_gdn_gating.py -v
"""
import gc
import pytest
import torch
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
SEED = 42
NUM_HEADS_VALUES = [4, 6, 8, 12, 16, 24, 32, 48, 64, 128]
BATCH_SIZES = [1, 7, 37, 128, 512, 4096, 16384]
DTYPES = [torch.bfloat16, torch.float16]
PARAM_DTYPES = [torch.float32, torch.bfloat16, torch.float16]
DTYPE_COMBINATIONS = [(dtype, param_dtype) for dtype in DTYPES for param_dtype in PARAM_DTYPES]
# ---------------------------------------------------------------------------
# Golden reference (CPU, pure PyTorch)
# ---------------------------------------------------------------------------
def _golden_fused_gdn_gating(
A_log: torch.Tensor,
a: torch.Tensor,
b: torch.Tensor,
dt_bias: torch.Tensor,
beta: float = 1.0,
threshold: float = 20.0,
) -> tuple[torch.Tensor, torch.Tensor]:
"""CPU golden reference for fused_gdn_gating.
Uses the same softplus threshold semantics as the Triton kernel:
where(beta * x <= threshold, log(1 + exp(beta * x)) / beta, x)
Returns:
g: [1, batch, num_heads], fp32.
beta_output: [1, batch, num_heads], original dtype.
"""
batch, num_heads = a.shape
compute_dtype = torch.float32
A_log_f = A_log.to(compute_dtype)
a_f = a.to(compute_dtype)
b_f = b.to(compute_dtype)
dt_bias_f = dt_bias.to(compute_dtype)
A_log_expanded = A_log_f.unsqueeze(0).expand(batch, -1)
dt_bias_expanded = dt_bias_f.unsqueeze(0).expand(batch, -1)
x = a_f + dt_bias_expanded
beta_x = beta * x
softplus_o = torch.where(
beta_x <= threshold,
torch.log1p(torch.exp(beta_x)) / beta,
x,
)
g = -torch.exp(A_log_expanded) * softplus_o
g = g.unsqueeze(0)
beta_output = torch.sigmoid(b_f).to(b.dtype)
beta_output = beta_output.unsqueeze(0)
return g, beta_output
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_inputs(
num_heads: int,
batch: int,
dtype: torch.dtype,
param_dtype: torch.dtype = torch.float32,
seed: int = SEED,
):
"""Build random tensors on CPU for both golden and NPU execution."""
torch.manual_seed(seed)
A_log = torch.randn(num_heads, dtype=param_dtype)
dt_bias = torch.randn(num_heads, dtype=param_dtype)
a = torch.randn(batch, num_heads, dtype=dtype)
b = torch.randn(batch, num_heads, dtype=dtype)
return A_log, a, b, dt_bias
def _force_softplus_threshold_cases(
a: torch.Tensor,
dt_bias: torch.Tensor,
beta: float,
threshold: float,
) -> None:
"""Force beta * (a + dt_bias) to cover threshold and non-threshold paths."""
if a.shape[0] < 4 or a.shape[1] < 4:
return
dt_bias[:4] = 0
boundary = threshold / beta
a[0, 0] = boundary + 2.0 # linear branch
a[1, 1] = boundary # softplus branch at equality
a[2, 2] = boundary - 0.5 # softplus branch below threshold
a[3, 3] = -boundary - 2.0 # negative softplus input
def _npu_op_exec(
A_log: torch.Tensor,
a: torch.Tensor,
b: torch.Tensor,
dt_bias: torch.Tensor,
beta: float = 1.0,
threshold: float = 20.0,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Execute the AscendC operator on NPU and return CPU tensors."""
# Ensure contiguity for the NPU operator.
if not A_log.is_contiguous():
A_log = A_log.contiguous()
if not dt_bias.is_contiguous():
dt_bias = dt_bias.contiguous()
g, beta_output = torch.ops._C_ascend.npu_fused_gdn_gating(
A_log.npu(),
a.npu(),
b.npu(),
dt_bias.npu(),
float(beta),
float(threshold),
)
return g.cpu(), beta_output.cpu()
def _assert_close(actual_g, actual_beta, ref_g, ref_beta, rtol=3e-3, atol=1e-2):
torch.testing.assert_close(
actual_g.to(torch.float32),
ref_g.to(torch.float32),
rtol=rtol,
atol=atol,
equal_nan=True,
)
torch.testing.assert_close(
actual_beta.to(torch.float32),
ref_beta.to(torch.float32),
rtol=rtol,
atol=atol,
equal_nan=True,
)
# ---------------------------------------------------------------------------
# Tests: core correctness
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("num_heads", NUM_HEADS_VALUES)
@pytest.mark.parametrize("batch", BATCH_SIZES)
@pytest.mark.parametrize("dtype", DTYPES)
def test_fused_gdn_gating_vs_reference(num_heads, batch, dtype):
A_log, a, b, dt_bias = _make_inputs(num_heads, batch, dtype)
ref_g, ref_beta = _golden_fused_gdn_gating(A_log, a, b, dt_bias)
npu_g, npu_beta = _npu_op_exec(A_log, a, b, dt_bias)
_assert_close(npu_g, npu_beta, ref_g, ref_beta)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.parametrize("num_heads", [16, 32, 64])
@pytest.mark.parametrize("batch", [1, 37])
def test_fused_gdn_gating_non_default_params(num_heads, batch):
A_log, a, b, dt_bias = _make_inputs(num_heads, batch, torch.bfloat16)
_force_softplus_threshold_cases(a, dt_bias, beta=0.5, threshold=1.0)
ref_g, ref_beta = _golden_fused_gdn_gating(
A_log,
a,
b,
dt_bias,
beta=0.5,
threshold=1.0,
)
npu_g, npu_beta = _npu_op_exec(
A_log,
a,
b,
dt_bias,
beta=0.5,
threshold=1.0,
)
_assert_close(npu_g, npu_beta, ref_g, ref_beta)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.parametrize(("dtype", "param_dtype"), DTYPE_COMBINATIONS)
def test_fused_gdn_gating_dtype_matrix(dtype, param_dtype):
A_log, a, b, dt_bias = _make_inputs(
32,
37,
dtype,
param_dtype=param_dtype,
)
_force_softplus_threshold_cases(a, dt_bias, beta=1.0, threshold=2.0)
ref_g, ref_beta = _golden_fused_gdn_gating(
A_log,
a,
b,
dt_bias,
threshold=2.0,
)
npu_g, npu_beta = _npu_op_exec(
A_log,
a,
b,
dt_bias,
threshold=2.0,
)
_assert_close(npu_g, npu_beta, ref_g, ref_beta)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
def test_fused_gdn_gating_output_shapes():
A_log, a, b, dt_bias = _make_inputs(32, 17, torch.bfloat16)
npu_g, npu_beta = _npu_op_exec(A_log, a, b, dt_bias)
assert npu_g.shape == (1, 17, 32), f"unexpected g shape: {npu_g.shape}"
assert npu_g.dtype == torch.float32, f"unexpected g dtype: {npu_g.dtype}"
assert npu_beta.shape == (1, 17, 32), f"unexpected beta shape: {npu_beta.shape}"
assert npu_beta.dtype == torch.bfloat16, f"unexpected beta dtype: {npu_beta.dtype}"
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
# ---------------------------------------------------------------------------
# Tests: multi-row processing and Bulk DMA
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("num_heads", [8, 16, 32, 64])
@pytest.mark.parametrize("batch", [64, 256, 1024, 4096])
def test_fused_gdn_gating_large_batch_multi_row(num_heads, batch):
A_log, a, b, dt_bias = _make_inputs(num_heads, batch, torch.bfloat16)
ref_g, ref_beta = _golden_fused_gdn_gating(A_log, a, b, dt_bias)
npu_g, npu_beta = _npu_op_exec(A_log, a, b, dt_bias)
_assert_close(npu_g, npu_beta, ref_g, ref_beta)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.parametrize("num_heads", [16, 32, 48])
def test_fused_gdn_gating_bulk_dma_alignment(num_heads):
"""Bulk DMA fast path for nh % 16 == 0."""
batch = 512
A_log, a, b, dt_bias = _make_inputs(num_heads, batch, torch.bfloat16)
ref_g, ref_beta = _golden_fused_gdn_gating(A_log, a, b, dt_bias)
npu_g, npu_beta = _npu_op_exec(A_log, a, b, dt_bias)
_assert_close(npu_g, npu_beta, ref_g, ref_beta)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.parametrize("num_heads", [6, 12, 24])
def test_fused_gdn_gating_non_bulk_dma_fallback(num_heads):
"""Per-row fallback for head counts that don't satisfy Bulk DMA alignment."""
batch = 256
A_log, a, b, dt_bias = _make_inputs(num_heads, batch, torch.bfloat16)
ref_g, ref_beta = _golden_fused_gdn_gating(A_log, a, b, dt_bias)
npu_g, npu_beta = _npu_op_exec(A_log, a, b, dt_bias)
_assert_close(npu_g, npu_beta, ref_g, ref_beta)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
def test_fused_gdn_gating_small_batch_optimization():
"""Small batch < rows_per_iter; adaptive UB budgeting."""
num_heads = 32
batch = 8
A_log, a, b, dt_bias = _make_inputs(num_heads, batch, torch.bfloat16)
ref_g, ref_beta = _golden_fused_gdn_gating(A_log, a, b, dt_bias)
npu_g, npu_beta = _npu_op_exec(A_log, a, b, dt_bias)
_assert_close(npu_g, npu_beta, ref_g, ref_beta)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
def test_fused_gdn_gating_extreme_large_batch():
"""Extreme large batch stress test."""
num_heads = 32
batch = 65536
A_log, a, b, dt_bias = _make_inputs(num_heads, batch, torch.bfloat16)
ref_g, ref_beta = _golden_fused_gdn_gating(A_log, a, b, dt_bias)
npu_g, npu_beta = _npu_op_exec(A_log, a, b, dt_bias)
_assert_close(npu_g, npu_beta, ref_g, ref_beta)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,361 @@
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# Copyright 2023 The vLLM team.
#
# 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.
# SPDX-License-Identifier: Apache-2.0
# This file is a part of the vllm-ascend project.
# Adapted from vllm/tests/kernels/test_moe.py
"""Tests for the MOE layers.
Run `pytest tests/ops/test_fused_moe.py`.
"""
import gc
from unittest.mock import MagicMock, patch
import pytest
import torch
import torch.nn.functional as F
import torch_npu
from vllm_ascend.ops.fused_moe.experts_selector import check_npu_moe_gating_top_k, select_experts
from vllm_ascend.ops.fused_moe.moe_mlp import unified_apply_mlp
from vllm_ascend.ops.fused_moe.moe_runtime_args import (
MoEQuantParams,
MoERoutingParams,
MoETokenDispatchInput,
build_fused_experts_input,
build_mlp_compute_input,
)
from vllm_ascend.ops.fused_moe.token_dispatcher import TokenDispatcherWithAllGather
from vllm_ascend.quantization.quant_type import QuantType
NUM_EXPERTS = [8, 64]
EP_SIZE = [1]
TOP_KS = [2, 6]
DEVICE = ["npu"]
class SiluAndMul:
"""SwiGLU activation function: silu(x[:d]) * x[d:] where d = x.shape[-1] // 2"""
def __call__(self, x: torch.Tensor) -> torch.Tensor:
d = x.shape[-1] // 2
return F.silu(x[..., :d]) * x[..., d:]
def apply_mlp(
hidden_states: torch.Tensor,
w1: torch.Tensor,
w2: torch.Tensor,
group_list: torch.Tensor,
group_list_type: int = 1,
) -> torch.Tensor:
w1 = w1.transpose(1, 2)
hidden_states = torch_npu.npu_grouped_matmul(
x=[hidden_states],
weight=[w1],
split_item=2,
group_list_type=group_list_type,
group_type=0,
group_list=group_list,
)[0]
hidden_states = torch_npu.npu_swiglu(hidden_states)
w2 = w2.transpose(1, 2)
hidden_states = torch_npu.npu_grouped_matmul(
x=[hidden_states],
weight=[w2],
split_item=2,
group_list_type=group_list_type,
group_type=0,
group_list=group_list,
)[0]
return hidden_states
def torch_moe(a, w1, w2, topk_weights, topk_ids, topk, expert_map):
B, D = a.shape
a = a.view(B, -1, D).repeat(1, topk, 1).reshape(-1, D)
out = torch.zeros(B * topk, w2.shape[1], dtype=a.dtype, device=a.device)
topk_weights = topk_weights.view(-1)
topk_ids = topk_ids.view(-1)
if expert_map is not None:
topk_ids = expert_map[topk_ids]
for i in range(w1.shape[0]):
mask = topk_ids == i
if mask.sum():
out[mask] = SiluAndMul()(a[mask] @ w1[i].transpose(0, 1)) @ w2[i].transpose(0, 1)
return (out.view(B, -1, w2.shape[1]) * topk_weights.view(B, -1, 1).to(out.dtype)).sum(dim=1)
@pytest.mark.skip("Probabilistic failure, need zengiant after fix")
@pytest.mark.parametrize("m", [1, 1024 * 128])
@pytest.mark.parametrize("n", [128, 2048])
@pytest.mark.parametrize("k", [128, 1024])
@pytest.mark.parametrize("e", NUM_EXPERTS)
@pytest.mark.parametrize("topk", TOP_KS)
@pytest.mark.parametrize("ep_size", EP_SIZE)
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
@pytest.mark.parametrize("device", DEVICE)
def test_token_dispatcher_with_all_gather(
m: int,
n: int,
k: int,
e: int,
topk: int,
ep_size: int,
dtype: torch.dtype,
device: str,
):
a = torch.randn((m, k), device=device, dtype=dtype) / 10
w1 = torch.randn((e, 2 * n, k), device=device, dtype=dtype) / 10
w2 = torch.randn((e, k, n), device=device, dtype=dtype) / 10
score = torch.randn((m, e), device=device, dtype=dtype)
expert_map = None
local_e = e
w1_local = w1
w2_local = w2
score = torch.softmax(score, dim=-1, dtype=dtype)
topk_weights, topk_ids = torch.topk(score, topk)
topk_ids = topk_ids.to(torch.int32)
dispatcher_kwargs = {
"num_experts": e,
"top_k": topk,
"num_local_experts": local_e,
}
dispatcher = TokenDispatcherWithAllGather(**dispatcher_kwargs)
apply_router_weight_on_input = False
token_dispatch_output = dispatcher.token_dispatch(
token_dispatch_input=MoETokenDispatchInput(
hidden_states=a,
topk_weights=topk_weights,
topk_ids=topk_ids,
routing=MoERoutingParams(
expert_map=expert_map,
global_redundant_expert_num=0,
mc2_mask=None,
apply_router_weight_on_input=apply_router_weight_on_input,
),
quant=MoEQuantParams(quant_type=QuantType.NONE),
)
)
sorted_hidden_states = token_dispatch_output.hidden_states
group_list = token_dispatch_output.group_list
group_list_type = token_dispatch_output.group_list_type
combine_metadata = token_dispatch_output.combine_metadata
expert_output = apply_mlp(
hidden_states=sorted_hidden_states,
w1=w1_local,
w2=w2_local,
group_list=group_list,
group_list_type=group_list_type,
)
combined_output = dispatcher.token_combine(
hidden_states=expert_output, combine_metadata=combine_metadata, bias=None
)
torch_output = torch_moe(a, w1, w2, topk_weights, topk_ids, topk, expert_map)
torch.testing.assert_close(combined_output, torch_output, atol=4e-2, rtol=1)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.skip("Probabilistic failure, need zengiant after fix")
@pytest.mark.parametrize("m", [1, 33, 64])
@pytest.mark.parametrize("n", [128, 1024, 2048])
@pytest.mark.parametrize("k", [128, 511, 1024])
@pytest.mark.parametrize("e", NUM_EXPERTS)
@pytest.mark.parametrize("topk", TOP_KS)
@pytest.mark.parametrize("ep_size", EP_SIZE)
@pytest.mark.parametrize("dtype", [torch.bfloat16])
@pytest.mark.parametrize("device", DEVICE)
def test_token_dispatcher_with_all_gather_quant(
m: int,
n: int,
k: int,
e: int,
topk: int,
ep_size: int,
dtype: torch.dtype,
device: str,
):
a = torch.randn((m, k), device=device, dtype=dtype) / 10
w1 = torch.randn((e, k, 2 * n), device=device, dtype=torch.int8)
w1_scale = torch.empty((e, 2 * n), device=device, dtype=dtype)
w2 = torch.randn((e, n, k), device=device, dtype=torch.int8)
w2_scale = torch.empty((e, k), device=device, dtype=dtype)
score = torch.randn((m, e), device=device, dtype=dtype)
expert_map = None
local_e = e
score = torch.softmax(score, dim=-1, dtype=dtype)
topk_weights, topk_ids = torch.topk(score, topk)
topk_ids = topk_ids.to(torch.int32)
dispatcher_kwargs = {
"num_experts": e,
"top_k": topk,
"num_local_experts": local_e,
}
dispatcher = TokenDispatcherWithAllGather(**dispatcher_kwargs)
apply_router_weight_on_input = False
token_dispatch_output = dispatcher.token_dispatch(
token_dispatch_input=MoETokenDispatchInput(
hidden_states=a,
topk_weights=topk_weights,
topk_ids=topk_ids,
routing=MoERoutingParams(
expert_map=expert_map,
global_redundant_expert_num=0,
mc2_mask=None,
apply_router_weight_on_input=apply_router_weight_on_input,
),
quant=MoEQuantParams(quant_type=QuantType.W8A8),
)
)
combine_metadata = token_dispatch_output.combine_metadata
mlp_compute_input = build_mlp_compute_input(
fused_experts_input=build_fused_experts_input(
hidden_states=a,
topk_weights=topk_weights,
topk_ids=topk_ids,
w1=w1,
w2=w2,
quant_type=QuantType.W8A8,
dynamic_eplb=False,
expert_map=expert_map,
w1_scale=w1_scale,
w2_scale=w2_scale,
),
token_dispatch_output=token_dispatch_output,
use_fusion_ops=False,
)
expert_output = unified_apply_mlp(mlp_compute_input=mlp_compute_input)
combined_output = dispatcher.token_combine(
hidden_states=expert_output, combine_metadata=combine_metadata, bias=None
)
assert combined_output.shape == (m, k)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.skip("Probabilistic failure, need zengiant after fix")
@pytest.mark.parametrize("m", [1, 33, 64])
@pytest.mark.parametrize("n", [128, 2048])
@pytest.mark.parametrize("e", NUM_EXPERTS)
@pytest.mark.parametrize("topk", TOP_KS)
@pytest.mark.parametrize("scoring_func", ["softmax", "sigmoid"])
@pytest.mark.parametrize("use_grouped_topk", [True, False])
@pytest.mark.parametrize("renormalize", [True, False])
@pytest.mark.parametrize("with_e_correction", [True, False])
@pytest.mark.parametrize("custom_routing", [True, False])
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
@pytest.mark.parametrize("device", DEVICE)
def test_select_experts(
m: int,
n: int,
e: int,
topk: int,
scoring_func: str,
use_grouped_topk: bool,
renormalize: bool,
with_e_correction: bool,
custom_routing: bool,
dtype: torch.dtype,
device: str,
):
topk_group = 4 if use_grouped_topk else None
num_expert_group = e // 4 if use_grouped_topk else None
hidden_states = torch.randn(m, n, device=device, dtype=dtype)
router_logits = torch.randn(m, e, device=device, dtype=dtype)
e_score_correction_bias = torch.randn(e, device=device, dtype=dtype) if with_e_correction else None
custom_routing_function = None
if custom_routing:
custom_routing_function = MagicMock()
mock_weights = torch.randn(m, topk, device=device, dtype=dtype)
mock_ids = torch.randint(0, e, (m, topk), device=device, dtype=torch.int32)
custom_routing_function.return_value = (mock_weights, mock_ids)
with (
patch("vllm_ascend.ops.fused_moe.experts_selector._native_grouped_topk") as mock_native_grouped_topk,
patch("vllm_ascend.ops.fused_moe.experts_selector.get_weight_prefetch_method", return_value=MagicMock()),
):
mock_native_grouped_topk.side_effect = lambda x, num_groups, k: torch.randn_like(x)
topk_weights, topk_ids = select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=topk,
use_grouped_topk=use_grouped_topk,
renormalize=renormalize,
topk_group=topk_group,
num_expert_group=num_expert_group,
custom_routing_function=custom_routing_function,
scoring_func=scoring_func,
e_score_correction_bias=e_score_correction_bias,
)
call_moe_gatingtopk = check_npu_moe_gating_top_k(
hidden_states, topk, renormalize, topk_group, num_expert_group, scoring_func, custom_routing_function
)
if not call_moe_gatingtopk and use_grouped_topk:
mock_native_grouped_topk.assert_called_once()
else:
mock_native_grouped_topk.assert_not_called()
assert topk_weights.shape == (m, topk)
assert topk_ids.shape == (m, topk)
assert topk_ids.dtype == torch.int32
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.skip("Probabilistic failure, need zengiant after fix")
@pytest.mark.parametrize("device", DEVICE)
def test_select_experts_invalid_scoring_func(device: str):
with (
patch("vllm_ascend.ops.fused_moe.experts_selector.get_weight_prefetch_method", return_value=MagicMock()),
pytest.raises(ValueError, match="Unsupported scoring function: invalid"),
):
select_experts(
hidden_states=torch.randn(1, 128, device=device),
router_logits=torch.randn(1, 8, device=device),
top_k=2,
use_grouped_topk=False,
renormalize=False,
scoring_func="invalid",
)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,40 @@
import gc
import pytest
import torch
import torch_npu
@pytest.mark.parametrize(
"B",
[1, 16, 64, 128, 32768],
)
@pytest.mark.parametrize(
"D",
[8, 16, 32, 64, 128],
)
@pytest.mark.parametrize(
"top_k",
[1, 2, 4, 8],
)
@pytest.mark.parametrize(
"dtype, atol, rtol",
[
(torch.float16, 1e-3, 1e-3),
(torch.bfloat16, 1e-3, 1e-3),
],
)
def test_quant_fpx_linear(B: int, D: int, top_k: int, dtype, atol, rtol):
x = torch.rand((B, D), dtype=dtype).to("npu")
# finished = torch.randint(1, size=(B,), dtype=torch.bool).to("npu")
finished = None
y, expert_idx, row_idx = torch_npu.npu_moe_gating_top_k_softmax(x, finished, k=top_k)
topk_weights = x.softmax(dim=-1)
topk_weights, topk_ids = topk_weights.topk(top_k, dim=-1)
topk_ids = topk_ids.to(torch.int32)
torch.allclose(y, topk_weights, atol=atol, rtol=rtol)
torch.allclose(expert_idx, topk_ids, atol=atol, rtol=rtol)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,140 @@
import gc
import torch
import torch_npu
from vllm_ascend.utils import enable_custom_op
# enable internal format
torch_npu.npu.config.allow_internal_format = True
# enable vllm-ascend custom ops
enable_custom_op()
def gmm_swiglu_quant(
x: torch.Tensor, weight: torch.Tensor, perChannelScale: torch.Tensor, perTokenScale: torch.Tensor, m: int
):
"""
Perform quantized GMM (Grouped Matrix Multiplication) operation with SwiGLU activation function.
Parameters:
x (torch.Tensor): Input tensor with shape (m, k).
weight (torch.Tensor): Weight tensor with shape (k, n).
perChannelScale (torch.Tensor): Per-channel scaling factor with shape (n,).
perTokenScale (torch.Tensor): Per-token scaling factor with shape (m,).
m (int): Number of tokens (rows of x).
Returns:
quantOutput (torch.Tensor): Quantized output tensor with shape (m, k // 2).
quantScaleOutput (torch.Tensor): Quantization scaling factor with shape (m,).
"""
# Perform matrix multiplication with int32 precision
c_temp1 = torch.matmul(x.to(torch.int32), weight.to(torch.int32))
c_temp1 = c_temp1.to(torch.float32) # Convert back to float32 for scaling
# Apply per-channel and per-token scaling
c_temp2 = torch.mul(c_temp1, perChannelScale)
c_temp3 = torch.mul(c_temp2, perTokenScale.reshape(m, 1))
# Split the result into two parts to apply SwiGLU activation function
c_temp4, gate = c_temp3.chunk(2, dim=-1)
c_temp5 = c_temp4 * torch.sigmoid(c_temp4) # SwiGLU activation
c_temp6 = c_temp5 * gate # Element-wise multiplication with gating values
# Quantize the output
max = torch.max(torch.abs(c_temp6), -1).values # Find maximum absolute value to calculate scaling factor
quantScaleOutput = 127 / max # Calculate quantization scaling factor
quantOutput = torch.round(c_temp6 * quantScaleOutput.reshape(m, 1)).to(torch.int8) # Quantize to int8
quantScaleOutput = 1 / quantScaleOutput # Inverse quantization scaling factor for subsequent dequantization
return quantOutput, quantScaleOutput
def process_groups(
x: torch.Tensor,
weight: torch.Tensor,
perChannelScale: torch.Tensor,
perTokenScale: torch.Tensor,
groupList: torch.Tensor,
):
"""
Process input data by groups and call GMM_Swiglu_quant function for quantized computation.
Parameters:
x (torch.Tensor): Input tensor with shape (M, K).
weight (torch.Tensor): List of weight tensors, each with shape (E, K, N).
perChannelScale (torch.Tensor): List of per-channel scaling factors, each with shape (E, N).
perTokenScale (torch.Tensor): Per-token scaling factor with shape (M,).
groupList (list): List defining the number of tokens in each group.
Returns:
quantOutput (torch.Tensor): Quantized output tensor with shape (M, N // 2).
quantScaleOutput (torch.Tensor): Quantization scaling factor with shape (M,).
"""
M, N = x.shape[0], weight.shape[2] # Get the shape of the input tensor
quantOutput = torch.zeros(M, N // 2).to(torch.int8) # Initialize quantized output tensor
quantScaleOutput = torch.zeros(M).to(torch.float32) # Initialize quantization scaling factor tensor
start_idx = 0 # Starting index
preV = 0 # Number of tokens in the previous group
groupList = groupList.tolist()
# Iterate through groupList to process data by groups
for i, v in enumerate(groupList):
currV = v
tempV = currV - preV # Calculate number of tokens in the current group
preV = currV # Update number of tokens in the previous group
if tempV > 0:
# Call GMM_Swiglu_quant to process the current group
quantOutput[start_idx : start_idx + tempV], quantScaleOutput[start_idx : start_idx + tempV] = (
gmm_swiglu_quant(
x[start_idx : start_idx + tempV],
weight[i],
perChannelScale[i],
perTokenScale[start_idx : start_idx + tempV],
tempV,
)
)
start_idx += tempV # Update starting index to process the next group
return quantOutput, quantScaleOutput
@torch.inference_mode()
def test_gmm_swiglu_quant_weight_nz_tensor_list():
M, K, E, N = 8192, 7168, 4, 4096
# x (M, K) - int8
x = torch.randint(-128, 127, (M, K), dtype=torch.int8)
# weight (E, N, K) - int8
weight = torch.randint(-128, 127, size=(E, K, N), dtype=torch.int8)
# weight_scale (E, N) - float32
weight_scale = torch.rand(E, N) * 0.9 + 0.1 # uniform(0.1, 1.0)
weight_scale = weight_scale.to(torch.float32)
weight_nz_npu = []
weight_scale_npu = []
for i in range(E):
weight_nz_npu.append(torch_npu.npu_format_cast(weight[i].npu(), 29))
weight_scale_npu.append(weight_scale[i].npu())
# x_scale (M,) - float32
x_scale = torch.rand(M) * 0.9 + 0.1 # uniform(0.1, 1.0)
x_scale = x_scale.to(torch.float32)
group_list = torch.tensor([2048, 4096, 6144, 8192], dtype=torch.int64)
output_cpu, output_scale_cpu = process_groups(x, weight, weight_scale, x_scale, group_list)
output_npu, output_scale_npu, _ = torch.ops._C_ascend.grouped_matmul_swiglu_quant_weight_nz_tensor_list(
x.npu(), weight_nz_npu, weight_scale_npu, x_scale.npu(), group_list.npu()
)
output_npu_valid = output_npu[: group_list[-1], :]
output_scale_npu_valid = output_scale_npu[: group_list[-1]]
torch.testing.assert_close(output_npu_valid.cpu(), output_cpu, atol=1, rtol=2**-13)
torch.testing.assert_close(output_scale_npu_valid.cpu(), output_scale_cpu, atol=1e-9, rtol=1e-6)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,161 @@
import gc
import numpy as np
import torch
import torch_npu
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
def x_int8_to_x_int4(x: torch.Tensor):
m, k = x.shape
x_high_4bit = torch.floor(x.to(torch.float16) // 16).to(torch.int8)
x_low_4bit = torch.bitwise_and(x.view(torch.int16), 0x0F0F).view(torch.int8) - 8
x_int4 = torch.empty((2 * m, k), dtype=torch.int8)
x_int4[::2, :] = x_high_4bit
x_int4[1::2, :] = x_low_4bit
return x_int4
def custom_mm(x: torch.Tensor, weight: torch.Tensor, weight_scale: torch.Tensor, m: int):
"""
Performing Quantized GMM (General Matrix Multiplication) Operation
Parameters:
x (torch.Tensor): Input tensor with shape (m, k).
weight (torch.Tensor): Weight tensor with shape (k, n).
weight_scale (torch.Tensor): Scaling factor for each channel.
- In perGroup scenario: Shape is (k_group_num, n). Note: When k_group_num == 1, it is a perChannel scenario.
- In perChannel scenario: Shape is (n).
m (int): Number of tokens (number of rows in x).
Returns:
mm_out(fp16): Result of MatMul + perGroup or perChannel dequantization.
"""
# Perform matrix multiplication with int32 precision
k, n = weight.shape
mm_out = torch.zeros((m, n), dtype=torch.float16)
# perGroup scenario
if len(weight_scale.shape) == 2 and weight_scale.shape[0] != 1:
k_group = weight_scale.shape[0]
per_group_ele = k // k_group
x_grouped = x.view(-1, k_group, per_group_ele).transpose(0, 1)
weight_grouped = weight.view(k_group, per_group_ele, n)
c_temp = torch.bmm(x_grouped.to(torch.int32), weight_grouped.to(torch.int32)).to(torch.float16)
for k_idx in range(k_group):
mm_out += (c_temp[k_idx] * weight_scale[k_idx].view(1, -1).to(torch.float16)).to(torch.float16)
# perChannel scenario
elif len(weight_scale.shape) == 1 or (len(weight_scale.shape) == 2 and weight_scale.shape[0] == 1):
c_temp = torch.matmul(x.to(torch.int32), weight.to(torch.int32)).to(torch.float32)
mm_out = c_temp * weight_scale.view(1, -1).to(torch.float16)
return mm_out.to(torch.float32)
def gmm_swiglu_quant_golden_a8_w4(
x: torch.Tensor,
weight: torch.Tensor,
weight_scale: torch.Tensor,
per_token_scale: torch.Tensor,
bias: torch.Tensor,
group_list: torch.Tensor,
):
"""
Process the input data by group and call the GMM_Swiglu_quant function for quantization computation.
Parameters:
x (torch.Tensor): Input tensor with shape (M, K), type INT8.
weight (torch.Tensor): List of weight tensors, each with shape (E, K, N),
data type INT8 but data range INT4, representing INT4 values.
weight_scale (torch.Tensor): Scaling factor for each channel.
- In perGroup scenario: shape (E, k_group_num, N).
- In perChannel scenario: shape (E, N).
per_token_scale (torch.Tensor): Scaling factor for each token, shape (M, ).
bias: torch.Tensor,
group_list (list): List defining the number of tokens in each group.
Returns:
quant_output (torch.Tensor): Quantized output tensor with shape (M, N // 2).
quant_scale_output (torch.Tensor): Quantization scaling factor, shape (M, ).
"""
M, N = x.shape[0], weight.shape[2]
quant_output = torch.zeros(M, N // 2).to(torch.int8)
quant_scale_output = torch.zeros(M).to(torch.float32)
# Preprocessing X_INT8 -> X_INT4
x_int4 = x_int8_to_x_int4(x)
start_idx = 0
# Number of tokens in the previous group
pre_v = 0
group_list = group_list.tolist()
# Traverse group_list and process data by group
for i, v in enumerate(group_list):
curr_v = v
# Calculate the number of tokens in the current group " * 2 " because 1 row of Int8--> 2 rows of Int4
temp_v = int((curr_v - pre_v) * 2)
# Update the number of tokens in the previous group
pre_v = curr_v
if temp_v > 0:
mm_out = custom_mm(x_int4[int(start_idx) : int(start_idx + temp_v)], weight[i], weight_scale[i], temp_v)
mm_num_concat = (mm_out[::2] * 16 + mm_out[1::2]) + bias[i].view(1, -1)
per_token_quant = mm_num_concat * per_token_scale[start_idx // 2 : (start_idx + temp_v) // 2].view(-1, 1)
swiglu, gate = per_token_quant.chunk(2, dim=-1)
temp = swiglu * torch.sigmoid(swiglu)
temp = temp * gate
max_value = torch.max(torch.abs(temp), dim=-1).values
quant_scale_output_temp = 127 / max_value
quant_output[start_idx // 2 : (start_idx + temp_v) // 2] = torch.round(
temp * quant_scale_output_temp.reshape(temp_v // 2, 1)
).to(torch.int8)
quant_scale_output[start_idx // 2 : (start_idx + temp_v) // 2] = 1 / quant_scale_output_temp
start_idx += temp_v
return quant_output, quant_scale_output
def generate_non_decreasing_sequence(length, upper_limit):
# Generate random increasing sequence
random_increments = torch.randint(0, 128, (length,))
sequence = torch.cumsum(random_increments, dim=0)
# Make sure the last value is less than the upper limit
if sequence[-1] >= upper_limit:
scale_factor = upper_limit / sequence[-1]
sequence = (sequence * scale_factor).to(torch.int64)
return sequence
@torch.inference_mode()
def test_grouped_matmul_swiglu_quant_kernel():
E = 16
M = 512
K = 7168
N = 4096
torch.npu.config.allow_internal_format = True
x = torch.randint(-5, 5, (M, K), dtype=torch.int8).npu()
weight_ori = torch.randint(-5, 5, (E, K, N), dtype=torch.int8)
weight_nz = torch_npu.npu_format_cast(weight_ori.npu().to(torch.float32), 29)
pack_weight = torch_npu.npu_quantize(weight_nz, torch.tensor([1.0], device="npu"), None, torch.quint4x2, -1, False)
weight_scale = torch.randn(E, 1, N)
scale_np = weight_scale.cpu().numpy()
scale_np.dtype = np.uint32
scale_uint64_tensor = torch.from_numpy(scale_np.astype(np.int64)).npu()
pertoken_scale = torch.randn(M).to(torch.float32).npu()
group_list = generate_non_decreasing_sequence(E, M).npu()
bias = torch.zeros((E, N), dtype=torch.float32, device="npu").uniform_(-5, 5)
output_golden, output_scale_golden = gmm_swiglu_quant_golden_a8_w4(
x.cpu(), weight_ori, weight_scale, pertoken_scale.cpu(), bias.cpu(), group_list.cpu()
)
output, output_scale, _ = torch.ops._C_ascend.grouped_matmul_swiglu_quant(
x=x,
weight=pack_weight,
bias=bias,
group_list=group_list,
weight_scale=scale_uint64_tensor,
x_scale=pertoken_scale,
)
torch.testing.assert_close(output_golden, output.cpu(), atol=1, rtol=0.005)
torch.testing.assert_close(output_scale_golden, output_scale.cpu(), atol=1, rtol=0.005)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,358 @@
import logging
import unittest
from unittest import TestCase
import npugraph_ex as nge
import numpy as np
import torch
import torch.nn as nn
import torch_npu
from npugraph_ex.core.utils import logger
from vllm_ascend.utils import enable_custom_op
logger.setLevel(logging.DEBUG)
torch._logging.set_logs(graph_code=True)
enable_custom_op()
torch_npu.npu.config.allow_internal_format = True
def run_tests():
unittest.main()
def unpackbits(x, dim=-1):
x_np = x.cpu().numpy()
unpacked = np.unpackbits(x_np, axis=dim)
return torch.from_numpy(unpacked).to(x.device).to(dtype=torch.float16)
def stable_topk_1d(input, k, largest=True, stable_index_order="ascending"):
values_with_indices = [(val.item(), idx) for idx, val in enumerate(input)]
if largest:
if stable_index_order == "ascending":
values_with_indices.sort(key=lambda x: (-x[0], x[1]))
else:
values_with_indices.sort(key=lambda x: (-x[0], -x[1]))
else:
if stable_index_order == "ascending":
values_with_indices.sort(key=lambda x: (x[0], x[1]))
else:
values_with_indices.sort(key=lambda x: (x[0], -x[1]))
values = torch.tensor([x[0] for x in values_with_indices[:k]])
indices = torch.tensor([x[1] for x in values_with_indices[:k]])
return values, indices
def torch_hamming_distance_all(hash_q_op, hash_k_op, block_table_op, seq_len_op, chunk_op, top_k_op, sink, recent):
print("torch_hamming_distance_all")
# hash q op
batch, num_head, max_q_seqlen, head_dim = hash_q_op.shape
assert num_head % 8 == 0
# hash k op
num_blocks, num_kv_head, block_size, head_dim = hash_k_op.shape
num_head_group = num_head // num_kv_head
ds_top_k_ = None
ds_top_k_idx_ = None
# batch loop
for per_batch in range(batch):
# num_kv_head loop
for single_num_kv_head in range(num_kv_head):
hash_q_op_per_batch = hash_q_op[
per_batch : per_batch + 1,
single_num_kv_head * num_head_group : (single_num_kv_head + 1) * num_head_group,
:,
:,
]
ds_top_k, sk_idx = torch_hamming_distance_topk(
hash_q_op_per_batch,
hash_k_op[:, single_num_kv_head : (single_num_kv_head + 1), :, :],
block_table_op[per_batch : (per_batch + 1), :],
seq_len_op[per_batch],
chunk_op[per_batch],
top_k_op[per_batch],
sink,
recent,
)
if ds_top_k_ is None:
ds_top_k_ = ds_top_k
else:
ds_top_k_ = torch.cat([ds_top_k_, ds_top_k], dim=-1)
if ds_top_k_idx_ is None:
ds_top_k_idx_ = sk_idx
else:
ds_top_k_idx_ = torch.cat([ds_top_k_idx_, sk_idx], dim=-1)
assert ds_top_k_ is not None
assert ds_top_k_idx_ is not None
top_k = top_k_op.max()
ds_top_k_idx_r = ds_top_k_idx_.reshape(batch, num_kv_head, top_k)
return ds_top_k_idx_r
def torch_hamming_distance_topk(hash_q_op, hash_k_op, block_table_op, seq_len, chunk_size, top_k, sink, recent):
# According to the current implementation of the Hamming operator,
# convert data into bits for matrix multiplication to calculate the Hamming distance.
# hash q op
batch, num_head, max_q_seqlen, head_dim = hash_q_op.shape # torch.Size([1, 16, 1, 72])
assert max_q_seqlen == 1
# hash k op
num_blocks, num_kv_head, block_size, head_dim = hash_k_op.shape # torch.Size([245, 1, 128, 72])
# Convert the query tensor into a binary bit representation; the current data type of a single element is uint8.
hash_q_op_unpack = unpackbits(hash_q_op, dim=-1) # torch.Size([1, 16, 1, 576])
hash_q_op_unpack = hash_q_op_unpack.squeeze(0).squeeze(1) # torch.Size([16, 576])
# Based on the current operator implementation, perform bit-wise conversion on the 8×16 query,
# with bit replacement rule: 1 → 1, 0 → -1.
hash_q_op_unpack = hash_q_op_unpack.where(hash_q_op_unpack == 1, torch.tensor(-1, device="cpu"))
# Accumulate values row-wise and output the last row.
hash_q_op_unpack_cumsum = torch.cumsum(hash_q_op_unpack, dim=0)[-1]
# Data range is constrained within [-8, 8].
if num_head > 8:
div = (num_head + 7) // 8
reciprocalDiv = 1.0 / div
reciprocalDiv = torch.tensor((reciprocalDiv,), dtype=torch.float16, device="cpu")
hash_q_op_unpack_cumsum = hash_q_op_unpack_cumsum * reciprocalDiv
# In the operator implementation, the computation casts half to int4b_t.
# When the value reaches 8, it is cast to 7 instead.
hash_q_op_unpack_cumsum = torch.where(
hash_q_op_unpack_cumsum == 8,
torch.tensor(7, dtype=hash_q_op_unpack_cumsum.dtype, device="cpu"),
hash_q_op_unpack_cumsum,
) # torch.Size([576])
# Extract non-zero elements from block_table_op.
block_table_op_origin = block_table_op
block_table_op = block_table_op[block_table_op != 0]
# Select the hash K data at the positions specified by block_table_op from hash_k_op.
key_selected = hash_k_op[block_table_op, :, :, :] # torch.Size([240, 1, 128, 72])
# Convert key tensor to binary bits.
key_selected_unpack = unpackbits(key_selected, dim=-1)
# Convert key to bits: 1 → 1, 0 → -1.
key_selected_unpack = key_selected_unpack.where(key_selected_unpack == 1, torch.tensor(-1, device="cpu"))
# The product of Q and K
q_k_mul = torch.matmul(key_selected_unpack, hash_q_op_unpack_cumsum) # torch.Size([240, 1, 128])
# Flatten into one dimension.
flat_q_k_mul = q_k_mul.view(-1) # (1, 33 * 16)
# Segment by chunkSize, and compute one maximum value within each segment.
last_len = seq_len % chunk_size
effect_len = seq_len
if last_len != 0:
effect_len = (seq_len // chunk_size + 1) * chunk_size
# number of chunks
reduce_max_chunk = flat_q_k_mul[:effect_len].view(1, -1, chunk_size).max(dim=-1).values
reduce_max_chunk[:, :sink] = 8192
# skip_tail_size
chunk_block_size = reduce_max_chunk.shape[1]
reduce_max_chunk[:, chunk_block_size - recent : chunk_block_size] = 8192
topk_values, topk_indices = stable_topk_1d(reduce_max_chunk.view(-1), k=top_k, largest=True)
topk_indices_new_sort, idx = torch.sort(topk_indices)
block_table_indices = block_table_op_origin.view(-1)[topk_indices_new_sort]
return topk_values, block_table_indices
class TestCustomHammingDistTopK(TestCase):
def setUp(self):
self.device = "cpu"
self.batch_size = 5
self.num_head = 16
self.num_kv_head = 1
self.head_dim = 128
self.compress_rate = 8
self.compressed_dim = self.head_dim // self.compress_rate
self.seqlen_q = 1
self.sparse_ratio = 0.2
self.chunk_size_value = 128
self.device_id = 0
self.DEVICE_ID = 0
self.seqlen_list = [30720] * self.batch_size
self.seqlen = torch.tensor(self.seqlen_list, dtype=torch.int32, device=self.device)
self.max_seq_len = max(self.seqlen_list)
self.chunk_size = torch.tensor([self.chunk_size_value] * self.batch_size, dtype=torch.int32, device=self.device)
self.top_k_list = [int(seq * self.sparse_ratio // self.chunk_size_value) for seq in self.seqlen_list]
self.top_k = torch.tensor(self.top_k_list, dtype=torch.int32, device=self.device)
print("self.top_k_list", self.top_k_list, self.top_k)
self.block_size = 128
self.num_blocks_per_seq = (self.seqlen + self.block_size - 1) // self.block_size
self.num_blocks = self.num_blocks_per_seq.sum().item() + 5
torch.manual_seed(42)
np.random.seed(42)
self.qhash = torch.randint(
255,
(self.batch_size, self.num_head, self.seqlen_q, self.compressed_dim),
dtype=torch.uint8,
device=self.device,
)
self.khash = torch.randint(
255,
(self.num_blocks, self.num_kv_head, self.block_size, self.compressed_dim),
dtype=torch.uint8,
device=self.device,
)
self.sink = 1
self.recent = 4
max_num_blocks_per_seq = (max(self.seqlen_list) + self.block_size - 1) // self.block_size + 5
self.block_table = torch.full(
(len(self.num_blocks_per_seq), max_num_blocks_per_seq), fill_value=0, dtype=torch.int32
)
start = 1
for i, n in enumerate(self.num_blocks_per_seq):
self.block_table[i, :n] = torch.arange(start, start + n, dtype=torch.int32)
start += n
self.block_table = self.block_table.to(device=self.device)
self.indices = torch.zeros([self.batch_size, self.num_kv_head, 128], dtype=torch.int32)
self.support_offload = 1
self.mask = torch.tensor([True, True, False, False, False])
torch.npu.set_device(self.device_id)
self.npu = f"npu:{self.device_id}"
def _run_eager_mode_without_support_offload_and_mask(self):
output_eager = torch.ops._C_ascend.npu_hamming_dist_top_k(
self.qhash.to(self.npu),
self.khash.to(self.npu),
None,
self.top_k.to(self.npu),
self.seqlen.to(self.npu),
self.chunk_size.to(self.npu),
self.max_seq_len,
self.sink,
self.recent,
None,
self.block_table.to(self.npu),
None,
self.indices.to(self.npu),
)
return output_eager
def _run_eager_mode_with_support_offload_and_mask(self):
output_eager = torch.ops._C_ascend.npu_hamming_dist_top_k(
self.qhash.to(self.npu),
self.khash.to(self.npu),
None,
self.top_k.to(self.npu),
self.seqlen.to(self.npu),
self.chunk_size.to(self.npu),
self.max_seq_len,
self.sink,
self.recent,
self.support_offload,
self.block_table.to(self.npu),
self.mask.to(self.npu),
self.indices.to(self.npu),
)
return output_eager
def _run_graph_mode(self):
class Network(nn.Module):
def __init__(self):
super().__init__()
def forward(
self,
qhash,
khash,
khash_rope,
top_k,
seqlen,
chunk_size,
max_seq_len,
sink,
recent,
support_offload,
block_table,
mask,
indices,
):
return torch.ops._C_ascend.npu_hamming_dist_top_k(
qhash,
khash,
None,
top_k,
seqlen,
chunk_size,
max_seq_len,
sink,
recent,
support_offload,
block_table,
mask,
indices,
)
npu_mode = Network().to(f"npu:{self.DEVICE_ID}")
config = nge.CompilerConfig()
npu_backend = nge.get_npu_backend(compiler_config=config)
npu_mode = torch.compile(npu_mode, backend=npu_backend, dynamic=False)
npu_out = npu_mode(
self.qhash.to(self.npu),
self.khash.to(self.npu),
None,
self.top_k.to(self.npu),
self.seqlen.to(self.npu),
self.chunk_size.to(self.npu),
self.max_seq_len,
self.sink,
self.recent,
self.support_offload,
self.block_table.to(self.npu),
self.mask.to(self.npu),
self.indices.to(self.npu),
)
return npu_out
def test_hamming_dist_top_k_compare(self):
output_eager = self._run_eager_mode_with_support_offload_and_mask()
output_graph = self._run_graph_mode()
print(f"===========output_eager {output_eager} ===================")
print(f"===========output_graph {output_graph} ===================")
self.assertEqual(output_eager.shape, output_graph.shape)
self.assertTrue(torch.allclose(output_eager.float(), output_graph.float(), atol=1e-05))
def test_hamming_dist_top_k(self, loop: int = 1, enable_assert: bool = True):
output_gd = torch_hamming_distance_all(
self.qhash, self.khash, self.block_table, self.seqlen, self.chunk_size, self.top_k, self.sink, self.recent
)
w = output_gd.shape[-1]
print(f"output_gd shape: {output_gd.shape}")
print(f"output_gd: {output_gd[:, :, :w]}")
output_op = self._run_eager_mode_without_support_offload_and_mask()
print(f"output_op shape: {output_op.shape}")
print(f"output_op: {output_op[:, :, :w]}")
assert torch.equal(output_op[:, :, :w].to("cpu"), output_gd[:, :, :w]), (
"Output from custom op does not match ground truth!"
)
if __name__ == "__main__":
run_tests()

View File

@@ -0,0 +1,243 @@
import gc
import torch
import torch_npu
from vllm_ascend.utils import enable_custom_op
torch_npu.npu.config.allow_internal_format = True
enable_custom_op()
BF16_ATOL = 6e-2
BF16_RTOL = 7.8125e-3
def _make_inputs():
torch.manual_seed(1024)
batch = 1
query_seq = 1
kv_seq = 4096
actual_kv_seq = 4096
query_heads = 64
kv_heads = 1
head_dim = 512
rope_head_dim = 64
tile_size = 128
block_size = 256
sparse_block_count = 2048
block_num = kv_seq // block_size
scale_value = (head_dim + rope_head_dim) ** -0.5
query_nope = (
torch.empty((batch, query_seq, query_heads, head_dim), dtype=torch.float32).uniform_(-10, 10).to(torch.bfloat16)
)
query_rope = (
torch.empty((batch, query_seq, query_heads, rope_head_dim), dtype=torch.float32)
.uniform_(-10, 10)
.to(torch.bfloat16)
)
query = torch.cat((query_nope, query_rope), dim=-1).npu()
key_nope = (
torch.empty((block_num, block_size, kv_heads, head_dim), dtype=torch.float32).uniform_(-5, 10).to(torch.int8)
)
value = key_nope.clone()
key_rope = (
torch.empty((block_num, block_size, kv_heads, rope_head_dim), dtype=torch.float32)
.uniform_(-10, 10)
.to(torch.bfloat16)
)
dequant_scale = torch.empty((block_num, block_size, kv_heads, head_dim // tile_size), dtype=torch.float32).uniform_(
0.1, 1.0
)
kv_cache = torch.cat(
(
key_nope,
key_rope.contiguous().view(torch.int8),
dequant_scale.contiguous().view(torch.int8),
),
dim=-1,
).npu()
sparse_indices = torch.randperm(actual_kv_seq, dtype=torch.int32)[:sparse_block_count].reshape(
batch, query_seq, kv_heads, sparse_block_count
)
block_table = torch.arange(block_num, dtype=torch.int32).reshape(batch, block_num)
actual_seq_lengths_query = torch.tensor([query_seq], dtype=torch.int32)
actual_seq_lengths_kv = torch.tensor([actual_kv_seq], dtype=torch.int32)
return {
"query": query,
"key": kv_cache,
"value": value.npu(),
"sparse_indices": sparse_indices.npu(),
"block_table": block_table.npu(),
"actual_seq_lengths_query": actual_seq_lengths_query.npu(),
"actual_seq_lengths_kv": actual_seq_lengths_kv.npu(),
"scale_value": scale_value,
"sparse_block_size": 1,
"layout_query": "BSND",
"layout_kv": "PA_BSND",
"sparse_mode": 3,
"attention_mode": 2,
"quant_scale_repo_mode": 1,
"tile_size": tile_size,
"rope_head_dim": rope_head_dim,
"key_quant_mode": 2,
"value_quant_mode": 2,
"cpu": {
"query_nope": query_nope,
"query_rope": query_rope,
"key_nope": key_nope,
"value": value,
"key_rope": key_rope,
"dequant_scale": dequant_scale,
"sparse_indices": sparse_indices,
"block_table": block_table,
"actual_seq_lengths_query": actual_seq_lengths_query,
"actual_seq_lengths_kv": actual_seq_lengths_kv,
"block_size": block_size,
},
}
def _sparse_token_indices(
sparse_indices,
sparse_block_size,
sparse_block_count,
query_idx,
actual_seq_query,
actual_seq_kv,
sparse_mode,
):
if sparse_mode == 0:
threshold = actual_seq_kv
elif sparse_mode == 3:
threshold = actual_seq_kv - actual_seq_query + query_idx + 1
else:
raise AssertionError(f"unsupported sparse_mode in test: {sparse_mode}")
valid_count = min(
sparse_block_count,
(threshold + sparse_block_size - 1) // sparse_block_size,
)
tokens: list[int] = []
for sparse_id in sparse_indices[:valid_count].tolist():
if sparse_id == -1:
break
begin = sparse_id * sparse_block_size
end = min(begin + sparse_block_size, actual_seq_kv)
if begin >= threshold:
continue
tokens.extend(range(begin, min(end, threshold)))
return torch.tensor(tokens, dtype=torch.long)
def _reference_attention(inputs):
cpu = inputs["cpu"]
query = torch.cat((cpu["query_nope"], cpu["query_rope"]), dim=-1).float()
dequant_scale = cpu["dequant_scale"].repeat_interleave(inputs["tile_size"], dim=-1)
key_nope = (cpu["key_nope"].float() * dequant_scale).to(torch.bfloat16)
value = (cpu["value"].float() * dequant_scale).to(torch.bfloat16).float()
key = torch.cat((key_nope, cpu["key_rope"]), dim=-1).float()
batch, query_seq, query_heads, _ = query.shape
kv_heads = cpu["key_nope"].shape[2]
group_size = query_heads // kv_heads
head_dim = cpu["value"].shape[-1]
output = torch.zeros((batch, query_seq, query_heads, head_dim), dtype=torch.float32)
block_table = cpu["block_table"].long()
block_size = cpu["block_size"]
sparse_block_count = cpu["sparse_indices"].shape[-1]
for batch_idx in range(batch):
actual_seq_query = int(cpu["actual_seq_lengths_query"][batch_idx])
actual_seq_kv = int(cpu["actual_seq_lengths_kv"][batch_idx])
for kv_head_idx in range(kv_heads):
head_begin = kv_head_idx * group_size
head_end = head_begin + group_size
for query_idx in range(actual_seq_query):
token_indices = _sparse_token_indices(
cpu["sparse_indices"][batch_idx, query_idx, kv_head_idx],
inputs["sparse_block_size"],
sparse_block_count,
query_idx,
actual_seq_query,
actual_seq_kv,
inputs["sparse_mode"],
)
if token_indices.numel() == 0:
continue
logical_blocks = token_indices // block_size
block_offsets = token_indices % block_size
physical_blocks = block_table[batch_idx, logical_blocks]
k_sparse = key[physical_blocks, block_offsets, kv_head_idx]
v_sparse = value[physical_blocks, block_offsets, kv_head_idx]
q_current = query[batch_idx, query_idx, head_begin:head_end]
scores = torch.matmul(q_current, k_sparse.T) * inputs["scale_value"]
probs = torch.softmax(scores, dim=-1)
output[batch_idx, query_idx, head_begin:head_end] = torch.matmul(
probs.to(torch.bfloat16).float(), v_sparse
)
return output
def _run_custom_op(inputs, return_softmax_lse=False):
return torch.ops._C_ascend.npu_kv_quant_sparse_flash_attention(
inputs["query"],
inputs["key"],
inputs["value"],
inputs["sparse_indices"],
inputs["scale_value"],
key_quant_mode=inputs["key_quant_mode"],
value_quant_mode=inputs["value_quant_mode"],
block_table=inputs["block_table"],
actual_seq_lengths_query=inputs["actual_seq_lengths_query"],
actual_seq_lengths_kv=inputs["actual_seq_lengths_kv"],
sparse_block_size=inputs["sparse_block_size"],
layout_query=inputs["layout_query"],
layout_kv=inputs["layout_kv"],
sparse_mode=inputs["sparse_mode"],
attention_mode=inputs["attention_mode"],
quant_scale_repo_mode=inputs["quant_scale_repo_mode"],
tile_size=inputs["tile_size"],
rope_head_dim=inputs["rope_head_dim"],
return_softmax_lse=return_softmax_lse,
)
@torch.inference_mode()
def test_kv_quant_sparse_flash_attention():
inputs = _make_inputs()
reference = _reference_attention(inputs)
output, softmax_max, softmax_sum = _run_custom_op(inputs)
assert output.shape == (1, 1, 64, 512)
assert output.dtype == torch.bfloat16
assert softmax_max.numel() == 0
assert softmax_sum.numel() == 0
assert torch.isfinite(output.cpu()).all()
torch.testing.assert_close(output.cpu().float(), reference, atol=BF16_ATOL, rtol=BF16_RTOL)
output_lse, softmax_max_lse, softmax_sum_lse = _run_custom_op(inputs, return_softmax_lse=True)
assert output_lse.shape == (1, 1, 64, 512)
assert softmax_max_lse.shape == (1, 1, 1, 64)
assert softmax_sum_lse.shape == (1, 1, 1, 64)
assert output_lse.dtype == torch.bfloat16
assert softmax_max_lse.dtype == torch.float32
assert softmax_sum_lse.dtype == torch.float32
assert torch.isfinite(output_lse.cpu()).all()
assert torch.isfinite(softmax_max_lse.cpu()).all()
assert torch.isfinite(softmax_sum_lse.cpu()).all()
torch.testing.assert_close(output_lse.cpu().float(), reference, atol=BF16_ATOL, rtol=BF16_RTOL)
torch.testing.assert_close(output_lse.cpu(), output.cpu(), atol=1e-2, rtol=1e-2)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,113 @@
import gc
import pytest
import torch
import torch_npu
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
@pytest.mark.skip(reason="Failure of an individual operator use case causes failures of other operators.")
@pytest.mark.parametrize("cache_mode", ["krope_ctkv", "nzcache"])
@torch.inference_mode()
def test_mla_preprocess_kernel(cache_mode: str):
token_num = 1
head_num = 2
N_7168 = 7168
block_num = 1
block_size = 128
dtype = torch.bfloat16
hidden_states = torch.randn((token_num, N_7168), dtype=dtype).npu()
quant_scale0 = torch.randn((1,), dtype=dtype).npu()
quant_offset0 = torch.randint(0, 7, (1,), dtype=torch.int8).npu()
wdqkv = torch.randint(0, 7, (1, 224, 2112, 32), dtype=torch.int8).npu()
wdqkv = torch_npu.npu_format_cast(wdqkv.contiguous(), 29)
de_scale0 = torch.rand((2112,), dtype=torch.float).npu()
bias0 = torch.randint(0, 7, (2112,), dtype=torch.int32).npu()
gamma1 = torch.randn((1536), dtype=dtype).npu()
beta1 = torch.randn((1536), dtype=dtype).npu()
quant_scale1 = torch.randn((1,), dtype=dtype).npu()
quant_offset1 = torch.randint(0, 7, (1,), dtype=torch.int8).npu()
wuq = torch.randint(0, 7, (1, 48, head_num * 192, 32), dtype=torch.int8).npu()
wuq = torch_npu.npu_format_cast(wuq.contiguous(), 29)
de_scale1 = torch.rand((head_num * 192,), dtype=torch.float).npu()
bias1 = torch.randint(0, 7, (head_num * 192,), dtype=torch.int32).npu()
gamma2 = torch.randn((512), dtype=dtype).npu()
cos = torch.randn((token_num, 64), dtype=dtype).npu()
sin = torch.randn((token_num, 64), dtype=dtype).npu()
wuk = torch.randn((head_num, 128, 512), dtype=dtype).npu()
wuk = torch_npu.npu_format_cast(wuk, 29)
kv_cache = torch.randint(0, 7, (block_num, head_num * 512 // 32, block_size, 32), dtype=dtype).npu()
kv_cache_rope = torch.randn((block_num, head_num * 64 // 16, block_size, 16), dtype=dtype).npu()
slotmapping = torch.randint(0, 7, (token_num,), dtype=torch.int32).npu()
ctkv_scale = torch.randn((1,), dtype=dtype).npu()
qnope_scale = torch.randn((head_num), dtype=dtype).npu()
q_nope_out = torch.empty(
(hidden_states.shape[0], wuk.shape[0], kv_cache.shape[-1]),
dtype=hidden_states.dtype,
device=hidden_states.device,
)
q_rope_out = torch.empty(
(hidden_states.shape[0], wuk.shape[0], kv_cache_rope.shape[-1]),
dtype=hidden_states.dtype,
device=hidden_states.device,
)
q_down = torch.empty(
(hidden_states.shape[0], 1536),
dtype=hidden_states.dtype,
device=hidden_states.device,
)
q_nope_old = q_nope_out.clone()
q_rope_old = q_rope_out.clone()
torch.ops._C_ascend.mla_preprocess(
hidden_states,
wdqkv,
de_scale0,
gamma1,
beta1,
wuq,
de_scale1,
gamma2,
cos,
sin,
wuk,
kv_cache,
kv_cache_rope,
slotmapping,
quant_scale0=quant_scale0,
quant_offset0=quant_offset0,
bias0=bias0,
quant_scale1=quant_scale1,
quant_offset1=quant_offset1,
bias1=bias1,
ctkv_scale=ctkv_scale,
q_nope_scale=qnope_scale,
cache_mode=cache_mode,
quant_mode="per_tensor_quant_asymm",
enable_inner_out=False,
q_out0=q_nope_out,
kv_cache_out0=kv_cache,
q_out1=q_rope_out,
kv_cache_out1=kv_cache_rope,
inner_out=q_down,
)
assert not torch.equal(q_nope_out, q_nope_old)
assert not torch.equal(q_rope_out, q_rope_old)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,97 @@
import gc
import pytest
import torch
import torch_npu
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
@pytest.mark.skip(reason="Failure of an individual operator use case causes failures of other operators.")
@torch.inference_mode()
def test_mla_preprocess_kernel():
token_num = 1
head_num = 2
N_7168 = 7168
block_num = 1
block_size = 128
dtype = torch.bfloat16
hidden_states = torch.randn((token_num, N_7168), dtype=dtype).npu()
wdqkv = torch.randint(0, 7, (1, 448, 2112, 16), dtype=dtype).npu()
wdqkv = torch_npu.npu_format_cast(wdqkv.contiguous(), 29)
gamma1 = torch.randn((1536), dtype=dtype).npu()
wuq = torch.randint(0, 7, (1, 96, head_num * 192, 16), dtype=dtype).npu()
wuq = torch_npu.npu_format_cast(wuq.contiguous(), 29)
gamma2 = torch.randn((512), dtype=dtype).npu()
cos = torch.randn((token_num, 64), dtype=dtype).npu()
sin = torch.randn((token_num, 64), dtype=dtype).npu()
wuk = torch.randn((head_num, 128, 512), dtype=dtype).npu()
# wuk = torch_npu.npu_format_cast(wuk, 29)
kv_cache = torch.randint(0, 7, (block_num, head_num * 512 // 32, block_size, 32), dtype=dtype).npu()
kv_cache_rope = torch.randn((block_num, head_num * 64 // 16, block_size, 16), dtype=dtype).npu()
slotmapping = torch.randint(0, 7, (token_num,), dtype=torch.int32).npu()
q_nope_out = torch.empty(
(hidden_states.shape[0], wuk.shape[0], kv_cache.shape[-1]),
dtype=hidden_states.dtype,
device=hidden_states.device,
)
q_rope_out = torch.empty(
(hidden_states.shape[0], wuk.shape[0], kv_cache_rope.shape[-1]),
dtype=hidden_states.dtype,
device=hidden_states.device,
)
q_down = torch.empty(
(hidden_states.shape[0], 1536),
dtype=hidden_states.dtype,
device=hidden_states.device,
)
q_nope_old = q_nope_out.clone()
q_rope_old = q_rope_out.clone()
torch.ops._C_ascend.mla_preprocess(
hidden_states,
wdqkv,
None,
gamma1,
None,
wuq,
None,
gamma2,
cos,
sin,
wuk,
kv_cache,
kv_cache_rope,
slotmapping,
None,
None,
None,
None,
None,
None,
None,
None,
cache_mode="krope_ctkv",
quant_mode="no_quant",
enable_inner_out=False,
q_out0=q_nope_out,
kv_cache_out0=kv_cache,
q_out1=q_rope_out,
kv_cache_out1=kv_cache_rope,
inner_out=q_down,
)
assert not torch.equal(q_nope_out, q_nope_old)
assert not torch.equal(q_rope_out, q_rope_old)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,116 @@
import gc
import pytest
import torch
import torch_npu
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
@pytest.mark.skip(reason="Failure of an individual operator use case causes failures of other operators.")
@pytest.mark.parametrize("cache_mode", ["krope_ctkv", "nzcache"])
@torch.inference_mode()
def test_mla_preprocess_kernel(cache_mode: str):
token_num = 1
head_num = 2
N_7168 = 7168
block_num = 1
block_size = 128
dtype = torch.bfloat16
hidden_states = torch.randn((token_num, N_7168), dtype=dtype).npu()
quant_scale0 = torch.randn((1,), dtype=dtype).npu()
quant_offset0 = torch.randint(0, 7, (1,), dtype=torch.int8).npu()
wdqkv = torch.randint(0, 7, (1, 224, 2112, 32), dtype=torch.int8).npu()
wdqkv = torch_npu.npu_format_cast(wdqkv.contiguous(), 29)
de_scale0 = torch.rand((2112,), dtype=torch.float).npu()
bias0 = torch.randint(0, 7, (2112,), dtype=torch.int32).npu()
gamma1 = torch.randn((1536), dtype=dtype).npu()
beta1 = torch.randn((1536), dtype=dtype).npu()
quant_scale1 = torch.randn((1,), dtype=dtype).npu()
quant_offset1 = torch.randint(0, 7, (1,), dtype=torch.int8).npu()
wuq = torch.randint(0, 7, (1, 48, head_num * 192, 32), dtype=torch.int8).npu()
wuq = torch_npu.npu_format_cast(wuq.contiguous(), 29)
de_scale1 = torch.rand((head_num * 192,), dtype=torch.float).npu()
bias1 = torch.randint(0, 7, (head_num * 192,), dtype=torch.int32).npu()
gamma2 = torch.randn((512), dtype=dtype).npu()
cos = torch.randn((token_num, 64), dtype=dtype).npu()
sin = torch.randn((token_num, 64), dtype=dtype).npu()
wuk = torch.randn((head_num, 128, 512), dtype=dtype).npu()
wuk = torch_npu.npu_format_cast(wuk, 29)
kv_cache = torch.randint(0, 7, (block_num, head_num * 512 // 32, block_size, 32), dtype=dtype).npu()
kv_cache_rope = torch.randn((block_num, head_num * 64 // 16, block_size, 16), dtype=dtype).npu()
slotmapping = torch.randint(0, 7, (token_num,), dtype=torch.int32).npu()
ctkv_scale = torch.randn((1,), dtype=dtype).npu()
qnope_scale = torch.randn((head_num), dtype=dtype).npu()
q_nope_out = torch.empty(
(hidden_states.shape[0], wuk.shape[0], kv_cache.shape[-1]),
dtype=hidden_states.dtype,
device=hidden_states.device,
)
q_rope_out = torch.empty(
(hidden_states.shape[0], wuk.shape[0], kv_cache_rope.shape[-1]),
dtype=hidden_states.dtype,
device=hidden_states.device,
)
q_down = torch.empty(
(hidden_states.shape[0], 1536),
dtype=hidden_states.dtype,
device=hidden_states.device,
)
q_nope_old = q_nope_out.clone()
q_rope_old = q_rope_out.clone()
q_down_old = q_down.clone()
torch.ops._C_ascend.mla_preprocess(
hidden_states,
wdqkv,
de_scale0,
gamma1,
beta1,
wuq,
de_scale1,
gamma2,
cos,
sin,
wuk,
kv_cache,
kv_cache_rope,
slotmapping,
quant_scale0=quant_scale0,
quant_offset0=quant_offset0,
bias0=bias0,
quant_scale1=quant_scale1,
quant_offset1=quant_offset1,
bias1=bias1,
ctkv_scale=ctkv_scale,
q_nope_scale=qnope_scale,
cache_mode=cache_mode,
quant_mode="per_tensor_quant_asymm",
enable_inner_out=True,
q_out0=q_nope_out,
kv_cache_out0=kv_cache,
q_out1=q_rope_out,
kv_cache_out1=kv_cache_rope,
inner_out=q_down,
)
assert not torch.equal(q_nope_out, q_nope_old)
assert not torch.equal(q_rope_out, q_rope_old)
assert not torch.equal(q_down, q_down_old)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,402 @@
import gc
import itertools
import random
import numpy as np
import torch
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
def adapter_capacity(sorted_row_idx, sorted_expert_idx, capacity):
count = 0
last = sorted_expert_idx[0]
for i, val in enumerate(sorted_expert_idx):
if last != val:
count = 1
last = val
else:
count += 1
if count > capacity:
sorted_expert_idx[i] = -1
sorted_row_idx[i] = -1
def moe_init_routing_golden(
x,
expert_idx,
scale,
offset,
active_num,
expert_capacity,
expert_num,
drop_pad_mode,
expert_tokens_num_type,
expert_tokens_num_flag,
active_expert_range,
quant_mode,
row_idx_type,
):
if drop_pad_mode == 1:
if expert_num <= 0:
print("expert num can not be 0")
return
expert_start = active_expert_range[0] if drop_pad_mode == 0 else 0
expert_end = active_expert_range[1] if drop_pad_mode == 0 else expert_num
num_rows = x.shape[0]
h = x.shape[1]
k = expert_idx.shape[-1]
expert_idx_in = expert_idx.copy().reshape(-1)
actual_expert_total_num: int = np.sum((expert_idx_in >= expert_start) & (expert_idx_in < expert_end))
expert_idx_in[(expert_idx_in < expert_start)] = np.int32(np.iinfo(np.int32).max)
sorted_expert_indices = np.argsort(expert_idx_in, axis=-1, kind="stable")
sorted_expert_idx = expert_idx_in[sorted_expert_indices]
if row_idx_type == 1:
expanded_row_idx = sorted_expert_indices[:actual_expert_total_num]
else:
expanded_row_idx = np.ones(num_rows * k).astype(np.int32) * -1
tmp_indices = np.arange(actual_expert_total_num)
expanded_row_idx[sorted_expert_indices[:actual_expert_total_num]] = tmp_indices
if not expert_tokens_num_flag:
expert_tokens_count = torch.tensor([0])
else:
if drop_pad_mode == 0:
if expert_tokens_num_type == 1:
expert_tokens_count = np.bincount(sorted_expert_idx[:actual_expert_total_num] - expert_start)
expert_tokens_count = np.concatenate(
[
expert_tokens_count,
np.zeros((expert_end - expert_start) - len(expert_tokens_count)).astype(np.int64),
]
)
elif expert_tokens_num_type == 0:
expert_tokens_count = np.bincount(sorted_expert_idx[:actual_expert_total_num] - expert_start)
expert_tokens_count = np.concatenate(
[
expert_tokens_count,
np.zeros((expert_end - expert_start) - len(expert_tokens_count)).astype(np.int64),
]
)
expert_tokens_count = np.cumsum(expert_tokens_count)
elif expert_tokens_num_type == 2:
expert_id, counts = np.unique(sorted_expert_idx[:actual_expert_total_num], return_counts=True)
expert_tokens_count = np.column_stack((expert_id, counts))
if expert_tokens_count.shape[0] < expert_num:
expert_tokens_count = np.concatenate(
(
expert_tokens_count,
[
[0, 0],
],
),
axis=0,
)
else:
expert_tokens_count = np.bincount(sorted_expert_idx[:actual_expert_total_num] - expert_start)
zeros_array = np.zeros((expert_end - expert_start) - len(expert_tokens_count), dtype=np.int64)
expert_tokens_count = np.concatenate([expert_tokens_count, zeros_array])
expert_tokens_count = expert_tokens_count.astype(np.int64)
if drop_pad_mode == 0:
if active_num == 0:
active_num = actual_expert_total_num
else:
active_num = min(active_num, actual_expert_total_num)
expanded_scale = None
expanded_x = x[sorted_expert_indices[:active_num] // k, :]
if scale is not None and quant_mode == -1:
expanded_scale = scale[sorted_expert_indices[:active_num] // k]
else:
adapter_capacity(sorted_expert_indices, sorted_expert_idx, expert_capacity)
sort_row_tmp = np.full((expert_num * expert_capacity), -1, dtype=int)
offset_tmp = 0
lastExpertId = 0
for i, val in enumerate(sorted_expert_indices):
if val != -1:
if lastExpertId != sorted_expert_idx[i]:
offset_tmp = 0
lastExpertId = sorted_expert_idx[i]
sort_row_tmp[sorted_expert_idx[i] * expert_capacity + offset_tmp] = sorted_expert_indices[i]
offset_tmp = offset_tmp + 1
expanded_row_idx = np.full(sorted_expert_indices.shape, -1)
for i, val in enumerate(sort_row_tmp):
if val != -1:
expanded_row_idx[val] = i
expanded_x_mask = np.full((expert_num * expert_capacity, h), 1, dtype=int)
expanded_x = np.full((expert_num * expert_capacity, h), 0, dtype=x.dtype)
for i, val in enumerate(sort_row_tmp):
if val != -1:
expanded_x[i] = x[val // k]
expanded_x_mask[i] = np.full((h,), 0, dtype=int)
if quant_mode == -1:
expanded_x = expanded_x
expanded_row_idx = expanded_row_idx
if scale is not None and drop_pad_mode == 1:
expanded_scale = np.full((expert_num * expert_capacity,), 0, dtype=scale.dtype)
for i, val in enumerate(sort_row_tmp):
if val != -1:
expanded_scale[i] = scale[val // k]
if scale is None:
expanded_scale = None
if quant_mode == 0:
expanded_scale = None
expanded_x_fp16 = expanded_x.astype(np.float16)
if scale is not None:
scale_val = scale.astype(np.float16)
else:
raise ValueError("scale cannot be None when quant_mode is 0")
if offset is not None:
offset_val = offset.astype(np.float16)
else:
raise ValueError("offset cannot be None when quant_mode is 0")
scale_rst = expanded_x_fp16 * scale_val[0]
add_offset = scale_rst + offset_val[0]
round_data = np.rint(add_offset)
round_data = np.clip(round_data, -128, 127)
expanded_x = round_data.astype(np.int8)
if quant_mode == 1:
x_final = expanded_x.astype(np.float32)
if scale is None:
x_abs = np.abs(x_final)
x_max = np.max(x_abs, axis=-1, keepdims=True)
expanded_scale = x_max / 127
expanded_x = x_final / expanded_scale
expanded_x = np.round(expanded_x).astype(np.int8)
else:
if scale.shape[0] == 1:
x_final = x_final * scale
else:
if drop_pad_mode == 0:
x_final = x_final * scale[sorted_expert_idx[:active_num] - expert_start]
else:
for i, val in enumerate(sort_row_tmp):
if val != -1:
x_final[i] = x_final[i] * scale[i // expert_capacity]
x_abs = np.abs(x_final)
x_max = np.max(x_abs, axis=-1, keepdims=True)
expanded_scale = x_max / 127
expanded_x = x_final / expanded_scale
expanded_x = np.round(expanded_x).astype(np.int8)
if x.dtype == np.int8:
expanded_scale = None
if drop_pad_mode == 1:
expanded_x = np.ma.array(expanded_x, mask=expanded_x_mask).filled(0)
expanded_x = expanded_x.reshape(expert_num, expert_capacity, h)
return expanded_x, expanded_row_idx, expert_tokens_count, expanded_scale
def npu_pta(
x,
expert_idx,
scale,
offset,
active_num,
expert_capacity,
expert_num,
drop_pad_mode,
expert_tokens_num_type,
expert_tokens_num_flag,
quant_mode,
active_expert_range,
row_idx_type,
):
expanded_x, expanded_row_idx, expert_token_cumsum_or_count, expanded_scale = (
torch.ops._C_ascend.npu_moe_init_routing_custom(
x,
expert_idx,
scale=scale,
offset=offset,
active_num=active_num,
expert_capacity=expert_capacity,
expert_num=expert_num,
drop_pad_mode=drop_pad_mode,
expert_tokens_num_type=expert_tokens_num_type,
expert_tokens_num_flag=expert_tokens_num_flag,
quant_mode=quant_mode,
active_expert_range=active_expert_range,
row_idx_type=row_idx_type,
)
)
return expanded_x, expanded_row_idx, expert_token_cumsum_or_count, expanded_scale
def cmp_out_golden(x_golden, x_out, dtype):
if dtype == "int8":
cmp = np.isclose(x_out.cpu().numpy()[: len(x_golden)], x_golden, atol=1)
else:
cmp = np.isclose(x_out.cpu().numpy()[: len(x_golden)], x_golden, rtol=1e-05, atol=1e-05)
return np.all(cmp)
def run_moe_npu_case(
x,
expert_idx,
scale,
offset,
active_num,
expert_capacity,
expert_num,
drop_pad_mode,
expert_tokens_num_type,
expert_tokens_num_flag,
quant_mode,
active_expert_range,
row_idx_type,
):
x_npu = x.npu()
expert_idx_npu = expert_idx.npu()
scale_npu = scale.npu() if scale is not None else None
offset_npu = offset.npu() if offset is not None else None
x_numpy = x.numpy()
expert_idx_numpy = expert_idx.numpy()
scale_numpy = scale.numpy() if scale is not None else None
offset_numpy = offset.numpy() if offset is not None else None
expanded_x_golden, expanded_row_idx_golden, expert_token_cumsum_or_count_golden, expanded_scale_golden = (
moe_init_routing_golden(
x_numpy,
expert_idx_numpy,
scale_numpy,
offset_numpy,
active_num,
expert_capacity,
expert_num,
drop_pad_mode,
expert_tokens_num_type,
expert_tokens_num_flag,
active_expert_range,
quant_mode,
row_idx_type,
)
)
expanded_x, expanded_row_idx, expert_token_cumsum_or_count, expanded_scale = npu_pta(
x_npu,
expert_idx_npu,
scale_npu,
offset_npu,
active_num,
expert_capacity,
expert_num,
drop_pad_mode,
expert_tokens_num_type,
expert_tokens_num_flag,
quant_mode,
active_expert_range,
row_idx_type,
)
if quant_mode == -1:
expanded_x_result = cmp_out_golden(expanded_x_golden, expanded_x, "float32")
else:
expanded_x_result = cmp_out_golden(expanded_x_golden, expanded_x, "int8")
expanded_row_idx_result = cmp_out_golden(expanded_row_idx_golden, expanded_row_idx, "int32")
if expert_tokens_num_flag:
expert_tokens_result = cmp_out_golden(
expert_token_cumsum_or_count_golden, expert_token_cumsum_or_count, "int64"
)
else:
expert_tokens_result = True
if quant_mode == 1 or (quant_mode == -1 and scale is not None):
expand_scale_result = cmp_out_golden(expanded_scale_golden.flatten(), expanded_scale, "float32")
else:
expand_scale_result = True
compare_result = expanded_x_result and expanded_row_idx_result and expert_tokens_result and expand_scale_result
# print('=======case result=======: ', compare_result)
return compare_result
def test_moe_init_routing_custom():
failed_test_cnt = 0
drop_pad_mode = [0, 1]
expert_tokens_num_type = [0, 1, 2]
expert_tokens_num_flag = [True, False]
quant_mode = [0, 1, -1]
row_idx_type = [0, 1]
scale_type = [0, 1, 2]
product_result = itertools.product(
drop_pad_mode, expert_tokens_num_type, expert_tokens_num_flag, quant_mode, row_idx_type, scale_type
)
for idx, (
drop_pad_mode_,
expert_tokens_num_type_,
expert_tokens_num_flag_,
quant_mode_,
row_idx_type_,
scale_type_,
) in enumerate(product_result, 5):
expert_num_ = random.randint(2, 500)
expert_start = random.randint(0, expert_num_ - 1)
expert_end = random.randint(expert_start + 1, expert_num_)
active_expert_range_ = [expert_start, expert_end]
N = random.randint(1, 100)
H = random.randint(12, 100)
K = random.randint(1, 12)
x_ = torch.randn(N, H, dtype=torch.float16) * 5
expert_capacity_ = random.randint(1, N - 1) if N > 1 else 1
expert_idx_ = torch.randint(0, expert_num_ - 1, (N, K), dtype=torch.int32)
active_num_ = N * K
if drop_pad_mode_ == 1:
active_expert_range_ = [0, expert_num_]
expert_tokens_num_type_ = 1
row_idx_type_ = 0
if quant_mode_ == 0:
scale_ = torch.randn(1, dtype=torch.float)
offset_ = torch.randn(1, dtype=torch.float)
elif quant_mode_ == -1:
scale_ = None
offset_ = None
else:
if scale_type_ == 0:
scale_ = None
offset_ = None
elif scale_type_ == 1:
scale_ = torch.randn(1, H, dtype=torch.float)
offset_ = None
else:
scale_ = torch.randn(active_expert_range_[1] - active_expert_range_[0], H, dtype=torch.float)
offset_ = None
result_pta = run_moe_npu_case(
x_,
expert_idx_,
scale_,
offset_,
active_num_,
expert_capacity_,
expert_num_,
drop_pad_mode_,
expert_tokens_num_type_,
expert_tokens_num_flag_,
quant_mode_,
active_expert_range_,
row_idx_type_,
)
if not result_pta:
failed_test_cnt += 1
assert failed_test_cnt == 0
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,584 @@
"""E2E accuracy test for NgramSpecDecode custom operator.
Tests the Ascend C kernel against a CPU golden reference implementation
with parametrized test cases covering various configurations.
"""
import time
import numpy as np
import pytest
import torch
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
SEED = 42
PERF_WARMUP = 3
PERF_ITERS = 20
# ---------------------------------------------------------------------------
# Golden reference (CPU, pure Python/NumPy)
# ---------------------------------------------------------------------------
def golden_ngram_spec_decode(
token_ids: np.ndarray, # [B, M], int32,
num_tokens_no_spec: np.ndarray, # [B], int32
sampled_token_ids: np.ndarray, # [B, N], int32
discard_request_mask: np.ndarray, # [B], int32
vocab_size: int,
min_n: int,
max_n: int,
k: int,
):
"""CPU golden reference for NgramSpecDecode.
Returns:
(token_ids_modified, next_token_ids, draft_token_ids, num_valid_draft_tokens)
"""
B = token_ids.shape[0]
M = token_ids.shape[1]
next_token_ids = np.zeros(B, dtype=np.int32)
draft_token_ids = np.full((B, k), -1, dtype=np.int32)
num_valid_draft_tokens = np.zeros(B, dtype=np.int32)
for i in range(B):
seq_len = int(num_tokens_no_spec[i])
discard = int(discard_request_mask[i])
valid_count = 0
# Stage 1: sample token valid
backup_pos = max(seq_len - 1, 0)
backup_token = int(token_ids[i, backup_pos])
for j in range(sampled_token_ids.shape[1]):
val = int(sampled_token_ids[i, j])
if discard != 0:
sampled_token_ids[i, j] = -1
elif val != -1 and val < vocab_size:
valid_count += 1
else:
sampled_token_ids[i, j] = -1
avail_space = M - seq_len
if avail_space < 0:
avail_space = 0
if valid_count > avail_space:
valid_count = avail_space
if valid_count > 0:
next_token_ids[i] = int(sampled_token_ids[i, valid_count - 1])
else:
next_token_ids[i] = backup_token
# Stage 2: scatter sampled token to token_ids tail
nt = seq_len + valid_count
for j in range(valid_count):
token_ids[i, seq_len + j] = int(sampled_token_ids[i, j])
# Stage 3: suffix n-gram match
best_match_pos = -1
best_ngram_len = 0
if valid_count > 0 and nt >= min_n:
for ngram_len in range(min_n, max_n + 1):
if ngram_len > nt:
break
wc = nt - ngram_len
if wc <= 0:
break
suffix = token_ids[i, nt - ngram_len : nt].tolist()
found = False
for pos in range(wc):
window = token_ids[i, pos : pos + ngram_len].tolist()
if window == suffix:
best_match_pos = pos
best_ngram_len = ngram_len
found = True
break
if found:
break
# Stage 4: get draft tokens
if best_match_pos >= 0:
draft_start = best_match_pos + best_ngram_len
tokens_available = nt - draft_start
for j in range(k):
if j < tokens_available:
draft_token_ids[i, j] = int(token_ids[i, draft_start + j])
else:
draft_token_ids[i, j] = -1
# else: init to -1
# static valid draft token
valid_draft_count = 0
for j in range(k):
if draft_token_ids[i, j] != -1:
valid_draft_count += 1
else:
break
num_valid_draft_tokens[i] = valid_draft_count
return token_ids, next_token_ids, draft_token_ids, num_valid_draft_tokens
# ---------------------------------------------------------------------------
# inputs construct helper
# ---------------------------------------------------------------------------
def _make_inputs(
batch_size: int,
seq_len: int,
max_new_tokens: int,
k: int,
vocab_size: int = 32000,
min_n: int = 3,
max_n: int = 5,
discard_rate: float = 0.0,
invalid_rate: float = 0.0,
seed: int = SEED,
):
"""
Returns:
(token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask,
vocab_size, min_n, max_n, k)
"""
rng = np.random.RandomState(seed)
token_ids = rng.randint(0, vocab_size, size=(batch_size, seq_len), dtype=np.int32)
max_valid_tokens = seq_len - max_new_tokens
if max_valid_tokens < 1:
max_valid_tokens = 1
num_tokens_no_spec = rng.randint(1, max_valid_tokens + 1, size=(batch_size,), dtype=np.int32)
sampled_token_ids = rng.randint(0, vocab_size, size=(batch_size, max_new_tokens), dtype=np.int32)
if invalid_rate > 0:
invalid_mask = rng.rand(batch_size, max_new_tokens) < invalid_rate
sampled_token_ids[invalid_mask] = -1
# discard_request_mask
discard_request_mask = np.zeros(batch_size, dtype=np.int32)
if discard_rate > 0:
discard_mask = rng.rand(batch_size) < discard_rate
discard_request_mask[discard_mask] = 1
return token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k
def _run_npu(token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k):
token_ids_t = torch.from_numpy(token_ids).to("npu")
num_tokens_t = torch.from_numpy(num_tokens_no_spec).to("npu")
sampled_t = torch.from_numpy(sampled_token_ids).to("npu")
discard_t = torch.from_numpy(discard_request_mask).to("npu")
result = torch.ops._C_ascend.npu_ngram_spec_decode(
token_ids_t, num_tokens_t, sampled_t, discard_t, vocab_size, min_n, max_n, k
)
torch.npu.synchronize()
return result
def _measure_perf(token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k):
token_ids_t = torch.from_numpy(token_ids.copy()).to("npu")
num_tokens_t = torch.from_numpy(num_tokens_no_spec.copy()).to("npu")
sampled_t = torch.from_numpy(sampled_token_ids.copy()).to("npu")
discard_t = torch.from_numpy(discard_request_mask.copy()).to("npu")
for _ in range(PERF_WARMUP):
_ = torch.ops._C_ascend.npu_ngram_spec_decode(
token_ids_t, num_tokens_t, sampled_t, discard_t, vocab_size, min_n, max_n, k
)
torch.npu.synchronize()
t0 = time.perf_counter()
for _ in range(PERF_ITERS):
_ = torch.ops._C_ascend.npu_ngram_spec_decode(
token_ids_t, num_tokens_t, sampled_t, discard_t, vocab_size, min_n, max_n, k
)
torch.npu.synchronize()
elapsed_us = (time.perf_counter() - t0) * 1e6 / PERF_ITERS
print(
f" [perf] B={token_ids.shape[0]} M={token_ids.shape[1]} N={sampled_token_ids.shape[1]} "
f"k={k} min_n={min_n} max_n={max_n} -> {elapsed_us:.1f} us/call",
flush=True,
)
# ===========================================================================
# Group 1: basic - basic function
# ===========================================================================
@pytest.mark.parametrize(
"batch_size,seq_len,max_new_tokens,k",
[
(1, 16, 4, 3),
(4, 64, 8, 5),
(16, 128, 16, 5),
],
)
@torch.inference_mode()
def test_ngram_spec_decode_basic(batch_size, seq_len, max_new_tokens, k):
inputs = _make_inputs(batch_size, seq_len, max_new_tokens, k)
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k = inputs
# CPU golden
golden_ids, golden_next, golden_draft, golden_valid = golden_ngram_spec_decode(
token_ids.copy(),
num_tokens_no_spec.copy(),
sampled_token_ids.copy(),
discard_request_mask.copy(),
vocab_size,
min_n,
max_n,
k,
)
# NPU
result_ids, result_next, result_draft, result_valid = _run_npu(
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k
)
# compare
assert result_ids.cpu().numpy().tolist() == golden_ids.tolist(), f"token_ids mismatch: B={batch_size}"
assert result_next.cpu().numpy().tolist() == golden_next.tolist(), f"next_token_ids mismatch: B={batch_size}"
assert result_draft.cpu().numpy().tolist() == golden_draft.tolist(), f"draft_token_ids mismatch: B={batch_size}"
assert result_valid.cpu().numpy().tolist() == golden_valid.tolist(), (
f"num_valid_draft_tokens mismatch: B={batch_size}"
)
_measure_perf(token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k)
# ===========================================================================
# Group 2: padding / optional
# ===========================================================================
@pytest.mark.parametrize(
"batch_size,seq_len,max_new_tokens,k,invalid_rate",
[
(4, 64, 8, 5, 0.3),
(4, 64, 8, 5, 0.7),
(4, 128, 16, 5, 0.5),
],
)
@torch.inference_mode()
def test_ngram_spec_decode_padding(batch_size, seq_len, max_new_tokens, k, invalid_rate):
inputs = _make_inputs(batch_size, seq_len, max_new_tokens, k, invalid_rate=invalid_rate)
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k = inputs
golden_ids, golden_next, golden_draft, golden_valid = golden_ngram_spec_decode(
token_ids.copy(),
num_tokens_no_spec.copy(),
sampled_token_ids.copy(),
discard_request_mask.copy(),
vocab_size,
min_n,
max_n,
k,
)
result_ids, result_next, result_draft, result_valid = _run_npu(
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k
)
assert result_ids.cpu().numpy().tolist() == golden_ids.tolist()
assert result_next.cpu().numpy().tolist() == golden_next.tolist()
assert result_draft.cpu().numpy().tolist() == golden_draft.tolist()
assert result_valid.cpu().numpy().tolist() == golden_valid.tolist()
_measure_perf(token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k)
@pytest.mark.parametrize("discard_rate", [0.2, 0.5, 1.0])
@torch.inference_mode()
def test_ngram_spec_decode_discard(discard_rate):
inputs = _make_inputs(4, 64, 8, 5, discard_rate=discard_rate)
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k = inputs
golden_ids, golden_next, golden_draft, golden_valid = golden_ngram_spec_decode(
token_ids.copy(),
num_tokens_no_spec.copy(),
sampled_token_ids.copy(),
discard_request_mask.copy(),
vocab_size,
min_n,
max_n,
k,
)
result_ids, result_next, result_draft, result_valid = _run_npu(
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k
)
assert result_ids.cpu().numpy().tolist() == golden_ids.tolist()
assert result_next.cpu().numpy().tolist() == golden_next.tolist()
assert result_draft.cpu().numpy().tolist() == golden_draft.tolist()
assert result_valid.cpu().numpy().tolist() == golden_valid.tolist()
# ===========================================================================
# Group 3: min_n / max_n / k
# ===========================================================================
@pytest.mark.parametrize(
"min_n,max_n,k",
[
(1, 1, 3),
(2, 4, 5),
(1, 8, 10),
(3, 3, 1),
(5, 10, 8),
],
)
@torch.inference_mode()
def test_ngram_spec_decode_attrs(min_n, max_n, k):
inputs = _make_inputs(4, 128, 16, k, min_n=min_n, max_n=max_n)
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, _, _, _ = inputs
golden_ids, golden_next, golden_draft, golden_valid = golden_ngram_spec_decode(
token_ids.copy(),
num_tokens_no_spec.copy(),
sampled_token_ids.copy(),
discard_request_mask.copy(),
vocab_size,
min_n,
max_n,
k,
)
result_ids, result_next, result_draft, result_valid = _run_npu(
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k
)
assert result_ids.cpu().numpy().tolist() == golden_ids.tolist()
assert result_next.cpu().numpy().tolist() == golden_next.tolist()
assert result_draft.cpu().numpy().tolist() == golden_draft.tolist()
assert result_valid.cpu().numpy().tolist() == golden_valid.tolist()
_measure_perf(token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k)
# ===========================================================================
# Group 4: large scale
# ===========================================================================
@torch.inference_mode()
def test_ngram_spec_decode_prefill():
inputs = _make_inputs(1, 2048, 16, 5, vocab_size=32000)
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k = inputs
golden_ids, golden_next, golden_draft, golden_valid = golden_ngram_spec_decode(
token_ids.copy(),
num_tokens_no_spec.copy(),
sampled_token_ids.copy(),
discard_request_mask.copy(),
vocab_size,
min_n,
max_n,
k,
)
result_ids, result_next, result_draft, result_valid = _run_npu(
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k
)
assert result_ids.cpu().numpy().tolist() == golden_ids.tolist()
assert result_next.cpu().numpy().tolist() == golden_next.tolist()
assert result_draft.cpu().numpy().tolist() == golden_draft.tolist()
assert result_valid.cpu().numpy().tolist() == golden_valid.tolist()
_measure_perf(token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k)
@torch.inference_mode()
def test_ngram_spec_decode_decode():
inputs = _make_inputs(64, 32, 5, 3, vocab_size=32000)
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k = inputs
golden_ids, golden_next, golden_draft, golden_valid = golden_ngram_spec_decode(
token_ids.copy(),
num_tokens_no_spec.copy(),
sampled_token_ids.copy(),
discard_request_mask.copy(),
vocab_size,
min_n,
max_n,
k,
)
result_ids, result_next, result_draft, result_valid = _run_npu(
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k
)
assert result_ids.cpu().numpy().tolist() == golden_ids.tolist()
assert result_next.cpu().numpy().tolist() == golden_next.tolist()
assert result_draft.cpu().numpy().tolist() == golden_draft.tolist()
assert result_valid.cpu().numpy().tolist() == golden_valid.tolist()
_measure_perf(token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k)
# ===========================================================================
# Group 5: boundary
# ===========================================================================
@torch.inference_mode()
def test_ngram_spec_decode_minimal():
inputs = _make_inputs(1, 4, 1, 1, vocab_size=100)
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k = inputs
golden_ids, golden_next, golden_draft, golden_valid = golden_ngram_spec_decode(
token_ids.copy(),
num_tokens_no_spec.copy(),
sampled_token_ids.copy(),
discard_request_mask.copy(),
vocab_size,
min_n,
max_n,
k,
)
result_ids, result_next, result_draft, result_valid = _run_npu(
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k
)
assert result_ids.cpu().numpy().tolist() == golden_ids.tolist()
assert result_next.cpu().numpy().tolist() == golden_next.tolist()
assert result_draft.cpu().numpy().tolist() == golden_draft.tolist()
assert result_valid.cpu().numpy().tolist() == golden_valid.tolist()
@torch.inference_mode()
def test_ngram_spec_decode_no_valid_sampled():
inputs = _make_inputs(4, 64, 8, 5)
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k = inputs
sampled_token_ids[:] = -1
golden_ids, golden_next, golden_draft, golden_valid = golden_ngram_spec_decode(
token_ids.copy(),
num_tokens_no_spec.copy(),
sampled_token_ids.copy(),
discard_request_mask.copy(),
vocab_size,
min_n,
max_n,
k,
)
result_ids, result_next, result_draft, result_valid = _run_npu(
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k
)
assert result_ids.cpu().numpy().tolist() == golden_ids.tolist()
assert result_next.cpu().numpy().tolist() == golden_next.tolist()
assert result_draft.cpu().numpy().tolist() == golden_draft.tolist()
assert result_valid.cpu().numpy().tolist() == golden_valid.tolist()
@torch.inference_mode()
def test_ngram_spec_decode_exact_match():
k = 3
vocab_size = 1000
token_ids = np.array(
[
[1, 2, 3, 1, 2, 3, 4, 5, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0],
[10, 20, 10, 20, 30, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0],
],
dtype=np.int32,
)
num_tokens_no_spec = np.array([6, 5], dtype=np.int32)
# sampled tokens
sampled_token_ids = np.array(
[
[1, 2, 3, -1],
[10, 20, -1, -1],
],
dtype=np.int32,
)
discard_request_mask = np.array([0, 0], dtype=np.int32)
golden_ids, golden_next, golden_draft, golden_valid = golden_ngram_spec_decode(
token_ids.copy(),
num_tokens_no_spec.copy(),
sampled_token_ids.copy(),
discard_request_mask.copy(),
vocab_size,
3,
5,
k,
)
result_ids, result_next, result_draft, result_valid = _run_npu(
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, 3, 5, k
)
assert result_ids.cpu().numpy().tolist() == golden_ids.tolist()
assert result_next.cpu().numpy().tolist() == golden_next.tolist()
assert result_draft.cpu().numpy().tolist() == golden_draft.tolist()
assert result_valid.cpu().numpy().tolist() == golden_valid.tolist()
@torch.inference_mode()
def test_ngram_spec_decode_full_capacity():
inputs = _make_inputs(4, 48, 16, 5, min_n=2, max_n=4)
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k = inputs
num_tokens_no_spec[:] = token_ids.shape[1] - sampled_token_ids.shape[1]
golden_ids, golden_next, golden_draft, golden_valid = golden_ngram_spec_decode(
token_ids.copy(),
num_tokens_no_spec.copy(),
sampled_token_ids.copy(),
discard_request_mask.copy(),
vocab_size,
min_n,
max_n,
k,
)
result_ids, result_next, result_draft, result_valid = _run_npu(
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k
)
assert result_ids.cpu().numpy().tolist() == golden_ids.tolist()
assert result_next.cpu().numpy().tolist() == golden_next.tolist()
assert result_draft.cpu().numpy().tolist() == golden_draft.tolist()
assert result_valid.cpu().numpy().tolist() == golden_valid.tolist()
@torch.inference_mode()
def test_ngram_spec_decode_k1():
inputs = _make_inputs(4, 64, 8, 1, min_n=2, max_n=3)
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k = inputs
golden_ids, golden_next, golden_draft, golden_valid = golden_ngram_spec_decode(
token_ids.copy(),
num_tokens_no_spec.copy(),
sampled_token_ids.copy(),
discard_request_mask.copy(),
vocab_size,
min_n,
max_n,
k,
)
result_ids, result_next, result_draft, result_valid = _run_npu(
token_ids, num_tokens_no_spec, sampled_token_ids, discard_request_mask, vocab_size, min_n, max_n, k
)
assert result_ids.cpu().numpy().tolist() == golden_ids.tolist()
assert result_next.cpu().numpy().tolist() == golden_next.tolist()
assert result_draft.cpu().numpy().tolist() == golden_draft.tolist()
assert result_valid.cpu().numpy().tolist() == golden_valid.tolist()

View File

@@ -0,0 +1,84 @@
import gc
import torch
import torch_npu
from vllm_ascend.utils import enable_custom_op
torch_npu.npu.config.allow_internal_format = True
enable_custom_op()
HC_MULT = 4
DSV4_FLASH_HIDDEN_SIZE = 4096
HIDDEN_SIZE = DSV4_FLASH_HIDDEN_SIZE
EXTENDED_HIDDEN_SIZE = 7168
MIX_HC = 24
HC_SINKHORN_ITERS = 20
NORM_EPS = 1e-6
HC_EPS = 1e-6
def _make_hc_pre_inputs(shape: tuple[int, ...]):
torch.manual_seed(1024)
hidden_size = shape[-1]
x = torch.randn(shape, dtype=torch.bfloat16, device="npu")
hc_fn = (
torch.randn(
MIX_HC,
HC_MULT * hidden_size,
dtype=torch.float32,
device="npu",
)
* 0.01
)
hc_scale = torch.randn(3, dtype=torch.float32, device="npu") * 0.01
hc_base = torch.randn(MIX_HC, dtype=torch.float32, device="npu") * 0.01
return x, hc_fn, hc_scale, hc_base
def _compare_hc_pre_outputs(shape: tuple[int, ...]):
x, hc_fn, hc_scale, hc_base = _make_hc_pre_inputs(shape)
expected = torch.ops._C_ascend.npu_hc_pre(x, hc_fn, hc_scale, hc_base, HC_MULT, HC_SINKHORN_ITERS, NORM_EPS, HC_EPS)
actual = torch.ops._C_ascend.npu_hc_pre_v2(
x, hc_fn, hc_scale, hc_base, HC_MULT, HC_SINKHORN_ITERS, NORM_EPS, HC_EPS
)
for actual_tensor, expected_tensor in zip(actual, expected, strict=True):
torch.testing.assert_close(
actual_tensor.cpu(),
expected_tensor.cpu(),
atol=5e-2,
rtol=5e-2,
)
@torch.inference_mode()
def test_npu_hc_pre_v1_v2_bf16_3d_input():
_compare_hc_pre_outputs((2, HC_MULT, HIDDEN_SIZE))
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@torch.inference_mode()
def test_npu_hc_pre_v1_v2_bf16_4d_input():
_compare_hc_pre_outputs((1, 2, HC_MULT, HIDDEN_SIZE))
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@torch.inference_mode()
def test_npu_hc_pre_v1_v2_bf16_dsv4_flash_hidden_size():
_compare_hc_pre_outputs((4, HC_MULT, DSV4_FLASH_HIDDEN_SIZE))
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@torch.inference_mode()
def test_npu_hc_pre_v1_v2_bf16_extended_hidden_size():
_compare_hc_pre_outputs((2, HC_MULT, EXTENDED_HIDDEN_SIZE))
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,209 @@
import gc
import random
import numpy
import pytest
import torch
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
# Fix random seed to ensure test reproducibility
RTOL_TOLERANCE = 1e-5
ATOL_TOLERANCE = 1e-8
seed = 45
random.seed(seed)
numpy.random.seed(seed)
torch.manual_seed(seed)
def softmax_func(x, axis=None):
"""Softmax implementation (adapted for numpy calculation)"""
if "float16" in x.dtype.name:
x = x.astype(numpy.float32)
x_max = x.max(axis=axis, keepdims=True)
x_sub = x - x_max
y = numpy.exp(x_sub)
x_sum = y.sum(axis=axis, keepdims=True)
res = y / x_sum
return res, x_max, x_sum
def moe_gating_top_k_numpy_ref(
x: torch.Tensor,
k: int,
bias: torch.Tensor | None,
k_group: int = 1,
group_count: int = 1,
group_select_mode: int = 0,
renorm: int = 0,
norm_type: int = 0,
y2_flag: bool = False,
routed_scaling_factor: float = 1.0,
eps: float = 1e-20,
) -> tuple:
"""NumPy reference implementation of MOE Gating TopK.
For result comparison with NPU operator, ensure the consistency
between NPU kernel and baseline implementation.
Args:
x: Input tensor of shape (num_tokens, num_experts)
k: Number of top-k experts to select
bias: Bias tensor of shape (num_experts,) (optional)
k_group: Number of top-k groups to select
group_count: Number of expert groups
group_select_mode: Group selection mode (0: max, 1: top2 sum)
renorm: Whether to renormalize the output (0/1)
norm_type: Normalization type (0: softmax, 1: sigmoid)
y2_flag: Whether to output original x as y2
routed_scaling_factor: Scaling factor for routing weights
eps: Small epsilon to avoid division by zero
Returns:
tuple: (y, indices, y2)
- y: Top-k weights of shape (num_tokens, k)
- indices: Top-k expert indices of shape (num_tokens, k)
- y2: Original x if y2_flag is True, else None
"""
dtype = x.dtype
if dtype != torch.float32:
x = x.to(dtype=torch.float32)
if bias is not None:
bias = bias.to(dtype=torch.float32)
x = x.numpy()
if bias is not None:
bias = bias.numpy()
if norm_type == 0: # softmax normalization
x, _, _ = softmax_func(x, -1)
else: # sigmoid normalization
x = 1 / (1 + numpy.exp(-x))
original_x = x
if bias is not None:
x = x + bias
if group_count > 1:
x = x.reshape(x.shape[0], group_count, -1)
if group_select_mode == 0:
group_x = numpy.amax(x, axis=-1)
else:
group_x = numpy.partition(x, -2, axis=-1)[..., -2:].sum(axis=-1)
indices = numpy.argsort(-group_x, axis=-1, kind="stable")[:, :k_group]
mask = numpy.ones((x.shape[0], group_count), dtype=bool)
mask[numpy.arange(x.shape[0])[:, None], indices] = False
x = numpy.where(mask[..., None], float("-inf"), x)
x = x.reshape(x.shape[0], -1)
_, indices = torch.sort(torch.from_numpy(x), dim=-1, stable=True, descending=True)
indices = numpy.asarray(indices[:, :k])
y = numpy.take_along_axis(original_x, indices, axis=1)
if norm_type == 1 or renorm == 1:
y /= numpy.sum(y, axis=-1, keepdims=True) + eps
y *= routed_scaling_factor
y2 = original_x if y2_flag else None
y = torch.tensor(y, dtype=dtype)
return y, indices.astype(numpy.int32), y2
# pytest parameterized decorators (cover all test scenarios)
@pytest.mark.parametrize("group_select_mode", [0, 1])
@pytest.mark.parametrize("renorm", [1])
@pytest.mark.parametrize("norm_type", [0, 1])
@pytest.mark.parametrize("group_count", [1, 8])
@pytest.mark.parametrize("k_ranges", [4, 16, 32])
@pytest.mark.parametrize("x_dim0_range", [1, 8, 16])
@pytest.mark.parametrize("x_dim1_range", [64, 128, 256])
def test_npu_moe_gating_topk_compare(
group_select_mode: int,
renorm: int,
norm_type: int,
group_count: int,
k_ranges: int,
x_dim0_range: int,
x_dim1_range: int,
device: str = "npu",
):
"""Ascend NPU MOE Gating TopK operator test.
Compare NPU kernel results with NumPy reference implementation
to verify the correctness of Ascend custom op.
Args:
group_select_mode: Group selection mode (0: max, 1: top2 sum)
renorm: Whether to renormalize output (fixed to 1 in test)
norm_type: Normalization type (0: softmax, 1: sigmoid)
group_count: Number of expert groups
k_ranges: Number of top-k experts to select
x_dim0_range: First dimension of input tensor (num_tokens)
x_dim1_range: Second dimension of input tensor (num_experts)
device: Target device (fixed to "npu" in test)
"""
# Simplify parameter names for better readability
k = k_ranges
dim0 = x_dim0_range
dim1 = x_dim1_range
# Skip invalid cases: k cannot exceed num_experts per group
if k > dim1 // group_count:
return
# Construct test inputs
x = numpy.random.uniform(-2, 2, (dim0, dim1)).astype(numpy.float32)
bias = numpy.random.uniform(-2, 2, (dim1,)).astype(numpy.float32)
x_tensor = torch.tensor(x, dtype=torch.float32)
bias_tensor = torch.tensor(bias, dtype=torch.float32)
# Fix k_group value to avoid irreproducibility caused by random.randint
k_group = min(1, group_count)
out_flag = False
routed_scaling_factor = 1.0
eps = 1e-20
# Calculate NumPy reference results
y, expert_idx, out = moe_gating_top_k_numpy_ref(
x_tensor,
k=k,
bias=bias_tensor,
k_group=k_group,
group_count=group_count,
group_select_mode=group_select_mode,
renorm=renorm,
norm_type=norm_type,
y2_flag=out_flag,
routed_scaling_factor=routed_scaling_factor,
eps=eps,
)
# Calculate NPU operator results
y_npu, expert_idx_npu, out_npu = torch.ops._C_ascend.moe_gating_top_k(
x_tensor.npu(),
k=k,
k_group=k_group,
group_count=group_count,
group_select_mode=group_select_mode,
renorm=renorm,
norm_type=norm_type,
out_flag=out_flag,
routed_scaling_factor=routed_scaling_factor,
eps=eps,
bias_opt=bias_tensor.npu() if bias_tensor is not None else None,
)
# Verify consistency between NPU and NumPy results
assert numpy.allclose(y.cpu().numpy(), y_npu.cpu().numpy(), rtol=RTOL_TOLERANCE, atol=ATOL_TOLERANCE)
assert numpy.allclose(expert_idx, expert_idx_npu.cpu().numpy(), rtol=RTOL_TOLERANCE, atol=ATOL_TOLERANCE)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
if __name__ == "__main__":
# Execute pytest tests with verbose output
pytest.main([__file__, "-sv"])

View File

@@ -0,0 +1,269 @@
import gc
import random
import numpy as np
import pytest
import torch
import torch_npu
torch_npu.npu.set_compile_mode(jit_compile=False)
seed = 42
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
def golden_recurrent_gated_delta_rule(
query,
key,
value,
state,
beta,
scale,
actual_seq_lengths,
ssm_state_indices,
g,
num_accepted_tokens,
):
"""Pure torch/CPU golden implementation of recurrent gated delta rule.
Args:
query: [T, nk, dk]
key: [T, nk, dk]
value: [T, nv, dv]
state: [S, nv, dv, dk]
beta: [T, nv]
scale: float
actual_seq_lengths: [batch_size] per-sequence lengths
ssm_state_indices: [T] per-token state block index
g: [T, nv] or None
num_accepted_tokens: [batch_size] or None
Returns:
(output [T, nv, dv], updated_state [S, nv, dv, dk])
"""
q = query.to(torch.float32)
k = key.to(torch.float32)
v = value.to(torch.float32)
initial_state = state.clone().to(torch.float32)
T, n_heads_v, Dv = v.shape
n_heads_k = q.shape[-2]
g = torch.ones(T, n_heads_v).to(torch.float32) if g is None else g.to(torch.float32).exp()
beta = torch.ones(T, n_heads_v).to(torch.float32) if beta is None else beta.to(torch.float32)
o = torch.empty_like(v).to(torch.float32)
if scale is None:
scale = k.shape[-1] ** -0.5
q = q * scale
seq_start = 0
for i in range(len(actual_seq_lengths)):
if num_accepted_tokens is None:
init_state = initial_state[ssm_state_indices[seq_start]]
else:
init_state = initial_state[ssm_state_indices[seq_start + num_accepted_tokens[i] - 1]]
for head_id in range(n_heads_v):
S = init_state[head_id]
for slot_id in range(seq_start, seq_start + actual_seq_lengths[i]):
q_i = q[slot_id][head_id // (n_heads_v // n_heads_k)]
k_i = k[slot_id][head_id // (n_heads_v // n_heads_k)]
v_i = v[slot_id][head_id]
alpha_i = g[slot_id][head_id]
beta_i = beta[slot_id][head_id]
S = S * alpha_i
x = (S * k_i.unsqueeze(-2)).sum(dim=-1)
y = (v_i - x) * beta_i
S_ = y[:, None] * k_i[None, :]
S = S + S_
initial_state[ssm_state_indices[slot_id]][head_id] = S
o[slot_id][head_id] = (S * q_i.unsqueeze(-2)).sum(dim=-1)
seq_start += actual_seq_lengths[i]
return o.to(query.dtype), initial_state.to(query.dtype)
@pytest.mark.parametrize("batch_size", [1, 4, 8])
@pytest.mark.parametrize("mtp", [1, 2])
@pytest.mark.parametrize("headnum", [(4, 8), (8, 16), (16, 32)])
@pytest.mark.parametrize("headdim_k", [128])
@pytest.mark.parametrize("headdim_v", [128])
@pytest.mark.parametrize("state_dtype", [torch.bfloat16, torch.float32])
def test_recurrent_gated_delta_rule(
batch_size,
mtp,
headnum,
headdim_k,
headdim_v,
state_dtype,
):
torch.manual_seed(seed)
dtype = torch.bfloat16
headnum_k, headnum_v = headnum
seq_lengths = torch.ones(batch_size, dtype=torch.int32) * mtp
T = int(torch.sum(seq_lengths))
state = torch.rand((T, headnum_v, headdim_v, headdim_k)).to(state_dtype)
query = torch.nn.functional.normalize(
torch.rand((T, headnum_k, headdim_k)),
p=2,
dim=-1,
).to(dtype)
key = torch.nn.functional.normalize(
torch.rand((T, headnum_k, headdim_k)),
p=2,
dim=-1,
).to(dtype)
value = torch.rand((T, headnum_v, headdim_v)).to(dtype)
g = torch.rand((T, headnum_v), dtype=torch.float32)
beta = torch.rand((T, headnum_v)).to(dtype)
ssm_state_indices = torch.arange(T, dtype=torch.int32)
num_accepted_tokens = torch.randint(1, mtp + 1, (batch_size,), dtype=torch.int32)
scale = headdim_k**-0.5
out_golden, state_golden = golden_recurrent_gated_delta_rule(
query,
key,
value,
state,
beta,
scale,
seq_lengths,
ssm_state_indices,
g,
num_accepted_tokens,
)
out_golden = out_golden.to(torch.float32)
state_golden = state_golden.to(torch.float32)
# torch_npu op expects actual_seq_lengths = [start_pos, len1, len2, ..., lenB]
actual_seq_lengths_npu = torch.cat(
[
torch.zeros(1, dtype=torch.int32),
seq_lengths,
]
)
state_npu = state.npu()
npu_out = torch.ops._C_ascend.npu_recurrent_gated_delta_rule(
query=query.npu(),
key=key.npu(),
value=value.npu(),
g=g.npu(),
beta=beta.npu(),
state=state_npu,
scale=scale,
actual_seq_lengths=actual_seq_lengths_npu.npu(),
ssm_state_indices=ssm_state_indices.npu(),
num_accepted_tokens=num_accepted_tokens.npu(),
)
torch.testing.assert_close(
npu_out.to(torch.float32).cpu(),
out_golden,
rtol=3e-3,
atol=1e-2,
equal_nan=True,
)
torch.testing.assert_close(
state_npu.to(torch.float32).cpu(),
state_golden,
rtol=3e-3,
atol=1e-2,
equal_nan=True,
)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@pytest.mark.parametrize("batch_size", [1, 4, 8])
@pytest.mark.parametrize("mtp", [1, 2])
@pytest.mark.parametrize("headnum", [(4, 8), (8, 16)])
@pytest.mark.parametrize("headdim_k", [128])
@pytest.mark.parametrize("headdim_v", [128])
@pytest.mark.parametrize("state_dtype", [torch.bfloat16, torch.float32])
def test_recurrent_gated_delta_rule_no_accepted(
batch_size,
mtp,
headnum,
headdim_k,
headdim_v,
state_dtype,
):
torch.manual_seed(seed)
dtype = torch.bfloat16
headnum_k, headnum_v = headnum
seq_lengths = torch.ones(batch_size, dtype=torch.int32) * mtp
T = int(torch.sum(seq_lengths))
state = torch.rand((T, headnum_v, headdim_v, headdim_k)).to(state_dtype)
query = torch.nn.functional.normalize(
torch.rand((T, headnum_k, headdim_k)),
p=2,
dim=-1,
).to(dtype)
key = torch.nn.functional.normalize(
torch.rand((T, headnum_k, headdim_k)),
p=2,
dim=-1,
).to(dtype)
value = torch.rand((T, headnum_v, headdim_v)).to(dtype)
g = torch.rand((T, headnum_v), dtype=torch.float32)
beta = torch.rand((T, headnum_v)).to(dtype)
ssm_state_indices = torch.arange(T, dtype=torch.int32)
scale = headdim_k**-0.5
out_golden, state_golden = golden_recurrent_gated_delta_rule(
query,
key,
value,
state,
beta,
scale,
seq_lengths,
ssm_state_indices,
g,
None,
)
out_golden = out_golden.to(torch.float32)
state_golden = state_golden.to(torch.float32)
actual_seq_lengths_npu = torch.cat(
[
torch.zeros(1, dtype=torch.int32),
seq_lengths,
]
)
state_npu = state.npu()
npu_out = torch.ops._C_ascend.npu_recurrent_gated_delta_rule(
query=query.npu(),
key=key.npu(),
value=value.npu(),
g=g.npu(),
beta=beta.npu(),
state=state_npu,
scale=scale,
actual_seq_lengths=actual_seq_lengths_npu.npu(),
ssm_state_indices=ssm_state_indices.npu(),
)
torch.testing.assert_close(
npu_out.to(torch.float32).cpu(),
out_golden,
rtol=3e-3,
atol=1e-2,
equal_nan=True,
)
torch.testing.assert_close(
state_npu.to(torch.float32).cpu(),
state_golden,
rtol=3e-3,
atol=1e-2,
equal_nan=True,
)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,141 @@
import pytest
import torch
import torch_npu
from vllm_ascend.utils import enable_custom_op
from vllm_ascend.utils import is_310p as is_310p_hw
torch_npu.npu.set_compile_mode(jit_compile=False)
def npu_recurrent_gated_delta_rule_310(
query,
key,
value,
beta,
state,
actual_seq_lengths,
ssm_state_indices,
g=None,
gk=None,
num_accepted_tokens=None,
scale=1.0,
):
"""Call RecurrentGatedDeltaRule."""
out = torch.ops._C_ascend.npu_recurrent_gated_delta_rule_310(
query=query,
key=key,
value=value,
g=g,
gk=gk,
beta=beta,
state=state,
actual_seq_lengths=actual_seq_lengths,
ssm_state_indices=ssm_state_indices,
num_accepted_tokens=num_accepted_tokens,
scale_value=scale,
)
return out
def golden_recurrent_gated_delta_rule(
query, key, value, state, beta, scale, actual_seq_lengths, ssm_state_indices, g, num_accepted_tokens
):
q = query.to(torch.float32)
k = key.to(torch.float32)
v = value.to(torch.float32)
initial_state = state.clone().to(torch.float32)
T, n_heads_v, Dv = v.shape
n_heads_k = q.shape[-2]
g = torch.ones(T, n_heads_v).to(torch.float32) if g is None else g.to(torch.float32).exp()
beta = torch.ones(T, n_heads_v).to(torch.float32) if beta is None else beta.to(torch.float32)
o = torch.empty_like(v).to(torch.float32)
if scale is None:
scale = k.shape[-1] ** -0.5
q = q * scale
seq_start = 0
for i in range(len(actual_seq_lengths)):
if num_accepted_tokens is None:
init_state = initial_state[ssm_state_indices[seq_start]]
else:
init_state = initial_state[ssm_state_indices[seq_start + num_accepted_tokens[i] - 1]]
for head_id in range(n_heads_v):
S = init_state[head_id]
for slot_id in range(seq_start, seq_start + actual_seq_lengths[i]):
q_i = q[slot_id][head_id // (n_heads_v // n_heads_k)]
k_i = k[slot_id][head_id // (n_heads_v // n_heads_k)]
v_i = v[slot_id][head_id]
alpha_i = g[slot_id][head_id]
beta_i = beta[slot_id][head_id]
S = S * alpha_i
x = (S * k_i.unsqueeze(-2)).sum(dim=-1)
y = (v_i - x) * beta_i
S_ = y[:, None] * k_i[None, :]
S = S + S_
initial_state[ssm_state_indices[slot_id]][head_id] = S
o[slot_id][head_id] = (S * q_i.unsqueeze(-2)).sum(dim=-1)
seq_start += actual_seq_lengths[i]
return o.to(query.dtype), initial_state.to(query.dtype)
@pytest.mark.skipif(not is_310p_hw(), reason="Tested separately on a 310P machine.")
@pytest.mark.parametrize("batch_size", [1, 4, 8])
@pytest.mark.parametrize("mtp", [1, 2])
@pytest.mark.parametrize("headnum", [(4, 8), (8, 16), (16, 32)])
@pytest.mark.parametrize("headdim_k", [128])
@pytest.mark.parametrize("headdim_v", [128])
def test_fused_recurrent_gated_delta_rule_310(batch_size, mtp, headnum, headdim_k, headdim_v):
enable_custom_op()
dtype = torch.float16
headnum_k, headnum_v = headnum
actual_seq_lengths = torch.ones(batch_size, dtype=torch.int32) * mtp
T = int(torch.sum(actual_seq_lengths))
state = torch.rand((T, headnum_v, headdim_v, headdim_k)).to(dtype)
query = torch.nn.functional.normalize(torch.rand((T, headnum_k, headdim_k)), p=2, dim=-1).to(dtype)
key = torch.nn.functional.normalize(torch.rand((T, headnum_k, headdim_k)), p=2, dim=-1).to(dtype)
value = torch.rand((T, headnum_v, headdim_v)).to(dtype)
g = torch.rand((T, headnum_v), dtype=torch.float32)
beta = torch.rand((T, headnum_v)).to(dtype)
ssm_state_indices = torch.arange(T, dtype=torch.int32)
num_accepted_tokens = torch.randint(1, mtp + 1, (batch_size,), dtype=torch.int32)
scale = headdim_k**-0.5
out_golden, state_golden = golden_recurrent_gated_delta_rule(
query, key, value, state, beta, scale, actual_seq_lengths, ssm_state_indices, g, num_accepted_tokens
)
out_golden = out_golden.to(torch.float32)
state_golden = state_golden.to(torch.float32)
state_npu = state.npu()
out = npu_recurrent_gated_delta_rule_310(
query.npu(),
key.npu(),
value.npu(),
beta.npu(),
state_npu,
actual_seq_lengths.npu(),
ssm_state_indices.npu(),
g=g.npu(),
num_accepted_tokens=num_accepted_tokens.npu(),
scale=scale,
)
out = out.to(torch.float32).cpu()
torch.testing.assert_close(
out.to(torch.float32).cpu(),
out_golden.to(torch.float32).cpu(),
rtol=3e-3,
atol=1e-2,
equal_nan=True,
)
torch.testing.assert_close(
state_npu.to(torch.float32).cpu(),
state_golden.to(torch.float32).cpu(),
rtol=3e-3,
atol=1e-2,
equal_nan=True,
)

View File

@@ -0,0 +1,188 @@
import logging
import unittest
from unittest import TestCase
import npugraph_ex as nge
import numpy as np
import torch
import torch.nn as nn
from npugraph_ex.core.utils import logger
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
logger.setLevel(logging.DEBUG)
torch._logging.set_logs(graph_code=True)
def run_tests():
unittest.main()
class TestCustomReshapeAndCacheBnsd(TestCase):
def setUp(self):
torch.manual_seed(42)
np.random.seed(42)
self.device_id = 0
torch.npu.set_device(self.device_id)
self.npu = f"npu:{self.device_id}"
self.num_blocks = 2
self.num_kv_heads = 2
self.block_size = 4
self.head_size = 16
self.token_num = 4
self.bs = 1
self.hashk_op = torch.randn(self.token_num * self.num_kv_heads, self.head_size, dtype=torch.float16).npu()
self.hashk_cache_op = torch.randn(
self.num_blocks, self.num_kv_heads, self.block_size, self.head_size, dtype=torch.float16
).npu()
self.slot_mapping_op = torch.tensor([0, 1, 2, 3], dtype=torch.int32).npu()
self.seq_lens_op = torch.tensor([self.token_num], dtype=torch.int32).npu()
self.hashk_op_org = self.hashk_op.reshape(self.num_kv_heads, self.token_num, self.head_size)
def _run_eager_mode(self):
return torch.ops._C_ascend.npu_reshape_and_cache_bnsd(
self.hashk_op, self.hashk_cache_op, self.slot_mapping_op, self.seq_lens_op, self.hashk_cache_op
)
def _run_graph_mode(self):
class Network(nn.Module):
def __init__(self):
super().__init__()
def forward(self, hashk_op, hashk_cache_op, slot_mapping_op, seq_lens_op, k_cache_out):
return torch.ops._C_ascend.npu_reshape_and_cache_bnsd(
hashk_op, hashk_cache_op, slot_mapping_op, seq_lens_op, k_cache_out
)
npu_mode = Network().to(f"npu:{self.device_id}")
config = nge.CompilerConfig()
npu_backend = nge.get_npu_backend(compiler_config=config)
npu_mode = torch.compile(npu_mode, backend=npu_backend, dynamic=False)
return npu_mode(self.hashk_op, self.hashk_cache_op, self.slot_mapping_op, self.seq_lens_op, self.hashk_cache_op)
def _verify_results(self, output):
for token_id in range(self.token_num):
slot = self.slot_mapping_op[token_id].item()
block_idx = slot // self.block_size
block_offset = slot % self.block_size
for kv_head_id in range(self.num_kv_heads):
output_hashk = output[block_idx, kv_head_id, block_offset, :]
token_id = self.bs - 1
slot = self.slot_mapping_op[token_id].item() + 1
block_idx = slot // self.block_size
block_offset = slot % self.block_size
for kv_head_id in range(self.num_kv_heads):
output_hashk = output[block_idx, kv_head_id, block_offset, :]
self.assertIsNotNone(output_hashk)
def test_reshape_and_cache_bnsd_compare(self):
print("=========== test_reshape_and_cache_bnsd_compare begin ===================")
output_eager = self._run_eager_mode()
output_graph = self._run_graph_mode()
self.assertIsNotNone(output_eager, "output_eager should not be None")
self.assertIsNotNone(output_graph, "output_graph should not be None")
self.assertEqual(output_eager.shape, output_graph.shape)
self.assertTrue(torch.allclose(output_eager, output_graph, atol=1e-05))
print("=========== test_reshape_and_cache_bnsd_compare end ===================")
def test_reshape_and_cache_bnsd_with_expected_output(self):
print("=========== test_reshape_and_cache_bnsd_with_expected_output begin ===================")
num_blocks = 2
num_kv_heads = 2
block_size = 4
head_size = 8
bs = 1
seq_lens_list = [8]
token_num = sum(seq_lens_list)
key_in = torch.zeros(num_kv_heads, token_num, head_size, dtype=torch.uint8)
val = 0
for head_id in range(num_kv_heads):
for token_id in range(token_num):
for dim_id in range(head_size):
key_in[head_id, token_id, dim_id] = val
val += 1
key_in = key_in.reshape(num_kv_heads * token_num, head_size).npu()
key_cache_out = torch.randint(
100, 200, (num_blocks, num_kv_heads, block_size, head_size), dtype=torch.uint8
).npu()
slot_mapping = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7], dtype=torch.int32).npu()
seq_lens = torch.tensor(seq_lens_list, dtype=torch.int32).npu()
key_cache_out_cpu = key_cache_out.cpu().clone()
output = torch.ops._C_ascend.npu_reshape_and_cache_bnsd(
key_in, key_cache_out, slot_mapping, seq_lens, key_cache_out
)
expected = key_cache_out_cpu.clone()
key_in_cpu = key_in.cpu().view(num_kv_heads, token_num, head_size)
for token_id in range(token_num):
slot = slot_mapping[token_id].cpu().item()
block_idx = slot // block_size
block_offset = slot % block_size
for head_id in range(num_kv_heads):
src_data = key_in_cpu[head_id, token_id, :]
expected[block_idx, head_id, block_offset, :] = src_data
self.assertIsNotNone(output, "output should not be None")
self.assertEqual(output.shape, expected.shape, "output shape should match expected shape")
output_cpu = output.cpu()
self.assertTrue(
torch.equal(output_cpu, expected),
f"output should match expected values\noutput:\n{output_cpu}\nexpected:\n{expected}",
)
print(f"bs: {bs}, seq_lens: {seq_lens_list}, total_tokens: {token_num}")
print(f"key_in shape: {key_in.shape}")
print(f"key_cache_out shape: {key_cache_out.shape}")
print(f"key_cache_out (before):\n{key_cache_out_cpu}")
print(f"output:\n{output_cpu}")
print(f"expected:\n{expected}")
print("=========== test_reshape_and_cache_bnsd_with_expected_output end ===================")
def test_reshape_and_cache_bnsd_bf16_shape(self):
print("=========== test_reshape_and_cache_bnsd_bf16_shape begin ===================")
num_blocks = 1
num_kv_heads = 2
block_size = 4
head_size = 4
token_num = 4
key_in = torch.zeros(num_kv_heads, token_num, head_size, dtype=torch.bfloat16)
for head_id in range(num_kv_heads):
for token_id in range(token_num):
start_val = head_id * token_num * head_size + token_id * head_size
key_in[head_id, token_id, :] = torch.arange(start_val, start_val + head_size, dtype=torch.bfloat16)
key_in = key_in.reshape(num_kv_heads * token_num, head_size).npu()
key_cache_out = torch.zeros(num_blocks, num_kv_heads, block_size, head_size, dtype=torch.bfloat16).npu()
slot_mapping = torch.tensor([0, 1, 2, 3], dtype=torch.int32).npu()
seq_lens = torch.tensor([token_num], dtype=torch.int32).npu()
output = torch.ops._C_ascend.npu_reshape_and_cache_bnsd(
key_in, key_cache_out, slot_mapping, seq_lens, key_cache_out
)
expected_shape = (num_blocks, num_kv_heads, block_size, head_size)
self.assertIsNotNone(output, "output should not be None")
self.assertEqual(output.shape, expected_shape, "output shape should match expected shape")
print(f"output shape: {output.shape}")
print(f"output:\n{output.cpu()}")
print("=========== test_reshape_and_cache_bnsd_bf16_shape end ===================")
if __name__ == "__main__":
run_tests()

View File

@@ -0,0 +1,104 @@
import gc
import time
import numpy as np
import pytest
import torch
import torch_npu
from vllm_ascend.utils import enable_custom_op
torch.set_printoptions(threshold=np.inf)
enable_custom_op()
def cal_slot(key, key_cache, slot_mapping, block_size):
key_expect = key_cache.clone()
for i, slot in enumerate(slot_mapping):
if slot < 0:
continue
token_key = key[i]
block_index = slot // block_size
block_offset = slot % block_size
key_expect[block_index][block_offset] = token_key
return key_expect.npu()
def cal_scatternd(key, key_cache, slot_mapping, block_size):
key_expect = key_cache.clone()
for i, slot in enumerate(slot_mapping):
if slot < 0:
continue
token_key = key[i]
key_expect[slot] = token_key
return key_expect.npu()
@pytest.mark.parametrize("num_tokens", [16]) # 6398
@pytest.mark.parametrize("num_head", [1]) # 512
@pytest.mark.parametrize("block_size", [128]) # 128
@pytest.mark.parametrize("num_blocks", [1773]) # 1599
@pytest.mark.parametrize("count", [1])
def test_scatter(num_tokens, num_head, block_size, num_blocks, count):
head_size_k = 64
key = torch.randint(low=0, high=128, size=(num_tokens, num_head, head_size_k), dtype=torch.int8).npu()
key_cache = torch.randint(
low=0, high=128, size=(num_blocks * block_size, num_head, head_size_k), dtype=torch.int8
).npu()
slot_list = []
for i in range(0, num_tokens):
slot_list.append([2 + i])
# slot_list.append(6+i)
assert num_tokens == len(slot_list)
slot_list_np = np.array(slot_list)
slot_mapping_npu = torch.from_numpy(slot_list_np).to(torch.int32).npu()
key_expect = cal_scatternd(key, key_cache, slot_mapping_npu, block_size)
N = 101
for i in range(N):
torch_npu.npu_scatter_nd_update_(key_cache, slot_mapping_npu, key)
torch.testing.assert_close(key_expect, key_cache, atol=0.001, rtol=0.1)
@pytest.mark.parametrize("num_tokens", [52]) # 6398
@pytest.mark.parametrize("num_head", [1]) # 512
@pytest.mark.parametrize("block_size", [1]) # 128
@pytest.mark.parametrize("num_blocks", [1773]) # 1599
@pytest.mark.parametrize("count", [1])
def test_myops(num_tokens, num_head, block_size, num_blocks, count):
head_size_k = 2
key_cache = torch.randint(low=0, high=128, size=(num_blocks, block_size, num_head, head_size_k), dtype=torch.int8)
key_cache_npu = key_cache.npu()
slot_list = []
for i in range(0, num_tokens):
slot_list.append(2 + i)
slot_list_np = np.array(slot_list)
slot_mapping_npu = torch.from_numpy(slot_list_np).to(torch.int32).npu()
key = torch.randint(low=0, high=128, size=(num_tokens, head_size_k), dtype=torch.int8)
key_npu = key.npu()
key_expect = cal_slot(key_npu, key_cache_npu, slot_list_np, block_size)
time.sleep(0.1)
group_len = torch.empty(num_tokens, dtype=torch.int32).npu()
group_key_idx = torch.empty(num_tokens, dtype=torch.int32).npu()
group_key_cache_idx = torch.empty(num_tokens, dtype=torch.int32).npu()
torch.ops._C_ascend.store_kv_block_metadata(
slot_mapping_npu, group_len, group_key_idx, group_key_cache_idx, block_size
)
torch.ops._C_ascend.store_kv_block(
key_npu, key_cache_npu, group_len, group_key_idx, group_key_cache_idx, block_size
)
torch.testing.assert_close(key_expect, key_cache_npu, atol=0.001, rtol=0.1)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,155 @@
import gc
import unittest
import torch
import torch_npu
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
torch.set_printoptions(threshold=float("inf"))
def clone_kv_cache(k_caches, v_caches):
new_k_caches = [cache.clone() for cache in k_caches]
new_v_caches = [cache.clone() for cache in v_caches]
return new_k_caches, new_v_caches
class TestTransposeKvCacheByBlock(unittest.TestCase):
def compute_golden(
self, k_caches, v_caches, block_ids_tensor, block_size, num_kv_head, head_dim, num_need_pulls, layers, dtype
):
num_blocks = block_ids_tensor.shape[0]
block_ids_tensor = block_ids_tensor.to(dtype=torch.int32)
block_offsets = torch.arange(0, block_size, dtype=torch.int32).npu()
slot_mapping = block_offsets.reshape((1, block_size)) + block_ids_tensor.reshape((num_blocks, 1)) * block_size
slot_mapping = slot_mapping.flatten()
block_len = num_blocks * block_size
block_len_tensor = torch.tensor([block_len], dtype=torch.int32).npu()
block_table = block_ids_tensor.view(1, -1)
seq_start_tensor = torch.tensor([0], dtype=torch.int32).npu()
k = torch.empty(block_len, num_kv_head, head_dim, dtype=dtype).npu()
v = torch.empty(block_len, num_kv_head, head_dim, dtype=dtype).npu()
for layer in range(layers):
k_cache_layer = k_caches[layer]
v_cache_layer = v_caches[layer]
torch_npu.atb.npu_paged_cache_load(
k_cache_layer,
v_cache_layer,
block_table,
block_len_tensor,
seq_starts=seq_start_tensor,
key=k,
value=v,
)
k = k.view(num_blocks, num_need_pulls, block_size, -1)
k.transpose_(1, 2)
k = k.contiguous().view(block_len, num_kv_head, -1)
v = v.view(num_blocks, num_need_pulls, block_size, -1)
v.transpose_(1, 2)
v = v.contiguous().view(block_len, num_kv_head, -1)
torch_npu._npu_reshape_and_cache(
key=k,
value=v,
key_cache=k_cache_layer,
value_cache=v_cache_layer,
slot_indices=slot_mapping,
)
del k, v
def test_transpose_kv_cache_by_block(self):
# (layers, block_num, block_size, num_kv_head, head_dim, num_need_pulls)
test_cases = [
(16, 128, 128, 4, 128, 4),
(16, 128, 128, 4, 128, 2),
(16, 128, 128, 4, 128, 1),
(16, 128, 128, 8, 128, 8),
(16, 128, 128, 8, 128, 4),
(16, 128, 128, 8, 128, 2),
]
dtypes = [torch.float16, torch.bfloat16]
for dtype in dtypes:
for layers, block_num, block_size, num_kv_head, head_dim, num_need_pulls in test_cases:
with self.subTest(
dtype=dtype,
shape=f"({layers}, {block_num}, {block_size}, {num_kv_head}, {head_dim}, {num_need_pulls})",
):
k_caches = []
v_caches = []
block_id_num = 33
block_ids_tensor = torch.randperm(block_num, dtype=torch.int64, device="npu")[:block_id_num]
for i in range(layers):
kcache = torch.randn(block_num, block_size, num_kv_head, head_dim, dtype=dtype, device="npu")
vcache = torch.randn(block_num, block_size, num_kv_head, head_dim, dtype=dtype, device="npu")
k_caches.append(kcache)
v_caches.append(vcache)
cloned_k_caches, cloned_v_caches = clone_kv_cache(k_caches, v_caches)
self.compute_golden(
cloned_k_caches,
cloned_v_caches,
block_ids_tensor,
block_size,
num_kv_head,
head_dim,
num_need_pulls,
layers,
dtype,
)
torch.ops._C_ascend.transpose_kv_cache_by_block(
k_caches, v_caches, block_ids_tensor, block_size, num_kv_head, head_dim, num_need_pulls, layers
)
for i in range(layers):
self.assert_tensors_almost_equal(k_caches[i], cloned_k_caches[i], dtype)
self.assert_tensors_almost_equal(v_caches[i], cloned_v_caches[i], dtype)
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
def assert_tensors_almost_equal(self, actual, expected, dtype):
"""Check if two tensors are approximately equal (considering floating point errors)"""
self.assertEqual(actual.shape, expected.shape, "Shape mismatch")
# Check for NaN
self.assertFalse(torch.isnan(actual).any(), "Actual result contains NaN")
self.assertFalse(torch.isnan(expected).any(), "Expected result contains NaN")
# Check for Inf
self.assertFalse(torch.isinf(actual).any(), "Actual result contains Inf")
self.assertFalse(torch.isinf(expected).any(), "Expected result contains Inf")
# Set different tolerances based on data type
if dtype == torch.float16:
rtol, atol = 1e-5, 1e-5
else: # bfloat16
rtol, atol = 1.5e-5, 1.5e-5
# Compare values
diff = torch.abs(actual - expected)
max_diff = diff.max().item()
max_expected = torch.abs(expected).max().item()
# Check relative and absolute errors
if max_expected > 0:
relative_diff = max_diff / max_expected
self.assertLessEqual(
relative_diff,
rtol,
f"Relative error too large: {relative_diff} > {rtol}. Max difference: {max_diff}",
)
self.assertLessEqual(max_diff, atol, f"Absolute error too large: {max_diff} > {atol}")
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

@@ -0,0 +1,99 @@
import gc
import pytest
import torch
import torch_npu # noqa: F401
import vllm_ascend.platform # noqa: F401
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
# Test parameters
DTYPES = [torch.int32]
# SHAPES = [(100,), (5, 20), (3, 4, 5)] # Various tensor shapes
# SHAPES = [(3, 4, 8), (3, 4, 5)] # Various tensor shapes
SHAPES = [(3, 4, 3)]
DEVICES = [f"npu:{0}"]
SEEDS = [0]
def get_masked_input_and_mask_ref(
input_: torch.Tensor,
org_vocab_start_index: int,
org_vocab_end_index: int,
num_org_vocab_padding: int,
added_vocab_start_index: int,
added_vocab_end_index: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Reference implementation for verification"""
org_vocab_mask = (input_ >= org_vocab_start_index) & (input_ < org_vocab_end_index)
added_vocab_mask = (input_ >= added_vocab_start_index) & (input_ < added_vocab_end_index)
added_offset = added_vocab_start_index - (org_vocab_end_index - org_vocab_start_index) - num_org_vocab_padding
valid_offset = (org_vocab_start_index * org_vocab_mask) + (added_offset * added_vocab_mask)
vocab_mask = org_vocab_mask | added_vocab_mask
masked_input = vocab_mask * (input_ - valid_offset)
return masked_input, ~vocab_mask
@pytest.mark.parametrize("shape", SHAPES)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("device", DEVICES)
@pytest.mark.parametrize("seed", SEEDS)
@torch.inference_mode()
def test_get_masked_input_and_mask(
shape: tuple[int, ...],
dtype: torch.dtype,
device: str,
seed: int,
) -> None:
# Set random seed
torch.manual_seed(seed)
torch.set_default_device(device)
# Generate random input tensor
input_tensor = torch.randint(0, 1000, shape, dtype=dtype)
# Test parameters
test_case = {
"org_start": 100,
"org_end": 200,
"padding": 0,
"added_start": 300,
"added_end": 400,
}
# Get reference result
ref_masked_input, ref_mask = get_masked_input_and_mask_ref(
input_tensor,
test_case["org_start"],
test_case["org_end"],
test_case["padding"],
test_case["added_start"],
test_case["added_end"],
)
# Get custom op result
print("input_tensor:", input_tensor)
custom_masked_input, custom_mask = torch.ops._C_ascend.get_masked_input_and_mask(
input_tensor,
test_case["org_start"],
test_case["org_end"],
test_case["padding"],
test_case["added_start"],
test_case["added_end"],
)
ref_masked_input = ref_masked_input.to(dtype)
print("custom_masked_input:", custom_masked_input)
print("ref_masked_input:", ref_masked_input)
print("custom_mask:", custom_mask)
print("ref_mask:", ref_mask)
# Compare results
torch.testing.assert_close(
custom_masked_input, ref_masked_input, rtol=1e-5, atol=1e-5, msg=f"Masked input mismatch for case: {test_case}"
)
torch.testing.assert_close(custom_mask, ref_mask, rtol=1e-5, atol=1e-5, msg=f"Mask mismatch for case: {test_case}")
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

Some files were not shown because too many files have changed in this diff Show More