Signed-off-by: Joey-gvwal <joey_gvwal@yeah.net> Co-authored-by: R0CKSTAR <yeahdongcn@gmail.com>
325 lines
15 KiB
C++
325 lines
15 KiB
C++
/* Copyright 2025 SGLang Team. 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.
|
|
==============================================================================*/
|
|
|
|
#include <ATen/core/dispatch/Dispatcher.h>
|
|
#include <torch/library.h>
|
|
|
|
#include "sgl_kernel_ops.h"
|
|
#include "torch_musa/csrc/aten/musa/MUSAContext.h"
|
|
|
|
TORCH_LIBRARY_EXPAND(sgl_kernel, m) {
|
|
/*
|
|
* From csrc/allreduce
|
|
*/
|
|
m.def("get_graph_buffer_ipc_meta", &get_graph_buffer_ipc_meta);
|
|
m.def("register_graph_buffers", ®ister_graph_buffers);
|
|
m.def("dispose", &dispose);
|
|
m.def("meta_size", &meta_size);
|
|
m.def("register_buffer", ®ister_buffer);
|
|
|
|
m.def(
|
|
"init_custom_ar(int[] ipc_tensors, Tensor rank_data, "
|
|
"int rank, bool full_nvlink) -> int");
|
|
m.impl("init_custom_ar", torch::kMUSA, &init_custom_ar);
|
|
|
|
m.def(
|
|
"all_reduce(int fa, Tensor inp, Tensor! out, int reg_buffer, "
|
|
"int reg_buffer_sz_bytes) -> ()");
|
|
m.impl("all_reduce", torch::kMUSA, &all_reduce);
|
|
|
|
/*
|
|
* From csrc/attention
|
|
*/
|
|
m.def("merge_state_v2(Tensor v_a, Tensor s_a, Tensor v_b, Tensor s_b, Tensor! v_merged, Tensor! s_merged) -> ()");
|
|
m.impl("merge_state_v2", torch::kMUSA, &merge_state_v2);
|
|
|
|
/*
|
|
* From csrc/elementwise
|
|
*/
|
|
m.def("rmsnorm(Tensor! output, Tensor input, Tensor weight, float eps, bool enable_pdl) -> ()");
|
|
m.impl("rmsnorm", torch::kMUSA, &rmsnorm);
|
|
|
|
m.def("fused_add_rmsnorm(Tensor! input, Tensor! residual, Tensor weight, float eps, bool enable_pdl) -> ()");
|
|
m.impl("fused_add_rmsnorm", torch::kMUSA, &musa_fused_add_rms_norm);
|
|
|
|
m.def("gemma_rmsnorm(Tensor! output, Tensor input, Tensor weight, float eps, bool enable_pdl) -> ()");
|
|
m.impl("gemma_rmsnorm", torch::kMUSA, &gemma_rmsnorm);
|
|
|
|
m.def("gemma_fused_add_rmsnorm(Tensor! input, Tensor! residual, Tensor weight, float eps, bool enable_pdl) -> ()");
|
|
m.impl("gemma_fused_add_rmsnorm", torch::kMUSA, &gemma_fused_add_rmsnorm);
|
|
|
|
m.def("silu_and_mul(Tensor! out, Tensor input) -> ()");
|
|
m.impl("silu_and_mul", torch::kMUSA, &silu_and_mul);
|
|
|
|
m.def("gelu_tanh_and_mul(Tensor! out, Tensor input) -> ()");
|
|
m.impl("gelu_tanh_and_mul", torch::kMUSA, &gelu_tanh_and_mul);
|
|
|
|
m.def("gelu_and_mul(Tensor! out, Tensor input) -> ()");
|
|
m.impl("gelu_and_mul", torch::kMUSA, &gelu_and_mul);
|
|
|
|
m.def("concat_mla_k(Tensor! k, Tensor k_nope, Tensor k_rope) -> ()");
|
|
m.impl("concat_mla_k", torch::kMUSA, &concat_mla_k);
|
|
|
|
m.def(
|
|
"rotary_embedding(Tensor positions, Tensor! query,"
|
|
" Tensor!? key, int head_size,"
|
|
" Tensor cos_sin_cache, bool is_neox) -> ()");
|
|
m.impl("rotary_embedding", torch::kMUSA, &rotary_embedding);
|
|
|
|
/*
|
|
* From csrc/gemm
|
|
*/
|
|
m.def("awq_dequantize(Tensor qweight, Tensor scales, Tensor qzeros) -> Tensor");
|
|
m.impl("awq_dequantize", torch::kMUSA, &awq_dequantize);
|
|
|
|
m.def(
|
|
"sgl_per_token_group_quant_8bit(Tensor input, Tensor output_q, Tensor output_s, int group_size,"
|
|
" float eps, float fp8_min, float fp8_max, bool scale_ue8m0) -> ()");
|
|
m.impl("sgl_per_token_group_quant_8bit", torch::kMUSA, &sgl_per_token_group_quant_8bit);
|
|
|
|
m.def(
|
|
"sgl_per_token_group_quant_8bit_v2(Tensor input, Tensor output_q, Tensor output_s, int group_size,"
|
|
" float eps, float fp8_min, float fp8_max, bool scale_ue8m0, bool fuse_silu_and_mul, Tensor? masked_m) -> ()");
|
|
m.impl("sgl_per_token_group_quant_8bit_v2", torch::kMUSA, &sgl_per_token_group_quant_8bit_v2);
|
|
|
|
m.def("sgl_per_token_quant_fp8(Tensor input, Tensor output_q, Tensor output_s) -> ()");
|
|
m.impl("sgl_per_token_quant_fp8", torch::kMUSA, &sgl_per_token_quant_fp8);
|
|
|
|
m.def("dsv3_fused_a_gemm(Tensor! output, Tensor mat_a, Tensor mat_b) -> ()");
|
|
m.impl("dsv3_fused_a_gemm", torch::kMUSA, &dsv3_fused_a_gemm);
|
|
|
|
m.def("dsv3_router_gemm(Tensor! output, Tensor mat_a, Tensor mat_b) -> ()");
|
|
m.impl("dsv3_router_gemm", torch::kMUSA, &dsv3_router_gemm);
|
|
|
|
/*
|
|
* From csrc/moe
|
|
*/
|
|
m.def(
|
|
"moe_align_block_size(Tensor topk_ids, int num_experts, int block_size, Tensor! sorted_token_ids, Tensor! "
|
|
"experts_ids, Tensor! num_tokens_post_pad, Tensor! cumsum_buffer, bool "
|
|
"pad_sorted_token_ids) -> ()");
|
|
m.impl("moe_align_block_size", torch::kMUSA, &moe_align_block_size);
|
|
|
|
m.def(
|
|
"topk_softmax(Tensor! topk_weights, Tensor! topk_indices, Tensor gating_output, bool renormalize, float "
|
|
"moe_softcapping, Tensor? correction_bias) -> ()");
|
|
m.impl("topk_softmax", torch::kMUSA, &topk_softmax);
|
|
|
|
m.def("moe_sum_reduce(Tensor input, Tensor output, float routed_scaling_factor) -> ()");
|
|
m.impl("moe_sum_reduce", torch::kMUSA, &moe_sum_reduce);
|
|
|
|
m.def("moe_sum(Tensor input, Tensor! output) -> ()");
|
|
m.impl("moe_sum", torch::kMUSA, &moe_sum);
|
|
|
|
m.def(
|
|
"moe_fused_gate(Tensor input, Tensor bias, int num_expert_group, int topk_group, int topk, int "
|
|
"num_fused_shared_experts, float routed_scaling_factor, bool apply_routed_scaling_factor_on_output) -> "
|
|
"(Tensor[])");
|
|
m.impl("moe_fused_gate", torch::kMUSA, &moe_fused_gate);
|
|
|
|
m.def(
|
|
"kimi_k2_moe_fused_gate(Tensor input, Tensor bias, int topk, bool renormalize, "
|
|
"float routed_scaling_factor, bool apply_routed_scaling_factor_on_output) -> "
|
|
"(Tensor[])");
|
|
m.impl("kimi_k2_moe_fused_gate", torch::kMUSA, &kimi_k2_moe_fused_gate);
|
|
|
|
/*
|
|
* From csrc/speculative
|
|
*/
|
|
m.def(
|
|
"tree_speculative_sampling_target_only(Tensor! predicts, Tensor! accept_index, Tensor! accept_token_num, "
|
|
"Tensor candidates, Tensor retrive_index, Tensor retrive_next_token, Tensor retrive_next_sibling, "
|
|
"Tensor uniform_samples, Tensor uniform_samples_for_final_sampling, Tensor target_probs, Tensor draft_probs, "
|
|
"float threshold_single, float threshold_acc, "
|
|
"bool deterministic) -> ()");
|
|
m.impl("tree_speculative_sampling_target_only", torch::kMUSA, &tree_speculative_sampling_target_only);
|
|
|
|
m.def(
|
|
"verify_tree_greedy(Tensor! predicts, Tensor! accept_index, Tensor! accept_token_num, "
|
|
"Tensor candidates, Tensor retrive_index, Tensor retrive_next_token, Tensor retrive_next_sibling, "
|
|
"Tensor target_predict) -> ()");
|
|
m.impl("verify_tree_greedy", torch::kMUSA, &verify_tree_greedy);
|
|
|
|
m.def(
|
|
"reconstruct_indices_from_tree_mask(Tensor tree_mask, Tensor verified_seq_len, Tensor positions, "
|
|
"Tensor retrive_index, Tensor retrive_next_token, Tensor retrive_next_sibling, "
|
|
"int batch_size, int draft_token_num) -> ()");
|
|
m.impl("reconstruct_indices_from_tree_mask", torch::kMUSA, &reconstruct_indices_from_tree_mask);
|
|
|
|
m.def(
|
|
"build_tree_kernel_efficient(Tensor parent_list, Tensor selected_index, Tensor verified_seq_len, "
|
|
"Tensor! tree_mask, Tensor! positions, Tensor! retrive_index, Tensor! retrive_next_token, "
|
|
"Tensor! retrive_next_sibling, int topk, int depth, int draft_token_num, int tree_mask_mode) -> "
|
|
"()");
|
|
m.impl("build_tree_kernel_efficient", torch::kMUSA, &build_tree_kernel_efficient);
|
|
|
|
/*
|
|
* From csrc/grammar
|
|
*/
|
|
m.def("apply_token_bitmask_inplace_cuda(Tensor logits, Tensor bitmask, Tensor? indices=None) -> ()");
|
|
m.impl("apply_token_bitmask_inplace_cuda", &ApplyTokenBitmaskInplace);
|
|
|
|
/*
|
|
* From csrc/quantization/gguf
|
|
*/
|
|
m.def(
|
|
"ggml_dequantize(Tensor W, int type, SymInt m, SymInt n, ScalarType? "
|
|
"dtype) -> Tensor");
|
|
m.impl("ggml_dequantize", torch::kMUSA, &ggml_dequantize);
|
|
|
|
m.def(
|
|
"ggml_mul_mat_vec_a8(Tensor W, Tensor X, int type, SymInt row) "
|
|
"-> Tensor");
|
|
m.impl("ggml_mul_mat_vec_a8", torch::kMUSA, &ggml_mul_mat_vec_a8);
|
|
|
|
m.def("ggml_mul_mat_a8(Tensor W, Tensor X, int type, SymInt row) -> Tensor");
|
|
m.impl("ggml_mul_mat_a8", torch::kMUSA, &ggml_mul_mat_a8);
|
|
|
|
m.def(
|
|
"ggml_moe_a8(Tensor X, Tensor W, "
|
|
"Tensor sorted_token_ids, Tensor expert_ids, Tensor "
|
|
"num_tokens_post_padded, "
|
|
"int type, SymInt row, SymInt top_k, SymInt tokens) -> Tensor");
|
|
m.impl("ggml_moe_a8", torch::kMUSA, &ggml_moe_a8);
|
|
|
|
m.def(
|
|
"ggml_moe_a8_vec(Tensor X, Tensor W, "
|
|
"Tensor topk_ids, int top_k, "
|
|
"int type, SymInt row, SymInt tokens) -> Tensor");
|
|
m.impl("ggml_moe_a8_vec", torch::kMUSA, &ggml_moe_a8_vec);
|
|
|
|
m.def("ggml_moe_get_block_size(int type) -> int");
|
|
m.impl("ggml_moe_get_block_size", torch::kMUSA, &ggml_moe_get_block_size);
|
|
|
|
/*
|
|
* From csrc/kvcacheio
|
|
*/
|
|
m.def(
|
|
"transfer_kv_per_layer(Tensor src_k, Tensor dst_k, Tensor src_v, Tensor dst_v, Tensor src_indices, Tensor "
|
|
"dst_indices, int item_size, int block_quota, int num_warps_per_block) -> ()");
|
|
m.impl("transfer_kv_per_layer", torch::kMUSA, &transfer_kv_per_layer);
|
|
m.def(
|
|
"transfer_kv_per_layer_pf_lf(Tensor src_k, Tensor dst_k, Tensor src_v, Tensor dst_v, Tensor src_indices, Tensor "
|
|
"dst_indices, int layer_id, int item_size, int src_layout_dim, int block_quota, int num_warps_per_block) -> ()");
|
|
m.impl("transfer_kv_per_layer_pf_lf", torch::kMUSA, &transfer_kv_per_layer_pf_lf);
|
|
m.def(
|
|
"transfer_kv_per_layer_ph_lf(Tensor src_k, Tensor dst_k, Tensor src_v, Tensor dst_v, Tensor src_indices, Tensor "
|
|
"dst_indices, int layer_id, int item_size, int src_layout_dim, int page_size, int head_num, int block_quota, int "
|
|
"num_warps_per_block) -> ()");
|
|
m.impl("transfer_kv_per_layer_ph_lf", torch::kMUSA, &transfer_kv_per_layer_ph_lf);
|
|
m.def(
|
|
"transfer_kv_all_layer(Tensor src_k_layers, Tensor dst_k_layers, Tensor src_v_layers, Tensor dst_v_layers, "
|
|
"Tensor src_indices, Tensor dst_indices, int item_size, int num_layers, int block_quota, int "
|
|
"num_warps_per_block) -> ()");
|
|
m.impl("transfer_kv_all_layer", torch::kMUSA, &transfer_kv_all_layer);
|
|
m.def(
|
|
"transfer_kv_all_layer_lf_pf(Tensor src_k_layers, Tensor dst_k, Tensor src_v_layers, Tensor dst_v, "
|
|
"Tensor src_indices, Tensor dst_indices, int item_size, int dst_layout_dim, int num_layers, int block_quota, int "
|
|
"num_warps_per_block) -> ()");
|
|
m.impl("transfer_kv_all_layer_lf_pf", torch::kMUSA, &transfer_kv_all_layer_lf_pf);
|
|
m.def(
|
|
"transfer_kv_all_layer_lf_ph(Tensor src_k_layers, Tensor dst_k, Tensor src_v_layers, Tensor dst_v, "
|
|
"Tensor src_indices, Tensor dst_indices, int item_size, int dst_layout_dim, int num_layers, int page_size, int "
|
|
"head_num, int block_quota, int num_warps_per_block) -> ()");
|
|
m.impl("transfer_kv_all_layer_lf_ph", torch::kMUSA, &transfer_kv_all_layer_lf_ph);
|
|
m.def(
|
|
"transfer_kv_per_layer_mla(Tensor src, Tensor dst, Tensor src_indices, Tensor dst_indices, int item_size, int "
|
|
"block_quota, int num_warps_per_block) -> ()");
|
|
m.impl("transfer_kv_per_layer_mla", torch::kMUSA, &transfer_kv_per_layer_mla);
|
|
m.def(
|
|
"transfer_kv_per_layer_mla_pf_lf(Tensor src, Tensor dst, Tensor src_indices, Tensor dst_indices, int layer_id, "
|
|
"int item_size, int src_layout_dim, int block_quota, int num_warps_per_block) -> ()");
|
|
m.impl("transfer_kv_per_layer_mla_pf_lf", torch::kMUSA, &transfer_kv_per_layer_mla_pf_lf);
|
|
m.def(
|
|
"transfer_kv_all_layer_mla(Tensor src_layers, Tensor dst_layers, Tensor src_indices, Tensor dst_indices, int "
|
|
"item_size, int num_layers, int block_quota, int num_warps_per_block) -> ()");
|
|
m.impl("transfer_kv_all_layer_mla", torch::kMUSA, &transfer_kv_all_layer_mla);
|
|
m.def(
|
|
"transfer_kv_all_layer_mla_lf_pf(Tensor src_layers, Tensor dst, Tensor src_indices, Tensor dst_indices, "
|
|
"int item_size, int dst_layout_dim, int num_layers, int block_quota, int num_warps_per_block) -> ()");
|
|
m.impl("transfer_kv_all_layer_mla_lf_pf", torch::kMUSA, &transfer_kv_all_layer_mla_lf_pf);
|
|
m.def(
|
|
"transfer_kv_direct(Tensor[] src_layers, Tensor[] dst_layers, Tensor src_indices, Tensor dst_indices, int "
|
|
"page_size) -> ()");
|
|
m.impl("transfer_kv_direct", torch::kMUSA, &transfer_kv_direct);
|
|
m.def(
|
|
"transfer_kv_per_layer_direct_pf_lf(Tensor[] src_ptrs, Tensor[] dst_ptrs, Tensor src_indices, "
|
|
"Tensor dst_indices, int layer_id, int page_size)->() ");
|
|
m.impl("transfer_kv_per_layer_direct_pf_lf", torch::kMUSA, &transfer_kv_per_layer_direct_pf_lf);
|
|
m.def(
|
|
"transfer_kv_all_layer_direct_lf_pf(Tensor[] src_ptrs, Tensor[] dst_ptrs, Tensor src_indices, "
|
|
"Tensor dst_indices, int page_size) ->() ");
|
|
m.impl("transfer_kv_all_layer_direct_lf_pf", torch::kMUSA, &transfer_kv_all_layer_direct_lf_pf);
|
|
|
|
/*
|
|
* From FlashInfer
|
|
*/
|
|
m.def(
|
|
"bmm_fp8(Tensor A, Tensor B, Tensor! D, Tensor A_scale, Tensor B_scale, Tensor workspace_buffer, "
|
|
"int cublas_handle) -> ()",
|
|
{at::Tag::needs_fixed_stride_order});
|
|
m.impl("bmm_fp8", torch::kMUSA, &bmm_fp8);
|
|
|
|
m.def("top_k_renorm_probs(Tensor probs, Tensor! renorm_probs, Tensor? maybe_top_k_arr, int top_k_val) -> ()");
|
|
m.impl("top_k_renorm_probs", torch::kMUSA, &top_k_renorm_probs);
|
|
|
|
m.def("top_p_renorm_probs(Tensor probs, Tensor! renorm_probs, Tensor? maybe_top_p_arr, float top_p_val) -> ()");
|
|
m.impl("top_p_renorm_probs", torch::kMUSA, &top_p_renorm_probs);
|
|
|
|
/*
|
|
* From csrc/musa
|
|
*/
|
|
m.def(
|
|
"musa_batched_rotary_embedding_contiguous(Tensor! positions, Tensor! query, Tensor! key, "
|
|
"int head_size, Tensor! cos_sin_cache, bool is_neox, int rot_dim, Tensor! cos_sin_cache_offsets) -> ()");
|
|
m.impl("musa_batched_rotary_embedding_contiguous", torch::kMUSA, &batched_rotary_embedding_contiguous);
|
|
|
|
m.def(
|
|
"musa_rotary_embedding_contiguous(Tensor! positions, Tensor! query, Tensor! key, "
|
|
"int head_size, Tensor! cos_sin_cache, bool is_neox) -> ()");
|
|
m.impl("musa_rotary_embedding_contiguous", torch::kMUSA, &rotary_embedding_contiguous);
|
|
|
|
m.def(
|
|
"musa_fused_moe_gemv(Tensor! A, Tensor! B, Tensor! C, Tensor? A_scale, Tensor? B_scale,"
|
|
"Tensor! topk_weights, Tensor! topk_ids, bool mul_routed_weight, int topk, bool use_int4_w4a16,"
|
|
"bool use_swigelu) -> ()");
|
|
m.impl("fused_moe_gemv", torch::kMUSA, &fused_moe_gemv);
|
|
|
|
m.def(
|
|
"musa_fused_gemv(Tensor! A, Tensor! B, Tensor! C, Tensor? A_scale, Tensor? B_scale,"
|
|
"bool use_int4_w4a16, bool use_swigelu, bool use_rms_norm, Tensor? gamma,"
|
|
"float eps) -> ()");
|
|
m.impl("musa_fused_gemv", torch::kMUSA, &musa_fused_gemv);
|
|
|
|
m.def(
|
|
"musa_fused_mul_add(Tensor! output, Tensor! self, Tensor! bias,"
|
|
"float scale) -> ()");
|
|
m.impl("musa_fused_mul_add", torch::kMUSA, &fused_mul_add);
|
|
|
|
m.def(
|
|
"musa_top_k_top_p_sampling_from_probs(Tensor probs, Tensor output, Tensor? maybe_indices, Tensor? "
|
|
"maybe_top_k_arr, "
|
|
"float top_k_val, Tensor? maybe_top_p_arr, float top_p_val, bool deterministic, Generator? gen) -> ()");
|
|
m.impl("musa_top_k_top_p_sampling_from_probs", torch::kMUSA, &musa_top_k_top_p_sampling_from_probs);
|
|
|
|
/*
|
|
* From csrc/memory
|
|
*/
|
|
m.def("weak_ref_tensor(Tensor tensor) -> Tensor");
|
|
m.impl("weak_ref_tensor", torch::kMUSA, &weak_ref_tensor);
|
|
}
|
|
|
|
REGISTER_EXTENSION(common_ops)
|