@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
66
tests/e2e/nightly/single_node/models/configs/GLM-4.7.yaml
Normal file
66
tests/e2e/nightly/single_node/models/configs/GLM-4.7.yaml
Normal 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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
91
tests/e2e/nightly/single_node/models/configs/Kimi-K2.5.yaml
Normal file
91
tests/e2e/nightly/single_node/models/configs/Kimi-K2.5.yaml
Normal 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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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:
|
||||
@@ -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:
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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`)
|
||||
16
tests/e2e/nightly/single_node/models/scripts/__init__.py
Normal file
16
tests/e2e/nightly/single_node/models/scripts/__init__.py
Normal 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.
|
||||
#
|
||||
@@ -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
|
||||
443
tests/e2e/nightly/single_node/models/scripts/test_single_node.py
Normal file
443
tests/e2e/nightly/single_node/models/scripts/test_single_node.py
Normal 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)
|
||||
0
tests/e2e/nightly/single_node/ops/__init__.py
Normal file
0
tests/e2e/nightly/single_node/ops/__init__.py
Normal file
61
tests/e2e/nightly/single_node/ops/conftest.py
Normal file
61
tests/e2e/nightly/single_node/ops/conftest.py
Normal 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}")
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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"])
|
||||
@@ -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()
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -0,0 +1,102 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Compare vllm_ascend.sample.penalties.apply_all_penalties (Triton-Ascend) with
|
||||
# vllm.v1.sample.ops.penalties.apply_all_penalties (PyTorch via model_executor).
|
||||
# Requires NPU and Triton-Ascend.
|
||||
|
||||
import gc
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from vllm.v1.sample.ops.penalties import apply_all_penalties as v1_apply_all_penalties
|
||||
|
||||
from vllm_ascend.sample.penalties import apply_all_penalties as ascend_apply_all_penalties
|
||||
|
||||
# Same scenario grid as test_apply_penalties_model_executor (equivalence + boundaries).
|
||||
APPLY_PENALTY_CASES = [
|
||||
pytest.param(0, 0, "mixed", id="empty-both"),
|
||||
pytest.param(0, 16, "mixed", id="empty-prompt"),
|
||||
pytest.param(32, 0, "mixed", id="empty-output"),
|
||||
pytest.param(1, 1, "mixed", id="single-token-each"),
|
||||
pytest.param(32, 16, "mixed", id="typical-small"),
|
||||
pytest.param(128, 64, "mixed", id="typical-large"),
|
||||
pytest.param(128, 64, "all_padding", id="all-padding"),
|
||||
]
|
||||
|
||||
|
||||
def _make_tokens(
|
||||
num_seqs: int,
|
||||
seq_len: int,
|
||||
vocab_size: int,
|
||||
mode: str,
|
||||
device: str,
|
||||
) -> torch.Tensor:
|
||||
if mode == "all_padding":
|
||||
return torch.full((num_seqs, seq_len), vocab_size, device=device, dtype=torch.int64)
|
||||
if seq_len == 0:
|
||||
return torch.empty((num_seqs, 0), device=device, dtype=torch.int64)
|
||||
tokens = torch.randint(0, vocab_size, (num_seqs, seq_len), device=device, dtype=torch.int64)
|
||||
pad_mask = torch.rand(num_seqs, seq_len, device=device) > 0.7
|
||||
tokens[pad_mask] = vocab_size
|
||||
return tokens
|
||||
|
||||
|
||||
@pytest.mark.skip("Probabilistic failure, need zengtian after fix")
|
||||
@pytest.mark.parametrize("num_seqs", [1, 8, 32, 128])
|
||||
@pytest.mark.parametrize("vocab_size", [5120, 151936])
|
||||
@pytest.mark.parametrize(
|
||||
"max_prompt_len,max_output_len,token_mode",
|
||||
APPLY_PENALTY_CASES,
|
||||
)
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
|
||||
@torch.inference_mode()
|
||||
def test_apply_all_penalties_v1_vs_ascend(
|
||||
num_seqs,
|
||||
vocab_size,
|
||||
max_prompt_len,
|
||||
max_output_len,
|
||||
token_mode,
|
||||
dtype,
|
||||
device="npu",
|
||||
seed=42,
|
||||
):
|
||||
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
|
||||
|
||||
init_device_properties_triton()
|
||||
torch.manual_seed(seed)
|
||||
|
||||
logits_v1 = torch.randn(num_seqs, vocab_size, device=device, dtype=dtype)
|
||||
logits_ascend = logits_v1.clone()
|
||||
|
||||
prompt_tokens = _make_tokens(num_seqs, max_prompt_len, vocab_size, token_mode, device)
|
||||
output_tokens = _make_tokens(num_seqs, max_output_len, vocab_size, token_mode, device)
|
||||
output_token_ids = [row.tolist() for row in output_tokens.cpu()]
|
||||
|
||||
presence_penalties = torch.rand(num_seqs, device=device, dtype=torch.float32) * 0.2
|
||||
frequency_penalties = torch.rand(num_seqs, device=device, dtype=torch.float32) * 0.2
|
||||
repetition_penalties = torch.rand(num_seqs, device=device, dtype=torch.float32) * 0.4 + 1.0
|
||||
|
||||
v1_apply_all_penalties(
|
||||
logits_v1,
|
||||
prompt_tokens,
|
||||
presence_penalties,
|
||||
frequency_penalties,
|
||||
repetition_penalties,
|
||||
output_token_ids,
|
||||
)
|
||||
ascend_apply_all_penalties(
|
||||
logits_ascend,
|
||||
prompt_tokens,
|
||||
presence_penalties,
|
||||
frequency_penalties,
|
||||
repetition_penalties,
|
||||
output_token_ids,
|
||||
)
|
||||
|
||||
atol = 1e-2 if dtype == torch.bfloat16 else 1e-3
|
||||
rtol = 1e-2 if dtype == torch.bfloat16 else 1e-3
|
||||
assert torch.allclose(logits_ascend.float(), logits_v1.float(), atol=atol, rtol=rtol), (
|
||||
f"Max diff: {(logits_ascend.float() - logits_v1.float()).abs().max().item()}"
|
||||
)
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
@@ -0,0 +1,202 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Test vllm_ascend.worker.v2.sample.bad_words.apply_bad_words (Triton-Ascend).
|
||||
# Requires NPU and Triton-Ascend.
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm_ascend.worker.v2.sample.bad_words import apply_bad_words
|
||||
|
||||
# Test cases for different input shapes
|
||||
BAD_WORDS_TEST_CASES = [
|
||||
pytest.param(512, 50257, 16, 3, 2, id="small-case"),
|
||||
pytest.param(1024, 50257, 32, 5, 3, id="medium-case"),
|
||||
pytest.param(2048, 50257, 64, 8, 4, id="large-case"),
|
||||
]
|
||||
|
||||
|
||||
def create_test_data(num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length, device):
|
||||
"""Create test data for testing"""
|
||||
# Create logits
|
||||
logits = torch.randn(num_tokens, vocab_size, dtype=torch.float32, device=device)
|
||||
|
||||
# Create expanded_idx_mapping (map each token to a request)
|
||||
expanded_idx_mapping = torch.randint(0, num_requests, (num_tokens,), dtype=torch.int32, device=device)
|
||||
|
||||
# Create bad_word_token_ids and bad_word_offsets
|
||||
MAX_BAD_WORDS_TOTAL_TOKENS = 1024
|
||||
MAX_NUM_BAD_WORDS = 128
|
||||
bad_word_token_ids = torch.zeros((num_requests, MAX_BAD_WORDS_TOTAL_TOKENS), dtype=torch.int32, device=device)
|
||||
bad_word_offsets = torch.zeros((num_requests, MAX_NUM_BAD_WORDS + 1), dtype=torch.int32, device=device)
|
||||
num_bad_words = torch.zeros(num_requests, dtype=torch.int32, device=device)
|
||||
|
||||
# Fill bad words data
|
||||
for req_idx in range(num_requests):
|
||||
offset = 0
|
||||
actual_bad_words = 0
|
||||
for bw_idx in range(num_bad_words_per_req):
|
||||
# Check if adding this bad word would exceed the token limit
|
||||
if offset + bad_word_length > MAX_BAD_WORDS_TOTAL_TOKENS:
|
||||
break
|
||||
# Create a bad word with specific tokens
|
||||
bad_word = torch.tensor([100 + req_idx * 10 + bw_idx] * bad_word_length, dtype=torch.int32, device=device)
|
||||
bad_word_token_ids[req_idx, offset : offset + bad_word_length] = bad_word
|
||||
bad_word_offsets[req_idx, bw_idx] = offset
|
||||
offset += bad_word_length
|
||||
actual_bad_words += 1
|
||||
bad_word_offsets[req_idx, actual_bad_words] = offset
|
||||
num_bad_words[req_idx] = actual_bad_words
|
||||
|
||||
# Create all_token_ids with some matching bad words
|
||||
max_seq_len = 1024
|
||||
all_token_ids = torch.randint(0, vocab_size, (num_requests, max_seq_len), dtype=torch.int32, device=device)
|
||||
|
||||
# Create prompt_len and total_len
|
||||
prompt_len = torch.tensor([50] * num_requests, dtype=torch.int32, device=device)
|
||||
total_len = torch.tensor([max_seq_len] * num_requests, dtype=torch.int32, device=device)
|
||||
|
||||
# Create input_ids with the same bad words, so they can be detected
|
||||
input_ids = torch.randint(0, vocab_size, (num_tokens,), dtype=torch.int32, device=device)
|
||||
# For each token, set input_ids to match the bad word for its request
|
||||
for token_idx in range(num_tokens):
|
||||
req_idx = expanded_idx_mapping[token_idx].item()
|
||||
if num_bad_words[req_idx] > 0:
|
||||
# Set input_ids to match the first bad word
|
||||
bad_word = bad_word_token_ids[req_idx, :bad_word_length]
|
||||
# For each position in the bad word, set input_ids accordingly
|
||||
for i in range(bad_word_length):
|
||||
if token_idx - i >= 0:
|
||||
input_ids[token_idx - i] = bad_word[bad_word_length - 1 - i]
|
||||
|
||||
# Create expanded_local_pos - set to bad_word_length - 1 so that effective_len = output_len + (bad_word_length - 1)
|
||||
# This ensures that we're checking the current token as the end of a bad word
|
||||
expanded_local_pos = torch.full((num_tokens,), bad_word_length - 1, dtype=torch.int32, device=device)
|
||||
|
||||
return (
|
||||
logits,
|
||||
expanded_idx_mapping,
|
||||
bad_word_token_ids,
|
||||
bad_word_offsets,
|
||||
num_bad_words,
|
||||
all_token_ids,
|
||||
prompt_len,
|
||||
total_len,
|
||||
input_ids,
|
||||
expanded_local_pos,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length", BAD_WORDS_TEST_CASES
|
||||
)
|
||||
@torch.inference_mode()
|
||||
def test_apply_bad_words_different_shapes(
|
||||
num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length, device="npu"
|
||||
):
|
||||
"""Test apply_bad_words with different input shapes"""
|
||||
test_data = create_test_data(num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length, device)
|
||||
|
||||
# Make a copy of logits to compare
|
||||
logits_before = test_data[0].clone()
|
||||
logits_after = test_data[0].clone()
|
||||
|
||||
# Apply bad words
|
||||
apply_bad_words(logits_after, *test_data[1:], num_bad_words_per_req)
|
||||
|
||||
# Verify that logits were modified
|
||||
assert not torch.allclose(logits_before, logits_after), "Logits should be modified when bad words are present"
|
||||
print(f"Test passed: tokens={num_tokens}, requests={num_requests}")
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def test_apply_bad_words_no_bad_words(device="npu"):
|
||||
"""Test apply_bad_words with no bad words"""
|
||||
num_tokens = 1024
|
||||
vocab_size = 50257
|
||||
num_requests = 32
|
||||
num_bad_words_per_req = 0
|
||||
bad_word_length = 3
|
||||
|
||||
test_data = create_test_data(num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length, device)
|
||||
|
||||
# Make a copy of logits to compare
|
||||
logits_before = test_data[0].clone()
|
||||
logits_after = test_data[0].clone()
|
||||
# Apply bad words
|
||||
apply_bad_words(logits_after, *test_data[1:], num_bad_words_per_req)
|
||||
|
||||
# Verify that logits were not modified
|
||||
assert torch.allclose(logits_before, logits_after), "Logits should not be modified when no bad words are present"
|
||||
print("No bad words test passed")
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def test_apply_bad_words_edge_cases(device="npu"):
|
||||
"""Test apply_bad_words with edge cases"""
|
||||
# Test with maximum bad words
|
||||
num_tokens = 1024
|
||||
vocab_size = 50257
|
||||
num_requests = 16
|
||||
num_bad_words_per_req = 128 # Maximum allowed
|
||||
bad_word_length = 2
|
||||
print("\nTesting edge case: maximum bad words")
|
||||
test_data = create_test_data(num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length, device)
|
||||
|
||||
# Make a copy of logits to compare
|
||||
logits_before = test_data[0].clone()
|
||||
logits_after = test_data[0].clone()
|
||||
|
||||
# Apply bad words
|
||||
apply_bad_words(logits_after, *test_data[1:], num_bad_words_per_req)
|
||||
|
||||
# Verify that logits were modified
|
||||
assert not torch.allclose(logits_before, logits_after), (
|
||||
"Logits should be modified when maximum bad words are present"
|
||||
)
|
||||
print("Maximum bad words test passed")
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def test_apply_bad_words_token_limit(device="npu"):
|
||||
"""Test apply_bad_words with token limit cases"""
|
||||
num_tokens = 1024
|
||||
vocab_size = 50257
|
||||
num_requests = 16
|
||||
|
||||
# Test case 1: Total tokens within limit
|
||||
print("\nTesting case: total tokens within limit")
|
||||
num_bad_words_per_req = 32
|
||||
bad_word_length = 32 # 32 * 32 = 1024 tokens (exactly at limit)
|
||||
|
||||
test_data = create_test_data(num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length, device)
|
||||
|
||||
# Make a copy of logits to compare
|
||||
logits_before = test_data[0].clone()
|
||||
logits_after = test_data[0].clone()
|
||||
|
||||
# Apply bad words
|
||||
apply_bad_words(logits_after, *test_data[1:], num_bad_words_per_req)
|
||||
|
||||
# Verify that logits were modified
|
||||
assert not torch.allclose(logits_before, logits_after), (
|
||||
"Logits should be modified when total tokens are within limit"
|
||||
)
|
||||
print("Total tokens within limit test passed")
|
||||
|
||||
# Test case 2: Total tokens exceeding limit (this should still work but only process up to limit)
|
||||
print("\nTesting case: total tokens exceeding limit")
|
||||
num_bad_words_per_req = 33
|
||||
bad_word_length = 32 # 33 * 32 = 1056 tokens (exceeding limit)
|
||||
|
||||
test_data = create_test_data(num_tokens, vocab_size, num_requests, num_bad_words_per_req, bad_word_length, device)
|
||||
|
||||
# Make a copy of logits to compare
|
||||
logits_before = test_data[0].clone()
|
||||
logits_after = test_data[0].clone()
|
||||
|
||||
# Apply bad words
|
||||
apply_bad_words(logits_after, *test_data[1:], num_bad_words_per_req)
|
||||
|
||||
# Verify that logits were modified (even though we exceed the limit)
|
||||
assert not torch.allclose(logits_before, logits_after), "Logits should be modified when total tokens exceed limit"
|
||||
print("Total tokens exceeding limit test passed")
|
||||
@@ -0,0 +1,35 @@
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm_ascend.ops.triton.batch_memcpy import batch_memcpy_kernel
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float32])
|
||||
def test_batch_memcpy(dtype):
|
||||
element_size = 2 if dtype == torch.bfloat16 else 4
|
||||
device = "npu:0"
|
||||
# this is a typical case when used in mamba states copy.
|
||||
sizes = torch.tensor([24576, 262144, 24576, 262144], device=device, dtype=torch.int32)
|
||||
|
||||
src_tensors_list = []
|
||||
src_addr_list = []
|
||||
dst_tensors_list = []
|
||||
dst_addr_list = []
|
||||
for i in range(len(sizes)):
|
||||
src_tensors_list.append(torch.rand(sizes[i].item() // element_size, dtype=dtype, device=device))
|
||||
src_addr_list.append(src_tensors_list[-1].data_ptr())
|
||||
dst_tensors_list.append(torch.empty(sizes[i].item() // element_size, dtype=dtype, device=device))
|
||||
dst_addr_list.append(dst_tensors_list[-1].data_ptr())
|
||||
|
||||
src_addr_list = torch.tensor(src_addr_list, dtype=torch.int64, device=device)
|
||||
dst_addr_list = torch.tensor(dst_addr_list, dtype=torch.int64, device=device)
|
||||
|
||||
batch = sizes.shape[0]
|
||||
|
||||
grid = (batch,)
|
||||
# using larger block_size to accelerate copy.
|
||||
BLOCK_SIZE = 8192
|
||||
batch_memcpy_kernel[grid](src_addr_list, dst_addr_list, sizes, BLOCK_SIZE=BLOCK_SIZE)
|
||||
|
||||
for i in range(len(sizes)):
|
||||
torch.testing.assert_close(src_tensors_list[i], dst_tensors_list[i], rtol=0, atol=0)
|
||||
@@ -0,0 +1,125 @@
|
||||
import pytest
|
||||
import torch
|
||||
from vllm.triton_utils import triton
|
||||
|
||||
from vllm_ascend.worker.v2.sample.penalties import _bincount_kernel
|
||||
|
||||
|
||||
def torch_bincount(
|
||||
expanded_idx_mapping: torch.Tensor,
|
||||
all_token_ids: torch.Tensor,
|
||||
prompt_len: torch.Tensor,
|
||||
prefill_len: torch.Tensor,
|
||||
prompt_bin_mask: torch.Tensor,
|
||||
output_bin_counts: torch.Tensor,
|
||||
):
|
||||
req_indices = expanded_idx_mapping
|
||||
prompt_bin_mask[req_indices] = 0
|
||||
output_bin_counts[req_indices] = 0
|
||||
|
||||
for token_idx in range(expanded_idx_mapping.shape[0]):
|
||||
req_idx = expanded_idx_mapping[token_idx].item()
|
||||
|
||||
p_len = prompt_len[req_idx].item()
|
||||
pref_len = prefill_len[req_idx].item()
|
||||
|
||||
tokens = all_token_ids[req_idx]
|
||||
|
||||
for pos in range(p_len):
|
||||
token = tokens[pos].item()
|
||||
bin_idx = token // 32
|
||||
bit_idx = token % 32
|
||||
prompt_bin_mask[req_idx, bin_idx] |= 1 << bit_idx
|
||||
|
||||
for pos in range(p_len, pref_len):
|
||||
token = tokens[pos].item()
|
||||
output_bin_counts[req_idx, token] += 1
|
||||
|
||||
|
||||
@pytest.mark.skip(reason="atomic_or operator hangs in current npu_ir version")
|
||||
def test_bincount_kernel():
|
||||
"""
|
||||
Compute the prompt binary mask and token bincount using the Triton kernel.
|
||||
|
||||
Args:
|
||||
expanded_idx_mapping: Tensor containing the indices of requests to process.
|
||||
all_token_ids: Batch of input token IDs for all requests.
|
||||
prompt_len: Tensor storing the prompt length for each request.
|
||||
prefill_len: Tensor storing the prefill length for each request.
|
||||
prompt_bin_mask: Output binary mask tensor to mark prompt tokens.
|
||||
output_bin_counts: Output tensor to store token frequency counts.
|
||||
max_prefill_len: Maximum prefill length to limit kernel processing.
|
||||
"""
|
||||
|
||||
torch.manual_seed(42)
|
||||
|
||||
expanded_idx_mapping = torch.tensor([63], dtype=torch.int32).npu()
|
||||
all_token_ids = torch.randint(
|
||||
low=0,
|
||||
high=10,
|
||||
size=(64, 40960),
|
||||
dtype=torch.int32,
|
||||
).npu()
|
||||
|
||||
prompt_len = torch.randint(
|
||||
low=0,
|
||||
high=10,
|
||||
size=(64,),
|
||||
dtype=torch.int32,
|
||||
).npu()
|
||||
|
||||
prefill_len = torch.randint(
|
||||
low=0,
|
||||
high=10,
|
||||
size=(64,),
|
||||
dtype=torch.int32,
|
||||
).npu()
|
||||
|
||||
prompt_bin_mask = torch.zeros(size=(64, 4748), dtype=torch.int32).npu()
|
||||
output_bin_counts = torch.zeros(size=(64, 151936), dtype=torch.int32).npu()
|
||||
|
||||
ref_prompt_bin_mask = torch.zeros(size=(64, 4748), dtype=torch.int32).npu()
|
||||
ref_output_bin_counts = torch.zeros(size=(64, 151936), dtype=torch.int32).npu()
|
||||
|
||||
max_prefill_len = 10
|
||||
|
||||
prompt_bin_mask[expanded_idx_mapping] = 0
|
||||
output_bin_counts[expanded_idx_mapping] = 0
|
||||
num_tokens = expanded_idx_mapping.shape[0]
|
||||
BLOCK_SIZE = 1024
|
||||
num_blocks = triton.cdiv(max_prefill_len, BLOCK_SIZE)
|
||||
|
||||
_bincount_kernel[(num_tokens, num_blocks)](
|
||||
expanded_idx_mapping,
|
||||
all_token_ids,
|
||||
all_token_ids.stride(0),
|
||||
prompt_len,
|
||||
prefill_len,
|
||||
prompt_bin_mask,
|
||||
prompt_bin_mask.stride(0),
|
||||
output_bin_counts,
|
||||
output_bin_counts.stride(0),
|
||||
BLOCK_SIZE=BLOCK_SIZE,
|
||||
)
|
||||
|
||||
torch_bincount(
|
||||
expanded_idx_mapping,
|
||||
all_token_ids,
|
||||
prompt_len,
|
||||
prefill_len,
|
||||
ref_prompt_bin_mask,
|
||||
ref_output_bin_counts,
|
||||
)
|
||||
|
||||
# ========== Verify results ==========
|
||||
assert torch.equal(prompt_bin_mask, ref_prompt_bin_mask), (
|
||||
f"prompt_bin_mask triton output differs from torch reference.\n"
|
||||
f"Max diff: {torch.max(torch.abs(prompt_bin_mask - ref_prompt_bin_mask))}\n"
|
||||
f"Mean diff: {torch.mean(torch.abs(prompt_bin_mask - ref_prompt_bin_mask))}"
|
||||
)
|
||||
|
||||
assert torch.equal(output_bin_counts, ref_output_bin_counts), (
|
||||
f"output_bin_counts triton output differs from torch reference.\n"
|
||||
f"Max diff: {torch.max(torch.abs(output_bin_counts - ref_output_bin_counts))}\n"
|
||||
f"Mean diff: {torch.mean(torch.abs(output_bin_counts - ref_output_bin_counts))}"
|
||||
)
|
||||
@@ -0,0 +1,315 @@
|
||||
import gc
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm_ascend._310p.ops.causal_conv1d import causal_conv1d_fn as causal_conv1d_fn_ref
|
||||
from vllm_ascend._310p.ops.causal_conv1d import causal_conv1d_update as causal_conv1d_update_ref
|
||||
from vllm_ascend.ops.triton.mamba.causal_conv1d import PAD_SLOT_ID, causal_conv1d_fn
|
||||
from vllm_ascend.ops.triton.mamba.causal_conv1d import causal_conv1d_update_npu as causal_conv1d_update
|
||||
from vllm_ascend.utils import enable_custom_op
|
||||
|
||||
|
||||
def validate_cmp(y_cal, y_ref, dtype, device="npu"):
|
||||
y_cal = y_cal.to(device)
|
||||
y_ref = y_ref.to(device)
|
||||
if dtype == torch.float16:
|
||||
torch.testing.assert_close(y_ref, y_cal, rtol=3e-03, atol=1e-02, equal_nan=True)
|
||||
elif dtype == torch.bfloat16:
|
||||
torch.testing.assert_close(y_ref, y_cal, rtol=1e-02, atol=1e-02, equal_nan=True)
|
||||
elif dtype == torch.float32:
|
||||
torch.testing.assert_close(y_ref, y_cal, rtol=1e-03, atol=4e-03, equal_nan=True)
|
||||
elif (
|
||||
dtype == torch.int32
|
||||
or dtype == torch.int64
|
||||
or dtype == torch.int16
|
||||
or dtype == torch.int8
|
||||
or dtype == torch.uint32
|
||||
or dtype == torch.bool
|
||||
):
|
||||
assert torch.equal(y_cal, y_ref)
|
||||
else:
|
||||
raise ValueError('Invalid parameter "dtype" is found : {}'.format(dtype))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("has_initial_state", [False, True])
|
||||
@pytest.mark.parametrize("itype", [torch.bfloat16])
|
||||
@pytest.mark.parametrize("silu_activation", [True])
|
||||
@pytest.mark.parametrize("has_bias", [True])
|
||||
@pytest.mark.parametrize("seq_len", [[128, 1024, 2048, 4096]])
|
||||
@pytest.mark.parametrize("extra_state_len", [0, 2])
|
||||
@pytest.mark.parametrize("width", [4])
|
||||
@pytest.mark.parametrize("dim", [2048])
|
||||
def test_ascend_causal_conv1d(
|
||||
dim, width, extra_state_len, seq_len, has_bias, silu_activation, itype, has_initial_state
|
||||
):
|
||||
torch.random.manual_seed(0)
|
||||
enable_custom_op()
|
||||
device = "npu"
|
||||
cu_seqlen, num_seq = sum(seq_len), len(seq_len)
|
||||
state_len = width - 1 + extra_state_len
|
||||
|
||||
x = torch.randn(cu_seqlen, dim, device=device, dtype=itype).transpose(0, 1)
|
||||
weight = torch.randn(dim, width, device=device, dtype=itype) #
|
||||
query_start_loc = torch.cumsum(torch.tensor([0] + seq_len, device=device, dtype=torch.int32), dim=0).to(
|
||||
dtype=torch.int32
|
||||
)
|
||||
cache_indices = torch.arange(num_seq, device=device, dtype=torch.int32)
|
||||
has_initial_state_tensor = torch.tensor([has_initial_state] * num_seq, device=device, dtype=torch.bool)
|
||||
activation = None if not silu_activation else "silu"
|
||||
|
||||
if has_initial_state:
|
||||
conv_states = torch.randn((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2)
|
||||
conv_states_ref = (
|
||||
torch.randn((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2).copy_(conv_states)
|
||||
)
|
||||
else:
|
||||
conv_states = torch.zeros((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2)
|
||||
conv_states_ref = torch.zeros((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2)
|
||||
|
||||
if has_bias:
|
||||
bias = torch.randn(dim, device=device, dtype=itype)
|
||||
else:
|
||||
bias = None
|
||||
|
||||
out_ref = causal_conv1d_fn_ref(
|
||||
x,
|
||||
weight,
|
||||
bias=bias,
|
||||
activation=activation,
|
||||
conv_states=conv_states_ref,
|
||||
has_initial_state=has_initial_state_tensor,
|
||||
cache_indices=cache_indices,
|
||||
query_start_loc=query_start_loc,
|
||||
)
|
||||
# out = causal_conv1d_fn(x,
|
||||
# weight,
|
||||
# bias=bias,
|
||||
# activation=activation,
|
||||
# conv_states=conv_states,
|
||||
# has_initial_state=has_initial_state_tensor,
|
||||
# cache_indices=cache_indices,
|
||||
# query_start_loc=query_start_loc)
|
||||
x_origin = x.transpose(-1, -2)
|
||||
weight_origin = weight.transpose(-1, -2)
|
||||
conv_states_origin = conv_states.transpose(-1, -2)
|
||||
activation_num = 1 if activation else 0
|
||||
out = torch.empty_like(x_origin)
|
||||
torch.ops._C_ascend.npu_causal_conv1d_custom(
|
||||
out,
|
||||
x_origin,
|
||||
weight_origin,
|
||||
conv_state=conv_states_origin,
|
||||
bias_opt=bias,
|
||||
query_start_loc_opt=query_start_loc,
|
||||
cache_indices_opt=cache_indices,
|
||||
initial_state_mode_opt=has_initial_state_tensor,
|
||||
num_accepted_tokens_opt=None,
|
||||
activation_mode=activation_num,
|
||||
pad_slot_id=PAD_SLOT_ID,
|
||||
run_mode=0,
|
||||
)
|
||||
out = out.transpose(-1, -2)
|
||||
validate_cmp(out, out_ref, itype)
|
||||
validate_cmp(conv_states, conv_states_ref, itype)
|
||||
|
||||
|
||||
@pytest.mark.skip(
|
||||
reason="To use this tirton ops:causal_conv1d_fn, you need to set `get_forward_context`. After\
|
||||
the model side dumps the data, Zeng Tian has made the necessary fixes."
|
||||
)
|
||||
@pytest.mark.parametrize("has_initial_state", [False, True])
|
||||
@pytest.mark.parametrize("itype", [torch.bfloat16])
|
||||
@pytest.mark.parametrize("silu_activation", [True])
|
||||
@pytest.mark.parametrize("has_bias", [True])
|
||||
@pytest.mark.parametrize("seq_len", [[128, 1024, 2048, 4096]])
|
||||
@pytest.mark.parametrize("extra_state_len", [0, 2])
|
||||
@pytest.mark.parametrize("width", [2, 4])
|
||||
@pytest.mark.parametrize("dim", [4160])
|
||||
def test_causal_conv1d(dim, width, extra_state_len, seq_len, has_bias, silu_activation, itype, has_initial_state):
|
||||
torch.random.manual_seed(0)
|
||||
|
||||
device = "npu"
|
||||
cu_seqlen, num_seq = sum(seq_len), len(seq_len)
|
||||
state_len = width - 1 + extra_state_len
|
||||
|
||||
x = torch.randn(cu_seqlen, dim, device=device, dtype=itype).transpose(0, 1)
|
||||
weight = torch.randn(dim, width, device=device, dtype=itype)
|
||||
query_start_loc = torch.cumsum(torch.tensor([0] + seq_len, device=device, dtype=torch.int32), dim=0)
|
||||
cache_indices = torch.arange(num_seq, device=device, dtype=torch.int32)
|
||||
has_initial_state_tensor = torch.tensor([has_initial_state] * num_seq, device=device, dtype=torch.bool)
|
||||
activation = None if not silu_activation else "silu"
|
||||
|
||||
if has_initial_state:
|
||||
conv_states = torch.randn((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2)
|
||||
conv_states_ref = (
|
||||
torch.randn((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2).copy_(conv_states)
|
||||
)
|
||||
else:
|
||||
conv_states = torch.zeros((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2)
|
||||
conv_states_ref = torch.zeros((num_seq, state_len, dim), device=device, dtype=itype).transpose(-1, -2)
|
||||
|
||||
if has_bias:
|
||||
bias = torch.randn(dim, device=device, dtype=itype)
|
||||
else:
|
||||
bias = None
|
||||
|
||||
out_ref = causal_conv1d_fn_ref(
|
||||
x,
|
||||
weight,
|
||||
bias=bias,
|
||||
activation=activation,
|
||||
conv_states=conv_states_ref,
|
||||
has_initial_state=has_initial_state_tensor,
|
||||
cache_indices=cache_indices,
|
||||
query_start_loc=query_start_loc,
|
||||
)
|
||||
out = causal_conv1d_fn(
|
||||
x,
|
||||
weight,
|
||||
bias=bias,
|
||||
activation=activation,
|
||||
conv_states=conv_states,
|
||||
has_initial_state=has_initial_state_tensor,
|
||||
cache_indices=cache_indices,
|
||||
query_start_loc=query_start_loc,
|
||||
)
|
||||
|
||||
validate_cmp(out, out_ref, itype)
|
||||
validate_cmp(conv_states, conv_states_ref, itype)
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
|
||||
@pytest.mark.skip(
|
||||
reason="In this scenario, using tirton ops:causal_conv1d_update will cause an overflow. \
|
||||
Later, Zeng Tian was responsible for fixing this issue."
|
||||
)
|
||||
@pytest.mark.parametrize("itype", [torch.bfloat16])
|
||||
@pytest.mark.parametrize("silu_activation", [True])
|
||||
@pytest.mark.parametrize("has_bias", [False, True])
|
||||
@pytest.mark.parametrize("seqlen", [1, 3])
|
||||
@pytest.mark.parametrize("width", [3, 4])
|
||||
@pytest.mark.parametrize("dim", [2048 + 16, 4096])
|
||||
# tests correctness in case subset of the sequences are padded
|
||||
@pytest.mark.parametrize("with_padding", [True, False])
|
||||
@pytest.mark.parametrize("batch_size", [3, 64])
|
||||
def test_causal_conv1d_update_with_batch_gather(
|
||||
batch_size, with_padding, dim, width, seqlen, has_bias, silu_activation, itype
|
||||
):
|
||||
device = "npu"
|
||||
rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (3e-3, 5e-3)
|
||||
if itype == torch.bfloat16:
|
||||
rtol, atol = 1e-2, 5e-2
|
||||
|
||||
padding = 5 if with_padding else 0
|
||||
padded_batch_size = batch_size + padding
|
||||
# total_entries = number of cache line
|
||||
total_entries = 10 * batch_size
|
||||
|
||||
# x will be (batch, dim, seqlen) with contiguous along dim-axis
|
||||
x = torch.randn(padded_batch_size, seqlen, dim, device=device, dtype=itype).transpose(1, 2)
|
||||
|
||||
x_ref = x.clone()
|
||||
|
||||
conv_state_indices = torch.randperm(total_entries)[:batch_size].to(dtype=torch.int32, device=device)
|
||||
unused_states_bool = torch.ones(total_entries, dtype=torch.bool, device=device)
|
||||
unused_states_bool[conv_state_indices] = False
|
||||
padded_state_indices = torch.concat(
|
||||
[
|
||||
conv_state_indices,
|
||||
torch.as_tensor([PAD_SLOT_ID] * padding, dtype=torch.int32, device=device),
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
|
||||
# conv_state will be (cache_lines, dim, state_len)
|
||||
# with contiguous along dim-axis
|
||||
conv_state = torch.randn(total_entries, width - 1, dim, device=device, dtype=itype).transpose(1, 2)
|
||||
|
||||
conv_state_for_padding_test = conv_state.clone()
|
||||
|
||||
weight = torch.randn(dim, width, device=device, dtype=itype)
|
||||
bias = torch.randn(dim, device=device, dtype=itype) if has_bias else None
|
||||
conv_state_ref = conv_state[conv_state_indices, :].detach().clone()
|
||||
activation = None if not silu_activation else "silu"
|
||||
|
||||
out = causal_conv1d_update(
|
||||
x,
|
||||
conv_state,
|
||||
weight,
|
||||
bias,
|
||||
activation=activation,
|
||||
conv_state_indices=padded_state_indices,
|
||||
pad_slot_id=PAD_SLOT_ID,
|
||||
)
|
||||
out_ref = causal_conv1d_update_ref(
|
||||
x_ref[:batch_size].transpose(1, 2), conv_state_ref, weight, bias, activation=activation
|
||||
).transpose(1, 2)
|
||||
|
||||
assert torch.equal(conv_state[conv_state_indices, :], conv_state_ref)
|
||||
assert torch.equal(conv_state[unused_states_bool], conv_state_for_padding_test[unused_states_bool])
|
||||
assert torch.allclose(out[:batch_size], out_ref, rtol=rtol, atol=atol)
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
|
||||
@pytest.mark.skip("Probabilistic failure, need zengtian after fix")
|
||||
def test_causal_conv1d_update_qwen3_next_shape():
|
||||
device = "npu"
|
||||
itype = torch.bfloat16
|
||||
rtol, atol = (3e-4, 1e-3) if itype == torch.float32 else (3e-3, 5e-3)
|
||||
if itype == torch.bfloat16:
|
||||
rtol, atol = 1e-2, 5e-2
|
||||
|
||||
total_tokens = 192
|
||||
dim = 4096
|
||||
kernel_size = 4
|
||||
batch_size = 96
|
||||
num_states = 929
|
||||
|
||||
x = torch.randn(total_tokens, dim, dtype=itype, device=device)
|
||||
conv_state = torch.randn(num_states, dim, kernel_size, dtype=itype, device=device)
|
||||
weight = torch.randn(dim, kernel_size, dtype=itype, device=device)
|
||||
bias = None
|
||||
conv_state_indices = torch.randint(0, num_states, (batch_size,), dtype=torch.int32, device=device)
|
||||
num_accepted_tokens = torch.ones(total_tokens, dtype=torch.int32, device=device)
|
||||
query_start_loc = torch.arange(0, total_tokens + 1, dtype=torch.int32, device=device)
|
||||
|
||||
activation = "silu"
|
||||
max_query_len = 2
|
||||
pad_slot_id = -1
|
||||
validate_data = False
|
||||
|
||||
block_idx_last_scheduled_token = None
|
||||
initial_state_idx = None
|
||||
|
||||
out = causal_conv1d_update(
|
||||
x,
|
||||
conv_state,
|
||||
weight,
|
||||
bias,
|
||||
activation,
|
||||
conv_state_indices,
|
||||
num_accepted_tokens,
|
||||
query_start_loc,
|
||||
max_query_len,
|
||||
pad_slot_id,
|
||||
block_idx_last_scheduled_token,
|
||||
initial_state_idx,
|
||||
validate_data,
|
||||
)
|
||||
|
||||
x_ref = x.clone()
|
||||
conv_state_ref = conv_state[conv_state_indices, :].detach().clone()
|
||||
out_ref = causal_conv1d_update_ref(
|
||||
x_ref[:batch_size].transpose(1, 2), conv_state_ref, weight, bias, activation=activation
|
||||
).transpose(1, 2)
|
||||
|
||||
assert torch.allclose(out[:batch_size], out_ref, rtol=rtol, atol=atol)
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
@@ -0,0 +1,83 @@
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from tests.ut.base import PytestBase
|
||||
from vllm_ascend._310p.ops.fla.chunk_gated_delta_rule import chunk_gated_delta_rule_pytorch
|
||||
from vllm_ascend.ops.triton.fla.chunk import chunk_gated_delta_rule
|
||||
|
||||
|
||||
class TestChunkGatedDeltaRule(PytestBase):
|
||||
def test_triton_fusion_ops(self):
|
||||
mock_attn_metadata = MagicMock()
|
||||
mock_attn_metadata.num_decodes = 1
|
||||
mock_forward_context = MagicMock()
|
||||
mock_forward_context.attn_metadata = mock_attn_metadata
|
||||
|
||||
q = torch.randn(1, 17, 4, 128, dtype=torch.bfloat16).npu()
|
||||
k = torch.randn(1, 17, 4, 128, dtype=torch.bfloat16).npu()
|
||||
v = torch.randn(1, 17, 8, 128, dtype=torch.bfloat16).npu()
|
||||
g = torch.randn(1, 17, 8, dtype=torch.float32).npu()
|
||||
beta = torch.randn(1, 17, 8, dtype=torch.bfloat16).npu()
|
||||
initial_state = torch.randn(3, 8, 128, 128, dtype=torch.bfloat16).npu()
|
||||
q_start_loc = torch.range(0, 3, dtype=torch.int).npu()
|
||||
|
||||
mock_pcp_group = MagicMock()
|
||||
mock_pcp_group.world_size = 1
|
||||
with (
|
||||
patch("vllm_ascend.ops.triton.fla.chunk.get_forward_context", return_value=mock_forward_context),
|
||||
patch("vllm_ascend.ops.triton.fla.chunk.get_pcp_group", return_value=mock_pcp_group),
|
||||
):
|
||||
(
|
||||
core_attn_out_non_spec,
|
||||
last_recurrent_state,
|
||||
) = chunk_gated_delta_rule(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
g=g,
|
||||
beta=beta,
|
||||
initial_state=initial_state,
|
||||
output_final_state=True,
|
||||
cu_seqlens=q_start_loc,
|
||||
head_first=False,
|
||||
use_qk_l2norm_in_kernel=True,
|
||||
)
|
||||
|
||||
assert core_attn_out_non_spec.shape == (1, 17, 8, 128)
|
||||
assert last_recurrent_state.shape == (3, 8, 128, 128)
|
||||
|
||||
|
||||
def test_chunk_gated_delta_rule_310_state_layout_matches_vllm():
|
||||
q = torch.tensor([[[[1.0, 0.0]]]], dtype=torch.float32)
|
||||
k = torch.tensor([[[[1.0, 0.0]]]], dtype=torch.float32)
|
||||
v = torch.tensor([[[[10.0, 20.0, 30.0]]]], dtype=torch.float32)
|
||||
g = torch.zeros(1, 1, 1, dtype=torch.float32)
|
||||
beta = torch.ones(1, 1, 1, dtype=torch.float32)
|
||||
initial_state = torch.tensor(
|
||||
[[[[1.0, 2.0], [4.0, 8.0], [16.0, 32.0]]]],
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
out, final_state = chunk_gated_delta_rule_pytorch(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
g=g,
|
||||
beta=beta,
|
||||
initial_state=initial_state,
|
||||
output_final_state=True,
|
||||
cu_seqlens=None,
|
||||
head_first=False,
|
||||
use_qk_l2norm_in_kernel=False,
|
||||
)
|
||||
|
||||
expected_out = torch.tensor([[[[10.0, 20.0, 30.0]]]], dtype=torch.float32) / (2.0**0.5)
|
||||
expected_state = torch.tensor(
|
||||
[[[[10.0, 2.0], [20.0, 8.0], [30.0, 32.0]]]],
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
torch.testing.assert_close(out, expected_out, rtol=1e-5, atol=1e-5)
|
||||
assert final_state is not None
|
||||
torch.testing.assert_close(final_state, expected_state, rtol=1e-5, atol=1e-5)
|
||||
@@ -0,0 +1,162 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
# mypy: ignore-errors
|
||||
"""Precision tests for vllm's chunk_kda Triton operator on NPU.
|
||||
|
||||
Compares chunk_kda against a naive recurrent reference (float32).
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torch_npu # noqa: F401
|
||||
|
||||
from vllm_ascend.ops.triton.kda.kda import chunk_kda
|
||||
|
||||
DEVICE = "npu"
|
||||
|
||||
NPU_RMSE_RATIO_O = 0.005
|
||||
NPU_RMSE_RATIO_HT = 0.005
|
||||
|
||||
|
||||
def reference_l2norm(x: torch.Tensor, eps: float = 1e-6) -> torch.Tensor:
|
||||
dtype = x.dtype
|
||||
x = x.to(torch.float32)
|
||||
return (x * torch.rsqrt(torch.sum(x * x, dim=-1, keepdim=True) + eps)).to(dtype)
|
||||
|
||||
|
||||
def naive_recurrent_kda(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
g: torch.Tensor,
|
||||
beta: torch.Tensor,
|
||||
scale: float | None = None,
|
||||
initial_state: torch.Tensor | None = None,
|
||||
output_final_state: bool = False,
|
||||
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||
"""Naive recurrent KDA reference, ported from FLA's naive.py."""
|
||||
dtype = v.dtype
|
||||
B, T, H, K, V = *q.shape, v.shape[-1]
|
||||
if scale is None:
|
||||
scale = K**-0.5
|
||||
|
||||
q, k, v, g, beta = map(lambda x: x.to(torch.float), [q, k, v, g, beta])
|
||||
q = q * scale
|
||||
|
||||
S = k.new_zeros(B, H, K, V).to(q)
|
||||
if initial_state is not None:
|
||||
S += initial_state
|
||||
o = torch.zeros_like(v)
|
||||
for i in range(T):
|
||||
q_i, k_i, v_i, g_i, b_i = q[:, i], k[:, i], v[:, i], g[:, i], beta[:, i]
|
||||
S = S * g_i[..., None].exp()
|
||||
S = S + torch.einsum(
|
||||
"bhk,bhv->bhkv",
|
||||
b_i[..., None] * k_i,
|
||||
v_i - (k_i[..., None] * S).sum(-2),
|
||||
)
|
||||
o[:, i] = torch.einsum("bhk,bhkv->bhv", q_i, S)
|
||||
if not output_final_state:
|
||||
S = None
|
||||
return o.to(dtype), S
|
||||
|
||||
|
||||
def assert_close(
|
||||
name: str,
|
||||
ref: torch.Tensor,
|
||||
tri: torch.Tensor,
|
||||
ratio: float,
|
||||
err_atol: float = 1e-6,
|
||||
):
|
||||
"""RMSE-based relative error comparison."""
|
||||
abs_err = (ref.detach() - tri.detach()).flatten().abs().max().item()
|
||||
rmse_diff = (ref.detach() - tri.detach()).flatten().square().mean().sqrt().item()
|
||||
rmse_base = ref.detach().flatten().square().mean().sqrt().item()
|
||||
rel_err = rmse_diff / (rmse_base + 1e-8)
|
||||
print(f"{name:>4} | abs={abs_err:.6f} | rmse={rel_err:.6f} | thr={ratio}")
|
||||
if abs_err <= err_atol:
|
||||
return
|
||||
assert not torch.isnan(ref).any(), f"{name}: NaN detected in ref"
|
||||
assert not torch.isnan(tri).any(), f"{name}: NaN detected in tri"
|
||||
assert rel_err < ratio, f"{name}: max abs err {abs_err:.6f}, rmse ratio {rel_err:.6f} >= {ratio}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("H", "D", "cu_seqlens", "dtype"),
|
||||
[
|
||||
pytest.param(
|
||||
*test,
|
||||
id="H{}-D{}-cu{}-{}".format(*test),
|
||||
)
|
||||
for test in [
|
||||
(32, 128, [0, 64], torch.float16),
|
||||
(32, 128, [0, 1024], torch.float16),
|
||||
(32, 128, [0, 15], torch.float16),
|
||||
(32, 128, [0, 256, 512, 768, 1024], torch.float16),
|
||||
(32, 128, [0, 15, 100, 300, 1200], torch.float16),
|
||||
(64, 128, [0, 256, 500, 1000], torch.float16),
|
||||
(32, 128, [0, 8192], torch.float16),
|
||||
(32, 128, [0, 256, 500, 1000], torch.bfloat16),
|
||||
(32, 128, [0, 4096], torch.float16),
|
||||
]
|
||||
],
|
||||
)
|
||||
@pytest.mark.skip_global_cleanup
|
||||
@torch.inference_mode()
|
||||
def test_chunk_kda(
|
||||
H: int,
|
||||
D: int,
|
||||
cu_seqlens: list[int],
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
T = cu_seqlens[-1]
|
||||
torch.manual_seed(42)
|
||||
B = 1
|
||||
cu_seqlens_t = torch.LongTensor(cu_seqlens).to(DEVICE)
|
||||
N = len(cu_seqlens) - 1
|
||||
|
||||
q = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
|
||||
k = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
|
||||
v = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
|
||||
g = F.logsigmoid(torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE)).to(dtype)
|
||||
beta = torch.rand(B, T, H, dtype=dtype, device=DEVICE).sigmoid()
|
||||
h0 = torch.randn(N, H, D, D, dtype=torch.float32, device=DEVICE)
|
||||
|
||||
ref_outputs = []
|
||||
ref_states = []
|
||||
for i in range(N):
|
||||
s, e = cu_seqlens[i], cu_seqlens[i + 1]
|
||||
q_i = reference_l2norm(q[:, s:e].contiguous())
|
||||
k_i = reference_l2norm(k[:, s:e].contiguous())
|
||||
o_i, ht_i = naive_recurrent_kda(
|
||||
q_i,
|
||||
k_i,
|
||||
v[:, s:e],
|
||||
g[:, s:e],
|
||||
beta[:, s:e],
|
||||
initial_state=h0[i],
|
||||
output_final_state=True,
|
||||
)
|
||||
ref_outputs.append(o_i)
|
||||
ref_states.append(ht_i)
|
||||
ref_o = torch.cat(ref_outputs, dim=1)
|
||||
ref_ht = torch.cat(ref_states, dim=0)
|
||||
|
||||
# h0 transposed to (V, K) layout for the kernel; naive uses (K, V)
|
||||
tri_o, tri_ht = chunk_kda(
|
||||
q=q.clone(),
|
||||
k=k.clone(),
|
||||
v=v.clone(),
|
||||
g=g.clone(),
|
||||
beta=beta.clone(),
|
||||
initial_state=h0.transpose(-1, -2).contiguous().clone(),
|
||||
output_final_state=True,
|
||||
cu_seqlens=cu_seqlens_t,
|
||||
use_qk_l2norm_in_kernel=True,
|
||||
)
|
||||
|
||||
assert not torch.isnan(tri_o).any(), "Triton output o contains NaN"
|
||||
assert not torch.isnan(tri_ht).any(), "Triton output ht contains NaN"
|
||||
assert_close("o", ref_o, tri_o, NPU_RMSE_RATIO_O)
|
||||
assert_close("ht", ref_ht, tri_ht.transpose(-1, -2).contiguous(), NPU_RMSE_RATIO_HT)
|
||||
@@ -0,0 +1,37 @@
|
||||
import gc
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm_ascend.ops.triton.fla.utils import clear_ssm_states
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.bfloat16])
|
||||
@pytest.mark.parametrize(
|
||||
"state_shape",
|
||||
[
|
||||
(6, 3, 5, 7),
|
||||
(4, 5, 25, 41),
|
||||
],
|
||||
)
|
||||
def test_clear_ssm_states_ref_parity(state_shape, dtype):
|
||||
torch.manual_seed(0)
|
||||
device = "npu"
|
||||
|
||||
ssm_states = torch.randn(*state_shape, device=device, dtype=dtype)
|
||||
has_initial_state = torch.tensor(
|
||||
[True, False, True, False, False, True][: state_shape[0]],
|
||||
device=device,
|
||||
dtype=torch.bool,
|
||||
)
|
||||
|
||||
ssm_states_ref = ssm_states.clone()
|
||||
ssm_states_ref[~has_initial_state, ...] = 0
|
||||
|
||||
clear_ssm_states(ssm_states, has_initial_state)
|
||||
|
||||
torch.testing.assert_close(ssm_states, ssm_states_ref)
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
@@ -0,0 +1,109 @@
|
||||
import torch
|
||||
from vllm.v1.worker.gpu.block_table import _compute_slot_mappings_kernel as ref_compute_slot_mappings_kernel
|
||||
|
||||
from vllm_ascend.worker.v2.block_table import _compute_slot_mappings_kernel as ascend_compute_slot_mappings_kernel
|
||||
|
||||
|
||||
def test_compute_slot_mapping_npu_kernel():
|
||||
"""
|
||||
Computes the physical slot IDs in KV cache for each token in the current batch.
|
||||
This function maps the logical positions of tokens to their actual storage locations
|
||||
in the block-managed KV cache, which is critical for efficient memory access in LLM inference.
|
||||
|
||||
Input:
|
||||
- max_num_batched_tokens (int): Maximum preallocated batched tokens in KV cache (memory limit)
|
||||
- idx_mapping (torch.Tensor): [num_reqs], int32 → Virtual-to-actual request index mapping
|
||||
- query_start_loc (torch.Tensor): [num_reqs+1], int32 → Batch-level token start positions per request
|
||||
- positions (torch.Tensor): [num_tokens], int64 → Per-token logical sequence positions in requests
|
||||
- block_table_ptrs (torch.Tensor): [num_kv_cache_groups], int32 → Pointers to block tables (virtual→physical)
|
||||
- block_table_strides (torch.Tensor): [num_kv_cache_groups], int32 → Stride for block table addressing
|
||||
- block_sizes_tensor (torch.Tensor): [num_kv_cache_groups], int32 → Token capacity per KV cache block
|
||||
- slot_mappings (torch.Tensor): [num_kv_cache_groups, max_num_batched_tokens], int32 → Output slot ID tensor
|
||||
- slot_mappings_stride0 (int): Stride of the first dimension of slot_mappings (memory layout)
|
||||
- cp_rank (int): Current device rank in column-parallel (CP) group
|
||||
- CP_SIZE (int): Total devices in CP parallel group
|
||||
- CP_INTERLEAVE (bool): Enable interleaved CP computation (memory access optimization)
|
||||
- PAD_ID (int): Padding value for invalid slot IDs (-1)
|
||||
- TRITON_BLOCK_SIZE (int): Block size for Triton kernel execution (hardware optimization),
|
||||
'TOTAL_BLOCK_SIZE' must be greater than the 'position / (block_size * CP_SIZE) + 1024'
|
||||
|
||||
Output:
|
||||
- slot_mappings (torch.Tensor): [num_kv_cache_groups, max_num_batched_tokens], int32 → Output slot ID tensor
|
||||
"""
|
||||
|
||||
torch.manual_seed(42)
|
||||
|
||||
device = "npu" if torch.npu.is_available() else "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
max_num_batched_tokens = 8192
|
||||
idx_mapping = torch.tensor([63], dtype=torch.int32, device=device)
|
||||
query_start_loc = torch.tensor([0, 5], dtype=torch.int32, device=device)
|
||||
positions = torch.tensor([0, 1, 2, 3, 4, 0, 0, 0], dtype=torch.int64, device=device)
|
||||
|
||||
num_kv_cache_groups = 1
|
||||
max_num_reqs = 64
|
||||
max_num_blocks = 320
|
||||
block_tables: list[torch.Tensor] = []
|
||||
for i in range(num_kv_cache_groups):
|
||||
block_table = torch.randint(0, 320, (max_num_reqs, max_num_blocks), dtype=torch.int32, device=device)
|
||||
block_tables.append(block_table)
|
||||
block_table_ptrs = torch.tensor([t.data_ptr() for t in block_table], dtype=torch.uint64, device=device)
|
||||
block_table_strides = torch.tensor([320], dtype=torch.int32, device=device)
|
||||
|
||||
block_sizes_tensor = torch.tensor([128], dtype=torch.int32, device=device)
|
||||
slot_mappings = torch.zeros(size=(1, 8192), dtype=torch.int64, device=device)
|
||||
ref_slot_mappings = torch.zeros(size=(1, 8192), dtype=torch.int64, device=device)
|
||||
cp_rank = 0
|
||||
cp_size = 1
|
||||
cp_interleave = 1
|
||||
num_reqs = query_start_loc.shape[0] - 1
|
||||
num_groups = num_kv_cache_groups
|
||||
|
||||
try:
|
||||
ascend_compute_slot_mappings_kernel[(num_groups, num_reqs + 1)](
|
||||
max_num_batched_tokens,
|
||||
idx_mapping,
|
||||
query_start_loc,
|
||||
positions,
|
||||
block_table_ptrs,
|
||||
block_table_strides,
|
||||
block_sizes_tensor,
|
||||
slot_mappings,
|
||||
slot_mappings.stride(0),
|
||||
cp_rank,
|
||||
CP_SIZE=cp_size,
|
||||
CP_INTERLEAVE=cp_interleave,
|
||||
PAD_ID=-1,
|
||||
TRITON_BLOCK_SIZE=1024, # type: ignore
|
||||
TOTAL_BLOCK_SIZE=4096,
|
||||
)
|
||||
|
||||
ref_compute_slot_mappings_kernel[(num_groups, num_reqs + 1)](
|
||||
max_num_batched_tokens,
|
||||
idx_mapping,
|
||||
query_start_loc,
|
||||
positions,
|
||||
block_table_ptrs,
|
||||
block_table_strides,
|
||||
block_sizes_tensor,
|
||||
ref_slot_mappings,
|
||||
ref_slot_mappings.stride(0),
|
||||
cp_rank,
|
||||
CP_SIZE=cp_size,
|
||||
CP_INTERLEAVE=cp_interleave,
|
||||
PAD_ID=-1,
|
||||
TRITON_BLOCK_SIZE=1024, # type: ignore
|
||||
)
|
||||
|
||||
# ========== Verify results ==========
|
||||
assert torch.equal(slot_mappings, ref_slot_mappings), (
|
||||
f"ascend output differs from gpu reference.\n"
|
||||
f"Max diff: {torch.max(torch.abs(slot_mappings - ref_slot_mappings))}\n"
|
||||
f"Mean diff: {torch.mean(torch.abs(slot_mappings - ref_slot_mappings).float())}"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during executionm: {e}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
@@ -0,0 +1,227 @@
|
||||
import random
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm_ascend.worker.v2.sample.logprob import compute_token_logprobs
|
||||
|
||||
|
||||
def torch_compute_token_logprobs(logits: torch.Tensor, token_ids: torch.Tensor) -> torch.Tensor:
|
||||
"""Pure PyTorch reference implementation of topk log softmax.
|
||||
|
||||
Computes log_softmax for the entire logits tensor, then gathers
|
||||
the values at the specified token_ids positions.
|
||||
|
||||
Args:
|
||||
logits: Tensor of shape (batch_size, vocab_size) containing the logits.
|
||||
token_ids: Tensor of shape (batch_size, topk) containing the token indices.
|
||||
|
||||
Returns:
|
||||
Tensor of shape (batch_size, topk) containing log probabilities.
|
||||
"""
|
||||
# Compute log_softmax along the vocab dimension
|
||||
# log_softmax(x) = x - log(sum(exp(x))) = x - max(x) - log(sum(exp(x - max(x))))
|
||||
log_probs = torch.nn.functional.log_softmax(logits.float(), dim=-1)
|
||||
|
||||
# Gather the log probabilities at the specified token positions
|
||||
token_ids = token_ids.to(torch.int64)
|
||||
result = torch.gather(log_probs, dim=1, index=token_ids)
|
||||
|
||||
return result.to(torch.float32)
|
||||
|
||||
|
||||
# Common vocab sizes from mainstream models
|
||||
VOCAB_SIZES = [
|
||||
32000, # LLaMA / LLaMA2 / Mistral
|
||||
50257, # GPT-2
|
||||
65024, # ChatGLM
|
||||
128256, # LLaMA3
|
||||
151936, # Qwen2
|
||||
]
|
||||
|
||||
# Different topk values to test
|
||||
TOPK_VALUES = [1, 2, 5, 10, 32, 64]
|
||||
|
||||
|
||||
@pytest.mark.skip("UB overflow, zengtian needs to fix it later")
|
||||
@pytest.mark.parametrize(
|
||||
"batch_size, vocab_size, topk",
|
||||
[(random.randint(1, 64), vocab_size, topk) for vocab_size in VOCAB_SIZES for topk in TOPK_VALUES],
|
||||
)
|
||||
def test_topk_log_softmax_kernel(batch_size, vocab_size, topk):
|
||||
"""
|
||||
Test the Triton _topk_log_softmax_kernel against a pure PyTorch reference.
|
||||
|
||||
The kernel computes log_softmax and gathers values at specified token positions.
|
||||
|
||||
Args:
|
||||
batch_size: Number of requests (rows) in the logits tensor.
|
||||
vocab_size: Vocabulary size (columns) in the logits tensor.
|
||||
topk: Number of token positions to compute log probabilities for.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
|
||||
device = "npu"
|
||||
|
||||
# Build input tensors
|
||||
logits = torch.randn((batch_size, vocab_size), dtype=torch.float32, device=device)
|
||||
|
||||
# Generate random token indices within vocab_size
|
||||
token_ids = torch.randint(0, vocab_size, (batch_size, topk), dtype=torch.int64, device=device)
|
||||
|
||||
# ========== Run Triton kernel ==========
|
||||
logprobs_triton = compute_token_logprobs(logits, token_ids)
|
||||
|
||||
# ========== Run PyTorch reference ==========
|
||||
logprobs_ref = torch_compute_token_logprobs(logits, token_ids)
|
||||
|
||||
# ========== Verify results ==========
|
||||
max_diff = torch.max(torch.abs(logprobs_triton - logprobs_ref)).item()
|
||||
mean_diff = torch.mean(torch.abs(logprobs_triton - logprobs_ref)).item()
|
||||
|
||||
assert torch.allclose(logprobs_triton, logprobs_ref, atol=1e-4, rtol=1e-5), (
|
||||
f"Triton topk_log_softmax kernel output differs from torch reference.\n"
|
||||
f"batch_size={batch_size}, vocab_size={vocab_size}, topk={topk}\n"
|
||||
f"Max diff: {max_diff}\n"
|
||||
f"Mean diff: {mean_diff}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skip("UB overflow, zengtian needs to fix it later")
|
||||
@pytest.mark.parametrize("vocab_size", VOCAB_SIZES)
|
||||
def test_topk_log_softmax_edge_cases(vocab_size):
|
||||
"""
|
||||
Test edge cases for the topk_log_softmax kernel.
|
||||
|
||||
Args:
|
||||
vocab_size: Vocabulary size to test.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
|
||||
device = "npu"
|
||||
|
||||
# Test case 1: Single batch, single topk
|
||||
logits = torch.randn((1, vocab_size), dtype=torch.float32, device=device)
|
||||
token_ids = torch.randint(0, vocab_size, (1, 1), dtype=torch.int64, device=device)
|
||||
|
||||
logprobs_triton = compute_token_logprobs(logits, token_ids)
|
||||
logprobs_ref = torch_compute_token_logprobs(logits, token_ids)
|
||||
|
||||
assert torch.allclose(logprobs_triton, logprobs_ref, atol=1e-4, rtol=1e-5), (
|
||||
f"Edge case (1,1) failed for vocab_size={vocab_size}"
|
||||
)
|
||||
|
||||
# Test case 2: Logits with extreme values
|
||||
logits_extreme = torch.randn((4, vocab_size), dtype=torch.float32, device=device)
|
||||
logits_extreme[0, 0] = 100.0 # Very large positive
|
||||
logits_extreme[1, 0] = -100.0 # Very large negative
|
||||
logits_extreme[2, :] = 0.0 # All zeros
|
||||
logits_extreme[3, :] = 1.0 # All ones
|
||||
|
||||
token_ids = torch.zeros((4, 5), dtype=torch.int64, device=device)
|
||||
token_ids[:, 0] = 0 # Include the extreme value position
|
||||
for i in range(1, 5):
|
||||
token_ids[:, i] = torch.randint(1, vocab_size, (4,))
|
||||
|
||||
logprobs_triton = compute_token_logprobs(logits_extreme, token_ids)
|
||||
logprobs_ref = torch_compute_token_logprobs(logits_extreme, token_ids)
|
||||
|
||||
assert torch.allclose(logprobs_triton, logprobs_ref, atol=1e-4, rtol=1e-5), (
|
||||
f"Extreme values test failed for vocab_size={vocab_size}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skip("UB overflow, zengtian needs to fix it later")
|
||||
@pytest.mark.parametrize(
|
||||
"batch_size, vocab_size, topk",
|
||||
[
|
||||
(16, 32000, 10),
|
||||
(32, 50257, 5),
|
||||
(64, 128256, 20),
|
||||
],
|
||||
)
|
||||
def test_topk_log_softmax_deterministic(batch_size, vocab_size, topk):
|
||||
"""
|
||||
Test that the kernel produces deterministic results across multiple runs.
|
||||
|
||||
Args:
|
||||
batch_size: Number of requests.
|
||||
vocab_size: Vocabulary size.
|
||||
topk: Number of token positions.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
|
||||
device = "npu"
|
||||
|
||||
logits = torch.randn((batch_size, vocab_size), dtype=torch.float32, device=device)
|
||||
token_ids = torch.randint(0, vocab_size, (batch_size, topk), dtype=torch.int64, device=device)
|
||||
|
||||
# Run multiple times and check consistency
|
||||
results = []
|
||||
for _ in range(3):
|
||||
result = compute_token_logprobs(logits, token_ids)
|
||||
results.append(result.clone())
|
||||
|
||||
for i in range(1, len(results)):
|
||||
assert torch.equal(results[0], results[i]), f"Non-deterministic results detected in run {i}"
|
||||
|
||||
|
||||
@pytest.mark.skip("UB overflow, zengtian needs to fix it later")
|
||||
@pytest.mark.parametrize("dtype", [torch.float32, torch.float16])
|
||||
def test_topk_log_softmax_dtypes(dtype):
|
||||
"""
|
||||
Test the kernel with different input dtypes.
|
||||
|
||||
Args:
|
||||
dtype: Input tensor dtype.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
|
||||
device = "npu"
|
||||
|
||||
batch_size = 8
|
||||
vocab_size = 32000
|
||||
topk = 10
|
||||
|
||||
logits = torch.randn((batch_size, vocab_size), dtype=dtype, device=device)
|
||||
token_ids = torch.randint(0, vocab_size, (batch_size, topk), dtype=torch.int64, device=device)
|
||||
|
||||
logprobs_triton = compute_token_logprobs(logits, token_ids)
|
||||
logprobs_ref = torch_compute_token_logprobs(logits.float(), token_ids)
|
||||
|
||||
# Use slightly larger tolerance for float16 due to precision loss
|
||||
atol = 1e-3 if dtype == torch.float16 else 1e-4
|
||||
|
||||
assert torch.allclose(logprobs_triton, logprobs_ref, atol=atol, rtol=1e-4), f"dtype {dtype} test failed"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run a quick sanity check
|
||||
print("Running quick sanity check...")
|
||||
|
||||
device = "npu"
|
||||
print(f"Using device: {device}")
|
||||
|
||||
torch.manual_seed(42)
|
||||
|
||||
batch_size = 4
|
||||
vocab_size = 32000
|
||||
topk = 5
|
||||
|
||||
logits = torch.randn((batch_size, vocab_size), dtype=torch.float32, device=device)
|
||||
token_ids = torch.randint(0, vocab_size, (batch_size, topk), dtype=torch.int64, device=device)
|
||||
|
||||
logprobs_triton = compute_token_logprobs(logits, token_ids)
|
||||
logprobs_ref = torch_compute_token_logprobs(logits, token_ids)
|
||||
|
||||
max_diff = torch.max(torch.abs(logprobs_triton - logprobs_ref)).item()
|
||||
mean_diff = torch.mean(torch.abs(logprobs_triton - logprobs_ref)).item()
|
||||
|
||||
print(f"Max diff: {max_diff}")
|
||||
print(f"Mean diff: {mean_diff}")
|
||||
print(f"All close (atol=1e-4): {torch.allclose(logprobs_triton, logprobs_ref, atol=1e-4, rtol=1e-5)}")
|
||||
|
||||
print("\nTriton output (first row):", logprobs_triton[0])
|
||||
print("PyTorch output (first row):", logprobs_ref[0])
|
||||
|
||||
print("\nSanity check passed!")
|
||||
@@ -0,0 +1,61 @@
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
|
||||
from vllm_ascend.worker.v2.sample.logprob import compute_topk_logprobs
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"batch_size,vocab_size,num_logprobs",
|
||||
[
|
||||
(48, 1024, 5),
|
||||
(96, 1024, 0),
|
||||
(24, 1519, 1),
|
||||
(1, 320, 10),
|
||||
],
|
||||
)
|
||||
def test_compute_topk_logprobs(batch_size, vocab_size, num_logprobs):
|
||||
"""Test compute_topk_logprobs for correctness of IDs, logprobs, and ranks.
|
||||
Args:
|
||||
batch_size: Number of sequences in the batch
|
||||
vocab_size: Size of the vocabulary
|
||||
num_logprobs: Number of top-k logprobs to return (excluding the sampled token)
|
||||
"""
|
||||
init_device_properties_triton()
|
||||
# ========== 1. Setup test data ==========
|
||||
torch.manual_seed(42)
|
||||
device = "npu"
|
||||
|
||||
logits = torch.randn(batch_size, vocab_size, device=device, dtype=torch.float32)
|
||||
sampled_token_ids = torch.randint(0, vocab_size, (batch_size,), device=device, dtype=torch.int64)
|
||||
|
||||
# ========== 2. Execute Triton implementation ==========
|
||||
triton_output = compute_topk_logprobs(logits, num_logprobs, sampled_token_ids)
|
||||
torch.npu.synchronize()
|
||||
|
||||
# ========== 3. Compute reference values using PyTorch ==========
|
||||
if num_logprobs == 0:
|
||||
ref_token_ids = sampled_token_ids.unsqueeze(-1)
|
||||
else:
|
||||
topk_indices = torch.topk(logits, num_logprobs, dim=-1).indices
|
||||
ref_token_ids = torch.cat((sampled_token_ids.unsqueeze(-1), topk_indices), dim=1)
|
||||
|
||||
ref_all_logprobs = torch.log_softmax(logits, dim=-1)
|
||||
ref_logprobs = torch.gather(ref_all_logprobs, dim=1, index=ref_token_ids)
|
||||
|
||||
sampled_logits = torch.gather(logits, 1, sampled_token_ids.unsqueeze(-1))
|
||||
ref_ranks = (logits > sampled_logits).sum(dim=1).to(torch.int64)
|
||||
|
||||
# ========== 4. Verify results ==========
|
||||
assert torch.equal(triton_output.logprob_token_ids, ref_token_ids), (
|
||||
"Token IDs (Sampled + TopK) do not match between Triton and PyTorch."
|
||||
)
|
||||
|
||||
assert torch.equal(triton_output.selected_token_ranks, ref_ranks), (
|
||||
f"Token Ranks do not match.\nTriton: {triton_output.selected_token_ranks}\nPyTorch: {ref_ranks}"
|
||||
)
|
||||
|
||||
assert torch.allclose(triton_output.logprobs, ref_logprobs, rtol=1e-4, atol=1e-4), (
|
||||
f"Logprobs values differ between Triton and PyTorch.\n"
|
||||
f"Max diff: {torch.max(torch.abs(triton_output.logprobs - ref_logprobs))}"
|
||||
)
|
||||
@@ -0,0 +1,51 @@
|
||||
import torch
|
||||
|
||||
from vllm_ascend._310p.ops.fla.fused_gdn_gating import fused_gdn_gating_pytorch
|
||||
from vllm_ascend.ops.triton.fused_gdn_gating import fused_gdn_gating_patch
|
||||
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
|
||||
|
||||
|
||||
def test_fused_gdn_gating_310p_parity_precision():
|
||||
init_device_properties_triton()
|
||||
torch.manual_seed(0)
|
||||
device = "npu"
|
||||
|
||||
num_tokens = 37
|
||||
num_heads = 8
|
||||
|
||||
A_log = torch.randn(num_heads, dtype=torch.float16, device=device)
|
||||
dt_bias = torch.randn(num_heads, dtype=torch.float16, device=device)
|
||||
a = torch.randn(num_tokens, num_heads, dtype=torch.float16, device=device)
|
||||
b = torch.randn(num_tokens, num_heads, dtype=torch.float16, device=device)
|
||||
|
||||
triton_g, triton_beta = fused_gdn_gating_patch(
|
||||
A_log=A_log,
|
||||
a=a,
|
||||
b=b,
|
||||
dt_bias=dt_bias,
|
||||
beta=1.0,
|
||||
threshold=20.0,
|
||||
)
|
||||
ref_g, ref_beta = fused_gdn_gating_pytorch(
|
||||
A_log=A_log,
|
||||
a=a,
|
||||
b=b,
|
||||
dt_bias=dt_bias,
|
||||
beta=1.0,
|
||||
threshold=20.0,
|
||||
)
|
||||
|
||||
torch.testing.assert_close(
|
||||
triton_g.to(torch.float32).cpu(),
|
||||
ref_g.to(torch.float32).cpu(),
|
||||
rtol=1e-2,
|
||||
atol=1e-2,
|
||||
equal_nan=True,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
triton_beta.to(torch.float32).cpu(),
|
||||
ref_beta.to(torch.float32).cpu(),
|
||||
rtol=1e-2,
|
||||
atol=1e-2,
|
||||
equal_nan=True,
|
||||
)
|
||||
@@ -0,0 +1,88 @@
|
||||
import gc
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from vllm.model_executor.layers.mamba.gdn.base import GatedDeltaNetAttention # type: ignore[import-not-found]
|
||||
|
||||
from vllm_ascend.ops.triton.fla.fused_qkvzba_split_reshape import fused_qkvzba_split_reshape_cat
|
||||
|
||||
|
||||
def validate_cmp(y_cal, y_ref, dtype, device="npu"):
|
||||
y_cal = y_cal.to(device)
|
||||
y_ref = y_ref.to(device)
|
||||
if dtype == torch.float16 or dtype == torch.bfloat16:
|
||||
torch.testing.assert_close(y_ref, y_cal, rtol=5e-03, atol=5e-03, equal_nan=True)
|
||||
elif dtype == torch.float32:
|
||||
torch.testing.assert_close(y_ref, y_cal, rtol=1e-03, atol=1e-03, equal_nan=True)
|
||||
elif (
|
||||
dtype == torch.int32
|
||||
or dtype == torch.int64
|
||||
or dtype == torch.int16
|
||||
or dtype == torch.int8
|
||||
or dtype == torch.uint32
|
||||
or dtype == torch.bool
|
||||
):
|
||||
assert torch.equal(y_cal, y_ref)
|
||||
else:
|
||||
raise ValueError('Invalid parameter "dtype" is found : {}'.format(dtype))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("seq_len", [1, 64, 1024, 2048])
|
||||
@pytest.mark.parametrize("num_heads_qk", [2, 4, 8])
|
||||
@pytest.mark.parametrize("num_heads_v", [8])
|
||||
@pytest.mark.parametrize("head_qk_dim", [256])
|
||||
@pytest.mark.parametrize("head_v_dim", [128])
|
||||
@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16])
|
||||
def test_fused_qkvzba_split_reshape_cat(
|
||||
seq_len,
|
||||
num_heads_qk,
|
||||
num_heads_v,
|
||||
head_qk_dim,
|
||||
head_v_dim,
|
||||
dtype,
|
||||
):
|
||||
if num_heads_v % num_heads_qk != 0:
|
||||
pytest.skip("num_heads_v must be divisible by num_heads_qk")
|
||||
|
||||
torch.random.manual_seed(0)
|
||||
device = "npu"
|
||||
|
||||
projected_states_qkvz = torch.randn(
|
||||
seq_len, 2 * head_qk_dim * num_heads_qk + 2 * head_v_dim * num_heads_v, dtype=dtype, device=device
|
||||
)
|
||||
|
||||
projected_states_ba = torch.randn(seq_len, 2 * num_heads_v, dtype=dtype, device=device)
|
||||
|
||||
projected_states_qkvz_copy = projected_states_qkvz.clone()
|
||||
projected_states_ba_copy = projected_states_ba.clone()
|
||||
|
||||
mixed_qkv, z, b, a = fused_qkvzba_split_reshape_cat(
|
||||
projected_states_qkvz_copy,
|
||||
projected_states_ba_copy,
|
||||
num_heads_qk,
|
||||
num_heads_v,
|
||||
head_qk_dim,
|
||||
head_v_dim,
|
||||
)
|
||||
|
||||
gdn = GatedDeltaNetAttention.__new__(GatedDeltaNetAttention)
|
||||
gdn.num_k_heads = num_heads_qk
|
||||
gdn.num_v_heads = num_heads_v
|
||||
gdn.head_k_dim = head_qk_dim
|
||||
gdn.head_v_dim = head_v_dim
|
||||
gdn.tp_size = 1
|
||||
|
||||
query, key, value, z_ref, b_ref, a_ref = gdn.fix_query_key_value_ordering(
|
||||
mixed_qkvz=projected_states_qkvz, mixed_ba=projected_states_ba
|
||||
)
|
||||
query, key, value = map(lambda x: rearrange(x, "l p d -> l (p d)"), (query, key, value))
|
||||
mixed_qkv_ref = torch.cat((query, key, value), dim=-1)
|
||||
|
||||
validate_cmp(mixed_qkv, mixed_qkv_ref, dtype)
|
||||
validate_cmp(z, z_ref, dtype)
|
||||
validate_cmp(b, b_ref, dtype)
|
||||
validate_cmp(a, a_ref, dtype)
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
@@ -0,0 +1,114 @@
|
||||
import pytest
|
||||
import torch
|
||||
from vllm.model_executor.layers.fla.ops import fused_recurrent_gated_delta_rule
|
||||
|
||||
from vllm_ascend._310p.ops.fla.fused_recurrent_gated_delta_rule import fused_recurrent_gated_delta_rule_pytorch
|
||||
|
||||
|
||||
@pytest.mark.skip("Probabilistic failure, need zengtian after fix")
|
||||
def test_fused_recurrent_gated_delta_rule_310p_parity_precision():
|
||||
torch.manual_seed(0)
|
||||
device = "npu"
|
||||
|
||||
bsz = 1
|
||||
total_tokens = 9
|
||||
num_qk_heads = 2
|
||||
num_v_heads = 4
|
||||
kdim = 64
|
||||
vdim = 48
|
||||
|
||||
q = torch.randn(bsz, total_tokens, num_qk_heads, kdim, dtype=torch.float16, device=device)
|
||||
k = torch.randn(bsz, total_tokens, num_qk_heads, kdim, dtype=torch.float16, device=device)
|
||||
v = torch.randn(bsz, total_tokens, num_v_heads, vdim, dtype=torch.float16, device=device)
|
||||
g = torch.randn(bsz, total_tokens, num_v_heads, dtype=torch.float32, device=device)
|
||||
beta = torch.sigmoid(torch.randn(bsz, total_tokens, num_v_heads, dtype=torch.float32, device=device)).to(
|
||||
torch.float16
|
||||
)
|
||||
|
||||
initial_state = torch.randn(2, num_v_heads, vdim, kdim, dtype=torch.float16, device=device)
|
||||
cu_seqlens = torch.tensor([0, 4, 9], dtype=torch.long, device=device)
|
||||
# For inplace_final_state=True, Ascend triton kernel expects explicit per-token state indices.
|
||||
# seq0 (len=4) -> state 0, seq1 (len=5) -> state 1.
|
||||
ssm_state_indices = torch.tensor(
|
||||
[
|
||||
[0, 0, 0, 0, 0],
|
||||
[1, 1, 1, 1, 1],
|
||||
],
|
||||
dtype=torch.long,
|
||||
device=device,
|
||||
)
|
||||
|
||||
triton_out, triton_state = fused_recurrent_gated_delta_rule(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
g=g,
|
||||
beta=beta,
|
||||
initial_state=initial_state.clone(),
|
||||
inplace_final_state=True,
|
||||
cu_seqlens=cu_seqlens,
|
||||
ssm_state_indices=ssm_state_indices,
|
||||
use_qk_l2norm_in_kernel=True,
|
||||
)
|
||||
ref_out, ref_state = fused_recurrent_gated_delta_rule_pytorch(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
g=g,
|
||||
beta=beta,
|
||||
initial_state=initial_state.clone(),
|
||||
inplace_final_state=True,
|
||||
cu_seqlens=cu_seqlens,
|
||||
ssm_state_indices=ssm_state_indices,
|
||||
use_qk_l2norm_in_kernel=True,
|
||||
)
|
||||
|
||||
torch.testing.assert_close(
|
||||
triton_out.to(torch.float32).cpu(),
|
||||
ref_out.to(torch.float32).cpu(),
|
||||
rtol=1e-2,
|
||||
atol=1e-2,
|
||||
equal_nan=True,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
triton_state.to(torch.float32).cpu(),
|
||||
ref_state.to(torch.float32).cpu(),
|
||||
rtol=1e-2,
|
||||
atol=1e-2,
|
||||
equal_nan=True,
|
||||
)
|
||||
|
||||
|
||||
def test_fused_recurrent_gated_delta_rule_310_state_layout_matches_vllm():
|
||||
q = torch.tensor([[[[1.0, 0.0]]]], dtype=torch.float32)
|
||||
k = torch.tensor([[[[1.0, 0.0]]]], dtype=torch.float32)
|
||||
v = torch.tensor([[[[10.0, 20.0, 30.0]]]], dtype=torch.float32)
|
||||
g = torch.zeros(1, 1, 1, dtype=torch.float32)
|
||||
beta = torch.ones(1, 1, 1, dtype=torch.float32)
|
||||
initial_state = torch.tensor(
|
||||
[[[[1.0, 2.0], [4.0, 8.0], [16.0, 32.0]]]],
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
out, final_state = fused_recurrent_gated_delta_rule_pytorch(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
g=g,
|
||||
beta=beta,
|
||||
initial_state=initial_state,
|
||||
inplace_final_state=False,
|
||||
cu_seqlens=None,
|
||||
ssm_state_indices=None,
|
||||
num_accepted_tokens=None,
|
||||
use_qk_l2norm_in_kernel=False,
|
||||
)
|
||||
|
||||
expected_out = torch.tensor([[[[10.0, 20.0, 30.0]]]], dtype=torch.float32) / (2.0**0.5)
|
||||
expected_state = torch.tensor(
|
||||
[[[[10.0, 2.0], [20.0, 8.0], [30.0, 32.0]]]],
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
torch.testing.assert_close(out, expected_out, rtol=1e-5, atol=1e-5)
|
||||
torch.testing.assert_close(final_state, expected_state, rtol=1e-5, atol=1e-5)
|
||||
@@ -0,0 +1,369 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
# mypy: ignore-errors
|
||||
"""Precision tests for vllm's fused_recurrent_kda Triton operator on NPU.
|
||||
|
||||
Tests the recurrent-mode (decode) kernel against a naive PyTorch recurrent
|
||||
reference implementation. Both ref and kernel use the same token-by-token
|
||||
recurrence algorithm, so errors are purely from FP accumulation differences
|
||||
on NPU triton-ascend.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torch_npu # noqa: F401
|
||||
|
||||
from vllm_ascend.ops.triton.kda.kda import fused_recurrent_kda
|
||||
|
||||
DEVICE = "npu"
|
||||
|
||||
# Both ref and kernel use the same recurrent algorithm; errors come from
|
||||
# FP accumulation differences on NPU triton-ascend.
|
||||
NPU_RMSE_RATIO_O = 0.005
|
||||
NPU_RMSE_RATIO_HT = 0.005
|
||||
|
||||
|
||||
def reference_l2norm(x: torch.Tensor, eps: float = 1e-6) -> torch.Tensor:
|
||||
dtype = x.dtype
|
||||
x = x.to(torch.float32)
|
||||
return (x * torch.rsqrt(torch.sum(x * x, dim=-1, keepdim=True) + eps)).to(dtype)
|
||||
|
||||
|
||||
def naive_recurrent_kda(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
g: torch.Tensor,
|
||||
beta: torch.Tensor,
|
||||
scale: float | None = None,
|
||||
initial_state: torch.Tensor | None = None,
|
||||
output_final_state: bool = False,
|
||||
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||
"""Naive recurrent KDA reference (pure PyTorch, runs on any device).
|
||||
|
||||
Ported from flash-linear-attention/fla/ops/kda/naive.py.
|
||||
"""
|
||||
dtype = v.dtype
|
||||
B, T, H, K, V = *q.shape, v.shape[-1]
|
||||
if scale is None:
|
||||
scale = K**-0.5
|
||||
|
||||
q, k, v, g, beta = map(lambda x: x.to(torch.float), [q, k, v, g, beta])
|
||||
q = q * scale
|
||||
|
||||
S = k.new_zeros(B, H, K, V).to(q)
|
||||
if initial_state is not None:
|
||||
S += initial_state
|
||||
o = torch.zeros_like(v)
|
||||
for i in range(T):
|
||||
q_i, k_i, v_i, g_i, b_i = q[:, i], k[:, i], v[:, i], g[:, i], beta[:, i]
|
||||
S = S * g_i[..., None].exp()
|
||||
S = S + torch.einsum(
|
||||
"bhk,bhv->bhkv",
|
||||
b_i[..., None] * k_i,
|
||||
v_i - (k_i[..., None] * S).sum(-2),
|
||||
)
|
||||
o[:, i] = torch.einsum("bhk,bhkv->bhv", q_i, S)
|
||||
if not output_final_state:
|
||||
S = None
|
||||
return o.to(dtype), S
|
||||
|
||||
|
||||
def assert_close(
|
||||
name: str,
|
||||
ref: torch.Tensor,
|
||||
tri: torch.Tensor,
|
||||
ratio: float,
|
||||
err_atol: float = 1e-6,
|
||||
):
|
||||
"""RMSE-based relative error comparison (same logic as FLA's assert_close)."""
|
||||
abs_err = (ref.detach() - tri.detach()).flatten().abs().max().item()
|
||||
rmse_diff = (ref.detach() - tri.detach()).flatten().square().mean().sqrt().item()
|
||||
rmse_base = ref.detach().flatten().square().mean().sqrt().item()
|
||||
rel_err = rmse_diff / (rmse_base + 1e-8)
|
||||
print(f"{name:>8} | max abs err: {abs_err:.6f} | rmse ratio: {rel_err:.6f} | threshold: {ratio}")
|
||||
if abs_err <= err_atol:
|
||||
return
|
||||
assert not torch.isnan(ref).any(), f"{name}: NaN detected in ref"
|
||||
assert not torch.isnan(tri).any(), f"{name}: NaN detected in tri"
|
||||
assert rel_err < ratio, f"{name}: max abs err {abs_err:.6f}, rmse ratio {rel_err:.6f} >= {ratio}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 1: Non-inplace varlen (clean output / state comparison)
|
||||
# ---------------------------------------------------------------------------
|
||||
@pytest.mark.parametrize(
|
||||
("H", "D", "cu_seqlens", "dtype"),
|
||||
[
|
||||
pytest.param(
|
||||
*test,
|
||||
id="H{}-D{}-cu{}-{}".format(*test),
|
||||
)
|
||||
for test in [
|
||||
# Decode: single token per sequence
|
||||
(32, 128, [0, 1], torch.float16),
|
||||
(32, 128, [0, 1, 2, 3, 4], torch.float16),
|
||||
(32, 128, [0, 1, 2, 3, 4, 5, 6, 7, 8], torch.float16),
|
||||
# Short sequences (multi-token recurrent)
|
||||
(32, 128, [0, 16], torch.float16),
|
||||
(32, 128, [0, 8, 24], torch.float16),
|
||||
(32, 128, [0, 4, 8, 16], torch.float16),
|
||||
(32, 128, [0, 64], torch.float16),
|
||||
# Different head count
|
||||
(64, 128, [0, 1, 2, 3, 4], torch.float16),
|
||||
# BFloat16
|
||||
(32, 128, [0, 1, 2, 3, 4], torch.bfloat16),
|
||||
(32, 128, [0, 8, 24], torch.bfloat16),
|
||||
]
|
||||
],
|
||||
)
|
||||
@pytest.mark.skip_global_cleanup
|
||||
@torch.inference_mode()
|
||||
def test_fused_recurrent_kda(
|
||||
H: int,
|
||||
D: int,
|
||||
cu_seqlens: list[int],
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
"""Non-inplace varlen mode — easiest to verify output and per-token state."""
|
||||
T = cu_seqlens[-1]
|
||||
N = len(cu_seqlens) - 1
|
||||
B = 1
|
||||
|
||||
torch.manual_seed(42)
|
||||
cu_seqlens_t = torch.LongTensor(cu_seqlens).to(DEVICE)
|
||||
|
||||
q = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
|
||||
k = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
|
||||
v = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
|
||||
g = F.logsigmoid(torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE)).to(dtype)
|
||||
beta = torch.rand(B, T, H, dtype=dtype, device=DEVICE).sigmoid()
|
||||
# Kernel layout: [T, H, V, K] = [T, H, D, D].
|
||||
# For varlen without ssm_state_indices, seq i reads from h0[cu_seqlens[i]].
|
||||
h0 = torch.randn(T, H, D, D, dtype=torch.float32, device=DEVICE)
|
||||
|
||||
# --- naive reference per sequence ---
|
||||
ref_outputs = []
|
||||
ref_states = []
|
||||
for i in range(N):
|
||||
s, e = cu_seqlens[i], cu_seqlens[i + 1]
|
||||
q_i = reference_l2norm(q[:, s:e].contiguous())
|
||||
k_i = reference_l2norm(k[:, s:e].contiguous())
|
||||
# Kernel state [H, V, K] -> naive [H, K, V]
|
||||
init_state_i = h0[s].transpose(-1, -2).unsqueeze(0)
|
||||
o_i, ht_i = naive_recurrent_kda(
|
||||
q_i,
|
||||
k_i,
|
||||
v[:, s:e],
|
||||
g[:, s:e],
|
||||
beta[:, s:e],
|
||||
initial_state=init_state_i,
|
||||
output_final_state=True,
|
||||
)
|
||||
ref_outputs.append(o_i)
|
||||
ref_states.append(ht_i)
|
||||
ref_o = torch.cat(ref_outputs, dim=1)
|
||||
|
||||
# --- Triton kernel ---
|
||||
tri_o, tri_ht = fused_recurrent_kda(
|
||||
q=q.clone(),
|
||||
k=k.clone(),
|
||||
v=v.clone(),
|
||||
g=g.clone(),
|
||||
beta=beta.clone(),
|
||||
initial_state=h0.clone(),
|
||||
inplace_final_state=False,
|
||||
use_qk_l2norm_in_kernel=True,
|
||||
cu_seqlens=cu_seqlens_t,
|
||||
)
|
||||
|
||||
assert not torch.isnan(tri_o).any(), "Triton output o contains NaN"
|
||||
assert not torch.isnan(tri_ht).any(), "Triton output ht contains NaN"
|
||||
|
||||
assert_close("o", ref_o, tri_o, NPU_RMSE_RATIO_O)
|
||||
# Compare final state per sequence: tri_ht[eos-1] in kernel layout [H,V,K]
|
||||
for i in range(N):
|
||||
e = cu_seqlens[i + 1]
|
||||
tri_state = tri_ht[e - 1].transpose(-1, -2).unsqueeze(0)
|
||||
assert_close(f"ht_{i}", ref_states[i], tri_state, NPU_RMSE_RATIO_HT)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 2: Inplace decode with ssm_state_indices (vllm actual pattern)
|
||||
# ---------------------------------------------------------------------------
|
||||
@pytest.mark.parametrize(
|
||||
("H", "D", "N", "dtype"),
|
||||
[
|
||||
pytest.param(
|
||||
*test,
|
||||
id="H{}-D{}-N{}-{}".format(*test),
|
||||
)
|
||||
for test in [
|
||||
(32, 128, 1, torch.float16),
|
||||
(32, 128, 4, torch.float16),
|
||||
(32, 128, 16, torch.float16),
|
||||
(64, 128, 4, torch.float16),
|
||||
(32, 128, 4, torch.bfloat16),
|
||||
]
|
||||
],
|
||||
)
|
||||
@pytest.mark.skip_global_cleanup
|
||||
@torch.inference_mode()
|
||||
def test_fused_recurrent_kda_decode_inplace(
|
||||
H: int,
|
||||
D: int,
|
||||
N: int,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
"""Decode with inplace state update + ssm_state_indices — vllm usage."""
|
||||
B = 1
|
||||
T = N # one token per sequence
|
||||
cu_seqlens = list(range(N + 1)) # [0, 1, 2, ..., N]
|
||||
|
||||
torch.manual_seed(42)
|
||||
cu_seqlens_t = torch.LongTensor(cu_seqlens).to(DEVICE)
|
||||
|
||||
q = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
|
||||
k = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
|
||||
v = torch.randn(B, T, H, D, dtype=dtype, device=DEVICE)
|
||||
g = F.logsigmoid(torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE)).to(dtype)
|
||||
beta = torch.rand(B, T, H, dtype=dtype, device=DEVICE).sigmoid()
|
||||
|
||||
# State buffer: slot 0 is NULL (invalid), slots 1..N are valid
|
||||
max_slots = N + 1
|
||||
state_buf = torch.randn(max_slots, H, D, D, dtype=torch.float32, device=DEVICE)
|
||||
state_buf[0] = 0 # NULL slot
|
||||
|
||||
# ssm_state_indices: seq i -> slot (i + 1), all valid (> 0)
|
||||
ssm_state_indices = torch.arange(1, N + 1, dtype=torch.long, device=DEVICE)
|
||||
|
||||
# --- naive reference per sequence ---
|
||||
ref_outputs = []
|
||||
ref_states = []
|
||||
for i in range(N):
|
||||
slot = i + 1
|
||||
q_i = reference_l2norm(q[:, i : i + 1].contiguous())
|
||||
k_i = reference_l2norm(k[:, i : i + 1].contiguous())
|
||||
init_state_i = state_buf[slot].transpose(-1, -2).unsqueeze(0)
|
||||
o_i, ht_i = naive_recurrent_kda(
|
||||
q_i,
|
||||
k_i,
|
||||
v[:, i : i + 1],
|
||||
g[:, i : i + 1],
|
||||
beta[:, i : i + 1],
|
||||
initial_state=init_state_i,
|
||||
output_final_state=True,
|
||||
)
|
||||
ref_outputs.append(o_i)
|
||||
ref_states.append(ht_i)
|
||||
ref_o = torch.cat(ref_outputs, dim=1)
|
||||
|
||||
# --- Triton kernel — inplace updates state_buf ---
|
||||
state_buf_tri = state_buf.clone()
|
||||
tri_o, _ = fused_recurrent_kda(
|
||||
q=q.clone(),
|
||||
k=k.clone(),
|
||||
v=v.clone(),
|
||||
g=g.clone(),
|
||||
beta=beta.clone(),
|
||||
initial_state=state_buf_tri,
|
||||
inplace_final_state=True,
|
||||
use_qk_l2norm_in_kernel=True,
|
||||
cu_seqlens=cu_seqlens_t,
|
||||
ssm_state_indices=ssm_state_indices,
|
||||
)
|
||||
|
||||
assert not torch.isnan(tri_o).any(), "Triton output o contains NaN"
|
||||
assert_close("o", ref_o, tri_o, NPU_RMSE_RATIO_O)
|
||||
|
||||
# Verify inplace state update at each slot
|
||||
for i in range(N):
|
||||
slot = i + 1
|
||||
tri_state = state_buf_tri[slot].transpose(-1, -2).unsqueeze(0)
|
||||
assert_close(f"ht_{i}", ref_states[i], tri_state, NPU_RMSE_RATIO_HT)
|
||||
|
||||
# Verify NULL slot was not modified
|
||||
assert torch.all(state_buf_tri[0] == 0), "NULL slot (0) should not be modified"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 3: Float32 — isolate algorithmic error from dtype precision
|
||||
# ---------------------------------------------------------------------------
|
||||
@pytest.mark.parametrize(
|
||||
("H", "D", "cu_seqlens"),
|
||||
[
|
||||
pytest.param(
|
||||
*test,
|
||||
id="H{}-D{}-cu{}".format(*test),
|
||||
)
|
||||
for test in [
|
||||
(32, 128, [0, 1]),
|
||||
(32, 128, [0, 1, 2, 3, 4]),
|
||||
(32, 128, [0, 16]),
|
||||
(32, 128, [0, 8, 24]),
|
||||
]
|
||||
],
|
||||
)
|
||||
@pytest.mark.skip_global_cleanup
|
||||
@torch.inference_mode()
|
||||
def test_fused_recurrent_kda_fp32(
|
||||
H: int,
|
||||
D: int,
|
||||
cu_seqlens: list[int],
|
||||
):
|
||||
"""Float32 test to isolate algorithmic error from dtype precision."""
|
||||
T = cu_seqlens[-1]
|
||||
N = len(cu_seqlens) - 1
|
||||
B = 1
|
||||
|
||||
torch.manual_seed(42)
|
||||
cu_seqlens_t = torch.LongTensor(cu_seqlens).to(DEVICE)
|
||||
|
||||
q = torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE)
|
||||
k = torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE)
|
||||
v = torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE)
|
||||
g = F.logsigmoid(torch.randn(B, T, H, D, dtype=torch.float32, device=DEVICE))
|
||||
beta = torch.rand(B, T, H, dtype=torch.float32, device=DEVICE).sigmoid()
|
||||
h0 = torch.randn(T, H, D, D, dtype=torch.float32, device=DEVICE)
|
||||
|
||||
ref_outputs = []
|
||||
ref_states = []
|
||||
for i in range(N):
|
||||
s, e = cu_seqlens[i], cu_seqlens[i + 1]
|
||||
q_i = reference_l2norm(q[:, s:e].contiguous())
|
||||
k_i = reference_l2norm(k[:, s:e].contiguous())
|
||||
init_state_i = h0[s].transpose(-1, -2).unsqueeze(0)
|
||||
o_i, ht_i = naive_recurrent_kda(
|
||||
q_i,
|
||||
k_i,
|
||||
v[:, s:e],
|
||||
g[:, s:e],
|
||||
beta[:, s:e],
|
||||
initial_state=init_state_i,
|
||||
output_final_state=True,
|
||||
)
|
||||
ref_outputs.append(o_i)
|
||||
ref_states.append(ht_i)
|
||||
ref_o = torch.cat(ref_outputs, dim=1)
|
||||
|
||||
tri_o, tri_ht = fused_recurrent_kda(
|
||||
q=q.clone(),
|
||||
k=k.clone(),
|
||||
v=v.clone(),
|
||||
g=g.clone(),
|
||||
beta=beta.clone(),
|
||||
initial_state=h0.clone(),
|
||||
inplace_final_state=False,
|
||||
use_qk_l2norm_in_kernel=True,
|
||||
cu_seqlens=cu_seqlens_t,
|
||||
)
|
||||
|
||||
assert not torch.isnan(tri_o).any(), "Triton output o contains NaN"
|
||||
assert not torch.isnan(tri_ht).any(), "Triton output ht contains NaN"
|
||||
|
||||
assert_close("o", ref_o, tri_o, NPU_RMSE_RATIO_O)
|
||||
for i in range(N):
|
||||
e = cu_seqlens[i + 1]
|
||||
tri_state = tri_ht[e - 1].transpose(-1, -2).unsqueeze(0)
|
||||
assert_close(f"ht_{i}", ref_states[i], tri_state, NPU_RMSE_RATIO_HT)
|
||||
@@ -0,0 +1,58 @@
|
||||
import gc
|
||||
|
||||
import torch
|
||||
from vllm.model_executor.layers.fla.ops import fused_recurrent_gated_delta_rule
|
||||
|
||||
from vllm_ascend.ops.triton.fla.sigmoid_gating import fused_sigmoid_gating_delta_rule_update
|
||||
from vllm_ascend.ops.triton.fused_gdn_gating import fused_gdn_gating_patch
|
||||
|
||||
|
||||
def test_triton_fusion_ops():
|
||||
q = torch.randn(1, 1, 4, 128, dtype=torch.bfloat16).npu()
|
||||
k = torch.randn(1, 1, 4, 128, dtype=torch.bfloat16).npu()
|
||||
v = torch.randn(1, 1, 8, 128, dtype=torch.bfloat16).npu()
|
||||
a = torch.tensor([[-2.6094, -0.2617, -0.3848, 2.2656, 3.6250, -0.7383, -1.0938, -0.0505]]).bfloat16().npu()
|
||||
b = torch.tensor([[0.4277, 0.8906, 1.6875, 2.3750, 4.1562, 0.3809, 1.0625, 3.6719]]).bfloat16().npu()
|
||||
non_spec_state_indices_tensor = torch.tensor([2]).int().npu()
|
||||
non_spec_query_start_loc = torch.tensor([0, 1]).int().npu()
|
||||
a_log = torch.tensor([-2.6875, -3.2031, -3.3438, -2.7812, -3.0625, -4.0312, -5.3750, 5.7188]).bfloat16().npu()
|
||||
dt_bias = torch.tensor([-4.7812, -5.0938, -5.5000, 9.4375, 7.6250, -4.3750, -3.0938, 0.9688]).bfloat16().npu()
|
||||
ssm_state1 = torch.ones(1, 8, 128, 128, dtype=torch.bfloat16).npu()
|
||||
core_attn_out_non_spec_fused = fused_sigmoid_gating_delta_rule_update(
|
||||
A_log=a_log.contiguous(),
|
||||
dt_bias=dt_bias.contiguous(),
|
||||
q=q.contiguous(),
|
||||
k=k.contiguous(),
|
||||
v=v.contiguous(),
|
||||
a=a.contiguous(),
|
||||
b=b.contiguous(),
|
||||
initial_state_source=ssm_state1,
|
||||
initial_state_indices=non_spec_state_indices_tensor,
|
||||
cu_seqlens=non_spec_query_start_loc,
|
||||
use_qk_l2norm_in_kernel=True,
|
||||
softplus_beta=1.0,
|
||||
softplus_threshold=20.0,
|
||||
)
|
||||
|
||||
ssm_state2 = torch.ones(1, 8, 128, 128, dtype=torch.bfloat16).npu()
|
||||
g, beta = fused_gdn_gating_patch(a_log, a, b, dt_bias)
|
||||
g_non_spec = g
|
||||
beta_non_spec = beta
|
||||
core_attn_out_non_spec_split, last_recurrent_state = fused_recurrent_gated_delta_rule(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
g=g_non_spec,
|
||||
beta=beta_non_spec,
|
||||
initial_state=ssm_state2,
|
||||
inplace_final_state=True,
|
||||
cu_seqlens=non_spec_query_start_loc,
|
||||
ssm_state_indices=non_spec_state_indices_tensor,
|
||||
use_qk_l2norm_in_kernel=True,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
core_attn_out_non_spec_fused, core_attn_out_non_spec_split, rtol=1e-02, atol=1e-02, equal_nan=True
|
||||
)
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
@@ -0,0 +1,39 @@
|
||||
import gc
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from vllm_ascend.ops.triton.fla.l2norm import l2norm_fwd
|
||||
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("B", "T", "H", "D", "dtype"),
|
||||
[
|
||||
pytest.param(*test, id="B{}-T{}-H{}-D{}-{}".format(*test))
|
||||
for test in [
|
||||
(1, 63, 1, 60, torch.float),
|
||||
(2, 500, 4, 64, torch.float),
|
||||
(2, 1000, 2, 100, torch.float),
|
||||
(3, 1024, 4, 128, torch.float),
|
||||
]
|
||||
],
|
||||
)
|
||||
def test_l2norm(B: int, T: int, H: int, D: int, dtype: torch.dtype):
|
||||
torch.manual_seed(42)
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
rtol, atol = (3e-4, 1e-3) if dtype == torch.float32 else (3e-3, 5e-3)
|
||||
if dtype == torch.bfloat16:
|
||||
rtol, atol = 1e-2, 5e-2
|
||||
x = torch.randn(B, T, H, D, dtype=dtype).to(device).requires_grad_(True)
|
||||
x = x * 0.5 + 0.3
|
||||
|
||||
ref = F.normalize(x, dim=-1, p=2)
|
||||
tri = l2norm_fwd(x)
|
||||
|
||||
assert torch.allclose(tri, ref, rtol=rtol, atol=atol)
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
@@ -0,0 +1,930 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
#
|
||||
|
||||
"""Tests for lightning attention triton kernels.
|
||||
|
||||
Covers the following 4 triton kernels:
|
||||
- ``_fwd_diag_kernel``: diagonal block causal attention
|
||||
- ``_fwd_kv_parallel``: key-value outer product per block
|
||||
- ``_fwd_kv_reduce``: prefix-sum reduction of KV across blocks
|
||||
- ``_fwd_none_diag_kernel``: non-diagonal block attention
|
||||
|
||||
All kernels are exercised through the public APIs:
|
||||
- ``lightning_attention_npu_`` (``_attention.apply``, single d-chunk)
|
||||
- ``lightning_attention_npu`` (full function with d-dimension chunking)
|
||||
- ``AscendLightningAttentionKernel.jit_linear_forward_prefix``
|
||||
|
||||
The naive reference implementations replicate the exact triton tiling
|
||||
algorithm (diagonal blocks -> KV parallel -> KV reduce -> non-diagonal)
|
||||
to ensure an apples-to-apples numerical comparison. Low-precision inputs use
|
||||
bounded random values and mirror the kernel's output stores to avoid comparing
|
||||
against an unrealistically precise PyTorch recurrence.
|
||||
|
||||
The production BailingMoE path promotes QKV to float32 before calling these
|
||||
kernels and commonly uses a float32 Mamba cache. Therefore larger accuracy
|
||||
cases use float32, while bf16/fp16 cases are kept as bounded compatibility
|
||||
coverage for the raw operator entry points.
|
||||
"""
|
||||
|
||||
import gc
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from einops import rearrange
|
||||
|
||||
from vllm_ascend.ops.triton.mamba.lightning_attn import (
|
||||
AscendLightningAttentionKernel,
|
||||
lightning_attention_npu,
|
||||
lightning_attention_npu_,
|
||||
)
|
||||
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Naive reference implementations (pure PyTorch, no Triton)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# The triton lightning attention algorithm processes the sequence in BLOCK-sized
|
||||
# tiles and follows these steps:
|
||||
# 1. _fwd_diag_kernel : causal attention within each diagonal block
|
||||
# 2. _fwd_kv_parallel : per-block KV outer product, decayed to block end
|
||||
# 3. _fwd_kv_reduce : prefix-sum across blocks, updates kv_history
|
||||
# 4. _fwd_none_diag_kernel : non-diagonal attention using prefix KV
|
||||
#
|
||||
# IMPORTANT: The non-diagonal kernel applies decay as exp(-s * t_in_block).
|
||||
# The references below mirror that tiled kernel convention directly rather
|
||||
# than using an equivalent-looking recurrence with a different history state.
|
||||
|
||||
BLOCK_SIZE = 256
|
||||
DIAG_BLOCK_SIZE = 32
|
||||
KV_BLOCK_SIZE = 64
|
||||
TEST_INPUT_SCALE = 0.02
|
||||
TEST_DECAY_SCALE = 0.02
|
||||
SEMANTIC_INPUT_SCALE = 0.05
|
||||
FLOAT32_TOLERANCE = (1e-2, 1e-2)
|
||||
LOW_PRECISION_TOLERANCE = (5e-2, 5e-2)
|
||||
|
||||
|
||||
def _round_like_kernel_store(x, dtype):
|
||||
"""Mirror Triton stores to output dtype, then compare in float32."""
|
||||
if dtype == torch.float32:
|
||||
return x
|
||||
return x.to(dtype).float()
|
||||
|
||||
|
||||
def _randn(shape, dtype, device, scale=TEST_INPUT_SCALE):
|
||||
return (torch.randn(*shape, dtype=dtype, device=device) * scale).to(dtype)
|
||||
|
||||
|
||||
def _rand_decay(h, device, scale=TEST_DECAY_SCALE):
|
||||
return torch.rand(h, dtype=torch.float32, device=device) * scale
|
||||
|
||||
|
||||
def _naive_triton_lightning_attention(q, k, v, s, kv_history, block_size=BLOCK_SIZE):
|
||||
"""Step-by-step replication of the triton tiling algorithm for a single
|
||||
d-chunk.
|
||||
|
||||
This function mirrors the four-kernel pipeline used in
|
||||
``_attention.apply`` (a.k.a. ``lightning_attention_npu_``):
|
||||
1. Diagonal blocks (within-block causal attention)
|
||||
2. Per-block KV outer product (decayed to block end)
|
||||
3. Prefix-sum reduce across blocks (updates kv_history in place)
|
||||
4. Non-diagonal blocks (cross-block attention with prefix KV)
|
||||
|
||||
Args:
|
||||
q: [b, h, n, d] queries
|
||||
k: [b, h, n, d] keys
|
||||
v: [b, h, n, e] values
|
||||
s: [1, h, 1, 1] per-head decay rates
|
||||
kv_history: [b, h, d, e] accumulated KV state from previous steps
|
||||
|
||||
Returns:
|
||||
(o, kv_return) where
|
||||
o: [b, h, n, e] attention output
|
||||
kv_return: [b, h, num_blocks + 1, d, e] block KV plus final state
|
||||
"""
|
||||
output_dtype = q.dtype
|
||||
b, h, n, d = q.shape
|
||||
e_dim = v.shape[-1]
|
||||
|
||||
q = q.float()
|
||||
k = k.float()
|
||||
v = v.float()
|
||||
decay_rate = s.float().reshape(1, h, 1, 1)
|
||||
|
||||
num_blocks = (n + block_size - 1) // block_size
|
||||
|
||||
# ---- Step 1: Diagonal blocks ----
|
||||
# Mirror _fwd_diag_kernel's tiled matmul path. The recurrence is
|
||||
# mathematically equivalent, but fp16/bf16 accumulation order differs.
|
||||
o = torch.zeros(b, h, n, e_dim, dtype=torch.float32, device=q.device)
|
||||
for block_idx in range(num_blocks):
|
||||
blk_start = block_idx * block_size
|
||||
blk_end = min(blk_start + block_size, n)
|
||||
block_len = blk_end - blk_start
|
||||
for q_start in range(0, block_len, DIAG_BLOCK_SIZE):
|
||||
q_end = min(q_start + DIAG_BLOCK_SIZE, block_len)
|
||||
q_block = q[:, :, blk_start + q_start : blk_start + q_end, :]
|
||||
q_pos = torch.arange(q_start, q_end, dtype=torch.float32, device=q.device)
|
||||
q_len = q_end - q_start
|
||||
qkv = torch.zeros(b, h, q_len, e_dim, dtype=torch.float32, device=q.device)
|
||||
|
||||
for kv_start in range(0, q_start + DIAG_BLOCK_SIZE, DIAG_BLOCK_SIZE):
|
||||
kv_end = min(kv_start + DIAG_BLOCK_SIZE, block_len)
|
||||
if kv_start >= kv_end:
|
||||
continue
|
||||
|
||||
k_block = k[:, :, blk_start + kv_start : blk_start + kv_end, :]
|
||||
v_block = v[:, :, blk_start + kv_start : blk_start + kv_end, :]
|
||||
kv_pos = torch.arange(kv_start, kv_end, dtype=torch.float32, device=q.device)
|
||||
kv_len = kv_end - kv_start
|
||||
diff = q_pos[:, None] - kv_pos[None, :]
|
||||
causal_mask = (diff >= 0).reshape(1, 1, q_len, kv_len)
|
||||
decay = torch.exp(-decay_rate * diff.clamp_min(0).reshape(1, 1, q_len, kv_len))
|
||||
decay = decay * causal_mask.to(torch.float32)
|
||||
|
||||
qk = torch.matmul(q_block, k_block.transpose(-1, -2)) * decay
|
||||
qkv = qkv + torch.matmul(qk, v_block)
|
||||
|
||||
o[:, :, blk_start + q_start : blk_start + q_end, :] = _round_like_kernel_store(qkv, output_dtype)
|
||||
|
||||
# ---- Step 2: Per-block KV outer product ----
|
||||
# For each block, accumulate K^T @ V with each row decayed to the end
|
||||
# of that block. The last partial block uses the same left-shifted
|
||||
# CBLOCK layout as _fwd_kv_parallel.
|
||||
kv_block = torch.zeros(b, h, num_blocks, d, e_dim, dtype=torch.float32, device=q.device)
|
||||
for block_idx in range(num_blocks):
|
||||
blk_start = block_idx * block_size
|
||||
blk_end = min(blk_start + block_size, n)
|
||||
block_len = blk_end - blk_start
|
||||
num_kv_blocks = min((block_len + KV_BLOCK_SIZE - 1) // KV_BLOCK_SIZE, block_size // KV_BLOCK_SIZE)
|
||||
left_shift = num_kv_blocks * KV_BLOCK_SIZE - block_len
|
||||
decay_start = (block_size // KV_BLOCK_SIZE - num_kv_blocks) * KV_BLOCK_SIZE
|
||||
|
||||
for kv_block_idx in range(num_kv_blocks):
|
||||
row_offsets = torch.arange(KV_BLOCK_SIZE, device=q.device)
|
||||
source_pos = kv_block_idx * KV_BLOCK_SIZE - left_shift + row_offsets
|
||||
left_bound = (1 - kv_block_idx) * left_shift
|
||||
valid = (row_offsets >= left_bound) & (source_pos >= 0) & (source_pos < block_len)
|
||||
safe_pos = source_pos.clamp(0, max(block_len - 1, 0))
|
||||
k_block = k[:, :, blk_start + safe_pos, :] * valid.reshape(1, 1, KV_BLOCK_SIZE, 1)
|
||||
v_block = v[:, :, blk_start + safe_pos, :] * valid.reshape(1, 1, KV_BLOCK_SIZE, 1)
|
||||
decay_pos = decay_start + kv_block_idx * KV_BLOCK_SIZE + row_offsets
|
||||
k_decay = torch.exp(-decay_rate * (block_size - 1 - decay_pos).float().reshape(1, 1, KV_BLOCK_SIZE, 1))
|
||||
weighted_k = k_block * k_decay
|
||||
kv_block[:, :, block_idx, :, :] = kv_block[:, :, block_idx, :, :] + torch.matmul(
|
||||
weighted_k.transpose(-1, -2),
|
||||
v_block,
|
||||
)
|
||||
|
||||
# ---- Step 3: Prefix-sum reduce across blocks ----
|
||||
# Replicates _fwd_kv_reduce exactly:
|
||||
# kv_pre starts as the existing kv_history
|
||||
# For each block i:
|
||||
# block_decay = exp(-s * block_size)
|
||||
# store kv_pre into KV[i]
|
||||
# kv_pre = block_decay * kv_pre + kv_cur
|
||||
kv = kv_block.clone()
|
||||
kv_pre = kv_history.clone().float()
|
||||
for i in range(num_blocks):
|
||||
blk_size = min(n - i * block_size, block_size)
|
||||
block_decay = torch.exp(-decay_rate * blk_size) # [1, h, 1, 1]
|
||||
kv_cur = kv[:, :, i, :, :].clone()
|
||||
kv[:, :, i, :, :] = kv_pre
|
||||
kv_pre = block_decay * kv_pre + kv_cur
|
||||
kv_history_updated = kv_pre
|
||||
|
||||
# ---- Step 4: Non-diagonal blocks ----
|
||||
# O[t] += Q[t] @ kv[block_idx] * exp(-s * t_local)
|
||||
# Note: the triton kernel uses t_local (NOT t_local + 1),
|
||||
# matching _fwd_none_diag_kernel's q_decay = exp(-s * (off_c*CBLOCK + c)).
|
||||
for block_idx in range(num_blocks):
|
||||
blk_start = block_idx * block_size
|
||||
blk_end = min(blk_start + block_size, n)
|
||||
block_len = blk_end - blk_start
|
||||
for q_start in range(0, block_len, KV_BLOCK_SIZE):
|
||||
q_end = min(q_start + KV_BLOCK_SIZE, block_len)
|
||||
q_len = q_end - q_start
|
||||
q_block = q[:, :, blk_start + q_start : blk_start + q_end, :]
|
||||
q_pos = torch.arange(q_start, q_end, dtype=torch.float32, device=q.device)
|
||||
q_decay = torch.exp(-decay_rate * q_pos.reshape(1, 1, q_len, 1))
|
||||
nondiag = torch.matmul(q_block, kv[:, :, block_idx, :, :]) * q_decay
|
||||
out_slice = o[:, :, blk_start + q_start : blk_start + q_end, :] + nondiag
|
||||
o[:, :, blk_start + q_start : blk_start + q_end, :] = _round_like_kernel_store(out_slice, output_dtype)
|
||||
|
||||
return _round_like_kernel_store(o, output_dtype), torch.cat([kv, kv_history_updated.unsqueeze(2)], dim=2)
|
||||
|
||||
|
||||
def _naive_lightning_attention_npu(q, k, v, ed, block_size, kv_history):
|
||||
"""Naive reference that replicates ``lightning_attention_npu``.
|
||||
|
||||
Handles the d-dimension chunking in the same way as the real
|
||||
implementation so that the comparison is apples-to-apples.
|
||||
"""
|
||||
d = q.shape[-1]
|
||||
e_dim = v.shape[-1]
|
||||
|
||||
if ed.dim() == 1:
|
||||
ed = ed.view(1, -1, 1, 1)
|
||||
|
||||
m = 128 if d >= 128 else 64
|
||||
arr = [m * i for i in range(d // m + 1)]
|
||||
if arr[-1] != d:
|
||||
arr.append(d)
|
||||
|
||||
if kv_history is None:
|
||||
kv_history = torch.zeros(
|
||||
(q.shape[0], q.shape[1], d, e_dim),
|
||||
dtype=torch.float32,
|
||||
device=q.device,
|
||||
)
|
||||
else:
|
||||
kv_history = kv_history.clone().contiguous().float()
|
||||
|
||||
output = torch.zeros(
|
||||
q.shape[0],
|
||||
q.shape[1],
|
||||
q.shape[2],
|
||||
e_dim,
|
||||
dtype=torch.float32,
|
||||
device=q.device,
|
||||
)
|
||||
|
||||
kv_state = None
|
||||
for i in range(len(arr) - 1):
|
||||
s_idx = arr[i]
|
||||
e_idx = arr[i + 1]
|
||||
q1 = q[..., s_idx:e_idx]
|
||||
k1 = k[..., s_idx:e_idx]
|
||||
kv_history_chunk = kv_history[:, :, s_idx:e_idx, :]
|
||||
o, kv_state = _naive_triton_lightning_attention(
|
||||
q1,
|
||||
k1,
|
||||
v,
|
||||
ed,
|
||||
kv_history_chunk,
|
||||
block_size=block_size,
|
||||
)
|
||||
output = _round_like_kernel_store(output + o, q.dtype)
|
||||
kv_history[:, :, s_idx:e_idx, :] = kv_state[:, :, -1, :, :]
|
||||
|
||||
return output, kv_state
|
||||
|
||||
|
||||
def _naive_jit_linear_forward_prefix(q, k, v, kv_caches, slope_rate, block_size, layer_idx=None):
|
||||
"""Naive reference for ``AscendLightningAttentionKernel.jit_linear_forward_prefix``."""
|
||||
slope_rate = slope_rate.to(torch.float32)
|
||||
should_squeeze = q.dim() == 3
|
||||
if should_squeeze:
|
||||
q = q.unsqueeze(0)
|
||||
k = k.unsqueeze(0)
|
||||
v = v.unsqueeze(0)
|
||||
|
||||
b, h, n, d = q.shape
|
||||
e_dim = v.shape[-1]
|
||||
|
||||
if slope_rate.dim() == 1:
|
||||
ed = slope_rate.view(1, -1, 1, 1)
|
||||
else:
|
||||
ed = slope_rate
|
||||
|
||||
kv_history = kv_caches.reshape(1, h, d, e_dim).contiguous().float()
|
||||
output, kv_state = _naive_lightning_attention_npu(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
ed,
|
||||
block_size,
|
||||
kv_history,
|
||||
)
|
||||
|
||||
# The triton kernel updates kv_caches in-place with the final KV state
|
||||
kv_caches_out = kv_state[:, :, -1, :, :].reshape_as(kv_caches)
|
||||
|
||||
assert output.shape[0] == 1, "batch size must be 1"
|
||||
result = rearrange(output.squeeze(0), "h n d -> n (h d)")
|
||||
return result.float(), kv_caches_out
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tolerance helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_tolerances(dtype):
|
||||
"""Return (rtol, atol) appropriate for the given dtype."""
|
||||
if dtype == torch.float32:
|
||||
return FLOAT32_TOLERANCE
|
||||
elif dtype in (torch.float16, torch.bfloat16):
|
||||
return LOW_PRECISION_TOLERANCE
|
||||
return FLOAT32_TOLERANCE
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests for lightning_attention_npu_ (single d-chunk, exercises all 4 kernels)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("b", "h", "n", "d", "e", "dtype"),
|
||||
[
|
||||
pytest.param(*case, id=f"b{case[0]}-h{case[1]}-n{case[2]}-d{case[3]}-e{case[4]}-{str(case[5]).split('.')[-1]}")
|
||||
for case in [
|
||||
# Small seq, basic sanity
|
||||
(1, 4, 32, 128, 64, torch.bfloat16),
|
||||
# n < BLOCK (256), not aligned
|
||||
(1, 4, 100, 128, 128, torch.bfloat16),
|
||||
# float16 dtype
|
||||
(1, 4, 256, 128, 128, torch.float16),
|
||||
# float32 dtype
|
||||
(1, 4, 128, 128, 128, torch.float32),
|
||||
# batch > 1
|
||||
(2, 4, 128, 128, 64, torch.bfloat16),
|
||||
# d = 64 (smaller head dim)
|
||||
(1, 4, 128, 64, 64, torch.bfloat16),
|
||||
# e != d
|
||||
(1, 4, 128, 64, 128, torch.bfloat16),
|
||||
]
|
||||
],
|
||||
)
|
||||
def test_lightning_attention_npu_single_chunk(b, h, n, d, e, dtype):
|
||||
"""Test lightning_attention_npu_ against naive PyTorch reference.
|
||||
|
||||
This exercises all 4 triton kernels through the _attention.apply path.
|
||||
All n values are <= BLOCK (256) to ensure single-block correctness.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
rtol, atol = _get_tolerances(dtype)
|
||||
|
||||
q = _randn((b, h, n, d), dtype, device)
|
||||
k = _randn((b, h, n, d), dtype, device)
|
||||
v = _randn((b, h, n, e), dtype, device)
|
||||
ed = _rand_decay(h, device).view(1, h, 1, 1)
|
||||
kv_history = torch.zeros(b, h, d, e, dtype=torch.float32, device=device)
|
||||
|
||||
# NOTE: Must clone kv_history before the triton call because
|
||||
# _fwd_kv_reduce modifies it in-place. The naive reference must
|
||||
# receive the ORIGINAL (pre-modification) value.
|
||||
o_triton, kv_triton = lightning_attention_npu_(q, k, v, ed, kv_history.clone())
|
||||
o_ref, _ = _naive_triton_lightning_attention(q, k, v, ed, kv_history)
|
||||
|
||||
torch.testing.assert_close(
|
||||
o_triton.float().cpu(),
|
||||
o_ref.cpu(),
|
||||
rtol=rtol,
|
||||
atol=atol,
|
||||
)
|
||||
assert kv_triton.shape == (b, h, (n + BLOCK_SIZE - 1) // BLOCK_SIZE + 1, d, e)
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("b", "h", "n", "d", "e", "dtype"),
|
||||
[
|
||||
pytest.param(*case, id=f"b{case[0]}-h{case[1]}-n{case[2]}-d{case[3]}-e{case[4]}-{str(case[5]).split('.')[-1]}")
|
||||
for case in [
|
||||
(1, 4, 256, 128, 128, torch.bfloat16),
|
||||
(2, 4, 128, 128, 64, torch.bfloat16),
|
||||
]
|
||||
],
|
||||
)
|
||||
def test_lightning_attention_npu_single_chunk_with_kv_history(b, h, n, d, e, dtype):
|
||||
"""Test lightning_attention_npu_ with non-zero initial KV history."""
|
||||
torch.manual_seed(42)
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
rtol, atol = _get_tolerances(dtype)
|
||||
|
||||
q = _randn((b, h, n, d), dtype, device)
|
||||
k = _randn((b, h, n, d), dtype, device)
|
||||
v = _randn((b, h, n, e), dtype, device)
|
||||
ed = _rand_decay(h, device).view(1, h, 1, 1)
|
||||
kv_history = _randn((b, h, d, e), torch.float32, device)
|
||||
|
||||
o_triton, kv_triton = lightning_attention_npu_(q, k, v, ed, kv_history.clone())
|
||||
o_ref, kv_ref = _naive_triton_lightning_attention(q, k, v, ed, kv_history)
|
||||
|
||||
torch.testing.assert_close(
|
||||
o_triton.float().cpu(),
|
||||
o_ref.cpu(),
|
||||
rtol=rtol,
|
||||
atol=atol,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
kv_triton[:, :, -1, :, :].cpu(),
|
||||
kv_ref[:, :, -1, :, :].cpu(),
|
||||
rtol=rtol,
|
||||
atol=atol,
|
||||
)
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests for multi-block sequences (n > BLOCK)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("b", "h", "n", "d", "e", "dtype"),
|
||||
[
|
||||
pytest.param(*case, id=f"b{case[0]}-h{case[1]}-n{case[2]}-d{case[3]}-e{case[4]}-{str(case[5]).split('.')[-1]}")
|
||||
for case in [
|
||||
# n > BLOCK, not aligned
|
||||
(1, 4, 300, 128, 128, torch.bfloat16),
|
||||
# Larger n, production path promotes q/k/v to float32.
|
||||
(1, 4, 768, 128, 128, torch.float32),
|
||||
]
|
||||
],
|
||||
)
|
||||
def test_lightning_attention_npu_multi_block(b, h, n, d, e, dtype):
|
||||
"""Test lightning_attention_npu_ with multi-block sequences (n > BLOCK).
|
||||
|
||||
This exercises the _fwd_kv_parallel, _fwd_kv_reduce, and
|
||||
_fwd_none_diag_kernel in the multi-block path.
|
||||
Uses small decay rates so inter-block numerical differences
|
||||
remain within tolerance.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
rtol, atol = _get_tolerances(dtype)
|
||||
|
||||
q = _randn((b, h, n, d), dtype, device)
|
||||
k = _randn((b, h, n, d), dtype, device)
|
||||
v = _randn((b, h, n, e), dtype, device)
|
||||
# Use small decay rates to keep errors within tolerance
|
||||
ed = _rand_decay(h, device, scale=0.01)
|
||||
ed = ed.view(1, h, 1, 1)
|
||||
kv_history = torch.zeros(b, h, d, e, dtype=torch.float32, device=device)
|
||||
|
||||
o_triton, kv_triton = lightning_attention_npu_(q, k, v, ed, kv_history.clone())
|
||||
o_ref, kv_ref = _naive_triton_lightning_attention(q, k, v, ed, kv_history)
|
||||
|
||||
torch.testing.assert_close(
|
||||
o_triton.float().cpu(),
|
||||
o_ref.cpu(),
|
||||
rtol=rtol,
|
||||
atol=atol,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
kv_triton[:, :, -1, :, :].cpu(),
|
||||
kv_ref[:, :, -1, :, :].cpu(),
|
||||
rtol=rtol,
|
||||
atol=atol,
|
||||
)
|
||||
|
||||
# Also verify the output is well-formed
|
||||
assert not torch.isnan(o_triton).any(), "Output contains NaN values"
|
||||
assert not torch.isinf(o_triton).any(), "Output contains Inf values"
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests for lightning_attention_npu (full function with d-dimension chunking)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("b", "h", "n", "d", "e", "dtype"),
|
||||
[
|
||||
pytest.param(*case, id=f"b{case[0]}-h{case[1]}-n{case[2]}-d{case[3]}-e{case[4]}-{str(case[5]).split('.')[-1]}")
|
||||
for case in [
|
||||
# d=128 -> single chunk (m=128)
|
||||
(1, 4, 128, 128, 128, torch.bfloat16),
|
||||
# d=64 -> single chunk (m=64)
|
||||
(1, 4, 128, 64, 64, torch.bfloat16),
|
||||
# n not aligned to BLOCK, single-block
|
||||
(1, 4, 100, 128, 128, torch.bfloat16),
|
||||
# n == BLOCK, production path promotes q/k/v to float32.
|
||||
(1, 4, 256, 128, 128, torch.float32),
|
||||
]
|
||||
],
|
||||
)
|
||||
def test_lightning_attention_npu(b, h, n, d, e, dtype):
|
||||
"""Test lightning_attention_npu (with d-dimension chunking) against naive reference."""
|
||||
torch.manual_seed(42)
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
rtol, atol = _get_tolerances(dtype)
|
||||
|
||||
q = _randn((b, h, n, d), dtype, device)
|
||||
k = _randn((b, h, n, d), dtype, device)
|
||||
v = _randn((b, h, n, e), dtype, device)
|
||||
ed = _rand_decay(h, device)
|
||||
|
||||
# Triton output
|
||||
o_triton, kv_triton = lightning_attention_npu(q, k, v, ed, block_size=256, kv_history=None)
|
||||
|
||||
# Naive reference output
|
||||
o_ref, kv_ref = _naive_lightning_attention_npu(q, k, v, ed, block_size=256, kv_history=None)
|
||||
|
||||
torch.testing.assert_close(
|
||||
o_triton.float().cpu(),
|
||||
o_ref.cpu(),
|
||||
rtol=rtol,
|
||||
atol=atol,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
kv_triton[:, :, -1, :, :].cpu(),
|
||||
kv_ref[:, :, -1, :, :].cpu(),
|
||||
rtol=rtol,
|
||||
atol=atol,
|
||||
)
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("b", "h", "n", "d", "e", "dtype"),
|
||||
[
|
||||
pytest.param(*case, id=f"b{case[0]}-h{case[1]}-n{case[2]}-d{case[3]}-e{case[4]}-{str(case[5]).split('.')[-1]}")
|
||||
for case in [
|
||||
# Production path promotes q/k/v to float32 while cache is float32.
|
||||
(1, 4, 128, 128, 128, torch.float32),
|
||||
(1, 4, 256, 64, 128, torch.bfloat16),
|
||||
]
|
||||
],
|
||||
)
|
||||
def test_lightning_attention_npu_with_kv_history(b, h, n, d, e, dtype):
|
||||
"""Test lightning_attention_npu with pre-existing KV history."""
|
||||
torch.manual_seed(42)
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
rtol, atol = _get_tolerances(dtype)
|
||||
|
||||
q = _randn((b, h, n, d), dtype, device)
|
||||
k = _randn((b, h, n, d), dtype, device)
|
||||
v = _randn((b, h, n, e), dtype, device)
|
||||
ed = _rand_decay(h, device)
|
||||
kv_history = _randn((b, h, d, e), torch.float32, device)
|
||||
|
||||
o_triton, kv_triton = lightning_attention_npu(q, k, v, ed, block_size=256, kv_history=kv_history.clone())
|
||||
o_ref, kv_ref = _naive_lightning_attention_npu(q, k, v, ed, block_size=256, kv_history=kv_history)
|
||||
|
||||
torch.testing.assert_close(
|
||||
o_triton.float().cpu(),
|
||||
o_ref.cpu(),
|
||||
rtol=rtol,
|
||||
atol=atol,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
kv_triton[:, :, -1, :, :].cpu(),
|
||||
kv_ref[:, :, -1, :, :].cpu(),
|
||||
rtol=rtol,
|
||||
atol=atol,
|
||||
)
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests for AscendLightningAttentionKernel.jit_linear_forward_prefix
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("h", "n", "d", "e", "dtype"),
|
||||
[
|
||||
pytest.param(*case, id=f"h{case[0]}-n{case[1]}-d{case[2]}-e{case[3]}-{str(case[4]).split('.')[-1]}")
|
||||
for case in [
|
||||
(4, 128, 128, 128, torch.bfloat16),
|
||||
(8, 256, 128, 128, torch.float32),
|
||||
(4, 100, 128, 128, torch.bfloat16),
|
||||
(4, 256, 64, 64, torch.bfloat16),
|
||||
(4, 128, 128, 128, torch.float16),
|
||||
]
|
||||
],
|
||||
)
|
||||
def test_ascend_lightning_attention_kernel_prefix(h, n, d, e, dtype):
|
||||
"""Test AscendLightningAttentionKernel.jit_linear_forward_prefix."""
|
||||
torch.manual_seed(42)
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
rtol, atol = _get_tolerances(dtype)
|
||||
|
||||
# jit_linear_forward_prefix receives one sequence in [h, n, d] layout.
|
||||
q = _randn((h, n, d), dtype, device)
|
||||
k = _randn((h, n, d), dtype, device)
|
||||
v = _randn((h, n, e), dtype, device)
|
||||
|
||||
slope_rate = _rand_decay(h, device)
|
||||
kv_caches = torch.zeros(h, d, e, dtype=torch.float32, device=device)
|
||||
|
||||
# Triton output (clone all shared tensors)
|
||||
kv_caches_triton = kv_caches.clone()
|
||||
out_triton = AscendLightningAttentionKernel.jit_linear_forward_prefix(
|
||||
q.clone(), k.clone(), v.clone(), kv_caches_triton, slope_rate.clone(), block_size=256
|
||||
)
|
||||
|
||||
# Naive reference output
|
||||
out_ref, kv_caches_ref = _naive_jit_linear_forward_prefix(q, k, v, kv_caches, slope_rate, block_size=256)
|
||||
|
||||
torch.testing.assert_close(
|
||||
out_triton.float().cpu(),
|
||||
out_ref.cpu(),
|
||||
rtol=rtol,
|
||||
atol=atol,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
kv_caches_triton.cpu(),
|
||||
kv_caches_ref.cpu(),
|
||||
rtol=rtol,
|
||||
atol=atol,
|
||||
)
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("h", "n", "d", "e", "dtype"),
|
||||
[
|
||||
pytest.param(*case, id=f"h{case[0]}-n{case[1]}-d{case[2]}-e{case[3]}-{str(case[4]).split('.')[-1]}")
|
||||
for case in [
|
||||
# Production path promotes q/k/v to float32 while cache is float32.
|
||||
(4, 128, 128, 128, torch.float32),
|
||||
]
|
||||
],
|
||||
)
|
||||
def test_ascend_lightning_attention_kernel_prefix_with_history(h, n, d, e, dtype):
|
||||
"""Test AscendLightningAttentionKernel.jit_linear_forward_prefix with
|
||||
non-zero initial kv_caches."""
|
||||
torch.manual_seed(42)
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
rtol, atol = _get_tolerances(dtype)
|
||||
|
||||
q = _randn((h, n, d), dtype, device)
|
||||
k = _randn((h, n, d), dtype, device)
|
||||
v = _randn((h, n, e), dtype, device)
|
||||
slope_rate = _rand_decay(h, device)
|
||||
kv_caches = _randn((h, d, e), torch.float32, device)
|
||||
|
||||
kv_caches_triton = kv_caches.clone()
|
||||
out_triton = AscendLightningAttentionKernel.jit_linear_forward_prefix(
|
||||
q.clone(), k.clone(), v.clone(), kv_caches_triton, slope_rate.clone(), block_size=256
|
||||
)
|
||||
|
||||
out_ref, kv_caches_ref = _naive_jit_linear_forward_prefix(q, k, v, kv_caches, slope_rate, block_size=256)
|
||||
|
||||
torch.testing.assert_close(
|
||||
out_triton.float().cpu(),
|
||||
out_ref.cpu(),
|
||||
rtol=rtol,
|
||||
atol=atol,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
kv_caches_triton.cpu(),
|
||||
kv_caches_ref.cpu(),
|
||||
rtol=rtol,
|
||||
atol=atol,
|
||||
)
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("h", "n", "d", "e"),
|
||||
[
|
||||
pytest.param(*case, id=f"h{case[0]}-n{case[1]}-d{case[2]}-e{case[3]}")
|
||||
for case in [
|
||||
# Multi-block sequence via prefix kernel
|
||||
(4, 512, 128, 128),
|
||||
]
|
||||
],
|
||||
)
|
||||
def test_ascend_lightning_attention_kernel_prefix_multi_block(h, n, d, e):
|
||||
"""Smoke test the prefix wrapper with multi-block sequences."""
|
||||
torch.manual_seed(42)
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
dtype = torch.bfloat16
|
||||
|
||||
q = _randn((h, n, d), dtype, device)
|
||||
k = _randn((h, n, d), dtype, device)
|
||||
v = _randn((h, n, e), dtype, device)
|
||||
|
||||
# Small decay rates for multi-block tolerance
|
||||
slope_rate = _rand_decay(h, device, scale=0.01)
|
||||
kv_caches = torch.zeros(h, d, e, dtype=torch.float32, device=device)
|
||||
|
||||
kv_caches_triton = kv_caches.clone()
|
||||
out_triton = AscendLightningAttentionKernel.jit_linear_forward_prefix(
|
||||
q.clone(), k.clone(), v.clone(), kv_caches_triton, slope_rate.clone(), block_size=256
|
||||
)
|
||||
|
||||
assert out_triton.shape == (n, h * e)
|
||||
assert not torch.isnan(out_triton).any(), "Output contains NaN values"
|
||||
assert not torch.isinf(out_triton).any(), "Output contains Inf values"
|
||||
assert kv_caches_triton.shape == (h, d, e)
|
||||
assert not torch.isnan(kv_caches_triton).any(), "KV cache contains NaN values"
|
||||
assert not torch.isinf(kv_caches_triton).any(), "KV cache contains Inf values"
|
||||
assert not torch.allclose(kv_caches_triton, kv_caches), "KV cache should be updated"
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Output shape and basic sanity tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_lightning_attention_npu_output_shapes():
|
||||
"""Verify output tensor shapes match expected shapes."""
|
||||
torch.manual_seed(42)
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
|
||||
b, h, n, d, e = 1, 4, 256, 128, 128
|
||||
q = _randn((b, h, n, d), torch.bfloat16, device)
|
||||
k = _randn((b, h, n, d), torch.bfloat16, device)
|
||||
v = _randn((b, h, n, e), torch.bfloat16, device)
|
||||
ed = _rand_decay(h, device)
|
||||
|
||||
o, kv = lightning_attention_npu(q, k, v, ed, block_size=256, kv_history=None)
|
||||
|
||||
assert o.shape == (b, h, n, e), f"Expected output shape {(b, h, n, e)}, got {o.shape}"
|
||||
assert kv.dim() == 5, f"Expected kv_return to be 5D, got {kv.dim()}D"
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
|
||||
def test_ascend_kernel_prefix_output_shape():
|
||||
"""Verify AscendLightningAttentionKernel.jit_linear_forward_prefix output shape."""
|
||||
torch.manual_seed(42)
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
|
||||
h, n, d, e = 4, 256, 128, 128
|
||||
q = _randn((h, n, d), torch.bfloat16, device)
|
||||
k = _randn((h, n, d), torch.bfloat16, device)
|
||||
v = _randn((h, n, e), torch.bfloat16, device)
|
||||
slope_rate = _rand_decay(h, device)
|
||||
kv_caches = torch.zeros(h, d, e, dtype=torch.float32, device=device)
|
||||
|
||||
out = AscendLightningAttentionKernel.jit_linear_forward_prefix(q, k, v, kv_caches, slope_rate, block_size=256)
|
||||
|
||||
expected_shape = (n, h * e)
|
||||
assert out.shape == expected_shape, f"Expected output shape {expected_shape}, got {out.shape}"
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
|
||||
def test_lightning_attention_output_no_nan():
|
||||
"""Verify the output does not contain NaN or Inf values."""
|
||||
torch.manual_seed(42)
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
|
||||
b, h, n, d, e = 1, 4, 256, 128, 128
|
||||
q = _randn((b, h, n, d), torch.bfloat16, device)
|
||||
k = _randn((b, h, n, d), torch.bfloat16, device)
|
||||
v = _randn((b, h, n, e), torch.bfloat16, device)
|
||||
ed = _rand_decay(h, device).view(1, h, 1, 1)
|
||||
|
||||
o, _ = lightning_attention_npu_(q, k, v, ed, torch.zeros(b, h, d, e, dtype=torch.float32, device=device))
|
||||
|
||||
assert not torch.isnan(o).any(), "Output contains NaN values"
|
||||
assert not torch.isinf(o).any(), "Output contains Inf values"
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
|
||||
def test_lightning_attention_causal_property():
|
||||
"""Verify causal property: output at position t should not depend on
|
||||
keys/values at positions > t.
|
||||
|
||||
If we modify V at positions after t, the output at position t should
|
||||
remain unchanged. This tests the diagonal kernel's causal mask.
|
||||
Uses float32 for numerical precision.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
|
||||
b, h, n, d, e = 1, 4, 128, 64, 64
|
||||
q = _randn((b, h, n, d), torch.float32, device, scale=SEMANTIC_INPUT_SCALE)
|
||||
k = _randn((b, h, n, d), torch.float32, device, scale=SEMANTIC_INPUT_SCALE)
|
||||
v = _randn((b, h, n, e), torch.float32, device, scale=SEMANTIC_INPUT_SCALE)
|
||||
ed = _rand_decay(h, device).view(1, h, 1, 1)
|
||||
kv_history = torch.zeros(b, h, d, e, dtype=torch.float32, device=device)
|
||||
|
||||
# Output with original V
|
||||
o_orig, _ = lightning_attention_npu_(q, k, v, ed, kv_history.clone())
|
||||
|
||||
# Scramble V at positions after t=50
|
||||
v_scrambled = v.clone()
|
||||
v_scrambled[:, :, 50:, :] = _randn(v_scrambled[:, :, 50:, :].shape, torch.float32, device)
|
||||
|
||||
# Output with scrambled V (should be identical at positions 0..49)
|
||||
o_scrambled, _ = lightning_attention_npu_(q, k, v_scrambled, ed, kv_history.clone())
|
||||
|
||||
# Output at positions 0..49 should be unchanged
|
||||
torch.testing.assert_close(
|
||||
o_orig[:, :, :50, :].cpu(),
|
||||
o_scrambled[:, :, :50, :].cpu(),
|
||||
rtol=1e-5,
|
||||
atol=1e-5,
|
||||
)
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
|
||||
|
||||
def test_lightning_attention_decay_effect():
|
||||
"""Verify that different decay rates produce different outputs.
|
||||
|
||||
With a larger decay rate, the output is more dominated by recent tokens.
|
||||
This test verifies the decay mechanism is active, using the naive
|
||||
reference as ground truth since both should produce valid outputs.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
|
||||
b, h, n, d, e = 1, 4, 128, 64, 64
|
||||
q = _randn((b, h, n, d), torch.float32, device, scale=SEMANTIC_INPUT_SCALE)
|
||||
k = _randn((b, h, n, d), torch.float32, device, scale=SEMANTIC_INPUT_SCALE)
|
||||
v = _randn((b, h, n, e), torch.float32, device, scale=SEMANTIC_INPUT_SCALE)
|
||||
kv_history = torch.zeros(b, h, d, e, dtype=torch.float32, device=device)
|
||||
|
||||
# Small decay rate (slow decay, long memory)
|
||||
ed_small = torch.full((1, h, 1, 1), 0.01, dtype=torch.float32, device=device)
|
||||
o_small, _ = lightning_attention_npu_(q, k, v, ed_small, kv_history.clone())
|
||||
|
||||
# Large decay rate (fast decay, short memory)
|
||||
ed_large = torch.full((1, h, 1, 1), 1.0, dtype=torch.float32, device=device)
|
||||
o_large, _ = lightning_attention_npu_(q, k, v, ed_large, kv_history.clone())
|
||||
|
||||
# Verify both produce valid (non-NaN, non-Inf) outputs
|
||||
assert not torch.isnan(o_small).any(), "Small decay output contains NaN"
|
||||
assert not torch.isinf(o_small).any(), "Small decay output contains Inf"
|
||||
assert not torch.isnan(o_large).any(), "Large decay output contains NaN"
|
||||
assert not torch.isinf(o_large).any(), "Large decay output contains Inf"
|
||||
|
||||
# The outputs should differ because different decay rates produce
|
||||
# different attention distributions.
|
||||
assert not torch.allclose(o_small, o_large, rtol=1e-3, atol=1e-3), (
|
||||
"Different decay rates should produce different outputs"
|
||||
)
|
||||
|
||||
# Verify both outputs match naive reference
|
||||
o_ref_small, _ = _naive_triton_lightning_attention(q, k, v, ed_small, kv_history)
|
||||
o_ref_large, _ = _naive_triton_lightning_attention(q, k, v, ed_large, kv_history)
|
||||
|
||||
torch.testing.assert_close(
|
||||
o_small.cpu(),
|
||||
o_ref_small.cpu(),
|
||||
rtol=1e-3,
|
||||
atol=1e-3,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
o_large.cpu(),
|
||||
o_ref_large.cpu(),
|
||||
rtol=1e-3,
|
||||
atol=1e-3,
|
||||
)
|
||||
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
@@ -0,0 +1,64 @@
|
||||
import pytest
|
||||
import torch
|
||||
from vllm.triton_utils import triton
|
||||
|
||||
from vllm_ascend.worker.v2.sample.logprob import _topk_log_softmax_kernel
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"batch_size,vocab_size,num_logprobs",
|
||||
[
|
||||
(48, 102400, 50),
|
||||
(96, 102400, 1),
|
||||
(24, 151936, 8),
|
||||
],
|
||||
)
|
||||
def test_topk_log_softmax_kernel(batch_size, vocab_size, num_logprobs):
|
||||
"""Test _topk_log_softmax_kernel for computing log probabilities
|
||||
Args:
|
||||
batch_size: Number of sequences in the batch
|
||||
vocab_size: Size of the vocabulary
|
||||
num_logprobs: Number of tokens to compute log probabilities for
|
||||
"""
|
||||
# ========== Setup test data ==========
|
||||
torch.manual_seed(42)
|
||||
|
||||
# Generate random logits
|
||||
logits = torch.randn(batch_size, vocab_size, device="npu", dtype=torch.float32)
|
||||
|
||||
# Generate token_ids for which to compute logprobs
|
||||
token_ids = torch.randint(0, vocab_size, (batch_size, num_logprobs), device="npu", dtype=torch.int64)
|
||||
|
||||
# ========== Execute test ==========
|
||||
# Prepare output tensor
|
||||
triton_output = torch.empty(batch_size, num_logprobs, dtype=torch.float32, device="npu")
|
||||
|
||||
# Invoke Triton kernel
|
||||
_topk_log_softmax_kernel[(batch_size,)](
|
||||
triton_output,
|
||||
logits,
|
||||
logits.stride(0),
|
||||
token_ids,
|
||||
num_logprobs,
|
||||
vocab_size,
|
||||
BLOCK_SIZE=1024,
|
||||
PADDED_TOPK=max(triton.next_power_of_2(num_logprobs), 2),
|
||||
)
|
||||
torch.npu.synchronize()
|
||||
|
||||
# Compute reference values using PyTorch
|
||||
torch_logprobs = torch.log_softmax(logits, dim=-1)
|
||||
|
||||
# Extract logprobs for each batch and token_id
|
||||
ref_output = torch.zeros_like(triton_output)
|
||||
for i in range(batch_size):
|
||||
for j in range(num_logprobs):
|
||||
token_id = token_ids[i, j]
|
||||
ref_output[i, j] = torch_logprobs[i, token_id]
|
||||
|
||||
# ========== Verify results ==========
|
||||
assert torch.allclose(triton_output, ref_output, rtol=1e-3, atol=1e-3), (
|
||||
f"Triton output differs from PyTorch reference.\n"
|
||||
f"Max diff: {torch.max(torch.abs(triton_output - ref_output))}\n"
|
||||
f"Mean diff: {torch.mean(torch.abs(triton_output - ref_output))}"
|
||||
)
|
||||
@@ -0,0 +1,91 @@
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
|
||||
from vllm_ascend.worker.v2.sample.min_p import apply_min_p
|
||||
|
||||
|
||||
def torch_min_p_torch(
|
||||
logits: torch.Tensor,
|
||||
expanded_idx_mapping: torch.Tensor,
|
||||
min_p: torch.Tensor,
|
||||
):
|
||||
num_tokens, _ = logits.shape
|
||||
out = logits.clone()
|
||||
|
||||
for token_idx in range(num_tokens):
|
||||
req_state_idx = expanded_idx_mapping[token_idx].item()
|
||||
|
||||
min_p_val = min_p[req_state_idx].item()
|
||||
if min_p_val == 0.0:
|
||||
continue
|
||||
|
||||
token_logits = out[token_idx]
|
||||
max_val = token_logits.max()
|
||||
threshold = max_val + torch.log(torch.tensor(min_p_val, device=logits.device))
|
||||
token_logits = torch.where(token_logits < threshold, -torch.inf, token_logits)
|
||||
out[token_idx] = token_logits
|
||||
|
||||
return out
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"num_reqs,vocab_size",
|
||||
[
|
||||
(48, 102400),
|
||||
(96, 102400),
|
||||
(24, 151936),
|
||||
(1, 32000),
|
||||
],
|
||||
)
|
||||
def test_apply_min_p_kernel(num_reqs, vocab_size):
|
||||
"""Test apply_min_p for computing Min-P sampling mask
|
||||
Args:
|
||||
num_reqs: Number of sequences in the batch
|
||||
vocab_size: Size of the vocabulary
|
||||
"""
|
||||
|
||||
init_device_properties_triton()
|
||||
# ========== Setup test data ==========
|
||||
torch.manual_seed(42)
|
||||
|
||||
# Generate random logits (using float32 as specified in your kernel)
|
||||
device = "npu"
|
||||
|
||||
original_logits = torch.randn(num_reqs, vocab_size, device=device, dtype=torch.float32)
|
||||
|
||||
triton_logits = original_logits.clone()
|
||||
ref_logits = original_logits.clone()
|
||||
|
||||
expanded_idx_mapping = torch.arange(num_reqs - 1, -1, -1, device=device, dtype=torch.int32)
|
||||
|
||||
# Generate random min_p values (valid range is typically (0, 1.0])
|
||||
min_p = torch.empty(num_reqs, device=device, dtype=torch.float32).uniform_(0.01, 0.5)
|
||||
|
||||
# ========== Execute test ==========
|
||||
# 1. Invoke your Triton kernel wrapper
|
||||
apply_min_p(triton_logits, expanded_idx_mapping, min_p)
|
||||
torch.npu.synchronize()
|
||||
|
||||
# 2. Compute reference values using PyTorch
|
||||
ref_logits = torch_min_p_torch(
|
||||
ref_logits,
|
||||
expanded_idx_mapping,
|
||||
min_p,
|
||||
)
|
||||
|
||||
# ========== Verify results ==========
|
||||
triton_inf_mask = torch.isinf(triton_logits)
|
||||
ref_inf_mask = torch.isinf(ref_logits)
|
||||
|
||||
assert torch.equal(triton_inf_mask, ref_inf_mask), (
|
||||
"Masked positions (where logits == -inf) do not match between Triton and PyTorch."
|
||||
)
|
||||
|
||||
valid_triton_logits = triton_logits[~triton_inf_mask]
|
||||
valid_ref_logits = ref_logits[~ref_inf_mask]
|
||||
|
||||
assert torch.allclose(valid_triton_logits, valid_ref_logits, rtol=1e-4, atol=1e-4), (
|
||||
f"Logits values differ between Triton and PyTorch reference.\n"
|
||||
f"Max diff: {torch.max(torch.abs(valid_triton_logits - valid_ref_logits))}"
|
||||
)
|
||||
@@ -0,0 +1,166 @@
|
||||
import gc
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from vllm.model_executor.layers.rotary_embedding.mrope import triton_mrope
|
||||
|
||||
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
|
||||
|
||||
MROPE_SECTION = [[32, 32, 32]]
|
||||
DTYPES = [torch.bfloat16, torch.float16]
|
||||
HEAD_SIZES = [128]
|
||||
ROTARY_DIMS = [128]
|
||||
NUM_Q_HEADS = [64]
|
||||
NUM_K_HEADS = [1]
|
||||
NUM_TOKENS = [1, 4, 8, 16]
|
||||
SEEDS = [0]
|
||||
DEVICES = [f"npu:{0}"]
|
||||
DEFAULT_ATOL = 1e-3
|
||||
DEFAULT_RTOL = 1e-3
|
||||
|
||||
|
||||
def pytorch_forward_native(q, k, cos, sin, mrope_section, head_size, rotary_dim, mrope_interleaved):
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
|
||||
num_tokens = q.shape[0]
|
||||
n_q_head = q.shape[1] // head_size
|
||||
n_kv_head = k.shape[1] // head_size
|
||||
|
||||
q_reshaped = q.view(num_tokens, n_q_head, head_size)
|
||||
k_reshaped = k.view(num_tokens, n_kv_head, head_size)
|
||||
|
||||
cos_reshaped = cos.permute(1, 2, 0)
|
||||
sin_reshaped = sin.permute(1, 2, 0)
|
||||
|
||||
half_rd = rotary_dim // 2
|
||||
|
||||
for token_idx in range(num_tokens):
|
||||
token_cos = cos_reshaped[token_idx]
|
||||
token_sin = sin_reshaped[token_idx]
|
||||
|
||||
cos_row = torch.zeros(head_size // 2, device=q.device, dtype=q.dtype)
|
||||
sin_row = torch.zeros(head_size // 2, device=q.device, dtype=q.dtype)
|
||||
|
||||
if mrope_interleaved:
|
||||
cos_offsets = torch.arange(0, head_size // 2, device=q.device)
|
||||
h_mask = ((cos_offsets % 3) == 1) & (cos_offsets <= 3 * mrope_section[1])
|
||||
w_mask = ((cos_offsets % 3) == 2) & (cos_offsets <= 3 * mrope_section[2])
|
||||
t_mask = ~(h_mask | w_mask)
|
||||
|
||||
cos_row[t_mask] = token_cos[t_mask, 0]
|
||||
cos_row[h_mask] = token_cos[h_mask, 1]
|
||||
cos_row[w_mask] = token_cos[w_mask, 2]
|
||||
|
||||
sin_row[t_mask] = token_sin[t_mask, 0]
|
||||
sin_row[h_mask] = token_sin[h_mask, 1]
|
||||
sin_row[w_mask] = token_sin[w_mask, 2]
|
||||
else:
|
||||
t_end = mrope_section[0]
|
||||
h_end = t_end + mrope_section[1]
|
||||
|
||||
if t_end > 0:
|
||||
cos_row[:t_end] = token_cos[:t_end, 0]
|
||||
sin_row[:t_end] = token_sin[:t_end, 0]
|
||||
|
||||
if mrope_section[1] > 0:
|
||||
cos_row[t_end:h_end] = token_cos[t_end:h_end, 1]
|
||||
sin_row[t_end:h_end] = token_sin[t_end:h_end, 1]
|
||||
|
||||
if mrope_section[2] > 0:
|
||||
w_start = h_end
|
||||
cos_row[w_start:half_rd] = token_cos[w_start:half_rd, 2]
|
||||
sin_row[w_start:half_rd] = token_sin[w_start:half_rd, 2]
|
||||
|
||||
q_token = q_reshaped[token_idx]
|
||||
k_token = k_reshaped[token_idx]
|
||||
|
||||
q1 = q_token[:, :half_rd]
|
||||
q2 = q_token[:, half_rd:]
|
||||
k1 = k_token[:, :half_rd]
|
||||
k2 = k_token[:, half_rd:]
|
||||
|
||||
cos_half = cos_row.unsqueeze(0)
|
||||
sin_half = sin_row.unsqueeze(0)
|
||||
|
||||
new_q1 = q1 * cos_half - q2 * sin_half
|
||||
new_q2 = q2 * cos_half + q1 * sin_half
|
||||
|
||||
new_k1 = k1 * cos_half - k2 * sin_half
|
||||
new_k2 = k2 * cos_half + k1 * sin_half
|
||||
|
||||
q_reshaped[token_idx] = torch.cat([new_q1, new_q2], dim=1)
|
||||
k_reshaped[token_idx] = torch.cat([new_k1, new_k2], dim=1)
|
||||
|
||||
q_result = q_reshaped.view(num_tokens, -1)
|
||||
k_result = k_reshaped.view(num_tokens, -1)
|
||||
|
||||
return q_result, k_result
|
||||
|
||||
|
||||
def create_test_data(num_tokens, n_q_head, n_kv_head, rotary_dim, head_size, device, dtype):
|
||||
q = torch.randn(num_tokens, n_q_head * head_size, dtype=dtype, device=device)
|
||||
k = torch.randn(num_tokens, n_kv_head * head_size, dtype=dtype, device=device)
|
||||
|
||||
sin = torch.randn(3, num_tokens, rotary_dim // 2, dtype=dtype, device=device)
|
||||
cos = torch.randn(3, num_tokens, rotary_dim // 2, dtype=dtype, device=device)
|
||||
|
||||
norm = torch.sqrt(cos**2 + sin**2)
|
||||
cos = cos / norm
|
||||
sin = sin / norm
|
||||
|
||||
return q, k, cos, sin
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mrope_section", MROPE_SECTION)
|
||||
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
|
||||
@pytest.mark.parametrize("num_q_heads", NUM_Q_HEADS)
|
||||
@pytest.mark.parametrize("num_k_heads", NUM_K_HEADS)
|
||||
@pytest.mark.parametrize("head_size", HEAD_SIZES)
|
||||
@pytest.mark.parametrize("rotary_dim", ROTARY_DIMS)
|
||||
@pytest.mark.parametrize("dtype", DTYPES)
|
||||
@pytest.mark.parametrize("seed", SEEDS)
|
||||
@pytest.mark.parametrize("device", DEVICES)
|
||||
@torch.inference_mode()
|
||||
def test_mrotary_embedding_triton_kernel(
|
||||
mrope_section: list[int],
|
||||
num_tokens: int,
|
||||
num_q_heads: int,
|
||||
num_k_heads: int,
|
||||
head_size: int,
|
||||
rotary_dim: int,
|
||||
dtype: torch.dtype,
|
||||
seed: int,
|
||||
device: str,
|
||||
) -> None:
|
||||
torch.manual_seed(seed)
|
||||
torch.set_default_device(device)
|
||||
init_device_properties_triton()
|
||||
if rotary_dim == -1:
|
||||
rotary_dim = head_size
|
||||
|
||||
q_trt, k_trt, cos, sin = create_test_data(
|
||||
num_tokens=num_tokens,
|
||||
n_q_head=num_q_heads,
|
||||
n_kv_head=num_k_heads,
|
||||
head_size=head_size,
|
||||
rotary_dim=rotary_dim,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
q_gold, k_gold = q_trt.clone(), k_trt.clone()
|
||||
|
||||
q_trt, k_trt = triton_mrope(q_trt, k_trt, cos, sin, mrope_section, head_size, rotary_dim, True)
|
||||
|
||||
q_gold, k_gold = pytorch_forward_native(q_gold, k_gold, cos, sin, mrope_section, head_size, rotary_dim, True)
|
||||
atol = DEFAULT_ATOL
|
||||
rtol = DEFAULT_RTOL
|
||||
if dtype == torch.bfloat16:
|
||||
atol = 1e-02
|
||||
rtol = 1e-02
|
||||
# Compare the results.
|
||||
torch.testing.assert_close(q_trt.view(q_gold.size()), q_gold, atol=atol, rtol=rtol)
|
||||
torch.testing.assert_close(k_trt.view(k_gold.size()), k_gold, atol=atol, rtol=rtol)
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
@@ -0,0 +1,33 @@
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm_ascend.ops.triton.muls_add import muls_add_triton
|
||||
from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("shape", "dtype", "scale"),
|
||||
[
|
||||
((1, 2048), torch.float16, 1.25),
|
||||
((4000, 2048), torch.float16, 0.75),
|
||||
((4, 2048), torch.bfloat16, 1.0),
|
||||
],
|
||||
)
|
||||
@torch.inference_mode()
|
||||
def test_muls_add_triton_correctness(shape, dtype, scale):
|
||||
"""compare the correctness of muls_add_triton with the PyTorch baseline implementation."""
|
||||
init_device_properties_triton()
|
||||
device = "npu"
|
||||
|
||||
torch.manual_seed(0)
|
||||
x = torch.randn(*shape, dtype=dtype, device=device)
|
||||
y = torch.randn(*shape, dtype=dtype, device=device)
|
||||
|
||||
out_triton = muls_add_triton(x, y, scale)
|
||||
out_ref = x * scale + y
|
||||
|
||||
rtol, atol = 1e-3, 1e-3
|
||||
|
||||
assert out_triton.shape == out_ref.shape
|
||||
assert out_triton.dtype == out_ref.dtype
|
||||
assert torch.allclose(out_triton, out_ref, rtol=rtol, atol=atol)
|
||||
@@ -0,0 +1,246 @@
|
||||
import gc
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from vllm_ascend.worker.v2.sample.penalties import apply_penalties
|
||||
|
||||
NUM_TOKENS = [1, 4]
|
||||
VOCAB_SIZE = [1000]
|
||||
NUM_STATUS = [1, 4]
|
||||
NUM_SPECULATIVE_TOKENS = [0, 1, 3]
|
||||
DTYPES = [torch.bfloat16, torch.float16]
|
||||
SEEDS = [42]
|
||||
DEVICES = [f"npu:{0}"]
|
||||
|
||||
DEFAULT_ATOL = 1e-3
|
||||
DEFAULT_RTOL = 1e-3
|
||||
|
||||
|
||||
def pytorch_apply_penalties(
|
||||
logits: torch.Tensor,
|
||||
idx_mapping: torch.Tensor,
|
||||
token_ids: torch.Tensor,
|
||||
expanded_local_pos: torch.Tensor,
|
||||
repetition_penalty: torch.Tensor,
|
||||
frequency_penalty: torch.Tensor,
|
||||
presence_penalty: torch.Tensor,
|
||||
prompt_bin_mask: torch.Tensor,
|
||||
output_bin_counts: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Pytorch equivalent implementation
|
||||
"""
|
||||
num_tokens, vocab_size = logits.shape
|
||||
device = logits.device
|
||||
dtype = logits.dtype
|
||||
|
||||
logits_float = logits.float()
|
||||
|
||||
num_status = prompt_bin_mask.shape[0]
|
||||
num_packed = prompt_bin_mask.shape[1]
|
||||
|
||||
prompt_masks_unpacked = torch.zeros(num_status, vocab_size, dtype=torch.bool, device=device)
|
||||
|
||||
for state_idx in range(num_status):
|
||||
for packed_idx in range(num_packed):
|
||||
packed_val = prompt_bin_mask[state_idx, packed_idx].item()
|
||||
if packed_val == 0:
|
||||
continue
|
||||
start_idx = packed_idx * 32
|
||||
end_idx = min(start_idx + 32, vocab_size)
|
||||
|
||||
for bit_pos in range(end_idx - start_idx):
|
||||
if (packed_val >> bit_pos) & 1:
|
||||
prompt_masks_unpacked[state_idx, start_idx + bit_pos] = True
|
||||
|
||||
start_idx_in_batch = torch.arange(num_tokens, device=device) - expanded_local_pos
|
||||
|
||||
for token_idx in range(num_tokens):
|
||||
req_state_idx = idx_mapping[token_idx].item()
|
||||
|
||||
rep_penalty = repetition_penalty[req_state_idx].item()
|
||||
freq_penalty = frequency_penalty[req_state_idx].item()
|
||||
pres_penalty = presence_penalty[req_state_idx].item()
|
||||
|
||||
use_rep_penalty = rep_penalty != 1.0
|
||||
use_freq_penalty = freq_penalty != 0.0
|
||||
use_pres_penalty = pres_penalty != 0.0
|
||||
use_penalty = use_rep_penalty or use_freq_penalty or use_pres_penalty
|
||||
|
||||
if not use_penalty:
|
||||
continue
|
||||
|
||||
current_prompt_mask = prompt_masks_unpacked[req_state_idx]
|
||||
base_counts = output_bin_counts[req_state_idx].clone()
|
||||
|
||||
# Compute cumulative draft counts
|
||||
pos = expanded_local_pos[token_idx].item()
|
||||
draft_counts = torch.zeros(vocab_size, device=device, dtype=torch.int32)
|
||||
|
||||
for prev_pos in range(pos):
|
||||
prev_token_idx = start_idx_in_batch[token_idx] + prev_pos + 1
|
||||
if 0 <= prev_token_idx < num_tokens:
|
||||
prev_token = token_ids[prev_token_idx].item()
|
||||
if 0 <= prev_token < vocab_size:
|
||||
draft_counts[prev_token] += 1
|
||||
|
||||
# Total counts = base output counts + cumulative draft counts
|
||||
total_counts = base_counts + draft_counts
|
||||
output_bin_mask = total_counts > 0
|
||||
|
||||
if use_rep_penalty:
|
||||
need_scale = current_prompt_mask | output_bin_mask
|
||||
scale = torch.where(need_scale, rep_penalty, 1.0)
|
||||
|
||||
pos_mask = logits_float[token_idx] > 0
|
||||
scale_factor = torch.where(pos_mask, 1.0 / scale, scale)
|
||||
logits_float[token_idx] *= scale_factor
|
||||
|
||||
if use_freq_penalty:
|
||||
logits_float[token_idx] -= freq_penalty * total_counts.float()
|
||||
|
||||
if use_pres_penalty:
|
||||
logits_float[token_idx] -= pres_penalty * output_bin_mask.float()
|
||||
|
||||
return logits_float.to(dtype)
|
||||
|
||||
|
||||
def create_test_data(
|
||||
num_tokens: int = 8,
|
||||
vocab_size: int = 51200,
|
||||
num_status: int = 16,
|
||||
num_speculative_tokens: int = 3,
|
||||
device: str = "npu",
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
seed: int = 42,
|
||||
):
|
||||
"""Create test data for penalties"""
|
||||
torch.manual_seed(seed)
|
||||
|
||||
logits = torch.randn(num_tokens, vocab_size, device=device, dtype=dtype)
|
||||
|
||||
repetition_penalty = torch.ones(num_status, device=device, dtype=torch.float32)
|
||||
for i in range(num_status):
|
||||
if torch.rand(1) > 0.3:
|
||||
repetition_penalty[i] = torch.rand(1, device=device).item() * 0.8 + 0.6
|
||||
|
||||
frequency_penalty = torch.zeros(num_status, device=device, dtype=torch.float32)
|
||||
for i in range(num_status):
|
||||
if torch.rand(1) > 0.5:
|
||||
frequency_penalty[i] = torch.rand(1, device=device).item() * 0.2
|
||||
|
||||
presence_penalty = torch.zeros(num_status, device=device, dtype=torch.float32)
|
||||
for i in range(num_status):
|
||||
if torch.rand(1) > 0.5:
|
||||
presence_penalty[i] = torch.rand(1, device=device).item() * 0.2
|
||||
|
||||
idx_mapping = torch.randint(0, num_status, (num_tokens,), device=device, dtype=torch.int32)
|
||||
|
||||
# Create token_ids for speculative decoding
|
||||
token_ids = torch.randint(0, vocab_size, (num_tokens,), device=device, dtype=torch.int32)
|
||||
|
||||
# Create expanded_local_pos (position within speculative decoding window)
|
||||
expanded_local_pos = torch.zeros(num_tokens, device=device, dtype=torch.int32)
|
||||
for i in range(num_tokens):
|
||||
expanded_local_pos[i] = torch.randint(0, num_speculative_tokens + 1, (1,)).item()
|
||||
|
||||
num_packed = (vocab_size + 31) // 32
|
||||
prompt_bin_mask = torch.zeros(num_status, num_packed, device=device, dtype=torch.int32)
|
||||
|
||||
for state_idx in range(num_status):
|
||||
num_tokens_in_prompt = max(1, vocab_size // 20)
|
||||
prompt_tokens = torch.randperm(vocab_size, device=device)[:num_tokens_in_prompt]
|
||||
|
||||
for token_id in prompt_tokens:
|
||||
packed_idx = token_id // 32
|
||||
bit_pos = token_id % 32
|
||||
prompt_bin_mask[state_idx, packed_idx] |= 1 << bit_pos
|
||||
|
||||
output_bin_counts = torch.zeros(num_status, vocab_size, device=device, dtype=torch.int32)
|
||||
for state_idx in range(num_status):
|
||||
num_output_tokens = max(1, vocab_size // 20)
|
||||
output_tokens = torch.randint(0, vocab_size, (num_output_tokens,), device=device)
|
||||
counts = torch.randint(1, 10, (num_output_tokens,), device=device)
|
||||
|
||||
for token, count in zip(output_tokens, counts):
|
||||
output_bin_counts[state_idx, token] = count
|
||||
|
||||
return (
|
||||
logits,
|
||||
idx_mapping,
|
||||
token_ids,
|
||||
expanded_local_pos,
|
||||
repetition_penalty,
|
||||
frequency_penalty,
|
||||
presence_penalty,
|
||||
prompt_bin_mask,
|
||||
output_bin_counts,
|
||||
)
|
||||
|
||||
|
||||
class TestApplyPenalties:
|
||||
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
|
||||
@pytest.mark.parametrize("vocab_size", VOCAB_SIZE)
|
||||
@pytest.mark.parametrize("num_status", NUM_STATUS)
|
||||
@pytest.mark.parametrize("num_speculative_tokens", NUM_SPECULATIVE_TOKENS)
|
||||
@pytest.mark.parametrize("dtype", DTYPES)
|
||||
@pytest.mark.parametrize("seed", SEEDS)
|
||||
@pytest.mark.parametrize("device", DEVICES)
|
||||
@torch.inference_mode()
|
||||
def test_apply_penalties(self, num_tokens, vocab_size, num_status, num_speculative_tokens, dtype, seed, device):
|
||||
(
|
||||
logits_triton,
|
||||
idx_mapping,
|
||||
token_ids,
|
||||
expanded_local_pos,
|
||||
repetition_penalty,
|
||||
frequency_penalty,
|
||||
presence_penalty,
|
||||
prompt_bin_mask,
|
||||
output_bin_counts,
|
||||
) = create_test_data(
|
||||
num_tokens=num_tokens,
|
||||
vocab_size=vocab_size,
|
||||
num_status=num_status,
|
||||
num_speculative_tokens=num_speculative_tokens,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
seed=seed,
|
||||
)
|
||||
|
||||
logits_pytorch = logits_triton.clone()
|
||||
|
||||
apply_penalties(
|
||||
logits_triton,
|
||||
idx_mapping,
|
||||
token_ids,
|
||||
expanded_local_pos,
|
||||
repetition_penalty,
|
||||
frequency_penalty,
|
||||
presence_penalty,
|
||||
prompt_bin_mask,
|
||||
output_bin_counts,
|
||||
)
|
||||
|
||||
logits_pytorch_result = pytorch_apply_penalties(
|
||||
logits_pytorch,
|
||||
idx_mapping,
|
||||
token_ids,
|
||||
expanded_local_pos,
|
||||
repetition_penalty,
|
||||
frequency_penalty,
|
||||
presence_penalty,
|
||||
prompt_bin_mask,
|
||||
output_bin_counts,
|
||||
)
|
||||
|
||||
atol = DEFAULT_ATOL
|
||||
rtol = DEFAULT_RTOL
|
||||
if dtype == torch.bfloat16:
|
||||
atol = 1e-02
|
||||
rtol = 1e-02
|
||||
assert torch.allclose(logits_triton, logits_pytorch_result, atol=atol, rtol=rtol)
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
@@ -0,0 +1,105 @@
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from vllm.v1.worker.gpu.input_batch import post_update as post_update_gpu
|
||||
|
||||
from vllm_ascend.worker.v2.input_batch import post_update as post_update_npu
|
||||
|
||||
|
||||
def generate_test_data(
|
||||
num_reqs: int, max_num_reqs: int, vocab_size: int, num_speculative_steps: int, device: str
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Generate random test data.
|
||||
Return a dictionary containing all input tensors and the additional field 'expected_query_lens' for validation.
|
||||
"""
|
||||
|
||||
if num_reqs > max_num_reqs:
|
||||
raise ValueError("num_reqs cannot be larger than max_num_reqs")
|
||||
|
||||
idx_mapping = torch.arange(num_reqs, dtype=torch.int32, device=device)
|
||||
num_computed_tokens = torch.randint(0, 100, (max_num_reqs,), dtype=torch.int32, device=device)
|
||||
last_sampled_tokens = torch.randint(0, vocab_size, (max_num_reqs,), dtype=torch.int32, device=device)
|
||||
output_bin_counts = torch.randint(0, 10, (max_num_reqs, vocab_size), dtype=torch.int32, device=device)
|
||||
sampled_tokens = torch.randint(
|
||||
0, vocab_size, (num_reqs, num_speculative_steps + 1), dtype=torch.int32, device=device
|
||||
)
|
||||
num_sampled = torch.randint(1, num_speculative_steps + 2, (num_reqs,), dtype=torch.int32, device=device)
|
||||
num_rejected = torch.randint(0, num_speculative_steps + 1, (num_reqs,), dtype=torch.int32, device=device)
|
||||
num_rejected = torch.min(num_rejected, num_sampled - 1)
|
||||
|
||||
query_lengths = torch.randint(1, 20, (num_reqs,), dtype=torch.int32, device=device)
|
||||
query_start_loc = torch.cat(
|
||||
[torch.tensor([0], dtype=torch.int32, device=device), torch.cumsum(query_lengths, dim=0)]
|
||||
)
|
||||
total_len = torch.randint(50, 200, (max_num_reqs,), dtype=torch.int32, device=device)
|
||||
|
||||
max_model_len = 3000 # 或者可以从total_len的最大值获取
|
||||
all_token_ids = torch.randint(0, vocab_size, (max_num_reqs, max_model_len), dtype=torch.int32, device=device)
|
||||
|
||||
return {
|
||||
"idx_mapping": idx_mapping,
|
||||
"num_computed_tokens": num_computed_tokens,
|
||||
"last_sampled_tokens": last_sampled_tokens,
|
||||
"output_bin_counts": output_bin_counts,
|
||||
"sampled_tokens": sampled_tokens,
|
||||
"num_sampled": num_sampled,
|
||||
"num_rejected": num_rejected,
|
||||
"query_start_loc": query_start_loc,
|
||||
"all_token_ids": all_token_ids,
|
||||
"total_len": total_len,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"num_reqs,max_num_reqs,vocab_size,num_speculative_steps",
|
||||
[
|
||||
(36, 36, 200, 2),
|
||||
(48, 48, 32000, 5),
|
||||
(128, 128, 32000, 5),
|
||||
],
|
||||
)
|
||||
def test_post_update(num_reqs: int, max_num_reqs: int, vocab_size: int, num_speculative_steps: int):
|
||||
"""Test _topk_log_softmax_kernel for computing log probabilities
|
||||
Args:
|
||||
batch_size: Number of sequences in the batch
|
||||
vocab_size: Size of the vocabulary
|
||||
num_logprobs: Number of tokens to compute log probabilities for
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
|
||||
post_update_params = [
|
||||
"idx_mapping",
|
||||
"num_computed_tokens",
|
||||
"last_sampled_tokens",
|
||||
"output_bin_counts",
|
||||
"sampled_tokens",
|
||||
"num_sampled",
|
||||
"num_rejected",
|
||||
"query_start_loc",
|
||||
"all_token_ids",
|
||||
"total_len",
|
||||
]
|
||||
|
||||
data = generate_test_data(num_reqs, max_num_reqs, vocab_size, num_speculative_steps, device="npu")
|
||||
kernel_inputs_gpu = {k: data[k].clone() for k in post_update_params}
|
||||
kernel_inputs_npu = {k: data[k].clone() for k in post_update_params}
|
||||
|
||||
# Invoke Triton kernel
|
||||
post_update_gpu(**kernel_inputs_gpu)
|
||||
torch.npu.synchronize()
|
||||
|
||||
post_update_npu(**kernel_inputs_npu)
|
||||
torch.npu.synchronize()
|
||||
|
||||
# ========== Verify results ==========
|
||||
assert torch.allclose(
|
||||
kernel_inputs_gpu["output_bin_counts"], kernel_inputs_npu["output_bin_counts"], rtol=1e-3, atol=1e-3
|
||||
), (
|
||||
f"Triton output differs from PyTorch reference.\n"
|
||||
f"Max diff: "
|
||||
f"{torch.max(torch.abs(kernel_inputs_gpu['output_bin_counts'] - kernel_inputs_npu['output_bin_counts']))}\n"
|
||||
f"Mean diff: "
|
||||
f"{torch.mean(torch.abs(kernel_inputs_gpu['output_bin_counts'] - kernel_inputs_npu['output_bin_counts']))}"
|
||||
)
|
||||
@@ -0,0 +1,318 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from vllm.model_executor.layers.mamba.mamba_utils import (
|
||||
get_conv_copy_spec,
|
||||
get_temporal_copy_spec,
|
||||
)
|
||||
from vllm.v1.core.sched.output import CachedRequestData, SchedulerOutput
|
||||
from vllm.v1.kv_cache_interface import KVCacheConfig, KVCacheGroupSpec, MambaSpec
|
||||
from vllm.v1.worker.mamba_utils import (
|
||||
MambaCopyBuffers,
|
||||
MambaSpecDecodeGPUContext,
|
||||
collect_mamba_copy_meta,
|
||||
do_mamba_copy_block,
|
||||
)
|
||||
|
||||
import vllm_ascend.patch.worker.patch_mamba_utils # noqa: F401
|
||||
|
||||
MambaStateCopyFunc = Callable[..., Any]
|
||||
_COPY_FUNCS: tuple[MambaStateCopyFunc, ...] = (
|
||||
get_conv_copy_spec,
|
||||
get_temporal_copy_spec,
|
||||
)
|
||||
|
||||
|
||||
def postprocess_mamba(
|
||||
scheduler_output: SchedulerOutput,
|
||||
kv_cache_config: KVCacheConfig,
|
||||
input_batch: Any,
|
||||
requests: dict[str, Any],
|
||||
forward_context: dict[str, Any],
|
||||
mamba_state_copy_funcs: tuple[MambaStateCopyFunc, ...],
|
||||
copy_bufs: MambaCopyBuffers,
|
||||
):
|
||||
assert input_batch.mamba_state_idx_cpu is not None
|
||||
num_scheduled_tokens_dict = scheduler_output.num_scheduled_tokens
|
||||
scheduled_spec_decode_tokens_dict = scheduler_output.scheduled_spec_decode_tokens
|
||||
num_accepted_tokens_cpu = input_batch.num_accepted_tokens_cpu
|
||||
mamba_state_idx_cpu = input_batch.mamba_state_idx_cpu
|
||||
mamba_group_ids = copy_bufs.mamba_group_ids
|
||||
mamba_spec = copy_bufs.mamba_spec
|
||||
copy_bufs.offset = 0
|
||||
for i, req_id in enumerate(input_batch.req_ids):
|
||||
req_state = requests[req_id]
|
||||
num_computed_tokens = req_state.num_computed_tokens
|
||||
num_draft_tokens = len(scheduled_spec_decode_tokens_dict.get(req_id, []))
|
||||
num_scheduled_tokens = num_scheduled_tokens_dict[req_id]
|
||||
num_accepted_tokens = num_accepted_tokens_cpu[i]
|
||||
num_tokens_running_state = num_computed_tokens + num_scheduled_tokens - num_draft_tokens
|
||||
new_num_computed_tokens = num_tokens_running_state + num_accepted_tokens - 1
|
||||
aligned_new_computed_tokens = new_num_computed_tokens // mamba_spec.block_size * mamba_spec.block_size
|
||||
if aligned_new_computed_tokens >= num_tokens_running_state:
|
||||
accept_token_bias = aligned_new_computed_tokens - num_tokens_running_state
|
||||
src_block_idx = mamba_state_idx_cpu[i]
|
||||
dest_block_idx = aligned_new_computed_tokens // mamba_spec.block_size - 1
|
||||
collect_mamba_copy_meta(
|
||||
copy_bufs,
|
||||
kv_cache_config,
|
||||
mamba_state_copy_funcs,
|
||||
mamba_group_ids,
|
||||
src_block_idx,
|
||||
dest_block_idx,
|
||||
accept_token_bias,
|
||||
req_state,
|
||||
forward_context,
|
||||
)
|
||||
if src_block_idx == dest_block_idx:
|
||||
num_accepted_tokens_cpu[i] = 1
|
||||
do_mamba_copy_block(copy_bufs)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _TestConfig:
|
||||
block_size: int = 16
|
||||
num_blocks: int = 32
|
||||
num_layers: int = 2
|
||||
max_num_reqs: int = 8
|
||||
conv_width: int = 4
|
||||
conv_inner_dim: int = 64
|
||||
temporal_state_dim: int = 128
|
||||
dtype: torch.dtype = torch.float16
|
||||
|
||||
|
||||
class _MockCpuGpuBuffer:
|
||||
def __init__(self, size: int, dtype: torch.dtype, device: torch.device):
|
||||
self.cpu = torch.zeros(size, dtype=dtype, device="cpu")
|
||||
self.gpu = torch.zeros(size, dtype=dtype, device=device)
|
||||
self.np = self.cpu.numpy()
|
||||
|
||||
def copy_to_gpu(self, n: int | None = None) -> torch.Tensor:
|
||||
if n is None:
|
||||
return self.gpu.copy_(self.cpu, non_blocking=True)
|
||||
return self.gpu[:n].copy_(self.cpu[:n], non_blocking=True)
|
||||
|
||||
|
||||
def _make_scheduler_output(
|
||||
num_scheduled_tokens: dict[str, int],
|
||||
scheduled_spec_decode_tokens: dict[str, list] | None = None,
|
||||
) -> SchedulerOutput:
|
||||
cached = CachedRequestData.make_empty()
|
||||
return SchedulerOutput(
|
||||
scheduled_new_reqs=[],
|
||||
scheduled_cached_reqs=cached,
|
||||
num_scheduled_tokens=num_scheduled_tokens,
|
||||
total_num_scheduled_tokens=sum(num_scheduled_tokens.values()),
|
||||
scheduled_spec_decode_tokens=scheduled_spec_decode_tokens or {},
|
||||
scheduled_encoder_inputs={},
|
||||
num_common_prefix_blocks=[],
|
||||
finished_req_ids=set(),
|
||||
free_encoder_mm_hashes=[],
|
||||
preempted_req_ids=set(),
|
||||
)
|
||||
|
||||
|
||||
def _make_mock_attention(conv_state: torch.Tensor, temporal_state: torch.Tensor) -> MagicMock:
|
||||
attention = MagicMock()
|
||||
attention.kv_cache = [conv_state, temporal_state]
|
||||
return attention
|
||||
|
||||
|
||||
def _make_states(
|
||||
cfg: _TestConfig, layer_names: list[str], device: torch.device
|
||||
) -> tuple[
|
||||
list[torch.Tensor],
|
||||
list[torch.Tensor],
|
||||
list[torch.Tensor],
|
||||
list[torch.Tensor],
|
||||
dict[str, MagicMock],
|
||||
dict[str, MagicMock],
|
||||
]:
|
||||
conv_py = [
|
||||
torch.randn(cfg.num_blocks, cfg.conv_width, cfg.conv_inner_dim, dtype=cfg.dtype, device=device)
|
||||
for _ in layer_names
|
||||
]
|
||||
temporal_py = [
|
||||
torch.randn(cfg.num_blocks, cfg.temporal_state_dim, dtype=cfg.dtype, device=device) for _ in layer_names
|
||||
]
|
||||
conv_gpu = [s.clone() for s in conv_py]
|
||||
temporal_gpu = [s.clone() for s in temporal_py]
|
||||
fwd_py = {name: _make_mock_attention(c, t) for name, c, t in zip(layer_names, conv_py, temporal_py)}
|
||||
fwd_gpu = {name: _make_mock_attention(c, t) for name, c, t in zip(layer_names, conv_gpu, temporal_gpu)}
|
||||
return conv_py, temporal_py, conv_gpu, temporal_gpu, fwd_py, fwd_gpu
|
||||
|
||||
|
||||
def _make_kv_cache_config(cfg: _TestConfig, layer_names: list[str]) -> KVCacheConfig:
|
||||
mamba_spec = MambaSpec(
|
||||
block_size=cfg.block_size,
|
||||
shapes=((cfg.conv_width, cfg.conv_inner_dim), (cfg.temporal_state_dim,)),
|
||||
dtypes=(cfg.dtype, cfg.dtype),
|
||||
mamba_cache_mode="all",
|
||||
)
|
||||
return KVCacheConfig(
|
||||
num_blocks=cfg.num_blocks,
|
||||
kv_cache_tensors=[],
|
||||
kv_cache_groups=[KVCacheGroupSpec(layer_names=layer_names, kv_cache_spec=mamba_spec)],
|
||||
)
|
||||
|
||||
|
||||
def _make_input_batch(req_ids: list[str], num_accepted_tokens: list[int], mamba_state_idx: list[int]) -> MagicMock:
|
||||
batch = MagicMock()
|
||||
batch.req_ids = req_ids
|
||||
batch.req_id_to_index = {rid: i for i, rid in enumerate(req_ids)}
|
||||
batch.num_accepted_tokens_cpu = np.array(num_accepted_tokens, dtype=np.int32)
|
||||
batch.mamba_state_idx_cpu = np.array(mamba_state_idx, dtype=np.int32)
|
||||
return batch
|
||||
|
||||
|
||||
def _make_requests(
|
||||
req_ids: list[str],
|
||||
num_computed_tokens: list[int],
|
||||
block_ids_per_req: list[list[int]],
|
||||
) -> dict[str, MagicMock]:
|
||||
requests: dict[str, MagicMock] = {}
|
||||
for i, req_id in enumerate(req_ids):
|
||||
req = MagicMock()
|
||||
req.num_computed_tokens = num_computed_tokens[i]
|
||||
req.block_ids = {0: block_ids_per_req[i]}
|
||||
requests[req_id] = req
|
||||
return requests
|
||||
|
||||
|
||||
def _make_copy_bufs(cfg: _TestConfig, kv_cache_config: KVCacheConfig, device: torch.device) -> MambaCopyBuffers:
|
||||
return MambaCopyBuffers.create(
|
||||
max_num_reqs=cfg.max_num_reqs,
|
||||
kv_cache_config=kv_cache_config,
|
||||
copy_funcs=_COPY_FUNCS,
|
||||
make_buffer=lambda n, dtype: _MockCpuGpuBuffer(n, dtype, device),
|
||||
)
|
||||
|
||||
|
||||
def _make_gpu_ctx(cfg: _TestConfig, kv_cache_config: KVCacheConfig, device: torch.device) -> MambaSpecDecodeGPUContext:
|
||||
return MambaSpecDecodeGPUContext.create(
|
||||
max_num_reqs=cfg.max_num_reqs,
|
||||
kv_cache_config=kv_cache_config,
|
||||
num_state_types=2,
|
||||
device=device,
|
||||
make_buffer=lambda n, dtype: _MockCpuGpuBuffer(n, dtype, device),
|
||||
)
|
||||
|
||||
|
||||
def _run_gpu_postprocess(
|
||||
gpu_ctx: MambaSpecDecodeGPUContext,
|
||||
*,
|
||||
kv_cache_config: KVCacheConfig,
|
||||
forward_context: dict[str, Any],
|
||||
copy_funcs: tuple,
|
||||
block_table: torch.Tensor,
|
||||
req_ids: list[str],
|
||||
num_accepted_tokens: list[int],
|
||||
mamba_state_idx: list[int],
|
||||
num_scheduled_tokens: dict[str, int],
|
||||
num_computed_tokens: list[int],
|
||||
num_draft_tokens: dict[str, int],
|
||||
device: torch.device,
|
||||
) -> None:
|
||||
def t(values):
|
||||
return torch.tensor(values, dtype=torch.int32, device=device)
|
||||
|
||||
gpu_ctx.initialize_from_forward_context(kv_cache_config, forward_context, copy_funcs, [block_table])
|
||||
gpu_ctx.run_fused_postprocess(
|
||||
num_reqs=len(req_ids),
|
||||
num_accepted_tokens_gpu=t(num_accepted_tokens),
|
||||
mamba_state_idx_gpu=t(mamba_state_idx),
|
||||
num_scheduled_tokens_gpu=t([num_scheduled_tokens[r] for r in req_ids]),
|
||||
num_computed_tokens_gpu=t(num_computed_tokens),
|
||||
num_draft_tokens_gpu=t([num_draft_tokens.get(r, 0) for r in req_ids]),
|
||||
)
|
||||
torch.accelerator.synchronize()
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.npu.is_available(), reason="NPU required")
|
||||
def test_matches_python_postprocess_mamba():
|
||||
cfg = _TestConfig()
|
||||
device = torch.device("npu:0")
|
||||
torch.manual_seed(42)
|
||||
|
||||
req_ids = ["req_0", "req_1", "req_2", "req_3"]
|
||||
num_computed_tokens = [60, 30, 45, 10]
|
||||
num_scheduled_tokens = {"req_0": 5, "req_1": 3, "req_2": 8, "req_3": 6}
|
||||
num_draft_tokens = {"req_0": 2, "req_1": 0, "req_2": 3, "req_3": 0}
|
||||
num_accepted_tokens = [3, 2, 4, 2]
|
||||
mamba_state_idx = [3, 1, 2, 0]
|
||||
block_ids_per_req = [
|
||||
list(range(8)),
|
||||
list(range(8, 16)),
|
||||
list(range(16, 24)),
|
||||
list(range(24, 32)),
|
||||
]
|
||||
|
||||
layer_names = [f"layer_{i}" for i in range(cfg.num_layers)]
|
||||
kv_cache_config = _make_kv_cache_config(cfg, layer_names)
|
||||
|
||||
(
|
||||
conv_states_py,
|
||||
temporal_states_py,
|
||||
conv_states_gpu,
|
||||
temporal_states_gpu,
|
||||
forward_context_py,
|
||||
forward_context_gpu,
|
||||
) = _make_states(cfg, layer_names, device)
|
||||
|
||||
scheduler_output = _make_scheduler_output(
|
||||
num_scheduled_tokens,
|
||||
{k: [None] * v for k, v in num_draft_tokens.items() if v > 0},
|
||||
)
|
||||
input_batch_py = _make_input_batch(req_ids, num_accepted_tokens.copy(), mamba_state_idx.copy())
|
||||
requests = _make_requests(req_ids, num_computed_tokens, block_ids_per_req)
|
||||
copy_bufs = _make_copy_bufs(cfg, kv_cache_config, device)
|
||||
|
||||
postprocess_mamba(
|
||||
scheduler_output,
|
||||
kv_cache_config,
|
||||
input_batch_py,
|
||||
requests,
|
||||
forward_context_py,
|
||||
_COPY_FUNCS,
|
||||
copy_bufs,
|
||||
)
|
||||
torch.accelerator.synchronize()
|
||||
|
||||
gpu_ctx = _make_gpu_ctx(cfg, kv_cache_config, device)
|
||||
block_table_gpu = torch.zeros(len(req_ids), 8, dtype=torch.int32, device=device)
|
||||
for i, block_ids in enumerate(block_ids_per_req):
|
||||
block_table_gpu[i, : len(block_ids)] = torch.tensor(block_ids, dtype=torch.int32)
|
||||
|
||||
_run_gpu_postprocess(
|
||||
gpu_ctx,
|
||||
kv_cache_config=kv_cache_config,
|
||||
forward_context=forward_context_gpu,
|
||||
copy_funcs=_COPY_FUNCS,
|
||||
block_table=block_table_gpu,
|
||||
req_ids=req_ids,
|
||||
num_accepted_tokens=num_accepted_tokens,
|
||||
mamba_state_idx=mamba_state_idx,
|
||||
num_scheduled_tokens=num_scheduled_tokens,
|
||||
num_computed_tokens=num_computed_tokens,
|
||||
num_draft_tokens=num_draft_tokens,
|
||||
device=device,
|
||||
)
|
||||
|
||||
for i in range(cfg.num_layers):
|
||||
torch.testing.assert_close(conv_states_gpu[i], conv_states_py[i])
|
||||
torch.testing.assert_close(temporal_states_gpu[i], temporal_states_py[i])
|
||||
|
||||
expected_accepted = torch.tensor(
|
||||
input_batch_py.num_accepted_tokens_cpu[: len(req_ids)],
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
)
|
||||
torch.testing.assert_close(gpu_ctx.num_accepted_tokens_out[: len(req_ids)], expected_accepted)
|
||||
@@ -0,0 +1,76 @@
|
||||
import gc
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from vllm.triton_utils import triton
|
||||
|
||||
from vllm_ascend.ops.triton.spec_decode.utils import prepare_inputs_padded_kernel
|
||||
from vllm_ascend.ops.triton.triton_utils import get_vectorcore_num
|
||||
from vllm_ascend.spec_decode.llm_base_proposer import _PREPARE_INPUTS_BLOCK_SIZE as BLOCK_SIZE
|
||||
|
||||
|
||||
def prepare_inputs_padded_ref(
|
||||
cu_num_draft_tokens,
|
||||
valid_sampled_tokens_count,
|
||||
query_start_loc,
|
||||
):
|
||||
num_draft_tokens = torch.cat(
|
||||
[
|
||||
cu_num_draft_tokens[0:1],
|
||||
cu_num_draft_tokens[1:] - cu_num_draft_tokens[:-1],
|
||||
]
|
||||
)
|
||||
|
||||
num_rejected_tokens = torch.where(
|
||||
num_draft_tokens > 0,
|
||||
num_draft_tokens + 1 - valid_sampled_tokens_count,
|
||||
torch.zeros_like(num_draft_tokens),
|
||||
)
|
||||
|
||||
token_indices_to_sample = query_start_loc[1:] - 1 - num_rejected_tokens
|
||||
|
||||
return token_indices_to_sample.to(torch.int32)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("num_reqs", [1, 7, 32, 128, 2048])
|
||||
def test_prepare_inputs_padded(num_reqs):
|
||||
device = "npu"
|
||||
torch.manual_seed(0)
|
||||
|
||||
draft_lens = torch.randint(1, 6, (num_reqs,), device=device, dtype=torch.int32)
|
||||
|
||||
cu_num_draft_tokens = torch.cumsum(draft_lens, dim=0).to(torch.int32)
|
||||
|
||||
valid_sampled_tokens_count = torch.zeros_like(draft_lens)
|
||||
for i in range(num_reqs):
|
||||
valid_sampled_tokens_count[i] = torch.randint(0, draft_lens[i] + 2, (1,)).item()
|
||||
|
||||
seq_lens = draft_lens + 1
|
||||
query_start_loc = torch.zeros(num_reqs + 1, device=device, dtype=torch.int32)
|
||||
query_start_loc[1:] = torch.cumsum(seq_lens, dim=0)
|
||||
|
||||
# Run PyTorch reference
|
||||
out_ref = prepare_inputs_padded_ref(cu_num_draft_tokens, valid_sampled_tokens_count, query_start_loc)
|
||||
|
||||
# Run Triton kernel
|
||||
out_tri = torch.empty(num_reqs, dtype=torch.int32, device=device)
|
||||
num_rejected_tokens = torch.empty(num_reqs, dtype=torch.int32, device=device)
|
||||
num_blocks_needed = triton.cdiv(num_reqs, BLOCK_SIZE)
|
||||
num_vector_core = get_vectorcore_num()
|
||||
grid_size = min(num_blocks_needed, num_vector_core)
|
||||
grid = (grid_size,)
|
||||
|
||||
prepare_inputs_padded_kernel[grid](
|
||||
cu_num_draft_tokens,
|
||||
valid_sampled_tokens_count,
|
||||
query_start_loc,
|
||||
out_tri,
|
||||
num_rejected_tokens,
|
||||
num_reqs,
|
||||
BLOCK_SIZE=BLOCK_SIZE,
|
||||
)
|
||||
|
||||
torch.testing.assert_close(out_tri, out_ref)
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
torch.npu.reset_peak_memory_stats()
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user