diff --git a/sgl-kernel/csrc/allreduce/custom_all_reduce.cuh b/sgl-kernel/csrc/allreduce/custom_all_reduce.cuh index 4a2d62b28..e9b32bc67 100644 --- a/sgl-kernel/csrc/allreduce/custom_all_reduce.cuh +++ b/sgl-kernel/csrc/allreduce/custom_all_reduce.cuh @@ -17,7 +17,24 @@ namespace sglang { +#ifndef USE_MUSA constexpr int kMaxBlocks = 36; +constexpr int kDefaultThreads = 512; +constexpr int kDefaultBlockLimit = 36; +constexpr int kMaxThreadsPerBlock = 512; +#else +constexpr int kMaxBlocks = 60; +constexpr int kDefaultThreads = 1024; +constexpr int kDefaultBlockLimit = 60; +constexpr int kMaxThreadsPerBlock = 1024; +#endif + +// Allreduce algorithm selection thresholds +constexpr int kAllReduceGPUSmall = 4; +constexpr int kAllReduceGPULarge = 8; +constexpr size_t kAllReduceSmallThreshold = 512 * 1024; // 512KB +constexpr size_t kAllReduceLargeThreshold = 256 * 1024; // 256KB + // Counter may overflow, but it's fine since unsigned int overflow is // well-defined behavior. using FlagType = uint32_t; @@ -134,7 +151,9 @@ DINLINE O downcast(array_t val) { } static DINLINE void st_flag_release(FlagType* flag_addr, FlagType flag) { -#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700 +#ifdef USE_MUSA + volatile_store((uint32_t)flag, (uint32_t*)flag_addr); +#elif defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700 asm volatile("st.release.sys.global.u32 [%1], %0;" ::"r"(flag), "l"(flag_addr)); #else asm volatile("membar.sys; st.volatile.global.u32 [%1], %0;" ::"r"(flag), "l"(flag_addr)); @@ -142,6 +161,11 @@ static DINLINE void st_flag_release(FlagType* flag_addr, FlagType flag) { } static DINLINE FlagType ld_flag_acquire(FlagType* flag_addr) { +#ifdef USE_MUSA + flushInv_byp(); + return (uint32_t)volatile_load((uint32_t*)flag_addr); +#endif + FlagType flag; #if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700 asm volatile("ld.acquire.sys.global.u32 %0, [%1];" : "=r"(flag) : "l"(flag_addr)); @@ -167,7 +191,12 @@ static DINLINE FlagType ld_flag_volatile(FlagType* flag_addr) { // barrier. template DINLINE void multi_gpu_barrier(const RankSignals& sg, Signal* self_sg, int rank) { - if constexpr (!is_start) __syncthreads(); + if constexpr (!is_start) +#ifndef USE_MUSA + __syncthreads(); +#else + __syncthreads_lm(); +#endif static_assert(!(is_start && need_fence)); // Start barrier shouldn't need fence. if (threadIdx.x < ngpus) { // Increment the counter. Technically we only need one counter, but we use @@ -187,7 +216,12 @@ DINLINE void multi_gpu_barrier(const RankSignals& sg, Signal* self_sg, int rank) ; } } - if constexpr (is_start || need_fence) __syncthreads(); + if constexpr (is_start || need_fence) +#ifndef USE_MUSA + __syncthreads(); +#else + __syncthreads_lm(); +#endif } template @@ -201,7 +235,7 @@ DINLINE P packed_reduce(const P* ptrs[], int idx) { } template -__global__ void __launch_bounds__(512, 1) cross_device_reduce_1stage( +__global__ void __launch_bounds__(kMaxThreadsPerBlock, 1) cross_device_reduce_1stage( RankData* _dp, RankSignals sg, Signal* self_sg, T* __restrict__ result, int rank, int size) { using P = typename packed_t::P; using A = typename packed_t::A; @@ -221,8 +255,132 @@ DINLINE P* get_tmp_buf(Signal* sg) { return (P*)(((Signal*)sg) + 1); } +#ifdef USE_MUSA +template +DINLINE void shfl_reduce(float* res) { + if constexpr (nranks >= 4) { +#pragma unroll + for (int32_t i = 0; i < vlen; i++) { + res[i] += __shfl_xor_sync(0xffffffff, res[i], 16); + } + } +#pragma unroll + for (int32_t i = 0; i < vlen; i++) { + res[i] += __shfl_xor_sync(0xffffffff, res[i], 8); + } +} + +template +__global__ void __launch_bounds__(kMaxThreadsPerBlock, 1) custom_all_reduce_2shot( + RankData* _dp, RankSignals sg, Signal* self_sg, T* __restrict__ result, int32_t local_rank, int32_t size) { + constexpr int32_t nranks_sft = (nranks >> 1) - (nranks >> 3); // 8->3, 4->2, 2->1 + constexpr int32_t coalesce_num = 8; + constexpr int32_t coalesce_sft = 3; // 8 threads per rank in group + constexpr int32_t group_size = nranks << coalesce_sft; // tp 8 -> 64 threads, tp 4 -> 32 threads, tp 2 -> 16 threads + constexpr int32_t group_stride_sft = nranks_sft + coalesce_sft; + const int32_t tidx = threadIdx.x; + const int32_t bidx = blockIdx.x; + const int32_t thread_num = blockDim.x; + const int32_t lane_idx = tidx & 31; + const int32_t warp_idx = tidx >> 5; + const int32_t group_num = thread_num >> group_stride_sft; + const int32_t target_rank = (tidx >> coalesce_sft) & (nranks - 1); + const int32_t group_id = tidx >> group_stride_sft; + const int32_t coalesce_tid = tidx & (coalesce_num - 1); + + typedef int16_t Vec __attribute__((vector_size(16))); + + const int32_t stride = gridDim.x * thread_num; + int32_t idx_base = bidx * thread_num; + int32_t idx_in_blk = coalesce_tid + (local_rank << coalesce_sft) + (group_id << group_stride_sft); + + // first sync barrier + FlagType* target_barrier = nullptr; + FlagType* local_barrier = nullptr; + FlagType flag; + if (tidx < nranks) { + flag = atomicAdd(&(self_sg->self_counter[bidx][tidx]), 1); + target_barrier = &sg.signals[tidx]->peer_counter[flag & 1][bidx][local_rank]; + local_barrier = &self_sg->peer_counter[flag & 1][bidx][tidx]; + atomicExch(target_barrier, flag); + while (atomicAdd(local_barrier, 0) != flag) { + } + } + __syncthreads_lm(); + + // reduce scatter + Vec* target_ptr = (Vec*)_dp->ptrs[target_rank]; + Vec* buffer_ptr = get_tmp_buf(sg.signals[local_rank]); + do { + int32_t idx = idx_in_blk + idx_base; + float temp_res[vlen] = {0}; + if (idx < size) { + T* data = reinterpret_cast(&(target_ptr[idx])); +#pragma unroll + for (int32_t i = 0; i < vlen; i++) { + temp_res[i] = upcast_s(data[i]); + } + } + shfl_reduce(temp_res); + // reduce cross warp, only trigger when tp 8 + if constexpr (nranks == 8) { + __shared__ float smem[kMaxThreadsPerBlock << 1]; + if (lane_idx < coalesce_num) { +#pragma unroll + for (int32_t i = 0; i < vlen; i++) { + smem[warp_idx * vlen * coalesce_num + coalesce_tid * vlen + i] = temp_res[i]; + } + } + __syncthreads_lm(); +#pragma unroll + for (int32_t i = 0; i < vlen; i++) { + temp_res[i] += smem[(warp_idx ^ 1) * vlen * coalesce_num + coalesce_tid * vlen + i]; + } + } + + if (local_rank == target_rank && idx < size) { + Vec res; +#pragma unroll + for (int32_t i = 0; i < vlen; i++) { + reinterpret_cast(&res)[i] = downcast_s(temp_res[i]); + } + buffer_ptr[idx] = res; + } + idx_base += stride; + } while (idx_base < size); + // make sure buffer_ptr data ready + __musa_barrier_slc(); + __syncthreads_lm(); + if (tidx == 0) { + __threadfence_system_noflush(); + } + buffer_ptr = get_tmp_buf(sg.signals[target_rank]); + // second sync barrier + if (tidx < nranks) { + flag = atomicAdd(&(self_sg->self_counter[bidx][tidx]), 1); + target_barrier = &sg.signals[tidx]->peer_counter[flag & 1][bidx][local_rank]; + local_barrier = &self_sg->peer_counter[flag & 1][bidx][tidx]; + atomicExch(target_barrier, flag); + while (atomicAdd(local_barrier, 0) != flag) { + } + } + __syncthreads_lm(); + + // all gather + idx_in_blk = coalesce_tid + (target_rank << coalesce_sft) + (group_id << group_stride_sft); + idx_base = bidx * thread_num; + do { + int32_t idx = idx_in_blk + idx_base; + if (idx < size) { + reinterpret_cast(result)[idx] = buffer_ptr[idx]; + } + idx_base += stride; + } while (idx_base < size); +} +#endif // USE_MUSA + template -__global__ void __launch_bounds__(512, 1) cross_device_reduce_2stage( +__global__ void __launch_bounds__(kMaxThreadsPerBlock, 1) cross_device_reduce_2stage( RankData* _dp, RankSignals sg, Signal* self_sg, T* __restrict__ result, int rank, int size) { int tid = blockIdx.x * blockDim.x + threadIdx.x; int stride = gridDim.x * blockDim.x; @@ -414,7 +572,13 @@ class CustomAllreduce { * guess is that too many SMs will cause contention on NVLink bus. */ template - void allreduce(cudaStream_t stream, T* input, T* output, int size, int threads = 512, int block_limit = 36) { + void allreduce( + cudaStream_t stream, + T* input, + T* output, + int size, + int threads = kDefaultThreads, + int block_limit = kDefaultBlockLimit) { auto d = packed_t::P::size; if (size % d != 0) throw std::runtime_error( @@ -462,26 +626,43 @@ class CustomAllreduce { #define KL(ngpus, name) name<<>>(ptrs, sg_, self_sg_, output, rank_, size); // TODO(hanzhi713): Threshold is different for A100 and H100. // Add per device threshold. -#define REDUCE_CASE(ngpus) \ - case ngpus: { \ - if (force_1stage) { \ - KL(ngpus, cross_device_reduce_1stage); \ - } else if (force_2stage) { \ - KL(ngpus, cross_device_reduce_2stage); \ - } else { \ - if (world_size_ == 2) { \ - KL(ngpus, cross_device_reduce_1stage); \ - } else if (full_nvlink_) { \ - if ((world_size_ <= 4 && bytes < 512 * 1024) || (world_size_ <= 8 && bytes < 256 * 1024)) { \ - KL(ngpus, cross_device_reduce_1stage); \ - } else { \ - KL(ngpus, cross_device_reduce_2stage); \ - } \ - } \ - } \ - break; \ +#ifndef USE_MUSA +#define REDUCE_CASE(ngpus) \ + case ngpus: { \ + if (force_1stage) { \ + KL(ngpus, cross_device_reduce_1stage); \ + } else if (force_2stage) { \ + KL(ngpus, cross_device_reduce_2stage); \ + } else { \ + if (world_size_ == 2) { \ + KL(ngpus, cross_device_reduce_1stage); \ + } else if (full_nvlink_) { \ + if ((world_size_ <= kAllReduceGPUSmall && bytes < kAllReduceSmallThreshold) || \ + (world_size_ <= kAllReduceGPULarge && bytes < kAllReduceLargeThreshold)) { \ + KL(ngpus, cross_device_reduce_1stage); \ + } else { \ + KL(ngpus, cross_device_reduce_2stage); \ + } \ + } \ + } \ + break; \ } - +#else +#define REDUCE_CASE(ngpus) \ + case ngpus: { \ + if constexpr (!std::is_same::value) { \ + custom_all_reduce_2shot<<>>(ptrs, sg_, self_sg_, output, rank_, size); \ + } else { \ + if ((world_size_ <= kAllReduceGPUSmall && bytes < kAllReduceSmallThreshold) || \ + (world_size_ <= kAllReduceGPULarge && bytes < kAllReduceLargeThreshold)) { \ + KL(ngpus, cross_device_reduce_1stage); \ + } else { \ + KL(ngpus, cross_device_reduce_2stage); \ + } \ + } \ + break; \ + } +#endif switch (world_size_) { REDUCE_CASE(2) REDUCE_CASE(4) diff --git a/sgl-kernel/csrc/common_extension_musa.cc b/sgl-kernel/csrc/common_extension_musa.cc index 00a83f5b5..4cb3a5f04 100644 --- a/sgl-kernel/csrc/common_extension_musa.cc +++ b/sgl-kernel/csrc/common_extension_musa.cc @@ -20,13 +20,251 @@ limitations under the License. #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); + + /* + * 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( - "min_p_sampling_from_probs(Tensor probs, Tensor output, Tensor? maybe_indices, Tensor? maybe_min_p_arr, float " - "min_p_val, bool deterministic, Generator? gen) -> ()"); - m.impl("min_p_sampling_from_probs", torch::kMUSA, &min_p_sampling_from_probs); + "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); @@ -34,15 +272,11 @@ TORCH_LIBRARY_EXPAND(sgl_kernel, m) { 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); - m.def( - "top_p_sampling_from_probs(Tensor probs, Tensor output, Tensor? maybe_indices, Tensor? " - "maybe_top_p_arr, float top_p_val, bool deterministic, Generator? gen) -> ()"); - m.impl("top_p_sampling_from_probs", torch::kMUSA, &top_p_sampling_from_probs); - - m.def( - "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("top_k_top_p_sampling_from_probs", torch::kMUSA, &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) diff --git a/sgl-kernel/csrc/elementwise/fused_add_rms_norm_kernel.mu b/sgl-kernel/csrc/elementwise/fused_add_rms_norm_kernel.mu new file mode 100644 index 000000000..c010a2868 --- /dev/null +++ b/sgl-kernel/csrc/elementwise/fused_add_rms_norm_kernel.mu @@ -0,0 +1,529 @@ +/* Copyright @2020-2026 Moore Threads Technology Co., Ltd("Moore Threads"). All + * rights reserved. + * + * This software ("this software and its documentations" or "the software") is + * protected by Copyright and the information contained herein is confidential. + * + * The software contained herein is PROPRIETARY to Moore Threads and is being + * provided under the terms and conditions of a form of Moore Threads software + * license agreement by and between Moore Threads and Licensee ("License + * Agreement") or electronically accepted by Licensee. Notwithstanding any + * terms or conditions to the contrary in the License Agreement, copy or + * disclosure of the software to any third party without the express written + * consent of Moore Threads is prohibited. + * + * NOTWITHSTANDING ANY TERMS OR CONDITIONS TO THE CONTRARY IN THE LICENSE + * AGREEMENT, MOORE THREADS MAKES NO REPRESENTATION ABOUT ANY WARRANTIES, + * INCLUDING BUT NOT LIMITED TO THE SUITABILITY OF THE SOFTWARE FOR ANY + * PURPOSE. IT IS PROVIDED "AS IS" WITHOUT EXPRESS OR IMPLIED WARRANTY OF + * ANY KIND. MOORE THREADS DISCLAIMS ALL WARRANTIES WITH REGARD TO THE + * SOFTWARE, INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY, + * NONINFRINGEMENT, AND FITNESS FOR A PARTICULAR PURPOSE. + * NOTWITHSTANDING ANY TERMS OR CONDITIONS TO THE CONTRARY IN THE + * LICENSE AGREEMENT, IN NO EVENT SHALL MOORE THREADS BE LIABLE FOR ANY + * SPECIAL, INDIRECT, INCIDENTAL, OR CONSEQUENTIAL DAMAGES, OR ANY + * DAMAGES WHATSOEVER RESULTING FROM LOSS OF USE, DATA OR PROFITS, + * WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR OTHER TORTIOUS + * ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR PERFORMANCE + * OF THE SOFTWARE. + */ + + #include "musa.h" + #include + #include + #include + #include + #include + #include "musa_bf16.h" + #include + #include + #include "torch_musa/csrc/core/MUSAGuard.h" + #include "torch_musa/csrc/core/MUSAStream.h" + + typedef __half float16_t; + typedef __mt_bfloat16 bfloat16_t; + + #define DEVICE_INLINE __device__ __forceinline__ + template + __device__ __forceinline__ T mudnn_shfl_down_sync(T val, unsigned int delta) { + return __shfl_down_sync(0xffffffff, val, delta, width); + } + + __device__ __host__ __forceinline__ constexpr int ceil_div(int a, int b) { + return (a + b - 1) / b; + } + + __device__ __host__ __forceinline__ constexpr int64_t ceil_div(int64_t a, + int64_t b) { + return (a + b - 1) / b; + } + + #define WARP_THREADS 32 + #define SMEM_STOP (WARP_THREADS / 2) + #define SHFL_START min(WARP_THREADS / 2, BLOCK_X / 2) + #define __SYNCTHREADS_LM __syncthreads_lm() + #define MACRO_UNROLL _Pragma("unroll") + #define LD_BYP_SLC(_BITS, _BYTES) \ + VecType dst; \ + const BaseType* addr = ptr + idx; \ + asm volatile("LSU.LD.B" #_BITS " %0, %1, _, " #_BYTES \ + ", 1, 1, inner_persist=0, " \ + "outer_persist=2, chrnt=l2_l3, slc=byp, persist=0, " \ + "stride_add_first=0" \ + : "=R"(dst) \ + : "R"(addr)); \ + return dst; + + #define ATTR_ALIGNED(v) __attribute__((aligned(v))) + #define SELF_VEC_DEF(BASE_TYPE, VEC_TYPE_V2, VEC_TYPE_V4) \ + struct ATTR_ALIGNED(sizeof(BASE_TYPE) * 2) VEC_TYPE_V2 { \ + __device__ VEC_TYPE_V2() {} \ + __device__ VEC_TYPE_V2(const VEC_TYPE_V2& t) { \ + this->x = t.x; \ + this->y = t.y; \ + } \ + BASE_TYPE x, y; \ + }; \ + \ + __device__ __forceinline__ VEC_TYPE_V2 make_##VEC_TYPE_V2(BASE_TYPE x, \ + BASE_TYPE y) { \ + VEC_TYPE_V2 t; \ + t.x = x, t.y = y; \ + return t; \ + } \ + \ + struct ATTR_ALIGNED(sizeof(BASE_TYPE) * 4) VEC_TYPE_V4 { \ + __device__ VEC_TYPE_V4() {} \ + __device__ VEC_TYPE_V4(const VEC_TYPE_V4& t) { \ + this->x = t.x; \ + this->y = t.y; \ + this->z = t.z; \ + this->w = t.w; \ + } \ + BASE_TYPE x, y, z, w; \ + }; \ + \ + __device__ __forceinline__ VEC_TYPE_V4 make_##VEC_TYPE_V4( \ + BASE_TYPE x, BASE_TYPE y, BASE_TYPE z, BASE_TYPE w) { \ + VEC_TYPE_V4 t; \ + t.x = x, t.y = y, t.z = z, t.w = w; \ + return t; \ + } + + SELF_VEC_DEF(float16_t, Half2, Half4) + SELF_VEC_DEF(bfloat16_t, Bhalf2, Bhalf4) + + #define GEN_VECTYPE(_CTYPE, _VECTYPE, _BYTES, _VLEN) \ + struct ATTR_ALIGNED(_BYTES) _VECTYPE { \ + __device__ _VECTYPE() {} \ + __device__ _VECTYPE(const _VECTYPE& t) { \ + MACRO_UNROLL \ + for (int i = 0; i < _VLEN; i++) { \ + this->arr[i] = t.arr[i]; \ + } \ + } \ + _CTYPE arr[_VLEN]; \ + } + + GEN_VECTYPE(float16_t, Half8, 16, 8); + GEN_VECTYPE(bfloat16_t, Bhalf8, 16, 8); + template + class Dtype; + + #define INST(_type, _vec2, _vec4) \ + template <> \ + class Dtype<_type> { \ + public: \ + using Scalar = _type; \ + using Vec2 = _vec2; \ + using Vec4 = _vec4; \ + static __device__ __forceinline__ Vec2 make_vec2(_type x, _type y) { \ + return make_##_vec2(x, y); \ + } \ + static __device__ __forceinline__ Vec4 make_vec4(_type x, _type y, \ + _type z, _type w) { \ + return make_##_vec4(x, y, z, w); \ + } \ + } + + INST(float, float2, float4); + INST(bfloat16_t, Bhalf2, Bhalf4); + + template + struct VecType; + + template + struct DeduceVectorizedType { + using Type = T; + }; + + template <> + struct DeduceVectorizedType { + using Type = _Float16; + }; + + template <> + struct DeduceVectorizedType { + using Type = _Float16; + }; + + #define DEF_VECT(_CTYPE, _VECTYPE) \ + template <> \ + struct VecType<_CTYPE, sizeof(_VECTYPE) * 8> { \ + static constexpr int vec_bytes = sizeof(_VECTYPE); \ + static constexpr int bit_per_byte = 8; \ + using BaseType = _CTYPE; \ + using RobustTypePtr = __musa::robust_ptr<_CTYPE>; \ + using Ttype = _VECTYPE; \ + static constexpr int bits = vec_bytes * bit_per_byte; \ + static constexpr int vlen = bits / (sizeof(BaseType) * bit_per_byte); \ + using VectorizedType = typename DeduceVectorizedType::Type; \ + typedef VectorizedType VxTy __attribute__((vector_size(vec_bytes))); \ + template \ + static __device__ __forceinline__ VecType load(const BaseType* ptr, \ + OffsetType idx) { \ + return *(VecType*)(ptr + idx); \ + } \ + template \ + static __device__ __forceinline__ VecType \ + load_byp_slc(const BaseType* ptr, OffsetType idx) { \ + if constexpr (vec_bytes == 16) { \ + LD_BYP_SLC(128, 16); \ + } else if constexpr (vec_bytes == 8) { \ + LD_BYP_SLC(64, 8); \ + } else if constexpr (vec_bytes == 4) { \ + LD_BYP_SLC(32, 4); \ + } else if constexpr (vec_bytes == 2) { \ + LD_BYP_SLC(16, 2); \ + } else { \ + LD_BYP_SLC(8, 1); \ + } \ + } \ + template \ + static __device__ __forceinline__ VecType \ + robust_load(const RobustTypePtr ptr, OffsetType idx) { \ + return __musa::robust_load(ptr, idx); \ + } \ + \ + template \ + static __device__ __forceinline__ void store(BaseType* ptr, \ + OffsetType idx, \ + const VecType& dst) { \ + *(VecType*)(ptr + idx) = dst; \ + } \ + template \ + static __device__ __forceinline__ void robust_store(RobustTypePtr ptr, \ + OffsetType idx, \ + const VecType& dst) { \ + __musa::robust_store(dst, ptr, idx); \ + } \ + \ + __device__ VecType() { \ + MACRO_UNROLL \ + for (int i = 0; i < sizeof(Ttype) / sizeof(BaseType); i++) { \ + this->val_.elem[i] = 0; \ + } \ + } \ + __device__ VecType(const VecType& t) { \ + MACRO_UNROLL \ + for (int i = 0; i < sizeof(Ttype) / sizeof(BaseType); i++) { \ + this->val_.elem[i] = t.val_.elem[i]; \ + } \ + } \ + __device__ VecType& operator=(const VecType& t) { \ + MACRO_UNROLL \ + for (int i = 0; i < sizeof(Ttype) / sizeof(BaseType); i++) { \ + this->val_.elem[i] = t.val_.elem[i]; \ + } \ + return *this; \ + } \ + __device__ VecType(_CTYPE val) { \ + MACRO_UNROLL \ + for (int i = 0; i < sizeof(Ttype) / sizeof(BaseType); i++) { \ + this->val_.elem[i] = val; \ + } \ + } \ + template \ + friend __device__ VecType operator+(VecType lhs, const SrcVecType& rhs) { \ + MACRO_UNROLL \ + for (int i = 0; i < sizeof(Ttype) / sizeof(BaseType); i++) { \ + lhs.val_.elem[i] += static_cast(rhs.val_.elem[i]); \ + } \ + return lhs; \ + } \ + friend __device__ VecType operator+(VecType lhs, const _CTYPE& rhs) { \ + MACRO_UNROLL \ + for (int i = 0; i < sizeof(Ttype) / sizeof(BaseType); i++) { \ + lhs.val_.elem[i] += rhs; \ + } \ + return lhs; \ + } \ + friend __device__ VecType operator-(VecType lhs, const VecType& rhs) { \ + MACRO_UNROLL \ + for (int i = 0; i < sizeof(Ttype) / sizeof(BaseType); i++) { \ + lhs.val_.elem[i] -= rhs.val_.elem[i]; \ + } \ + return lhs; \ + } \ + friend __device__ VecType operator*(VecType lhs, const VecType& rhs) { \ + MACRO_UNROLL \ + for (int i = 0; i < sizeof(Ttype) / sizeof(BaseType); i++) { \ + lhs.val_.elem[i] *= rhs.val_.elem[i]; \ + } \ + return lhs; \ + } \ + template \ + __device__ VecType& apply() { \ + MACRO_UNROLL \ + for (int i = 0; i < sizeof(Ttype) / sizeof(BaseType); i++) { \ + this->val_.elem[i] = Func::apply(this->val_.elem[i]); \ + } \ + return *this; \ + } \ + template \ + static __device__ VecType cvt(const SrcVecType& src) { \ + VecType dst; \ + MACRO_UNROLL \ + for (int i = 0; i < sizeof(Ttype) / sizeof(BaseType); i++) { \ + dst.val_.elem[i] = (BaseType)(src.val_.elem[i]); \ + } \ + return dst; \ + } \ + union U { \ + __device__ U() { \ + MACRO_UNROLL \ + for (int i = 0; i < sizeof(Ttype) / sizeof(BaseType); i++) { \ + this->elem[i] = 0; \ + } \ + } \ + Ttype storage; \ + BaseType elem[sizeof(Ttype) / sizeof(BaseType)]; \ + VxTy vt_elem; \ + }; \ + U val_{}; \ + } + DEF_VECT(float16_t, float16_t); + DEF_VECT(float16_t, Half2); + DEF_VECT(bfloat16_t, bfloat16_t); + DEF_VECT(bfloat16_t, Bhalf2); + DEF_VECT(bfloat16_t, Bhalf8); + DEF_VECT(float16_t, Half8); + DEF_VECT(float, float4); + + enum class VarUpdateMode { WELFORD, WELFORD_ONLY_MEAN, CHAN, CHAN_ONLY_MEAN }; + + static __device__ __forceinline__ float fast_rcpf(float x) { + float y = __frcp_rn(x); + y = y * (2.0 - x * y); + return y; + } + + static __device__ __forceinline__ float fast_divf(float a, float b) { + return a * fast_rcpf(b); + } + + static __device__ __forceinline__ float fast_rsqrtf(float a) { + float x = 0.5 * a; + float y = __frsqrt_rn(a); + y = y * (1.5 - x * y * y); + return y; + } + + template + struct VarUpdate; + + template + struct VarUpdate { + DEVICE_INLINE void apply(T curr, T* mu, T* cnt) { + *cnt += 1; + T delta = curr - *mu; + *mu += fast_divf(delta, *cnt); + } + }; + + template + struct VarUpdate { + DEVICE_INLINE void apply(T mu_B, T cnt_B, T* mu, T* cnt) { + if (cnt_B > 0) { + T n_AB = cnt_B + (*cnt); + T delta = mu_B - (*mu); + *mu += delta * fast_divf(cnt_B, n_AB); + *cnt = n_AB; + } + } + }; + + template + struct AllReduceOp { + DEVICE_INLINE void apply(ComputeType* sum, int tx, int ty) { + __shared__ ComputeType __attribute__((aligned(16))) + smem[BLOCK_X * BLOCK_Y * Vlen]; + ComputeType* smem_sum = &smem[0]; + + static_assert(Vlen == 1, + "Axis COLUMN doesn't support vlen greater than 1"); + #pragma unroll + for (int offset = BLOCK_X / 2; offset > SMEM_STOP; offset /= 2) { + if (tx >= offset && tx < 2 * offset) { + smem_sum[ty * BLOCK_X + tx] = *sum; + } + __SYNCTHREADS_LM; + if (tx < offset) { + *sum += smem_sum[ty * BLOCK_X + tx + offset]; + } + } + #if ((defined __MUSA_ARCH__) && (__MUSA_ARCH__ >= 220)) + #pragma unroll + for (int offset = SHFL_START; offset > 0; offset /= 2) { + *sum += mudnn_shfl_down_sync(*sum, offset); + } + #endif + if (tx == 0) { + smem_sum[ty * BLOCK_X + tx] = *sum; + } + __SYNCTHREADS_LM; + *sum = smem_sum[ty * BLOCK_X]; + } + }; + + template + __global__ void LayerNormGlobalKernelVlen( + SrcDtype* input, SrcDtype* residual, const SrcDtype* weight, + const size_t M, const size_t N, const ComputeType eps) { + size_t tx = threadIdx.x; + size_t ty = threadIdx.y; + size_t m_idx = blockIdx.x * blockDim.y + ty; + size_t n_idx = tx * vlen; + size_t n_step = (size_t)blockDim.x * vlen; + + using SrcVec = VecType; + + ComputeType var = 0; + const SrcDtype* __restrict p_src = input + m_idx * N; + SrcDtype* __restrict p_res = residual + m_idx * N; // residual ptr + + // TODO(wuke): use robust_load, robust_store + bool m_valid = m_idx < M; + if (m_valid) { + for (size_t j = n_idx; j < N; j += n_step) { + SrcVec curr, res_vec, fused_vec; + #if ((defined __MUSA_ARCH__) && (__MUSA_ARCH__ == 220)) + curr = *(SrcVec *)(p_src+j); + res_vec = *(SrcVec *)(p_res+j); + #elif ((defined __MUSA_ARCH__) && (__MUSA_ARCH__ == 310)) + curr = SrcVec::load_byp_slc(p_src, j); + res_vec = SrcVec::load_byp_slc(p_res, j); + #endif + #pragma unroll + for (int k = 0; k < vlen; k++) { + fused_vec.val_.elem[k] = curr.val_.elem[k] + res_vec.val_.elem[k]; + var += (ComputeType)fused_vec.val_.elem[k] * (ComputeType)fused_vec.val_.elem[k]; + } + *(SrcVec*)(p_res + j) = fused_vec; + } + } + AllReduceOp all_reduce_op; + all_reduce_op.apply(&var, tx, ty); + if (m_valid) { + ComputeType inv_var = fast_rsqrtf(var / N + eps); + SrcDtype* __restrict p_dst = input + m_idx * N; + bool with_weight = (weight != NULL); + if (with_weight) { + for (size_t j = n_idx; j < N; j += n_step) { + SrcVec fused_vec, weight_val, dst; + #if ((defined __MUSA_ARCH__) && (__MUSA_ARCH__ == 220)) + fused_vec = *(SrcVec *)(p_res+j); + weight_val = *(SrcVec *)(weight+j); + #elif ((defined __MUSA_ARCH__) && (__MUSA_ARCH__ == 310)) + fused_vec = SrcVec::load_byp_slc(p_res, j); + weight_val = SrcVec::load_byp_slc(weight, j); + #endif + #pragma unroll + for (int k = 0; k < vlen; k++) { + dst.val_.elem[k] = (SrcDtype)((ComputeType)fused_vec.val_.elem[k] * inv_var * + (ComputeType)weight_val.val_.elem[k]); + } + *(SrcVec*)(p_dst + j) = dst; + } + } + } + } + + #define CALL_KERN(_SRC_DTYPE,_KERN, _BLKX, _BLKY, _VLEN) \ + { \ + const uint32_t block_x = _BLKX; \ + const uint32_t block_y = _BLKY; \ + const uint32_t nr_blocks = ceil_div(m, (size_t)block_y); \ + dim3 block_size{block_x, block_y, 1}; \ + dim3 grid_size{nr_blocks, 1, 1}; \ + LayerNorm##_KERN##KernelVlen<_SRC_DTYPE, float, \ + block_x, block_y, _VLEN> \ + <<>>( \ + static_cast<_SRC_DTYPE*>(input), \ + static_cast<_SRC_DTYPE*>(residual), \ + static_cast<_SRC_DTYPE*>(weight), \ + m, n, static_cast(epsilon)); \ + } + + #define DISPATCH_KERNEL(_KERN, _BLKX, _BLKY) \ + if constexpr (std::is_same_v) { \ + CALL_KERN(float16_t, _KERN, _BLKX, _BLKY, 8); \ + } else if constexpr (std::is_same_v) { \ + CALL_KERN(bfloat16_t, _KERN, _BLKX, _BLKY, 8); \ + } else if constexpr (std::is_same_v) { \ + CALL_KERN(float, _KERN, _BLKX, _BLKY, 4); \ + } + + template + void rms_fused_add_rms_norm(SrcDtype* input, SrcDtype* residual, SrcDtype* weight, int m, int n, double epsilon) { + auto stream = c10::musa::getCurrentMUSAStream().stream(); + DISPATCH_KERNEL(Global, 1024, 1); + } + +void musa_fused_add_rms_norm( + torch::Tensor &input, + torch::Tensor &residual, + torch::Tensor &weight, + double epsilon, + bool enable_pdl) { + + int m = input.size(0); + int n = input.size(1); + + const at::musa::OptionalMUSAGuard device_guard(device_of(input)); + + if (input.scalar_type() == at::ScalarType::BFloat16) + { + rms_fused_add_rms_norm<__mt_bfloat16>( + static_cast<__mt_bfloat16*>(input.data_ptr()), + static_cast<__mt_bfloat16*>(residual.data_ptr()), + static_cast<__mt_bfloat16*>(weight.data_ptr()), + m, + n, + epsilon); + } + else if (input.scalar_type() == at::ScalarType::Half) + { + rms_fused_add_rms_norm<__half>( + static_cast<__half*>(input.data_ptr()), + static_cast<__half*>(residual.data_ptr()), + static_cast<__half*>(weight.data_ptr()), + m, + n, + epsilon); + } + else if (input.scalar_type() == at::ScalarType::Float) + { + rms_fused_add_rms_norm( + static_cast(input.data_ptr()), + static_cast(residual.data_ptr()), + static_cast(weight.data_ptr()), + m, + n, + epsilon); + } + else + { + TORCH_CHECK(false, "only support Float32, Half and BFloat16 dtype"); + } +} diff --git a/sgl-kernel/csrc/elementwise/utils.cuh b/sgl-kernel/csrc/elementwise/utils.cuh index 3010a5452..cf46734b9 100644 --- a/sgl-kernel/csrc/elementwise/utils.cuh +++ b/sgl-kernel/csrc/elementwise/utils.cuh @@ -7,11 +7,18 @@ #include +#ifndef USE_MUSA __forceinline__ __device__ int get_lane_id() { int lane_id; asm("mov.s32 %0, %laneid;" : "=r"(lane_id)); return lane_id; } +#else +constexpr int WarpSize = 32; +__forceinline__ __device__ int get_lane_id() { + return threadIdx.x % WarpSize; +} +#endif int ceil_div(int a, int b) { return (a + b - 1) / b; diff --git a/sgl-kernel/csrc/gemm/dsv3_fused_a_gemm.cu b/sgl-kernel/csrc/gemm/dsv3_fused_a_gemm.cu index 37aff1b9a..4dc0a796a 100644 --- a/sgl-kernel/csrc/gemm/dsv3_fused_a_gemm.cu +++ b/sgl-kernel/csrc/gemm/dsv3_fused_a_gemm.cu @@ -652,7 +652,11 @@ void dsv3_fused_a_gemm(torch::Tensor& output, torch::Tensor const& mat_a, torch: TORCH_CHECK(output.scalar_type() == torch::kBFloat16, "Only BFloat16 output dtype is supported") auto const sm = getSMVersion(); +#ifndef USE_MUSA TORCH_CHECK(sm >= 90, "required CUDA ARCH >= SM_90"); +#else + TORCH_CHECK(sm >= 22, "required MUSA ARCH >= MP_22"); +#endif auto stream = at::cuda::getCurrentCUDAStream(mat_a.get_device()); if (num_tokens <= 8) { diff --git a/sgl-kernel/csrc/gemm/dsv3_router_gemm_entry.cu b/sgl-kernel/csrc/gemm/dsv3_router_gemm_entry.cu index 4f09e6cf4..a3b6b272c 100644 --- a/sgl-kernel/csrc/gemm/dsv3_router_gemm_entry.cu +++ b/sgl-kernel/csrc/gemm/dsv3_router_gemm_entry.cu @@ -121,7 +121,11 @@ void dsv3_router_gemm( output.dtype() == torch::kFloat32 || output.dtype() == torch::kBFloat16, "output must be float32 or bf16"); auto const sm = getSMVersion(); +#ifndef USE_MUSA TORCH_CHECK(sm >= 90, "required CUDA ARCH >= SM_90"); +#else + TORCH_CHECK(sm >= 22, "required MUSA ARCH >= MP_22"); +#endif const cudaStream_t stream = at::cuda::getCurrentCUDAStream(); diff --git a/sgl-kernel/csrc/gemm/per_token_group_quant_8bit_v2.cu b/sgl-kernel/csrc/gemm/per_token_group_quant_8bit_v2.cu index dea405663..c886569f9 100644 --- a/sgl-kernel/csrc/gemm/per_token_group_quant_8bit_v2.cu +++ b/sgl-kernel/csrc/gemm/per_token_group_quant_8bit_v2.cu @@ -90,16 +90,25 @@ __forceinline__ __device__ OUT_DTYPE_T extract_required_scale_format(float value } __device__ __forceinline__ void st_global(const int4* ptr, const int4& value) { +#ifndef USE_MUSA asm volatile( "st.global.v4.s32 [%0], {%1, %2, %3, %4};" ::"l"(ptr), "r"(value.x), "r"(value.y), "r"(value.z), "r"(value.w)); +#else + int4* p = const_cast(ptr); + *p = value; +#endif } __device__ __forceinline__ int4 ld_global_nc(const int4* ptr) { +#ifndef USE_MUSA int4 ret; asm volatile("ld.global.nc.v4.s32 {%0, %1, %2, %3}, [%4];" : "=r"(ret.x), "=r"(ret.y), "=r"(ret.z), "=r"(ret.w) : "l"(ptr)); return ret; +#else + return *ptr; +#endif } template diff --git a/sgl-kernel/csrc/kvcacheio/transfer.cu b/sgl-kernel/csrc/kvcacheio/transfer.cu index 230932c15..b4372a956 100644 --- a/sgl-kernel/csrc/kvcacheio/transfer.cu +++ b/sgl-kernel/csrc/kvcacheio/transfer.cu @@ -7,7 +7,7 @@ #include #include -#ifndef USE_ROCM +#if !defined(USE_ROCM) && !defined(USE_MUSA) #include #define WARP_SIZE 32 #include "pytorch_extension_utils.h" @@ -24,7 +24,7 @@ transfer_item_warp(int32_t lane_id, const void* src_addr, void* dst_addr, int64_ #pragma unroll for (int j = lane_id; j < total_chunks; j += WARP_SIZE) { -#ifndef USE_ROCM +#if !defined(USE_ROCM) && !defined(USE_MUSA) uint64_t tmp; asm volatile("ld.global.nc.b64 %0,[%1];" : "=l"(tmp) : "l"(src + j) : "memory"); asm volatile("st.global.cg.b64 [%0],%1;" ::"l"(dst + j), "l"(tmp) : "memory"); diff --git a/sgl-kernel/csrc/moe/moe_fused_gate_musa.cu b/sgl-kernel/csrc/moe/moe_fused_gate_musa.cu new file mode 100644 index 000000000..facc66e9e --- /dev/null +++ b/sgl-kernel/csrc/moe/moe_fused_gate_musa.cu @@ -0,0 +1,840 @@ +#include +#include +#include +#include +#include +#include + +#include +#include + +#include "torch_musa/csrc/aten/musa/MUSAContext.h" +template +using AlignedArray = mutlass::AlignedArray; +using bfloat16_t = mutlass::bfloat16_t; +using float16_t = mutlass::half_t; +using float32_t = float; + +constexpr float log2ef = 1.4426950408889634074f; + +static __device__ __forceinline__ float fast_expf(float a) { + return __musa_exp2_f(a * log2ef); +} + +static __device__ __forceinline__ float fast_rcpf(float x) { + float y = __frcp_rn(x); + y = y * (2.f - x * y); + return y; +} + +// QQ NOTE: to handle the case for at::Half, error: more than one operator ">" +// matches these operands: built-in operator "arithmetic > arithmetic" function +// "operator>(const __half &, const __half &)" +template +__device__ inline bool cmp_gt(const T& a, const T& b) { + if constexpr (std::is_same::value) { + // at::Half (or float16_t in our native case) causes ambiguity, so we cast + // to float. + return static_cast(a) > static_cast(b); + } else { + // For types like float, at::BFloat16, or mutlass::half_t / + // mutlass::bfloat16_t, assume operator> works as expected. + return a > b; + } +} + +template +__device__ inline bool cmp_eq(const T& a, const T& b) { + if constexpr (std::is_same::value) { + return static_cast(a) == static_cast(b); + } else { + return a == b; + } +} + +template +__device__ inline bool cmp_ge(const T& a, const T& b, const int& x, const int& y) { + return (x > y && a == b) || a < b; +} + +// Fixed constants common to both dynamic and static template versions: +static constexpr int WARP_SIZE = 32; +static constexpr int WARPS_PER_CTA = 16; +static constexpr int MAX_VPT = 32; // maximum VPT we support, > params.VPT = num_expert / num_expert_group + +// Create an alias for Array using AlignedArray +template +using Array = AlignedArray; +// QQ: NOTE expression must have a constant value, this has to be > params.VPT +template +using AccessType = AlignedArray; + +template +__device__ void moe_fused_gate_impl_dynamic( + void* input, + void* bias, + float* output_ptr, + int32_t* indices_ptr, + int64_t num_rows, + int64_t topk_group, + int64_t topk, + int64_t num_fused_shared_experts, + double routed_scaling_factor, + bool apply_routed_scaling_factor_on_output, + Params params) { + int tidx = threadIdx.x; + int64_t thread_row = + blockIdx.x * params.ROWS_PER_CTA + threadIdx.y * params.ROWS_PER_WARP + tidx / params.THREADS_PER_ROW; + // Calculate topk_excluding_share_expert_fusion from topk + int64_t topk_excluding_share_expert_fusion = topk - num_fused_shared_experts; + + // Cast pointers to type T: + auto* input_ptr = reinterpret_cast(input); + auto* bias_ptr = reinterpret_cast(bias); + auto* thread_row_ptr = input_ptr + thread_row * params.NUM_EXPERTS; + + int thread_group_idx = tidx % params.THREADS_PER_ROW; + int first_elt_read_by_thread = thread_group_idx * params.VPT; + + // Create local arrays for the row chunk and bias chunk and then reinterpret + // the address of row_chunk as a pointer to AccessType. + T* thread_read_ptr = thread_row_ptr + first_elt_read_by_thread; + Array row_chunk; + AccessType const* vec_thread_read_ptr = reinterpret_cast const*>(thread_read_ptr); + + T* bias_thread_read_ptr = bias_ptr + first_elt_read_by_thread; + Array bias_chunk; + AccessType const* vec_bias_thread_read_ptr = reinterpret_cast const*>(bias_thread_read_ptr); + + // QQ NOTE: doing the follow will be slower than loop assign and more + // importantly have misaligned address issue when params.VPT < 8 and mismatch + // with MAX_VPT AccessType* row_chunk_vec_ptr = + // reinterpret_cast*>(&row_chunk); row_chunk_vec_ptr[0] = + // vec_thread_read_ptr[0]; + if (thread_row < num_rows) { +#pragma unroll + for (int ii = 0; ii < params.VPT; ++ii) { + row_chunk[ii] = vec_thread_read_ptr[0][ii]; + bias_chunk[ii] = vec_bias_thread_read_ptr[0][ii]; + } + } + + ////////////////////// Sigmoid ////////////////////// + if (thread_row < num_rows) { +#pragma unroll + for (int ii = 0; ii < params.VPT; ++ii) { + row_chunk[ii] = static_cast(fast_rcpf(1.0f + fast_expf(-float(row_chunk[ii])))); + } + } + + ////////////////////// Add Bias ////////////////////// + if (thread_row < num_rows) { +#pragma unroll + for (int ii = 0; ii < params.VPT; ++ii) { + bias_chunk[ii] = row_chunk[ii] + bias_chunk[ii]; + } + } + + ////////////////////// Exclude Groups ////////////////////// + if (thread_row < num_rows) { +#pragma unroll + for (int k_idx = 0; k_idx < params.THREADS_PER_ROW - topk_group; + ++k_idx) { // QQ NOTE Here params.THREADS_PER_ROW = num_expert_group + int expert = first_elt_read_by_thread; + // local argmax + T max_val = static_cast(-FLT_MAX); + T max_val_second = static_cast(-FLT_MAX); +#pragma unroll + for (int ii = 0; ii < params.VPT; ++ii) { + T val = bias_chunk[ii]; + + if (cmp_gt(val, max_val)) { + max_val_second = max_val; + max_val = val; + } else if (cmp_gt(val, max_val_second)) { + max_val_second = val; + } + } + + // QQ NOTE: currently fixed to pick top2 sigmoid weight value in each + // expert group and sum them as the group weight to select expert groups + T max_sum = max_val + max_val_second; + +// argmin reduce +#pragma unroll + for (int mask = params.THREADS_PER_ROW / 2; mask > 0; mask /= 2) { + T other_max_sum = + static_cast(__shfl_xor_sync(0xFFFFFFFF, static_cast(max_sum), mask, params.THREADS_PER_ROW)); + int other_expert = __shfl_xor_sync(0xFFFFFFFF, expert, mask, params.THREADS_PER_ROW); + + // higher indices win + if (cmp_gt(max_sum, other_max_sum) || (cmp_eq(other_max_sum, max_sum) && other_expert > expert)) { + max_sum = other_max_sum; + expert = other_expert; + } + } + + // clear the max value in the thread + if (k_idx < params.THREADS_PER_ROW - topk_group) { + int const thread_to_clear_in_group = expert / params.VPT; + + if (thread_group_idx == thread_to_clear_in_group) { +#pragma unroll + for (int ii = 0; ii < params.VPT; ++ii) { + bias_chunk[ii] = static_cast(FLT_MAX); + } + } + } + } + } + + ////////////////////// Topk ////////////////////// + float output_sum = 0.0f; + for (int k_idx = 0; k_idx < topk_excluding_share_expert_fusion; ++k_idx) { + if (thread_row < num_rows) { + // local argmax + T max_val = bias_chunk[0]; + int expert = first_elt_read_by_thread; + + if (!cmp_eq(max_val, static_cast(FLT_MAX))) { +#pragma unroll + for (int ii = 1; ii < params.VPT; ++ii) { + T val = bias_chunk[ii]; + if (cmp_gt(val, max_val)) { + max_val = val; + expert = first_elt_read_by_thread + ii; + } + } + } else { + max_val = static_cast(-FLT_MAX); + } + + // argmax reduce +#pragma unroll + for (int mask = params.THREADS_PER_ROW / 2; mask > 0; mask /= 2) { + T other_max = + static_cast(__shfl_xor_sync(0xFFFFFFFF, static_cast(max_val), mask, params.THREADS_PER_ROW)); + int other_expert = __shfl_xor_sync(0xFFFFFFFF, expert, mask, params.THREADS_PER_ROW); + + // lower indices to win + if (cmp_gt(other_max, max_val) || (cmp_eq(other_max, max_val) && other_expert < expert)) { + max_val = other_max; + expert = other_expert; + } + } + + int thread_to_clear_in_group = expert / params.VPT; + int64_t idx = topk * thread_row + k_idx; + + if (thread_group_idx == thread_to_clear_in_group) { + int expert_to_clear_in_thread = expert % params.VPT; + +#pragma unroll + for (int v = 0; v < MAX_VPT; v++) { + if (v < params.VPT && expert_to_clear_in_thread == v) { + // clear the max value in the thread + bias_chunk[v] = static_cast(-FLT_MAX); + // store output + output_ptr[idx] = static_cast(row_chunk[v]); + } + } + indices_ptr[idx] = static_cast(expert); + } + + __threadfence_block(); + // accumulate sum for all elements + if (thread_group_idx == 0) { + output_sum += output_ptr[idx]; + } + } + } + + if (thread_row < num_rows) { + if (thread_group_idx == 0 && num_fused_shared_experts > 0) { + int64_t last_idx = topk * thread_row + topk_excluding_share_expert_fusion; + int64_t expert_offset = 0; + indices_ptr[last_idx] = static_cast(params.NUM_EXPERTS + expert_offset); + + // Set the weight to the sum of all weights divided by + // routed_scaling_factor + output_ptr[last_idx] = output_sum / routed_scaling_factor; + + if (num_fused_shared_experts > 1) { + for (int i = 1; i < num_fused_shared_experts; ++i) { + ++last_idx; + ++expert_offset; + indices_ptr[last_idx] = static_cast(params.NUM_EXPERTS + expert_offset); + // Set the weight to the sum of all weights divided by + // routed_scaling_factor + output_ptr[last_idx] = output_sum / routed_scaling_factor; + } + } + } + } + __threadfence_block(); + + ////////////////////// Rescale Output ////////////////////// + if (thread_row < num_rows) { + if (thread_group_idx == 0) { +#pragma unroll + for (int ii = 0; ii < topk; ++ii) { + int64_t const idx = topk * thread_row + ii; + output_ptr[idx] = output_ptr[idx] / output_sum; + if (apply_routed_scaling_factor_on_output) { + output_ptr[idx] *= routed_scaling_factor; + } + } + } + } +} + +template +__device__ void moe_fused_gate_impl_static( + void* input, + void* bias, + float* output_ptr, + int32_t* indices_ptr, + int64_t num_rows, + int64_t topk_group, + int64_t topk, + int64_t num_fused_shared_experts, + double routed_scaling_factor, + bool apply_routed_scaling_factor_on_output, + float last_val, + Params params) { + using ArrayVal = AlignedArray; + using ArrayIndex = AlignedArray; + + int tidx = threadIdx.x % (params.NUM_EXPERTS / Vlen) * Vlen; + int tidy = threadIdx.x / (params.NUM_EXPERTS / Vlen); + int64_t thread_row = blockIdx.x * params.ROWS_PER_CTA + tidy; + + constexpr int NR_EXPERTS = params.NUM_EXPERTS; + constexpr int NR_ROWS_PER_CTA = params.ROWS_PER_CTA; + constexpr int NR_EXPERT_GRPS = params.NUM_EXPERTS / params.VPT; + constexpr int NR_EXPERT_PER_GRP = params.VPT; + constexpr int NR_THREADS_PER_GRP = NR_EXPERT_PER_GRP / Vlen; + __shared__ int smem_grp_flag[NR_ROWS_PER_CTA * NR_EXPERT_GRPS]; + __shared__ float smem_grp_max_sum[NR_ROWS_PER_CTA * NR_EXPERT_GRPS]; + __shared__ T smem_score[NR_ROWS_PER_CTA * NR_EXPERTS]; + __shared__ int smem_idx[NR_ROWS_PER_CTA * NR_EXPERTS]; + __shared__ T smem_bias[NR_EXPERTS]; + + static_assert(Vlen <= NR_EXPERT_PER_GRP); + + // Calculate topk_excluding_share_expert_fusion from topk + int topk_excluding_share_expert_fusion = topk - num_fused_shared_experts; + + // Cast pointers to type T: + auto* input_ptr = reinterpret_cast(input); + auto* bias_ptr = reinterpret_cast(bias); + auto* thread_row_ptr = input_ptr + thread_row * params.NUM_EXPERTS; + + int grp_idx = tidx / NR_EXPERT_PER_GRP; + int exp_idx_in_grp = tidx % NR_EXPERT_PER_GRP; + + ArrayVal row_chunk; + ArrayVal bias_chunk; + ArrayIndex idx_chunk; + if (thread_row < num_rows) { + row_chunk = *(ArrayVal*)(thread_row_ptr + tidx); + bias_chunk = *(ArrayVal*)(bias_ptr + tidx); + } + +#pragma unroll + for (int v = 0; v < Vlen; v++) { + ////////////////////// Sigmoid ////////////////////// + row_chunk[v] = static_cast(fast_rcpf(1.0f + fast_expf(-float(row_chunk[v])))); + if (tidy == 0) { + smem_bias[tidx + v] = bias_chunk[v]; + } + bias_chunk[v] = row_chunk[v] + bias_chunk[v]; + idx_chunk[v] = tidx + v; + } + + int max_idx = exp_idx_in_grp; + T max_val = bias_chunk[0]; + float max_sum = 0.f; + + ////////////////////// top 1 ////////////////////// +#pragma unroll + for (int v = 1; v < Vlen; v++) { + // per-thread max + if (bias_chunk[v] > max_val) { + max_val = bias_chunk[v]; + max_idx = exp_idx_in_grp + v; + } + } +#pragma unroll + for (int mask = NR_THREADS_PER_GRP / 2; mask > 0; mask /= 2) { + T peer_max_val = static_cast(__shfl_xor_sync(0xFFFFFFFF, static_cast(max_val), mask, NR_THREADS_PER_GRP)); + int peer_idx = __shfl_xor_sync(0xFFFFFFFF, max_idx, mask, NR_THREADS_PER_GRP); + if (cmp_gt(peer_max_val, max_val)) { + max_val = peer_max_val; + max_idx = peer_idx; + } + } + int top1_max_idx = __shfl_sync(0xFFFFFFFF, static_cast(max_idx), 0, NR_THREADS_PER_GRP); + max_sum += max_val; + + ////////////////////// top 2 ////////////////////// + max_val = static_cast(-FLT_MAX); + for (int v = 0; v < Vlen; v++) { + // per-thread reset + if (bias_chunk[v] > max_val && exp_idx_in_grp + v != top1_max_idx) { + max_val = bias_chunk[v]; + } + } +#pragma unroll + for (int mask = NR_THREADS_PER_GRP / 2; mask > 0; mask /= 2) { + T peer_max_val = static_cast(__shfl_xor_sync(0xFFFFFFFF, static_cast(max_val), mask, NR_THREADS_PER_GRP)); + if (cmp_gt(peer_max_val, max_val)) { + max_val = peer_max_val; + } + } + max_sum += max_val; + + ////////////////////// sort groups by max_sum ////////////////////// + if (exp_idx_in_grp == 0) { + smem_grp_max_sum[tidy * NR_EXPERT_GRPS + grp_idx] = max_sum; + smem_grp_flag[tidy * NR_EXPERT_GRPS + grp_idx] = grp_idx; + } + __syncthreads_lm(); + int cur_grp_rank = 0; + if (exp_idx_in_grp == 0) { + float cur_grp_max = max_sum; +#pragma unroll + for (int i = 0; i < NR_EXPERT_GRPS; i++) { + float other_grp_max = smem_grp_max_sum[tidy * NR_EXPERT_GRPS + i]; + int other_grp_idx = smem_grp_flag[tidy * NR_EXPERT_GRPS + i]; + if (cmp_ge(cur_grp_max, other_grp_max, grp_idx, other_grp_idx)) { + cur_grp_rank++; + } + } + } + __syncthreads_lm(); + if (exp_idx_in_grp == 0) { + smem_grp_flag[tidy * NR_EXPERT_GRPS + grp_idx] = cur_grp_rank; + } + __syncthreads_lm(); + + ////////////////////// TopK experts ////////////////////// + cur_grp_rank = smem_grp_flag[tidy * NR_EXPERT_GRPS + grp_idx]; + +#pragma unroll + for (int v = 0; v < Vlen; v++) { + if (cur_grp_rank >= topk_group) { + bias_chunk[v] = static_cast(-FLT_MAX); + } + } + + float output_sum = 0.f; + for (int i = 0; i < topk_excluding_share_expert_fusion; i++) { + T thread_max_val = static_cast(-FLT_MAX); + int thread_max_idx = idx_chunk[0]; +#pragma unroll + for (int v = 0; v < Vlen; v++) { + if (bias_chunk[v] > thread_max_val) { + thread_max_val = bias_chunk[v]; + thread_max_idx = idx_chunk[v]; + } + } + +#pragma unroll + for (int mask = WARP_SIZE / 2; mask > 0; mask /= 2) { + T peer_max_val = static_cast(__shfl_xor_sync(0xFFFFFFFF, static_cast(thread_max_val), mask, WARP_SIZE)); + int peer_idx = __shfl_xor_sync(0xFFFFFFFF, thread_max_idx, mask, WARP_SIZE); + if (cmp_ge(thread_max_val, peer_max_val, thread_max_idx, peer_idx)) { + thread_max_val = peer_max_val; + thread_max_idx = peer_idx; + } + } + int warp_max_idx = __shfl_sync(0xFFFFFFFF, thread_max_idx, 0, WARP_SIZE); + + if (tidx == 0) { + // restore row_chunk + float restored_val = (float)thread_max_val - (float)smem_bias[thread_max_idx]; + output_sum += restored_val; + smem_score[tidy * NR_EXPERTS + i] = (T)restored_val; + smem_idx[tidy * NR_EXPERTS + i] = thread_max_idx; + } + +#pragma unroll + for (int v = 0; v < Vlen; v++) { + if (warp_max_idx == idx_chunk[v]) { + bias_chunk[v] = static_cast(-FLT_MAX); + } + } + } + + __syncthreads_lm(); + output_sum = __shfl_sync(0xFFFFFFFF, output_sum, 0, WARP_SIZE); + + ////////////////////// store output ////////////////////// + int64_t out_idx = thread_row * topk; + int tid_st_x = threadIdx.x % WARP_SIZE; + if (thread_row < num_rows) { + for (int i = tid_st_x; i < topk_excluding_share_expert_fusion; i += WARP_SIZE) { + float output_val = smem_score[tidy * NR_EXPERTS + i] * fast_rcpf(output_sum); + if (apply_routed_scaling_factor_on_output) { + output_val *= routed_scaling_factor; + } + output_ptr[out_idx + i] = output_val; + indices_ptr[out_idx + i] = smem_idx[tidy * NR_EXPERTS + i]; + } + } + + ////////////////////// handle shared experts ////////////////////// + if (thread_row < num_rows && tidx == 0 && num_fused_shared_experts > 0) { + int64_t last_idx = thread_row * topk + topk_excluding_share_expert_fusion; + int64_t expert_offset = 0; + // Set the weight to the sum of all weights divided by routed_scaling_factor + indices_ptr[last_idx] = static_cast(NR_EXPERTS + expert_offset); + output_ptr[last_idx] = last_val; + + if (num_fused_shared_experts > 1) { + for (int i = 1; i < num_fused_shared_experts; ++i) { + ++last_idx; + ++expert_offset; + indices_ptr[last_idx] = static_cast(NR_EXPERTS + expert_offset); + output_ptr[last_idx] = last_val; + } + } + } +} + +//------------------------------------------------------------------------------ +// Templated Kernel Version (using compile-time constants) +//------------------------------------------------------------------------------ +template +struct KernelParams { + static constexpr int VPT = VPT_; + static constexpr int NUM_EXPERTS = NUM_EXPERTS_; + static constexpr int ROWS_PER_CTA = ROWS_PER_CTA_; +}; + +template +__global__ void moe_fused_gate_kernel_static( + void* input, + void* bias, + float* output_ptr, + int32_t* indices_ptr, + int64_t num_rows, + int64_t topk_group, + int64_t topk, + int64_t num_fused_shared_experts, + double routed_scaling_factor, + bool apply_routed_scaling_factor_on_output, + float last_val) { + KernelParams params; + moe_fused_gate_impl_static, Vlen>( + input, + bias, + output_ptr, + indices_ptr, + num_rows, + topk_group, + topk, + num_fused_shared_experts, + routed_scaling_factor, + apply_routed_scaling_factor_on_output, + last_val, + params); +} + +// Macro to compute compile-time constants and launch the kernel. +#define LAUNCH_MOE_GATE_CONFIG(T, EXPERTS, EXPERT_GROUP) \ + do { \ + constexpr int vlen = EXPERTS / WARP_SIZE; \ + int block_x = num_experts / vlen; \ + int block_y = block_size / block_x; \ + int64_t num_blocks = (num_rows + block_y - 1) / block_y; \ + dim3 block_dim(block_size, 1, 1); \ + constexpr int VPT = (EXPERTS) / (EXPERT_GROUP); \ + constexpr int ROWS_PER_CTA = block_size / (EXPERTS / vlen); \ + moe_fused_gate_kernel_static<<>>( \ + input.data_ptr(), \ + bias.data_ptr(), \ + output.data_ptr(), \ + indices.data_ptr(), \ + num_rows, \ + topk_group, \ + topk, \ + num_fused_shared_experts, \ + routed_scaling_factor, \ + apply_routed_scaling_factor_on_output, \ + last_val); \ + dispatched = true; \ + } while (0); + +//------------------------------------------------------------------------------ +// Dynamic Kernel Version (parameters computed at runtime) +//------------------------------------------------------------------------------ +struct KernelParamsDynamic { + int VPT; + int NUM_EXPERTS; + int THREADS_PER_ROW; + int ROWS_PER_WARP; + int ROWS_PER_CTA; + int WARPS_PER_CTA; +}; + +template +__global__ void moe_fused_gate_kernel_dynamic( + void* input, + void* bias, + float* output_ptr, + int32_t* indices_ptr, + int64_t num_rows, + int64_t num_experts, + int64_t num_expert_group, + int64_t topk_group, + int64_t topk, + int64_t num_fused_shared_experts, + double routed_scaling_factor, + bool apply_routed_scaling_factor_on_output) { + KernelParamsDynamic params; + params.NUM_EXPERTS = num_experts; // e.g, for deepseek v3, this is 256 + params.VPT = num_experts / num_expert_group; // e.g., for deepseek v3, this is 256 / 8 = 32 + params.THREADS_PER_ROW = num_expert_group; // fixed as num_expert_group, e.g., for deepseek v3, + // this is 8 + params.WARPS_PER_CTA = WARPS_PER_CTA; // fixed as 6 + params.ROWS_PER_WARP = std::max(1, WARP_SIZE / num_expert_group); // WARP_SIZE is fixed as 32 + params.ROWS_PER_CTA = params.WARPS_PER_CTA * params.ROWS_PER_WARP; + + moe_fused_gate_impl_dynamic( + input, + bias, + output_ptr, + indices_ptr, + num_rows, + topk_group, + topk, + num_fused_shared_experts, + routed_scaling_factor, + apply_routed_scaling_factor_on_output, + params); +} + +void dispatch_moe_fuse_gate_dynamic( + at::Tensor& output, + at::Tensor& indices, + at::Tensor& input, + at::Tensor& bias, + int64_t num_rows, + int64_t num_experts, + int64_t num_expert_group, + int64_t topk_group, + int64_t topk, + int64_t num_fused_shared_experts, + double routed_scaling_factor, + bool apply_routed_scaling_factor_on_output) { + // Compute grid dimensions based on runtime value for num_expert_group. + int64_t rows_per_warp = std::max(1, WARP_SIZE / num_expert_group); + int64_t num_warps = (num_rows + rows_per_warp - 1) / rows_per_warp; + int64_t num_blocks = (num_warps + WARPS_PER_CTA - 1) / WARPS_PER_CTA; + const musaStream_t stream = at::musa::getCurrentMUSAStream(); + dim3 block_dim(WARP_SIZE, WARPS_PER_CTA); + + // Fallback to the dynamic kernel if none of the supported combinations match. + // currently only support num_experts / num_expert_group <= 32 for dynamic + // kernels + if (input.scalar_type() == at::kBFloat16) { + moe_fused_gate_kernel_dynamic<<>>( + input.data_ptr(), + bias.data_ptr(), + output.data_ptr(), + indices.data_ptr(), + num_rows, + num_experts, + num_expert_group, + topk_group, + topk, + num_fused_shared_experts, + routed_scaling_factor, + apply_routed_scaling_factor_on_output); + } else if (input.scalar_type() == at::kHalf) { + moe_fused_gate_kernel_dynamic<<>>( + input.data_ptr(), + bias.data_ptr(), + output.data_ptr(), + indices.data_ptr(), + num_rows, + num_experts, + num_expert_group, + topk_group, + topk, + num_fused_shared_experts, + routed_scaling_factor, + apply_routed_scaling_factor_on_output); + } else if (input.scalar_type() == at::kFloat) { + moe_fused_gate_kernel_dynamic<<>>( + input.data_ptr(), + bias.data_ptr(), + output.data_ptr(), + indices.data_ptr(), + num_rows, + num_experts, + num_expert_group, + topk_group, + topk, + num_fused_shared_experts, + routed_scaling_factor, + apply_routed_scaling_factor_on_output); + } else { + TORCH_CHECK(false, "Unsupported data type for moe_fused_gate"); + } +} + +bool dispatch_moe_fuse_gate_static( + at::Tensor& output, + at::Tensor& indices, + at::Tensor& input, + at::Tensor& bias, + int64_t num_rows, + int64_t num_experts, + int64_t num_expert_group, + int64_t topk_group, + int64_t topk, + int64_t num_fused_shared_experts, + double routed_scaling_factor, + bool apply_routed_scaling_factor_on_output) { + const musaStream_t stream = at::musa::getCurrentMUSAStream(); + bool dispatched = false; + float last_val = apply_routed_scaling_factor_on_output ? 1.f : 1.f / routed_scaling_factor; + // Dispatch to templated kernel for known compile-time configurations. + // We currently only support for: + // Case 1: 256 experts, with 8 or 16 groups. + // Case 2: 128 experts, with 4 or 8 groups. + // Case 3: other cases, require 8 <= num_experts / num_expert_group <= 32 + constexpr int block_size = 256; + switch (num_experts) { + case 256: + if (num_expert_group == 8) { + // This is deepseek v3 case. Here VPT = 256/8 = 32, ROWS_PER_WARP = 32/8 + // = 4, ROWS_PER_CTA = 6 * 4 = 24. + if (input.scalar_type() == at::kBFloat16) { + LAUNCH_MOE_GATE_CONFIG(bfloat16_t, 256, 8); + } else if (input.scalar_type() == at::kHalf) { + LAUNCH_MOE_GATE_CONFIG(float16_t, 256, 8); + } else if (input.scalar_type() == at::kFloat) { + LAUNCH_MOE_GATE_CONFIG(float32_t, 256, 8); + } + } else if (num_expert_group == 16) { + // Here VPT = 256/16 = 16, ROWS_PER_WARP = 32/16 = 2, ROWS_PER_CTA + // = 6 * 2 = 12. + if (input.scalar_type() == at::kBFloat16) { + LAUNCH_MOE_GATE_CONFIG(bfloat16_t, 256, 16); + } else if (input.scalar_type() == at::kHalf) { + LAUNCH_MOE_GATE_CONFIG(float16_t, 256, 16); + } else if (input.scalar_type() == at::kFloat) { + LAUNCH_MOE_GATE_CONFIG(float32_t, 256, 16); + } + } + break; + case 128: + if (num_expert_group == 4) { + // VPT = 128/4 = 32, ROWS_PER_WARP = 32/16 = 2, ROWS_PER_CTA = 6 * 2 + // = 12. + if (input.scalar_type() == at::kBFloat16) { + LAUNCH_MOE_GATE_CONFIG(bfloat16_t, 128, 4); + } else if (input.scalar_type() == at::kHalf) { + LAUNCH_MOE_GATE_CONFIG(float16_t, 128, 4); + } else if (input.scalar_type() == at::kFloat) { + LAUNCH_MOE_GATE_CONFIG(float32_t, 128, 4); + } + } else if (num_expert_group == 8) { + // VPT = 128/8 = 16, ROWS_PER_WARP = 32/8 = 4, ROWS_PER_CTA = 6 * 4 + // = 24. + if (input.scalar_type() == at::kBFloat16) { + LAUNCH_MOE_GATE_CONFIG(bfloat16_t, 128, 8); + } else if (input.scalar_type() == at::kHalf) { + LAUNCH_MOE_GATE_CONFIG(float16_t, 128, 8); + } else if (input.scalar_type() == at::kFloat) { + LAUNCH_MOE_GATE_CONFIG(float32_t, 128, 8); + } + } + break; + default: + break; + } + + return dispatched; +} + +#undef LAUNCH_MOE_GATE_CONFIG + +//------------------------------------------------------------------------------ +// Host Launcher Function +//------------------------------------------------------------------------------ +std::vector moe_fused_gate( + at::Tensor& input, + at::Tensor& bias, + int64_t num_expert_group, + int64_t topk_group, + int64_t topk, + int64_t num_fused_shared_experts, + double routed_scaling_factor, + bool apply_routed_scaling_factor_on_output) { + TORCH_CHECK(input.dtype() == bias.dtype(), "input and bias should have the same dtype"); + int64_t num_rows = input.size(0); + int32_t num_experts = input.size(1); + auto options = torch::TensorOptions().dtype(torch::kFloat32).device(input.device()); + auto output = torch::empty({num_rows, topk}, options); + auto indices = torch::empty({num_rows, topk}, options.dtype(torch::kInt32)); + + // Check 1: Ensure that num_experts is a power of 2. + TORCH_CHECK((num_experts & (num_experts - 1)) == 0, "num_experts must be a power of 2, but got ", num_experts); + + // Check 2: Ensure that num_experts is divisible by num_expert_group. (this + // also means num_expert_group is power of 2) + TORCH_CHECK( + num_experts % num_expert_group == 0, + "num_experts must be divisible by num_expert_group, but got ", + num_experts, + " / ", + num_expert_group); + + int computed_vpt = num_experts / num_expert_group; + // Check 3: Ensure that num_experts/num_expert_group does not exceed + // MAX_VPT=32. Maximum VPT indicate max value per threads we can process. + TORCH_CHECK( + computed_vpt <= MAX_VPT, + "Per group experts: num_experts / num_expert_group = (", + computed_vpt, + ") exceeds the maximum supported (", + MAX_VPT, + ")"); + + bool static_dispatched = dispatch_moe_fuse_gate_static( + output, + indices, + input, + bias, + num_rows, + num_experts, + num_expert_group, + topk_group, + topk, + num_fused_shared_experts, + routed_scaling_factor, + apply_routed_scaling_factor_on_output); + + if (!static_dispatched) { + dispatch_moe_fuse_gate_dynamic( + output, + indices, + input, + bias, + num_rows, + num_experts, + num_expert_group, + topk_group, + topk, + num_fused_shared_experts, + routed_scaling_factor, + apply_routed_scaling_factor_on_output); + } + + return {output, indices}; +} diff --git a/sgl-kernel/csrc/moe/moe_topk_softmax_kernels.cu b/sgl-kernel/csrc/moe/moe_topk_softmax_kernels.cu index ba12f4dd8..82f8b89fc 100644 --- a/sgl-kernel/csrc/moe/moe_topk_softmax_kernels.cu +++ b/sgl-kernel/csrc/moe/moe_topk_softmax_kernels.cu @@ -23,7 +23,9 @@ limitations under the License. #ifndef USE_ROCM #include #include +#ifndef USE_MUSA #include +#endif #else #include #include diff --git a/sgl-kernel/csrc/quantization/gguf/ggml-common.h b/sgl-kernel/csrc/quantization/gguf/ggml-common.h index f6fbe57aa..88c21a4ab 100644 --- a/sgl-kernel/csrc/quantization/gguf/ggml-common.h +++ b/sgl-kernel/csrc/quantization/gguf/ggml-common.h @@ -964,11 +964,11 @@ static __device__ __forceinline__ dst_t convert_from_half(half val) { template <> __device__ __forceinline__ c10::BFloat16 convert_from_half(half val) { -#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 || defined(USE_MUSA) return __float2bfloat16(__half2float(val)); #else return __half2float(val); -#endif // defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 +#endif // defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800 || defined(USE_MUSA) } template <> diff --git a/sgl-kernel/csrc/quantization/gguf/vecdotq.cuh b/sgl-kernel/csrc/quantization/gguf/vecdotq.cuh index 933c2aa51..08b4cd269 100644 --- a/sgl-kernel/csrc/quantization/gguf/vecdotq.cuh +++ b/sgl-kernel/csrc/quantization/gguf/vecdotq.cuh @@ -48,7 +48,7 @@ static __device__ __forceinline__ int get_int_from_uint8_aligned(const uint8_t* template static __device__ __forceinline__ float vec_dot_q4_0_q8_1_impl(const int* v, const int* u, const float& d4, const half2& ds8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA int sumi = 0; #pragma unroll @@ -74,7 +74,7 @@ vec_dot_q4_0_q8_1_impl(const int* v, const int* u, const float& d4, const half2& template static __device__ __forceinline__ float vec_dot_q4_1_q8_1_impl(const int* v, const int* u, const half2& dm4, const half2& ds8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA int sumi = 0; #pragma unroll @@ -102,7 +102,7 @@ vec_dot_q4_1_q8_1_impl(const int* v, const int* u, const half2& dm4, const half2 template static __device__ __forceinline__ float vec_dot_q5_0_q8_1_impl(const int* vl, const int* vh, const int* u, const float& d5, const half2& ds8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA int sumi = 0; #pragma unroll @@ -135,7 +135,7 @@ vec_dot_q5_0_q8_1_impl(const int* vl, const int* vh, const int* u, const float& template static __device__ __forceinline__ float vec_dot_q5_1_q8_1_impl(const int* vl, const int* vh, const int* u, const half2& dm5, const half2& ds8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA int sumi = 0; #pragma unroll @@ -170,7 +170,7 @@ vec_dot_q5_1_q8_1_impl(const int* vl, const int* vh, const int* u, const half2& template static __device__ __forceinline__ float vec_dot_q8_0_q8_1_impl(const int* v, const int* u, const float& d8_0, const float& d8_1) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA int sumi = 0; #pragma unroll @@ -185,7 +185,7 @@ vec_dot_q8_0_q8_1_impl(const int* v, const int* u, const float& d8_0, const floa template static __device__ __forceinline__ float vec_dot_q8_1_q8_1_impl(const int* v, const int* u, const half2& dm8, const half2& ds8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA int sumi = 0; @@ -214,7 +214,7 @@ static __device__ __forceinline__ float vec_dot_q2_K_q8_1_impl_mmvq( const uint8_t* __restrict__ scales, const half2& dm2, const float* __restrict__ d8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA float sumf_d = 0.0f; float sumf_m = 0.0f; @@ -245,7 +245,7 @@ static __device__ __forceinline__ float vec_dot_q2_K_q8_1_impl_mmq( const uint8_t* __restrict__ scales, const half2& dm2, const float& d8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA int sumi_d = 0; int sumi_m = 0; @@ -287,7 +287,7 @@ static __device__ __forceinline__ float vec_dot_q3_K_q8_1_impl_mmvq( const int& scale_offset, const float& d3, const float* __restrict__ d8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA float sumf = 0.0f; @@ -324,7 +324,7 @@ static __device__ __forceinline__ float vec_dot_q3_K_q8_1_impl_mmq( const int8_t* __restrict__ scales, const float& d3, const float& d8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA int sumi = 0; #pragma unroll @@ -353,7 +353,7 @@ static __device__ __forceinline__ float vec_dot_q4_K_q8_1_impl_vmmq( const uint8_t* __restrict__ m, const half2& dm4, const float* __restrict__ d8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA float sumf_d = 0.0f; float sumf_m = 0.0f; @@ -382,7 +382,7 @@ static __device__ __forceinline__ float vec_dot_q4_K_q8_1_impl_mmq( const uint8_t* __restrict__ m, const half2& dm4, const half2* __restrict__ ds8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA float sumf_d = 0.0f; float sumf_m = 0.0f; @@ -418,7 +418,7 @@ static __device__ __forceinline__ float vec_dot_q5_K_q8_1_impl_vmmq( const uint8_t* __restrict__ m, const half2& dm5, const float* __restrict__ d8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA float sumf_d = 0.0f; float sumf_m = 0.0f; @@ -453,7 +453,7 @@ static __device__ __forceinline__ float vec_dot_q5_K_q8_1_impl_mmq( const uint8_t* __restrict__ m, const half2& dm4, const half2* __restrict__ ds8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA float sumf_d = 0.0f; float sumf_m = 0.0f; @@ -489,7 +489,7 @@ static __device__ __forceinline__ float vec_dot_q6_K_q8_1_impl_mmvq( const int8_t* __restrict__ scales, const float& d, const float* __restrict__ d8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA float sumf = 0.0f; #pragma unroll @@ -512,7 +512,7 @@ static __device__ __forceinline__ float vec_dot_q6_K_q8_1_impl_mmq( const int8_t* __restrict__ sc, const float& d6, const float* __restrict__ d8) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA float sumf_d = 0.0f; #pragma unroll @@ -1807,7 +1807,7 @@ vec_dot_iq2_xs_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ static __device__ __forceinline__ float vec_dot_iq2_s_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ bq8_1, const int& iqs) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA const block_iq2_s* bq2 = (const block_iq2_s*)vbq; const int ib32 = iqs; @@ -1846,7 +1846,7 @@ vec_dot_iq2_s_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ static __device__ __forceinline__ float vec_dot_iq3_xxs_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ bq8_1, const int& iqs) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA const block_iq3_xxs* bq2 = (const block_iq3_xxs*)vbq; const int ib32 = iqs; @@ -1873,7 +1873,7 @@ vec_dot_iq3_xxs_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict_ static __device__ __forceinline__ float vec_dot_iq3_s_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ bq8_1, const int& iqs) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA const block_iq3_s* bq2 = (const block_iq3_s*)vbq; const int ib32 = iqs; @@ -1899,7 +1899,7 @@ vec_dot_iq3_s_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ static __device__ __forceinline__ float vec_dot_iq1_s_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ bq8_1, const int& iqs) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA const block_iq1_s* bq1 = (const block_iq1_s*)vbq; const int qs_packed = get_int_b2(bq1->qs, iqs); @@ -1931,7 +1931,7 @@ vec_dot_iq1_s_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ static __device__ __forceinline__ float vec_dot_iq1_m_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ bq8_1, const int& iqs) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA const block_iq1_m* bq1 = (const block_iq1_m*)vbq; @@ -1991,7 +1991,7 @@ get_int_from_table_16(const uint32_t& q4, const uint8_t* values, int& val1, int& static __device__ __forceinline__ float vec_dot_iq4_nl_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ bq8_1, const int& iqs) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA const block_iq4_nl* bq = (const block_iq4_nl*)vbq; @@ -2015,7 +2015,7 @@ vec_dot_iq4_nl_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ static __device__ __forceinline__ float vec_dot_iq4_xs_q8_1(const void* __restrict__ vbq, const block_q8_1* __restrict__ bq8_1, const int& iqs) { -#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM +#if defined __CUDA_ARCH__ && __CUDA_ARCH__ >= 610 || defined USE_ROCM || defined USE_MUSA const block_iq4_xs* bq4 = (const block_iq4_xs*)vbq; const uint8_t* values = (const uint8_t*)kvalues_iq4nl; diff --git a/sgl-kernel/csrc/speculative/eagle_utils.cu b/sgl-kernel/csrc/speculative/eagle_utils.cu index e8e306325..06524e773 100644 --- a/sgl-kernel/csrc/speculative/eagle_utils.cu +++ b/sgl-kernel/csrc/speculative/eagle_utils.cu @@ -17,7 +17,7 @@ #include #include -#ifndef USE_ROCM +#if !defined(USE_ROCM) && !defined(USE_MUSA) #include "pytorch_extension_utils.h" #else #include "pytorch_extension_utils_rocm.h" diff --git a/sgl-kernel/include/sgl_kernel_ops.h b/sgl-kernel/include/sgl_kernel_ops.h index 77068fb8d..3891dc900 100644 --- a/sgl-kernel/include/sgl_kernel_ops.h +++ b/sgl-kernel/include/sgl_kernel_ops.h @@ -127,6 +127,8 @@ int64_t cutlass_mla_get_workspace_size( void rmsnorm(at::Tensor& output, at::Tensor& input, at::Tensor& weight, double eps, bool enable_pdl); void sgl_fused_add_rmsnorm( torch::Tensor input, torch::Tensor residual, torch::Tensor weight, double eps, bool enable_pdl); +void musa_fused_add_rms_norm( + torch::Tensor& input, torch::Tensor& residual, torch::Tensor& weight, double epsilon, bool enable_pdl); void gemma_rmsnorm(at::Tensor& output, at::Tensor& input, at::Tensor& weight, double eps, bool enable_pdl); void gemma_fused_add_rmsnorm(at::Tensor& input, at::Tensor& residual, at::Tensor& weight, double eps, bool enable_pdl); void silu_and_mul(at::Tensor& out, at::Tensor& input); diff --git a/sgl-kernel/setup_musa.py b/sgl-kernel/setup_musa.py index 1abc0ea4d..cdfcf87b4 100644 --- a/sgl-kernel/setup_musa.py +++ b/sgl-kernel/setup_musa.py @@ -76,10 +76,37 @@ include_dirs = [ ] sources = [ + "csrc/allreduce/custom_all_reduce.cu", + "csrc/attention/merge_attn_states.cu", "csrc/common_extension_musa.cc", + "csrc/elementwise/activation.cu", + "csrc/elementwise/concat_mla.cu", + "csrc/elementwise/fused_add_rms_norm_kernel.mu", + "csrc/grammar/apply_token_bitmask_inplace_cuda.cu", + "csrc/moe/moe_align_kernel.cu", + "csrc/moe/moe_fused_gate_musa.cu", + "csrc/moe/kimi_k2_moe_fused_gate.cu", + "csrc/moe/moe_sum.cu", + "csrc/moe/moe_sum_reduce.cu", + "csrc/moe/moe_topk_softmax_kernels.cu", + "csrc/quantization/gguf/gguf_kernel.cu", + "csrc/speculative/eagle_utils.cu", + "csrc/speculative/ngram_utils.cu", + "csrc/speculative/packbit.cu", + "csrc/speculative/speculative_sampling.cu", + "csrc/kvcacheio/transfer.cu", + "csrc/gemm/awq_kernel.cu", + "csrc/gemm/bmm_fp8.cu", + "csrc/gemm/dsv3_fused_a_gemm.cu", + "csrc/gemm/dsv3_router_gemm_bf16_out.cu", + "csrc/gemm/dsv3_router_gemm_entry.cu", + "csrc/gemm/dsv3_router_gemm_float_out.cu", + "csrc/gemm/per_token_quant_fp8.cu", + "csrc/gemm/per_token_group_quant_8bit.cu", + "csrc/gemm/per_token_group_quant_8bit_v2.cu", + "csrc/memory/weak_ref_tensor.cpp", str(_FLASHINFER_REPO.source_dir / "csrc/norm.cu"), str(_FLASHINFER_REPO.source_dir / "csrc/renorm.cu"), - str(_FLASHINFER_REPO.source_dir / "csrc/sampling.cu"), ] cxx_flags = ["force_mcc"]