From 3ecf2c76ad1b5d54cb308f420a299a8889c13e51 Mon Sep 17 00:00:00 2001 From: blzheng Date: Fri, 29 May 2026 16:05:26 +0800 Subject: [PATCH] [CPU] Add GPT-OSS model optimization for CPU (#16775) Co-authored-by: mingfeima Co-authored-by: jianan-gu --- python/sglang/srt/layers/amx_utils.py | 2 + .../srt/layers/attention/intel_amx_backend.py | 6 + .../srt/layers/moe/fused_moe_triton/layer.py | 5 + .../srt/layers/quantization/__init__.py | 4 +- .../quantization/awq/schemes/awq_cpu.py | 4 + python/sglang/srt/layers/quantization/fp8.py | 4 + .../srt/layers/quantization/gptq_cpu.py | 4 + .../sglang/srt/layers/quantization/mxfp4.py | 60 ++++ .../sglang/srt/layers/quantization/unquant.py | 12 + .../srt/layers/quantization/w8a8_int8.py | 4 + python/sglang/srt/models/deepseek.py | 4 + python/sglang/srt/models/gpt_oss.py | 21 +- python/sglang/srt/server_args.py | 4 + sgl-kernel/csrc/cpu/common.h | 7 + sgl-kernel/csrc/cpu/decode.cpp | 65 +++- sgl-kernel/csrc/cpu/extend.cpp | 64 +++- sgl-kernel/csrc/cpu/flash_attn.cpp | 13 +- sgl-kernel/csrc/cpu/flash_attn.h | 170 ++++++----- sgl-kernel/csrc/cpu/gemm.h | 62 +++- sgl-kernel/csrc/cpu/gemm_fp8.cpp | 83 ++++++ sgl-kernel/csrc/cpu/moe.cpp | 254 +++++++++++++--- sgl-kernel/csrc/cpu/moe.h | 106 +++++++ sgl-kernel/csrc/cpu/moe_fp8.cpp | 145 +++++---- sgl-kernel/csrc/cpu/qkv_proj.cpp | 1 + sgl-kernel/csrc/cpu/torch_extension_cpu.cpp | 53 +++- sgl-kernel/csrc/cpu/vec.h | 14 + test/registered/cpu/test_decode.py | 158 ++++++++-- test/registered/cpu/test_extend.py | 177 +++++++++-- test/registered/cpu/test_gemm.py | 188 ++++-------- test/registered/cpu/test_mla.py | 2 + test/registered/cpu/test_moe.py | 281 ++++++++++++------ test/registered/cpu/utils.py | 216 ++++++++++++++ test/srt/cpu/test_decode.py | 158 ++++++++-- test/srt/cpu/test_extend.py | 177 +++++++++-- test/srt/cpu/test_mla.py | 2 + 35 files changed, 2000 insertions(+), 530 deletions(-) diff --git a/python/sglang/srt/layers/amx_utils.py b/python/sglang/srt/layers/amx_utils.py index b2be65058..18f7c52f6 100755 --- a/python/sglang/srt/layers/amx_utils.py +++ b/python/sglang/srt/layers/amx_utils.py @@ -15,6 +15,7 @@ class CPUQuantMethod(IntEnum): INT8_W8A8 = 1 FP8_W8A16 = 2 INT4_W4A8 = 3 + MXFP4 = 4 class CPUQuantAlgo(IntEnum): @@ -96,6 +97,7 @@ def dtype_is_supported(weight): return weight.dtype in [ torch.float16, torch.bfloat16, + torch.uint8, torch.int8, torch.float8_e4m3fn, ] diff --git a/python/sglang/srt/layers/attention/intel_amx_backend.py b/python/sglang/srt/layers/attention/intel_amx_backend.py index 2f5e4141b..fbf7a58dd 100644 --- a/python/sglang/srt/layers/attention/intel_amx_backend.py +++ b/python/sglang/srt/layers/attention/intel_amx_backend.py @@ -98,6 +98,7 @@ class IntelAMXAttnBackend(AttentionBackend): layer: RadixAttention, forward_batch: ForwardBatch, save_kv_cache=True, + sinks=None, ): if layer.qk_head_dim != layer.v_head_dim: o = q.new_empty((q.shape[0], layer.tp_q_head_num * layer.v_head_dim)) @@ -128,7 +129,9 @@ class IntelAMXAttnBackend(AttentionBackend): layer.scaling, layer.logit_cap, layer.is_cross_attention, + layer.sliding_window_size + 1, forward_batch.encoder_lens, + sinks, ) return o @@ -140,6 +143,7 @@ class IntelAMXAttnBackend(AttentionBackend): layer: RadixAttention, forward_batch: ForwardBatch, save_kv_cache=True, + sinks=None, ): attn_logits, _ = self.forward_metadata @@ -169,7 +173,9 @@ class IntelAMXAttnBackend(AttentionBackend): layer.scaling, layer.logit_cap, layer.is_cross_attention, + layer.sliding_window_size + 1, forward_batch.encoder_lens, + sinks, ) return o diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py index 04831e08e..59ca3f9cc 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -218,6 +218,7 @@ class FusedMoE(torch.nn.Module): self.use_presharded_weights = use_presharded_weights self.use_triton_kernels = get_moe_runner_backend().is_triton_kernels() + self.use_flashinfer_trtllm_moe = ( get_moe_runner_backend().is_flashinfer_trtllm() or get_moe_runner_backend().is_flashinfer_trtllm_routed() @@ -465,6 +466,8 @@ class FusedMoE(torch.nn.Module): start = 0 if self.use_padded_loading: + if _is_cpu and is_bias: + shard_dim = 1 expert_data, loaded_weight = narrow_padded_param_and_loaded_weight( expert_data, loaded_weight, @@ -534,6 +537,8 @@ class FusedMoE(torch.nn.Module): shard_size = expert_data.shape[shard_dim] if self.use_padded_loading: + if _is_cpu and is_bias: + shard_dim = 1 expert_data, loaded_weight = narrow_padded_param_and_loaded_weight( expert_data, loaded_weight, diff --git a/python/sglang/srt/layers/quantization/__init__.py b/python/sglang/srt/layers/quantization/__init__.py index bc9aa33af..8fdcb737e 100644 --- a/python/sglang/srt/layers/quantization/__init__.py +++ b/python/sglang/srt/layers/quantization/__init__.py @@ -53,6 +53,7 @@ from sglang.srt.layers.quantization.w8a8_fp8 import W8A8Fp8Config from sglang.srt.layers.quantization.w8a8_int8 import W8A8Int8Config from sglang.srt.utils import ( cpu_has_amx_support, + is_cpu, is_cuda, is_hip, is_mps, @@ -95,7 +96,7 @@ BASE_QUANTIZATION_METHODS: Dict[str, Type[QuantizationConfig]] = { } -if is_cuda() or (_is_mxfp_supported and is_hip()): +if is_cpu() or is_cuda() or (_is_mxfp_supported and is_hip()): BASE_QUANTIZATION_METHODS.update( { "mxfp4": Mxfp4Config, @@ -126,6 +127,7 @@ CPU_QUANTIZATION_METHODS = { "compressed-tensors": CompressedTensorsConfig, "awq": AWQCPUConfig, "gptq": CPUGPTQConfig, + "mxfp4": Mxfp4Config, } QUANTIZATION_METHODS = {**BASE_QUANTIZATION_METHODS} diff --git a/python/sglang/srt/layers/quantization/awq/schemes/awq_cpu.py b/python/sglang/srt/layers/quantization/awq/schemes/awq_cpu.py index 7f258685d..f1dc785ae 100644 --- a/python/sglang/srt/layers/quantization/awq/schemes/awq_cpu.py +++ b/python/sglang/srt/layers/quantization/awq/schemes/awq_cpu.py @@ -99,6 +99,10 @@ class AWQIntelAMXMoEKernel: layer.w13_qzeros, layer.w2_qzeros, None, # block_size + None, # w1 bias + None, # w3 bias + None, # alpha + None, # limit True, # is_vnni ) return StandardCombineInput(hidden_states=output) diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py index 3fe0b1884..6903104ac 100644 --- a/python/sglang/srt/layers/quantization/fp8.py +++ b/python/sglang/srt/layers/quantization/fp8.py @@ -1828,6 +1828,10 @@ class Fp8MoEMethod(FusedMoEMethodBase): None, # w1_zp None, # w2_zp self.quant_config.weight_block_size, # block_size + None, # w1 bias + None, # w3 bias + None, # alpha + None, # limit True, # is_vnni ) return StandardCombineInput(hidden_states=output) diff --git a/python/sglang/srt/layers/quantization/gptq_cpu.py b/python/sglang/srt/layers/quantization/gptq_cpu.py index 1bac4ce0e..e50f14fb6 100644 --- a/python/sglang/srt/layers/quantization/gptq_cpu.py +++ b/python/sglang/srt/layers/quantization/gptq_cpu.py @@ -370,6 +370,10 @@ class GPTQMoEIntelAMXMethod(FusedMoEMethodBase): layer.w13_qzeros, layer.w2_qzeros, None, # block_size + None, # w1 bias + None, # w3 bias + None, # alpha + None, # limit True, # is_vnni ) return StandardCombineInput(hidden_states=output) diff --git a/python/sglang/srt/layers/quantization/mxfp4.py b/python/sglang/srt/layers/quantization/mxfp4.py index fa1142b21..c45d44cc8 100644 --- a/python/sglang/srt/layers/quantization/mxfp4.py +++ b/python/sglang/srt/layers/quantization/mxfp4.py @@ -33,6 +33,10 @@ from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) +from sglang.srt.layers.amx_utils import ( + CPUQuantMethod, + _amx_process_weight_after_loading, +) from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.moe import MoeRunner, MoeRunnerBackend, MoeRunnerConfig from sglang.srt.layers.moe.moe_runner.marlin import MarlinMoeQuantInfo @@ -46,6 +50,8 @@ from sglang.srt.layers.quantization.base_config import ( from sglang.srt.layers.quantization.utils import is_layer_skipped from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( + cpu_has_amx_support, + is_cpu, is_flashinfer_available, is_gfx95_supported, is_hip, @@ -57,6 +63,7 @@ from sglang.srt.utils import ( next_power_of_2, round_up, set_weight_attrs, + use_intel_amx_backend, ) from sglang.srt.utils.common import get_bool_env_var from sglang.srt.utils.custom_op import register_custom_op @@ -138,9 +145,11 @@ if TYPE_CHECKING: StandardDispatchOutput, ) +_is_cpu = is_cpu() _is_hip = is_hip() _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip _is_shuffle_moe_mxfp4 = is_gfx95_supported() +_is_cpu_amx_available = cpu_has_amx_support() _sm120_mxfp4_min_warps_patched = False if _is_hip: @@ -849,6 +858,30 @@ class Mxfp4MoEMethod(FusedMoEMethodBase): self.w2_weight_triton_tensor = w2_weight del layer.w13_weight del layer.w2_weight + elif _is_cpu and _is_cpu_amx_available: + _amx_process_weight_after_loading(layer, ["w13_weight", "w2_weight"]) + if use_intel_amx_backend(layer): + packed_w13_weight_scale = torch.ops.sgl_kernel.convert_scale_packed( + layer.w13_weight_scale + ) + packed_w2_weight_scale = torch.ops.sgl_kernel.convert_scale_packed( + layer.w2_weight_scale + ) + layer.w13_weight_scale = Parameter( + packed_w13_weight_scale, requires_grad=False + ) + layer.w2_weight_scale = Parameter( + packed_w2_weight_scale, requires_grad=False + ) + if hasattr(layer, "w13_weight_bias"): + layer.w13_weight_bias = Parameter( + layer.w13_weight_bias.float(), requires_grad=False + ) + if hasattr(layer, "w2_weight_bias"): + layer.w2_weight_bias = Parameter( + layer.w2_weight_bias.float(), requires_grad=False + ) + return else: from triton_kernels.numerics_details.mxfp import upcast_from_mxfp @@ -1107,6 +1140,33 @@ class Mxfp4MoEMethod(FusedMoEMethodBase): x = dispatch_output.hidden_states topk_output = dispatch_output.topk_output + if use_intel_amx_backend(layer): + from sglang.srt.layers.moe.topk import apply_topk_weights_cpu + + topk_weights, topk_ids, _ = dispatch_output.topk_output + x, topk_weights = apply_topk_weights_cpu( + self.moe_runner_config.apply_router_weight_on_input, topk_weights, x + ) + output = torch.ops.sgl_kernel.fused_experts_cpu( + x, + layer.w13_weight, + layer.w2_weight, + topk_weights, + topk_ids, + False, # inplace See [Note] inplace should be False in fused_experts. + CPUQuantMethod.MXFP4, + layer.w13_weight_scale, # w1_scale + layer.w2_weight_scale, # w2_scale + None, # w1_zp + None, # w2_zp + None, # block_size + getattr(layer, "w13_weight_bias", None), + getattr(layer, "w2_weight_bias", None), + layer.moe_runner_config.gemm1_alpha, + layer.moe_runner_config.gemm1_clamp_limit, + True, # is_vnni + ) + return StandardCombineInput(hidden_states=output) if self.use_marlin: assert TopKOutputChecker.format_is_standard(topk_output) diff --git a/python/sglang/srt/layers/quantization/unquant.py b/python/sglang/srt/layers/quantization/unquant.py index 1a4ab64ae..4c252d093 100644 --- a/python/sglang/srt/layers/quantization/unquant.py +++ b/python/sglang/srt/layers/quantization/unquant.py @@ -268,6 +268,14 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp): # Pack weight for get better performance on CPU if _is_cpu and _is_cpu_amx_available: _amx_process_weight_after_loading(layer, ["w13_weight", "w2_weight"]) + if hasattr(layer, "w13_weight_bias"): + layer.w13_weight_bias = Parameter( + layer.w13_weight_bias.float(), requires_grad=False + ) + if hasattr(layer, "w2_weight_bias"): + layer.w2_weight_bias = Parameter( + layer.w2_weight_bias.float(), requires_grad=False + ) if ( self.use_deep_gemm @@ -579,6 +587,10 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp): None, # w1_zp None, # w2_zp None, # block_size + getattr(layer, "w13_weight_bias", None), + getattr(layer, "w2_weight_bias", None), + layer.moe_runner_config.gemm1_alpha, + layer.moe_runner_config.gemm1_clamp_limit, True, # is_vnni ) return StandardCombineInput(hidden_states=output) diff --git a/python/sglang/srt/layers/quantization/w8a8_int8.py b/python/sglang/srt/layers/quantization/w8a8_int8.py index 09de341de..f94c90b30 100644 --- a/python/sglang/srt/layers/quantization/w8a8_int8.py +++ b/python/sglang/srt/layers/quantization/w8a8_int8.py @@ -375,6 +375,10 @@ class W8A8Int8MoEMethod(FusedMoEMethodBase): None, # w1_zp None, # w2_zp None, # block_size + None, # w1 bias + None, # w3 bias + None, # alpha + None, # limit True, # is_vnni ) return StandardCombineInput(hidden_states=output) diff --git a/python/sglang/srt/models/deepseek.py b/python/sglang/srt/models/deepseek.py index 9f1b345eb..5e122a05c 100644 --- a/python/sglang/srt/models/deepseek.py +++ b/python/sglang/srt/models/deepseek.py @@ -207,6 +207,10 @@ class DeepseekMoE(nn.Module): None, # w1_zp None, # w2_zp None, # block_size + None, # w1_bias + None, # w2_bias + None, # alpha + None, # limit True, # is_vnni ) else: diff --git a/python/sglang/srt/models/gpt_oss.py b/python/sglang/srt/models/gpt_oss.py index 84f8890ee..b979a9264 100644 --- a/python/sglang/srt/models/gpt_oss.py +++ b/python/sglang/srt/models/gpt_oss.py @@ -81,6 +81,7 @@ from sglang.srt.utils import ( add_prefix, get_cuda_version, is_blackwell_supported, + is_cpu, is_cuda, is_flashinfer_available, is_npu, @@ -89,6 +90,7 @@ from sglang.srt.utils import ( ) from sglang.srt.utils.custom_op import register_custom_op +_is_cpu = is_cpu() _is_npu = is_npu() _is_cuda = is_cuda() _is_tinygemm_supported = ( @@ -881,7 +883,8 @@ class GptOssForCausalLM(nn.Module): moe_ep_rank_end = (moe_ep_rank + 1) * moe_num_local_experts for name, weight in weights: - weight = weight.cuda() + if _is_cuda: + weight = weight.cuda() if "gate_up_proj_blocks" in name: # Handle MLP gate and up projection weights @@ -1163,6 +1166,22 @@ class GptOssForCausalLM(nn.Module): param = params_dict[name] if "sinks" in name: start = get_attention_tp_rank() * param.numel() + tp_size = get_tensor_model_parallel_world_size() + full_shard_size = param.numel() * tp_size + # This handles TP padding: if the checkpoint dim is not divisible by tp_size, + # the last TP shard extends beyond `loaded_weight`, pad with zeros before slicing. + if ( + _is_cpu + and full_shard_size > loaded_weight.size(0) + and start + param.numel() >= loaded_weight.size(0) + ): + pad_size = start + param.numel() - loaded_weight.size(0) + pad_tensor = torch.zeros(pad_size).to( + loaded_weight.dtype + ) + loaded_weight = torch.cat( + [loaded_weight, pad_tensor], dim=0 + ).to(loaded_weight.dtype) param.data.copy_( loaded_weight[start : start + param.numel()] ) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 6fd058607..e604e3f9f 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -2059,6 +2059,8 @@ class ServerArgs: self.attention_backend = "trtllm_mha" elif is_sm90_supported(): self.attention_backend = "fa3" + elif is_cpu() and cpu_has_amx_support(): + self.attention_backend = "intel_amx" elif is_xpu(): self.attention_backend = "intel_xpu" elif is_hip(): @@ -2084,6 +2086,7 @@ class ServerArgs: "fa3", "fa4", "ascend", + "intel_amx", "intel_xpu", "aiter", ] @@ -2147,6 +2150,7 @@ class ServerArgs: self.ep_size == 1 and is_triton_kernels_available() and self.quantization is None + and not (is_cpu() and cpu_has_amx_support()) ): # The triton_kernels package segfaults on Blackwell (B200) # with NVIDIA driver >= 595. Fall back to triton backend. diff --git a/sgl-kernel/csrc/cpu/common.h b/sgl-kernel/csrc/cpu/common.h index 0da10f0bc..84eae8647 100644 --- a/sgl-kernel/csrc/cpu/common.h +++ b/sgl-kernel/csrc/cpu/common.h @@ -352,6 +352,13 @@ inline int get_cache_blocks(int chunk_size) { return std::min(MAX_CACHE_BLOCK_SIZE, cache_block_size); } +template <> +inline int get_cache_blocks(int chunk_size) { + // mxfp4 uses bf16 as accumulate type + int cache_block_size = get_cache_blocks(chunk_size); + return std::min(MAX_CACHE_BLOCK_SIZE, cache_block_size); +} + // 2d sequential loop in range : [mb0, mb1), [nb0, nb1) template inline void loop_2d(int64_t mb0, int64_t mb1, int64_t nb0, int64_t nb1, int64_t chunk_size, const func_t& f) { diff --git a/sgl-kernel/csrc/cpu/decode.cpp b/sgl-kernel/csrc/cpu/decode.cpp index b3ea06f10..e8ca4554e 100644 --- a/sgl-kernel/csrc/cpu/decode.cpp +++ b/sgl-kernel/csrc/cpu/decode.cpp @@ -981,16 +981,20 @@ template void decode_accumulate_kv_splits( scalar_t* __restrict__ output, float* __restrict__ attn_logits, + const scalar_t* __restrict__ sinks_ptr, int64_t batches, int64_t num_heads, int64_t head_size_v, int64_t num_kv_splits, int64_t l_stride1, - int64_t l_stride2) { + int64_t l_stride2, + bool has_sink) { using Vec = at::vec::Vectorized; // parallel on [batches, num_heads] at::parallel_for(0, batches * num_heads, 0, [&](int64_t begin, int64_t end) { + int64_t bi{0}, ni{0}; + data_index_init(begin, bi, batches, ni, num_heads); // NB: here we use logits[b][h][0] as acc, since // for the first kv split (kv_id == 0): // m_delta = std::exp(-inf) = 0 @@ -1022,8 +1026,12 @@ void decode_accumulate_kv_splits( s_prime = s_prime * m_delta + e_logic; m_prime = m_i; } - + if (has_sink) { + s_prime += std::exp(sinks_ptr[ni] - m_prime); + } copy_stub(output + i * head_size_v, acc, 1 / s_prime, head_size_v); + // move to the next index + data_index_step(bi, batches, ni, num_heads); } }); } @@ -1039,6 +1047,7 @@ void decode_attention_kernel_impl( const int64_t* __restrict__ req_pool_indices, const int64_t* __restrict__ seq_lens, const int64_t* __restrict__ encoder_lens, + const scalar_t* __restrict__ sinks, int64_t batches, int64_t num_heads, int64_t head_size, @@ -1055,8 +1064,10 @@ void decode_attention_kernel_impl( int64_t max_num_reqs, int64_t max_context_len, int64_t max_total_num_tokens, + int64_t sliding_window_size, bool is_cross_attn, - bool has_encoder_lens) { + bool has_encoder_lens, + bool has_sink) { using Vec = at::vec::Vectorized; // strides @@ -1083,6 +1094,10 @@ void decode_attention_kernel_impl( int64_t seq_len_kv = is_cross_attn ? encoder_lens[bs] : seq_lens[bs]; int64_t req_pool_id = req_pool_indices[bs]; int64_t kv_offset = (has_encoder_lens && (!is_cross_attn)) ? encoder_lens[bs] : 0; + if (sliding_window_size > 0 && seq_len_kv > sliding_window_size) { + kv_offset = seq_len_kv - sliding_window_size; + seq_len_kv = sliding_window_size; + } TORCH_CHECK(seq_len_kv <= max_context_len, "seq_len_kv out of scope!"); TORCH_CHECK(req_pool_id < max_num_reqs, "req_pool_id out of scope!"); @@ -1173,7 +1188,7 @@ void decode_attention_kernel_impl( }); decode_accumulate_kv_splits( - output, attn_logits, batches, num_heads, head_size_v, num_kv_splits, l_stride1, l_stride2); + output, attn_logits, sinks, batches, num_heads, head_size_v, num_kv_splits, l_stride1, l_stride2, has_sink); } // MHA template @@ -1187,6 +1202,7 @@ void decode_attention_mla_kernel_impl( const int64_t* __restrict__ req_pool_indices, const int64_t* __restrict__ seq_lens, scalar_t* __restrict__ buffer, + const scalar_t* __restrict__ sinks, int64_t batches, int64_t num_heads, int64_t head_size, @@ -1203,7 +1219,8 @@ void decode_attention_mla_kernel_impl( int64_t max_num_reqs, int64_t max_context_len, int64_t max_total_num_tokens, - int64_t buffer_size_per_thread) { + int64_t buffer_size_per_thread, + bool has_sink) { using Vec = at::vec::Vectorized; // block length for heads @@ -1369,7 +1386,7 @@ void decode_attention_mla_kernel_impl( }); decode_accumulate_kv_splits( - output, attn_logits, batches, num_heads, head_size_v, num_kv_splits, l_stride1, l_stride2); + output, attn_logits, sinks, batches, num_heads, head_size_v, num_kv_splits, l_stride1, l_stride2, has_sink); } // MLA template @@ -1383,6 +1400,7 @@ void decode_attention_grouped_kernel_impl( const int64_t* __restrict__ req_pool_indices, const int64_t* __restrict__ seq_lens, const int64_t* __restrict__ encoder_lens, + const scalar_t* __restrict__ sinks, int64_t batches, int64_t num_heads, int64_t num_heads_kv, @@ -1400,8 +1418,10 @@ void decode_attention_grouped_kernel_impl( int64_t max_num_reqs, int64_t max_context_len, int64_t max_total_num_tokens, + int64_t sliding_window_size, bool is_cross_attn, - bool has_encoder_lens) { + bool has_encoder_lens, + bool has_sink) { using Vec = at::vec::Vectorized; // block length for heads @@ -1447,7 +1467,10 @@ void decode_attention_grouped_kernel_impl( int64_t kv_offset = (has_encoder_lens && (!is_cross_attn)) ? encoder_lens[bs] : 0; TORCH_CHECK(seq_len_kv <= max_context_len, "seq_len_kv out of scope!"); TORCH_CHECK(req_pool_id < max_num_reqs, "req_pool_id out of scope!"); - + if (sliding_window_size > 0 && seq_len_kv > sliding_window_size) { + kv_offset = seq_len_kv - sliding_window_size; + seq_len_kv = sliding_window_size; + } const int64_t SPLIT_SIZE = div_up(seq_len_kv, num_kv_splits); const int64_t kv_start = kv_id * SPLIT_SIZE; const int64_t kv_end = std::min(kv_start + SPLIT_SIZE, seq_len_kv); @@ -1545,7 +1568,7 @@ void decode_attention_grouped_kernel_impl( }); decode_accumulate_kv_splits( - output, attn_logits, batches, num_heads, head_size_v, num_kv_splits, l_stride1, l_stride2); + output, attn_logits, sinks, batches, num_heads, head_size_v, num_kv_splits, l_stride1, l_stride2, has_sink); } // GQA/MQA } // anonymous namespace @@ -1559,7 +1582,7 @@ void decode_attention_grouped_kernel_impl( // req_pool_indices: [num_seqs] int64 // seq_lens: [num_seqs] int64 // encoder_lens: [num_seqs] int64 or None -// +// sinks: [num_heads] or None void decode_attention_cpu( at::Tensor& query, at::Tensor& k_buffer, @@ -1575,7 +1598,9 @@ void decode_attention_cpu( double sm_scale, double logit_cap, bool is_cross_attn, - std::optional encoder_lens) { + int64_t sliding_window_size, + std::optional encoder_lens, + std::optional sinks) { CHECK_LAST_DIM_CONTIGUOUS_INPUT(query); CHECK_LAST_DIM_CONTIGUOUS_INPUT(k_buffer); CHECK_LAST_DIM_CONTIGUOUS_INPUT(v_buffer); @@ -1642,6 +1667,10 @@ void decode_attention_cpu( encoder_lens_t = encoder_lens.value(); CHECK_EQ(encoder_lens_t.size(0), num_seqs); } + bool has_sink = sinks.has_value(); + at::Tensor sinks_tensor = has_sink ? sinks.value() : at::empty({num_heads}, query.options()); + CHECK_DIM(1, sinks_tensor); + CHECK_EQ(sinks_tensor.size(0), num_heads); AT_DISPATCH_REDUCED_FLOATING_TYPES(query.scalar_type(), "decode_attention_kernel", [&] { AT_DISPATCH_INDEX_TYPES(index_dtype, "decode_attention_indices", [&] { if (key.has_value()) { @@ -1693,6 +1722,7 @@ void decode_attention_cpu( req_pool_indices.data_ptr(), seq_lens.data_ptr(), encoder_lens_t.data_ptr(), + sinks_tensor.data_ptr(), num_seqs, num_heads, head_size, @@ -1709,8 +1739,10 @@ void decode_attention_cpu( max_num_reqs, max_context_len, max_total_num_tokens, + sliding_window_size, is_cross_attn, - has_encoder_lens); + has_encoder_lens, + has_sink); } else if (is_mla) { // MLA decode_attention_mla_kernel_impl( @@ -1723,6 +1755,7 @@ void decode_attention_cpu( req_pool_indices.data_ptr(), seq_lens.data_ptr(), buffer.data_ptr(), + sinks_tensor.data_ptr(), num_seqs, num_heads, head_size, @@ -1739,7 +1772,8 @@ void decode_attention_cpu( max_num_reqs, max_context_len, max_total_num_tokens, - size_per_thread); + size_per_thread, + has_sink); } else { // GQA/MQA decode_attention_grouped_kernel_impl( @@ -1752,6 +1786,7 @@ void decode_attention_cpu( req_pool_indices.data_ptr(), seq_lens.data_ptr(), encoder_lens_t.data_ptr(), + sinks_tensor.data_ptr(), num_seqs, num_heads, num_heads_kv, @@ -1769,8 +1804,10 @@ void decode_attention_cpu( max_num_reqs, max_context_len, max_total_num_tokens, + sliding_window_size, is_cross_attn, - has_encoder_lens); + has_encoder_lens, + has_sink); } }); }); diff --git a/sgl-kernel/csrc/cpu/extend.cpp b/sgl-kernel/csrc/cpu/extend.cpp index 9cc378602..01458f5fb 100644 --- a/sgl-kernel/csrc/cpu/extend.cpp +++ b/sgl-kernel/csrc/cpu/extend.cpp @@ -26,6 +26,7 @@ void extend_attention_kernel_impl( const index_t* __restrict__ extend_seq_lens, const index_t* __restrict__ extend_start_loc, const void* __restrict__ buffer, + const scalar_t* __restrict__ sinks, int batches, int num_heads, int num_heads_kv, @@ -47,9 +48,11 @@ void extend_attention_kernel_impl( int max_total_num_tokens, int max_len_extend, int buffer_size_per_thread, + int64_t sliding_window_size, bool is_prefix_skipped, bool is_cross_attn, - bool has_encoder_lens) { + bool has_encoder_lens, + bool has_sink) { // strides const int o_strideM = num_heads * head_size_v; const int o_strideH = head_size_v; @@ -69,18 +72,21 @@ void extend_attention_kernel_impl( data_index_init(begin, bs, batches, head_id, num_heads, mb, MB); int tid = at::get_thread_num(); - // s_i and s_delta: [BLOCK_M, BLOCK_N] + // s_i: [BLOCK_M, BLOCK_N] float* __restrict__ s_i = reinterpret_cast((char*)(buffer) + tid * buffer_size_per_thread); - scalar_t* __restrict__ s_delta = reinterpret_cast(s_i); // v_prime: [BLOCK_M, head_size_v] float* __restrict__ v_prime = s_i + BLOCK_M * BLOCK_N; + // s_delta: [BLOCK_M, BLOCK_N] + scalar_t* __restrict__ s_delta = reinterpret_cast(v_prime + BLOCK_M * head_size_v); + // Btmp: [BLOCK_N, max(head_size, head_size_v)] - scalar_t* __restrict__ Btmp = reinterpret_cast(v_prime + BLOCK_M * head_size_v); + scalar_t* __restrict__ Btmp = reinterpret_cast(s_delta + BLOCK_M * BLOCK_N); // init Btmp just once for each thread to prevent NaN fill_stub(Btmp, 0.f, BLOCK_N * ldb_tmp); + fill_stub(s_delta, 0.f, BLOCK_M * BLOCK_N); alignas(64) float s_prime[BLOCK_M]; alignas(64) float m_prime[BLOCK_M]; @@ -151,11 +157,20 @@ void extend_attention_kernel_impl( /* B */ Btmp, /* C */ s_i); - flash_attn_softmax::apply( - s_i, s_delta, v_prime, s_prime, m_prime, m_size, n_size, padded_n_size, head_size_v, sm_scale); + for (int row = 0; row < m_size; ++row) { + if (sliding_window_size > 0) { + int last_col = seq_len_prefix + row + m - sliding_window_size + 1; + if (last_col >= n + n_size) { + continue; + } + fill_stub(s_i + row * BLOCK_N, -std::numeric_limits::infinity(), last_col - n); + } + flash_attn_softmax::apply( + s_i, s_delta, v_prime, s_prime, m_prime, m_size, n_size, padded_n_size, head_size_v, sm_scale, row); + } // get value and pack - pack_vnni2( + pack_vnni2( /* dst */ Btmp, /* src */ v_buffer + head_kv_id * v_strideH, /* ind */ req_to_token + req_pool_id * max_context_len + n + kv_offset, @@ -233,8 +248,17 @@ void extend_attention_kernel_impl( } } - flash_attn_softmax::apply( - s_i, s_delta, v_prime, s_prime, m_prime, m_size, n_size, padded_n_size, head_size_v, sm_scale); + for (int row = 0; row < m_size; ++row) { + if (sliding_window_size > 0 && row + m + 1 >= n + sliding_window_size - 1 && + row + m + 1 < n + sliding_window_size + n_size) { + fill_stub( + s_i + row * BLOCK_N, -std::numeric_limits::infinity(), row + m - n - sliding_window_size + 1); + } else if (sliding_window_size > 0 && row + m + 1 >= n + sliding_window_size) { + continue; + } + flash_attn_softmax::apply( + s_i, s_delta, v_prime, s_prime, m_prime, m_size, n_size, padded_n_size, head_size_v, sm_scale, row); + } // get value and pack pack_vnni2( @@ -261,6 +285,9 @@ void extend_attention_kernel_impl( } scalar_t* __restrict__ out_ptr = o_extend + (seq_extend_start_loc + m) * o_strideM + head_id * o_strideH; for (int row = 0; row < m_size; ++row) { + if (has_sink) { + s_prime[row] += std::exp(sinks[head_id] - m_prime[row]); + } float s = 1 / s_prime[row]; copy_stub(out_ptr + row * o_strideM, v_prime + row * head_size_v, s, head_size_v); } @@ -280,6 +307,7 @@ inline int resize_buffer(at::Tensor& buffer, int num_threads, int head_size, int const int size_per_thread = /* s_i */ BLOCK_M * BLOCK_N * sizeof(float) + /* v_prime */ BLOCK_M * head_size_v * sizeof(float) + + /* s_delta */ BLOCK_M * BLOCK_N * sizeof(uint16_t) + /* Btmp */ BLOCK_N * std::max(head_size, head_size_v) * sizeof(uint16_t); buffer.resize_({num_threads, size_per_thread}); @@ -304,6 +332,7 @@ inline int resize_buffer(at::Tensor& buffer, int num_threads, int head_size, int extend_seq_lens.data_ptr(), \ extend_start_loc.data_ptr(), \ buffer.data_ptr(), \ + sinks_tensor.data_ptr(), \ num_seqs, \ num_heads, \ num_heads_kv, \ @@ -325,9 +354,11 @@ inline int resize_buffer(at::Tensor& buffer, int num_threads, int head_size, int max_total_num_tokens, \ max_len_extend, \ sz, \ + sliding_window_size, \ is_prefix_skipped, \ is_cross_attn, \ - has_encoder_lens); \ + has_encoder_lens, \ + has_sink); \ } while (0) // q_extend, k_extend, v_extend, o_extend: contiguous tensors @@ -344,8 +375,8 @@ inline int resize_buffer(at::Tensor& buffer, int num_threads, int head_size, int // seq_lens: [num_seqs] int64 // extend_seq_lens: [num_seqs] // extend_start_loc: [num_seqs] -// encoder_lens: [num_seqs] int64 -// +// encoder_lens: [num_seqs] int64 or None +// sinks: [num_heads] or None void extend_attention_cpu( at::Tensor& q_extend, const std::optional& k_extend_opt, @@ -362,7 +393,9 @@ void extend_attention_cpu( double sm_scale, double logit_cap, bool is_cross_attn, - std::optional encoder_lens) { + int64_t sliding_window_size, + std::optional encoder_lens, + std::optional sinks) { if (!is_cross_attn) { TORCH_CHECK( k_extend_opt.has_value() && v_extend_opt.has_value(), @@ -443,6 +476,11 @@ void extend_attention_cpu( encoder_lens_t = encoder_lens.value(); CHECK_EQ(encoder_lens_t.size(0), num_seqs); } + bool has_sink = sinks.has_value(); + at::Tensor sinks_tensor = has_sink ? sinks.value() : at::empty({num_heads}, q_extend.options()); + CHECK_DIM(1, sinks_tensor); + CHECK_EQ(sinks_tensor.size(0), num_heads); + AT_DISPATCH_REDUCED_FLOATING_TYPES(q_extend.scalar_type(), "extend_attention_kernel", [&] { AT_DISPATCH_INDEX_TYPES(index_dtype, "extend_attention_indices", [&] { if (max_len_extend <= 256) { diff --git a/sgl-kernel/csrc/cpu/flash_attn.cpp b/sgl-kernel/csrc/cpu/flash_attn.cpp index 152ac662a..4b3da3314 100644 --- a/sgl-kernel/csrc/cpu/flash_attn.cpp +++ b/sgl-kernel/csrc/cpu/flash_attn.cpp @@ -161,9 +161,10 @@ void flash_attn_kernel_impl( } } - flash_attn_softmax::apply( - s_i, s_delta, v_prime, s_prime, m_prime, m_size, n_size, padded_n_size, head_size_v, sm_scale); - + for (int row = 0; row < m_size; ++row) { + flash_attn_softmax::apply( + s_i, s_delta, v_prime, s_prime, m_prime, m_size, n_size, padded_n_size, head_size_v, sm_scale, row); + } // get value and pack pack_vnni2( /* dst */ Btmp, @@ -344,8 +345,10 @@ void flash_attn_varlen_kernel_impl( } } - flash_attn_softmax::apply( - s_i, s_delta, v_prime, s_prime, m_prime, m_size, n_size, padded_n_size, head_size_v, sm_scale); + for (int row = 0; row < m_size; ++row) { + flash_attn_softmax::apply( + s_i, s_delta, v_prime, s_prime, m_prime, m_size, n_size, padded_n_size, head_size_v, sm_scale, row); + } // get value and pack pack_vnni2( diff --git a/sgl-kernel/csrc/cpu/flash_attn.h b/sgl-kernel/csrc/cpu/flash_attn.h index 3aad15f44..ad90e989e 100644 --- a/sgl-kernel/csrc/cpu/flash_attn.h +++ b/sgl-kernel/csrc/cpu/flash_attn.h @@ -100,45 +100,43 @@ struct flash_attn_softmax { int n_size, int padded_n_size, int head_size_v, - const float sm_scale) { + const float sm_scale, + int row) { using Vec = at::vec::Vectorized; const Vec scale_vec = Vec(sm_scale); float* s_delta = s_i; - for (int row = 0; row < m_size; ++row) { - // s_i <- s_i * scale - at::vec::map( - [scale_vec](Vec x) { return x * scale_vec; }, s_i + row * BLOCK_N, s_i + row * BLOCK_N, n_size); + // s_i <- s_i * scale + at::vec::map([scale_vec](Vec x) { return x * scale_vec; }, s_i + row * BLOCK_N, s_i + row * BLOCK_N, n_size); - // m_i: max value per row - float m_i = at::vec::reduce_all( - [](Vec& x, Vec& y) { return at::vec::maximum(x, y); }, s_i + row * BLOCK_N, n_size); - m_i = std::max(m_i, m_prime[row]); + // m_i: max value per row + float m_i = + at::vec::reduce_all([](Vec& x, Vec& y) { return at::vec::maximum(x, y); }, s_i + row * BLOCK_N, n_size); + m_i = std::max(m_i, m_prime[row]); - // m_delta <- exp(m' - m_i) - float m_delta = std::exp(m_prime[row] - m_i); + // m_delta <- exp(m' - m_i) + float m_delta = std::exp(m_prime[row] - m_i); - // s_delta <- exp(s_i - m_i) - at::vec::map( - [m_i](Vec x) { return (x - Vec(m_i)).fexp_u20(); }, s_delta + row * BLOCK_N, s_i + row * BLOCK_N, n_size); + // s_delta <- exp(s_i - m_i) + at::vec::map( + [m_i](Vec x) { return (x - Vec(m_i)).fexp_u20(); }, s_delta + row * BLOCK_N, s_i + row * BLOCK_N, n_size); - // s' <- s' * m_delta + sum(s_delta) - s_prime[row] *= m_delta; - s_prime[row] += at::vec::reduce_all([](Vec& x, Vec& y) { return x + y; }, s_delta + row * BLOCK_N, n_size); + // s' <- s' * m_delta + sum(s_delta) + s_prime[row] *= m_delta; + s_prime[row] += at::vec::reduce_all([](Vec& x, Vec& y) { return x + y; }, s_delta + row * BLOCK_N, n_size); - m_prime[row] = m_i; + m_prime[row] = m_i; - // v' <- v' * m_delta - at::vec::map( - [m_delta](Vec x) { return x * Vec(m_delta); }, - v_prime + row * head_size_v, - v_prime + row * head_size_v, - head_size_v); + // v' <- v' * m_delta + at::vec::map( + [m_delta](Vec x) { return x * Vec(m_delta); }, + v_prime + row * head_size_v, + v_prime + row * head_size_v, + head_size_v); - // Keep s_delta row-major for the following brgemm(P @ V), and only - // convert the columns that brgemm will consume. - fill_stub(s_delta + row * BLOCK_N + n_size, 0.f, padded_n_size - n_size); - copy_stub(s_delta2 + row * BLOCK_N, s_delta + row * BLOCK_N, 1.f, padded_n_size); - } + // Keep s_delta row-major for the following brgemm(P @ V), and only + // convert the columns that brgemm will consume. + fill_stub(s_delta + row * BLOCK_N + n_size, 0.f, padded_n_size - n_size); + copy_stub(s_delta2 + row * BLOCK_N, s_delta + row * BLOCK_N, 1.f, padded_n_size); } }; @@ -155,7 +153,8 @@ struct flash_attn_softmax { int n_size, int padded_n_size, int head_size_v, - const float sm_scale) { + const float sm_scale, + int row) { float* s_delta = s_i; const __m512 vscale = _mm512_set1_ps(sm_scale); @@ -175,72 +174,71 @@ struct flash_attn_softmax { const __m512 vneg_inf = _mm512_set1_ps(NEG_INF); - for (int m = 0; m < m_size; ++m) { - vmax = vneg_inf; + int m = row; + vmax = vneg_inf; - // s_i <- s_i * scale - int n = 0; - for (; n <= n_size - 16; n += 16) { - va = _mm512_mul_ps(_mm512_loadu_ps(s_i + m * BLOCK_N + n), vscale); - vmax = _mm512_max_ps(va, vmax); - } - if (n_remainder > 0) { - va = _mm512_mul_ps(_mm512_mask_loadu_ps(vneg_inf, vmask, s_i + m * BLOCK_N + n), vscale); - vmax = _mm512_max_ps(va, vmax); - } + // s_i <- s_i * scale + int n = 0; + for (; n <= n_size - 16; n += 16) { + va = _mm512_mul_ps(_mm512_loadu_ps(s_i + m * BLOCK_N + n), vscale); + vmax = _mm512_max_ps(va, vmax); + } + if (n_remainder > 0) { + va = _mm512_mul_ps(_mm512_mask_loadu_ps(vneg_inf, vmask, s_i + m * BLOCK_N + n), vscale); + vmax = _mm512_max_ps(va, vmax); + } - // m_i: max value per row - float m_i = _mm512_reduce_max_ps(vmax); - m_i = std::max(m_i, m_prime[m]); - vmax = _mm512_set1_ps(m_i); + // m_i: max value per row + float m_i = _mm512_reduce_max_ps(vmax); + m_i = std::max(m_i, m_prime[m]); + vmax = _mm512_set1_ps(m_i); - // m_delta <- exp(m' - m_i) - float m_delta = std::exp(m_prime[m] - m_i); + // m_delta <- exp(m' - m_i) + float m_delta = std::exp(m_prime[m] - m_i); - // s_delta <- exp(s_i - m_i) - vsum = _mm512_setzero_ps(); - for (n = 0; n <= n_size - 16; n += 16) { - va = _mm512_mul_ps(_mm512_loadu_ps(s_i + m * BLOCK_N + n), vscale); - va = _mm512_fexp_u20_ps(_mm512_sub_ps(va, vmax)); - vsum = _mm512_add_ps(vsum, va); + // s_delta <- exp(s_i - m_i) + vsum = _mm512_setzero_ps(); + for (n = 0; n <= n_size - 16; n += 16) { + va = _mm512_mul_ps(_mm512_loadu_ps(s_i + m * BLOCK_N + n), vscale); + va = _mm512_fexp_u20_ps(_mm512_sub_ps(va, vmax)); + vsum = _mm512_add_ps(vsum, va); - vb = (__m256i)(_mm512_cvtneps_pbh(va)); - _mm256_storeu_si256(reinterpret_cast<__m256i*>(s_delta2 + m * BLOCK_N + n), vb); - } - if (n_remainder > 0) { - va = _mm512_mul_ps(_mm512_mask_loadu_ps(vneg_inf, vmask, s_i + m * BLOCK_N + n), vscale); - va = _mm512_fexp_u20_ps(_mm512_sub_ps(va, vmax)); - vsum = _mm512_add_ps(vsum, va); + vb = (__m256i)(_mm512_cvtneps_pbh(va)); + _mm256_storeu_si256(reinterpret_cast<__m256i*>(s_delta2 + m * BLOCK_N + n), vb); + } + if (n_remainder > 0) { + va = _mm512_mul_ps(_mm512_mask_loadu_ps(vneg_inf, vmask, s_i + m * BLOCK_N + n), vscale); + va = _mm512_fexp_u20_ps(_mm512_sub_ps(va, vmax)); + vsum = _mm512_add_ps(vsum, va); - vb = (__m256i)(_mm512_cvtneps_pbh(va)); - _mm256_mask_storeu_epi16(reinterpret_cast<__m256i*>(s_delta2 + m * BLOCK_N + n), vmask, vb); - } + vb = (__m256i)(_mm512_cvtneps_pbh(va)); + _mm256_mask_storeu_epi16(reinterpret_cast<__m256i*>(s_delta2 + m * BLOCK_N + n), vmask, vb); + } - // s' <- s' * m_delta + sum(s_delta) - s_prime[m] *= m_delta; - s_prime[m] += _mm512_reduce_add_ps(vsum); + // s' <- s' * m_delta + sum(s_delta) + s_prime[m] *= m_delta; + s_prime[m] += _mm512_reduce_add_ps(vsum); - m_prime[m] = m_i; + m_prime[m] = m_i; - // pad s_delta with 0, pad_size range from [0, 32) - int pad_size = padded_n_size - n_size; - if (pad_size > 0) { - const __m512i vzero = _mm512_setzero_si512(); - __mmask32 vmask2 = (1ULL << pad_size) - 1; - _mm512_mask_storeu_epi16(reinterpret_cast<__m512i*>(s_delta2 + m * BLOCK_N + n_size), vmask2, vzero); - } + // pad s_delta with 0, pad_size range from [0, 32) + int pad_size = padded_n_size - n_size; + if (pad_size > 0) { + const __m512i vzero = _mm512_setzero_si512(); + __mmask32 vmask2 = (1ULL << pad_size) - 1; + _mm512_mask_storeu_epi16(reinterpret_cast<__m512i*>(s_delta2 + m * BLOCK_N + n_size), vmask2, vzero); + } - // v' <- v' * m_delta - vmdelta = _mm512_set1_ps(m_delta); - int k = 0; - for (; k <= head_size_v - 16; k += 16) { - va = _mm512_mul_ps(_mm512_loadu_ps(v_prime + m * head_size_v + k), vmdelta); - _mm512_storeu_ps(reinterpret_cast<__m512*>(v_prime + m * head_size_v + k), va); - } - if (v_remainder > 0) { - va = _mm512_mul_ps(_mm512_maskz_loadu_ps(vmask1, v_prime + m * head_size_v + k), vmdelta); - _mm512_mask_storeu_ps(reinterpret_cast<__m512*>(v_prime + m * head_size_v + k), vmask1, va); - } + // v' <- v' * m_delta + vmdelta = _mm512_set1_ps(m_delta); + int k = 0; + for (; k <= head_size_v - 16; k += 16) { + va = _mm512_mul_ps(_mm512_loadu_ps(v_prime + m * head_size_v + k), vmdelta); + _mm512_storeu_ps(reinterpret_cast<__m512*>(v_prime + m * head_size_v + k), va); + } + if (v_remainder > 0) { + va = _mm512_mul_ps(_mm512_maskz_loadu_ps(vmask1, v_prime + m * head_size_v + k), vmdelta); + _mm512_mask_storeu_ps(reinterpret_cast<__m512*>(v_prime + m * head_size_v + k), vmask1, va); } } }; diff --git a/sgl-kernel/csrc/cpu/gemm.h b/sgl-kernel/csrc/cpu/gemm.h index bbfdccf45..102779a7a 100644 --- a/sgl-kernel/csrc/cpu/gemm.h +++ b/sgl-kernel/csrc/cpu/gemm.h @@ -67,7 +67,13 @@ inline int64_t get_row_size(int64_t K, bool use_int8_w8a8) { return use_int8_w8a8 ? K + sizeof(int32_t) : K; } -enum class CPUQuantMethod : int64_t { BF16 = 0, INT8_W8A8 = 1, FP8_W8A16 = 2, INT4_W4A8 = 3 }; +enum class CPUActMethod : int { + silu_and_mul = 0, + swiglu = 1, + gelu_and_mul = 2, +}; + +enum class CPUQuantMethod : int64_t { BF16 = 0, INT8_W8A8 = 1, FP8_W8A16 = 2, INT4_W4A8 = 3, MXFP4 = 4 }; constexpr bool operator==(CPUQuantMethod a, int64_t b) { return static_cast(a) == b; @@ -87,6 +93,17 @@ constexpr bool operator==(int64_t a, CPUQuantAlgo b) { return a == static_cast(b); } +inline int64_t get_row_size(CPUQuantMethod quant, int64_t K) { + switch (quant) { + case CPUQuantMethod::INT8_W8A8: + return K + sizeof(int32_t); + case CPUQuantMethod::MXFP4: + return K >> 1; + default: + return K; + } +} + inline int64_t get_4bit_block_k_size(int64_t group_size) { return group_size > 128 ? 128 : group_size; } @@ -124,9 +141,9 @@ void fused_experts_int8_kernel_impl( int64_t topk, int64_t num_tokens_post_pad); -// moe implementations for fp8 w8a16 -template -void fused_experts_fp8_kernel_impl( +// moe implementations for fp8 w8a16 and mxfp4 +template +void fused_experts_fp_kernel_impl( scalar_t* __restrict__ output, scalar_t* __restrict__ ic0, scalar_t* __restrict__ ic1, @@ -135,10 +152,12 @@ void fused_experts_fp8_kernel_impl( scalar_t* __restrict__ B_tmp, float* __restrict__ C_tmp, const scalar_t* __restrict__ input, - const at::Float8_e4m3fn* __restrict__ packed_w1, - const at::Float8_e4m3fn* __restrict__ packed_w2, - const float* __restrict__ w1s, - const float* __restrict__ w2s, + const packed_t* __restrict__ packed_w1, + const packed_t* __restrict__ packed_w2, + const float* __restrict__ w1_bias, + const float* __restrict__ w2_bias, + const param_t* __restrict__ w1s, + const param_t* __restrict__ w2s, int64_t block_size_N, int64_t block_size_K, const float* __restrict__ topk_weights, @@ -150,7 +169,11 @@ void fused_experts_fp8_kernel_impl( int64_t K, int64_t E, int64_t topk, - int64_t num_tokens_post_pad); + int64_t num_tokens_post_pad, + float alpha, + float limit, + CPUActMethod act_func, + bool with_bias); // shared expert implementation for int8 w8a8 template @@ -261,6 +284,7 @@ void tinygemm_kernel( scalar_t* __restrict__ C, scalar_t* __restrict__ Btmp, float* __restrict__ Ctmp, + const float* __restrict__ Bbias, const float* __restrict__ scale, int64_t M, int64_t N, @@ -289,6 +313,26 @@ void tinygemm_kernel( int64_t ldc, bool brg); +// mxfp4 +template +void tinygemm_kernel( + const scalar_t* __restrict__ A, + const uint8_t* __restrict__ B, + scalar_t* __restrict__ C, + scalar_t* __restrict__ Btmp, + float* __restrict__ Ctmp, + const float* __restrict__ Bbias, + const uint8_t* __restrict__ scale, + int64_t M, + int64_t N, + int64_t K, + int64_t lda, + int64_t ldb, + int64_t ldc, + bool brg, + int64_t block_size_K, + bool do_unpack = true); + template void tinygemm_kernel( scalar_t* C, diff --git a/sgl-kernel/csrc/cpu/gemm_fp8.cpp b/sgl-kernel/csrc/cpu/gemm_fp8.cpp index e4c195dd9..c176c418d 100644 --- a/sgl-kernel/csrc/cpu/gemm_fp8.cpp +++ b/sgl-kernel/csrc/cpu/gemm_fp8.cpp @@ -62,6 +62,23 @@ inline void copy_mul_stub(scalar_t* __restrict__ out, const float* __restrict__ } } +template <> +inline void +copy_add_stub(float* __restrict__ out, const float* __restrict__ input, const float* __restrict__ bias, int64_t size) { + using fVec = at::vec::Vectorized; + constexpr int kVecSize = fVec::size(); + + int64_t d; +#pragma GCC unroll 4 + for (d = 0; d <= size - kVecSize; d += kVecSize) { + fVec data = fVec::loadu(input + d) + fVec::loadu(bias + d); + data.store(out + d); + } + for (; d < size; ++d) { + out[d] = input[d] + bias[d]; + } +} + inline void unpack_B( at::BFloat16* __restrict__ Btmp, const at::Float8_e4m3fn* __restrict__ packed_B, @@ -913,6 +930,7 @@ void tinygemm_kernel( scalar_t* __restrict__ C, scalar_t* __restrict__ Btmp, float* __restrict__ Ctmp, + const float* __restrict__ Bbias, const float* __restrict__ scale, int64_t M, int64_t N, @@ -923,6 +941,11 @@ void tinygemm_kernel( bool brg, int64_t block_size_K, bool do_unpack) { + if (Bbias != nullptr) { + tinygemm_kernel( + A, B, C, Btmp, Ctmp, scale, Bbias, M, N, K, lda, ldb, ldc, brg, block_size_K, do_unpack); + return; + } tinygemm_kernel( A, B, C, Btmp, Ctmp, scale, nullptr, M, N, K, lda, ldb, ldc, brg, block_size_K, do_unpack); } @@ -952,6 +975,7 @@ void tinygemm_kernel( scalar_t* __restrict__ C, scalar_t* __restrict__ Btmp, float* __restrict__ Ctmp, + const float* __restrict__ Bbias, const uint8_t* __restrict__ scale, int64_t M, int64_t N, @@ -962,10 +986,68 @@ void tinygemm_kernel( bool brg, int64_t block_size_K, bool do_unpack) { + if (Bbias != nullptr) { + tinygemm_kernel( + A, B, C, Btmp, Ctmp, scale, Bbias, M, N, K, lda, ldb, ldc, brg, block_size_K, do_unpack); + return; + } tinygemm_kernel( A, B, C, Btmp, Ctmp, scale, nullptr, M, N, K, lda, ldb, ldc, brg, block_size_K, do_unpack); } +// tinygemm interface +template +void tinygemm_kernel( + const scalar_t* __restrict__ A, + const at::Float8_e4m3fn* __restrict__ B, + float* __restrict__ C, + scalar_t* __restrict__ Btmp, + const float* __restrict__ Bbias, + const float* __restrict__ scale, + int64_t M, + int64_t N, + int64_t K, + int64_t lda, + int64_t ldb, + int64_t ldc, + bool brg, + int64_t block_size_K, + bool do_unpack) { + if (Bbias != nullptr) { + tinygemm_kernel( + A, B, C, Btmp, scale, Bbias, M, N, K, lda, ldb, ldc, brg, block_size_K, do_unpack); + return; + } + tinygemm_kernel( + A, B, C, Btmp, scale, nullptr, M, N, K, lda, ldb, ldc, brg, block_size_K, do_unpack); +} + +template +void tinygemm_kernel( + const scalar_t* __restrict__ A, + const uint8_t* __restrict__ B, + float* __restrict__ C, + scalar_t* __restrict__ Btmp, + const float* __restrict__ Bbias, + const uint8_t* __restrict__ scale, + int64_t M, + int64_t N, + int64_t K, + int64_t lda, + int64_t ldb, + int64_t ldc, + bool brg, + int64_t block_size_K, + bool do_unpack) { + if (Bbias != nullptr) { + tinygemm_kernel( + A, B, C, Btmp, scale, Bbias, M, N, K, lda, ldb, ldc, brg, block_size_K, do_unpack); + return; + } + tinygemm_kernel( + A, B, C, Btmp, scale, nullptr, M, N, K, lda, ldb, ldc, brg, block_size_K, do_unpack); +} + #define INSTANTIATE_TINYGEMM_TEMPLATE(TYPE_A, TYPE_B, TYPE_S) \ template void tinygemm_kernel( \ const TYPE_A* __restrict__ A, \ @@ -973,6 +1055,7 @@ void tinygemm_kernel( TYPE_A* __restrict__ C, \ TYPE_A* __restrict__ Btmp, \ float* __restrict__ Ctmp, \ + const float* __restrict__ Bbias, \ const TYPE_S* __restrict__ scale, \ int64_t M, \ int64_t N, \ diff --git a/sgl-kernel/csrc/cpu/moe.cpp b/sgl-kernel/csrc/cpu/moe.cpp index d1740c981..797a6edc7 100644 --- a/sgl-kernel/csrc/cpu/moe.cpp +++ b/sgl-kernel/csrc/cpu/moe.cpp @@ -158,6 +158,56 @@ inline void silu_and_mul( } } +template +inline void clamp_sigmoid_and_mul( + scalar_t* __restrict__ output, + const float* __restrict__ input0, + int64_t m_size, + int64_t N, + const float alpha, + const float limit, + int64_t offset) { + using bVec = at::vec::Vectorized; + using fVec = at::vec::Vectorized; + + const fVec one = fVec(1.f); + const fVec zero = fVec(0.f); + const fVec limit_v = fVec(limit); + const fVec nlimit_v = fVec(-limit); + const fVec alpha_v = fVec(alpha); + + // no remainder + for (int64_t m = 0; m < m_size; ++m) { + scalar_t* __restrict__ out = output + m * N; + const float* __restrict__ cur_ptr = input0 + m * BLOCK_N; + for (int64_t d = 0; d < BLOCK_N; d += bVec::size()) { + float tmp_glu0[fVec::size()]; // 16 + float tmp_linear0[fVec::size()]; // 16 + + // interleaved: x[2i] = glu, x[2i+1] = linear + for (int j = 0; j < fVec::size(); ++j) { + // x0 [0,2,..30] + tmp_glu0[j] = cur_ptr[d + j * 2]; + // y0 [1,3,...31] + tmp_linear0[j] = cur_ptr[d + j * 2 + 1]; + } + fVec x0 = fVec::loadu(tmp_glu0); + fVec y0 = fVec::loadu(tmp_linear0); + + // clamp + x0 = at::vec::minimum(x0, limit_v); + y0 = at::vec::minimum(limit_v, at::vec::maximum(nlimit_v, y0)); + // x * sigmoid(x * alpha) + x0 = x0 / (one + (x0 * alpha_v).neg().exp_u20()); + // (y + 1) * x + y0 = y0 + one; + x0 = x0 * y0; + // // convert + convert_from_float_and_store(out + d / 2 + offset, x0); + } + } +} + template struct tinygemm_kernel_nn2 { static inline void apply( @@ -450,6 +500,8 @@ void fused_experts_kernel_impl( const scalar_t* __restrict__ input, const scalar_t* __restrict__ packed_w1, const scalar_t* __restrict__ packed_w2, + const float* __restrict__ w1_bias, + const float* __restrict__ w2_bias, const float* __restrict__ topk_weights, const int32_t* __restrict__ sorted_ids, const int32_t* __restrict__ expert_ids, @@ -459,7 +511,11 @@ void fused_experts_kernel_impl( int64_t K, int64_t E, int64_t topk, - int64_t num_tokens_post_pad) { + int64_t num_tokens_post_pad, + float alpha, + float limit, + CPUActMethod act_func, + bool with_bias) { // handle 2 tiles per block constexpr int64_t BLOCK_M = block_size_m(); constexpr int64_t BLOCK_N = block_size_n(); @@ -494,6 +550,8 @@ void fused_experts_kernel_impl( int32_t expert_id = expert_ids[mb]; const scalar_t* __restrict__ B0 = packed_w1 + expert_id * stride_e + nb_upper * BLOCK_N * stride_n; const scalar_t* __restrict__ B1 = packed_w1 + expert_id * stride_e + nb_lower * BLOCK_N * stride_n; + const float* __restrict__ B0_bias = w1_bias + expert_id * 2 * N + nb_upper * BLOCK_N; + const float* __restrict__ B1_bias = w1_bias + expert_id * 2 * N + nb_lower * BLOCK_N; int64_t m_size = offsets[mb + 1] - offsets[mb]; @@ -533,23 +591,58 @@ void fused_experts_kernel_impl( /* B */ B1, /* C */ C1); - // 1.d silu and mul - const int64_t offset = offsets[mb]; - silu_and_mul(ic1 + offset * N + nb * BLOCK_N, C0, C1, m_size, N); } else { - // fused 1.bcd: silu_and_mul(A @ B0, A @ B1) const int64_t offset = offsets[mb]; - tinygemm_kernel( - /* A */ A, - /* B0 */ B0, - /* B1 */ B1, - /* C */ ic1 + offset * N + nb * BLOCK_N, - /* M */ m_size, - /* N */ n_size, - /* K */ K, - /* lda */ K, - /* ldb */ n_size, - /* ldc */ N); + if (act_func == CPUActMethod::swiglu) { + tinygemm_kernel( + /* A */ A, + /* B */ B0, + /* C */ C0, + /* M */ m_size, + /* N */ n_size, + /* K */ K, + /* lda */ K, + /* ldb */ n_size, + /* ldc */ BLOCK_N); + tinygemm_kernel( + /* A */ A, + /* B */ B1, + /* C */ C1, + /* M */ m_size, + /* N */ n_size, + /* K */ K, + /* lda */ K, + /* ldb */ n_size, + /* ldc */ BLOCK_N); + } else { + // fused 1.bcd: silu_and_mul(A @ B0, A @ B1) + tinygemm_kernel( + /* A */ A, + /* B0 */ B0, + /* B1 */ B1, + /* C */ ic1 + offset * N + nb * BLOCK_N, + /* M */ m_size, + /* N */ n_size, + /* K */ K, + /* lda */ K, + /* ldb */ n_size, + /* ldc */ N); + } + } + if (with_bias) { + for (int64_t m = 0; m < m_size; ++m) { + add_bias_stub(C0 + m * BLOCK_N, B0_bias, n_size); + add_bias_stub(C1 + m * BLOCK_N, B1_bias, n_size); + } + } + // 1.d silu and mul + const int64_t offset = offsets[mb]; + if (act_func == CPUActMethod::silu_and_mul && use_brgemm) { + silu_and_mul(ic1 + offset * N + nb * BLOCK_N, C0, C1, m_size, N); + } else if (act_func == CPUActMethod::swiglu) { + clamp_sigmoid_and_mul(ic1 + offset * N, C0, m_size, N, alpha, limit, 0 + nb * BLOCK_N / 2); + clamp_sigmoid_and_mul( + ic1 + offset * N, C1, m_size, N, alpha, limit, N / 2 + nb * BLOCK_N / 2); } }); @@ -586,6 +679,7 @@ void fused_experts_kernel_impl( // B shape [IC, n_size] in vnni format int32_t expert_id = expert_ids[mb]; const scalar_t* __restrict__ B = packed_w2 + expert_id * stride_e2 + nb * BLOCK_N * stride_oc; + const float* __restrict__ B_bias = w2_bias + expert_id * OC + nb * BLOCK_N; // 2.a gemm: C = A @ B if (use_brgemm) { @@ -613,6 +707,11 @@ void fused_experts_kernel_impl( /* ldc */ BLOCK_N); } + if (with_bias) { + for (int64_t m = 0; m < m_size; ++m) { + add_bias_stub(C + m * BLOCK_N, B_bias, n_size); + } + } // 2.b copy from C to ic2 in original order // and also mul topk_weights in float32 for (int64_t m = 0; m < m_size; ++m) { @@ -804,21 +903,38 @@ void shared_expert_kernel_impl( } // anonymous namespace // common checks +template static inline void check_moe_scales( - bool use_int8_w8a8, - bool use_fp8_w8a16, const std::optional& w1_scale, const std::optional& w2_scale, const std::optional> block_size) { - if (use_int8_w8a8) { + if constexpr (quant == CPUQuantMethod::INT8_W8A8) { TORCH_CHECK(w1_scale.has_value(), "missing w1_scale for int8 w8a8."); TORCH_CHECK(w2_scale.has_value(), "missing w2_scale for int8 w8a8."); - } - if (use_fp8_w8a16) { + } else if constexpr (quant == CPUQuantMethod::FP8_W8A16) { TORCH_CHECK(w1_scale.has_value(), "missing w1_scale for fp8 w8a16."); TORCH_CHECK(w2_scale.has_value(), "missing w2_scale for fp8 w8a16."); TORCH_CHECK(block_size.has_value(), "missing block_size for fp8 w8a16."); TORCH_CHECK(block_size.value().size() == 2, "expect block_size.size() to be 2."); + } else if constexpr (quant == CPUQuantMethod::MXFP4) { + TORCH_CHECK(w1_scale.has_value(), "missing w1_scale for mxfp4."); + TORCH_CHECK(w2_scale.has_value(), "missing w2_scale for mxfp4."); + TORCH_CHECK(w1_scale.value().scalar_type() == at::kByte, "expect w1_scale to be uint8."); + TORCH_CHECK(w2_scale.value().scalar_type() == at::kByte, "expect w2_scale to be uint8."); + } +} + +static inline void check_moe_scales( + int64_t moe_comp_method, + const std::optional& w1_scale, + const std::optional& w2_scale, + const std::optional> block_size) { + if (moe_comp_method == CPUQuantMethod::INT8_W8A8) { + check_moe_scales(w1_scale, w2_scale, block_size); + } else if (moe_comp_method == CPUQuantMethod::FP8_W8A16) { + check_moe_scales(w1_scale, w2_scale, block_size); + } else if (moe_comp_method == CPUQuantMethod::MXFP4) { + check_moe_scales(w1_scale, w2_scale, block_size); } } @@ -834,8 +950,8 @@ static inline void check_moe_scales( TORCH_CHECK(w2s.size(DIM1) == div_up(N, block_size_K)) // hidden_states: [M, K] -// w1: [E, 2N, K] -// w2: [E, K, N] +// w1: [E, 2N, K] or [E, 2N, K / 2] for uint8 +// w2: [E, K, N] or [E, K, N / 2] for uint8 // topk_weights: [M, topk] // topk_ids: [M, topk] (int32_t) // @@ -853,6 +969,10 @@ at::Tensor fused_experts_cpu( const std::optional& w1_zero, const std::optional& w2_zero, const std::optional> block_size, + const std::optional& w1_bias, + const std::optional& w2_bias, + const std::optional& alpha, + const std::optional& limit, bool is_vnni) { auto packed_w1 = is_vnni ? w1 : convert_weight_packed(w1); auto packed_w2 = is_vnni ? w2 : convert_weight_packed(w2); @@ -895,8 +1015,8 @@ at::Tensor fused_experts_cpu( int64_t topk = topk_weights_.size(1); // we use int32_t compensation for int8 w8a8 - int64_t packed_K = get_row_size(K, moe_comp_method == CPUQuantMethod::INT8_W8A8); - int64_t packed_N = get_row_size(N, moe_comp_method == CPUQuantMethod::INT8_W8A8); + int64_t packed_K = get_row_size(static_cast(moe_comp_method), K); + int64_t packed_N = get_row_size(static_cast(moe_comp_method), N); // check weight shapes CHECK_EQ(w2.size(0), E); @@ -906,12 +1026,7 @@ at::Tensor fused_experts_cpu( CHECK_EQ(packed_w2.size(2), packed_N / (moe_comp_method == CPUQuantMethod::INT4_W4A8 ? 2 : 1)); } // check scales - check_moe_scales( - moe_comp_method == CPUQuantMethod::INT8_W8A8, - moe_comp_method == CPUQuantMethod::FP8_W8A16, - w1_scale, - w2_scale, - block_size); + check_moe_scales(moe_comp_method, w1_scale, w2_scale, block_size); at::Tensor out_hidden_states = inplace ? hidden_states : at::empty_like(hidden_states); @@ -963,7 +1078,7 @@ at::Tensor fused_experts_cpu( // 5. Aq_tmp : [M, K] or [M * topk, N] // 6. As_tmp : [M * topk] // - // for fp8 w8a16: + // for fp8 w8a16 and mxfp4: // 7. intermediate_cache0 : [M * topk, 2N] // 8. B_tmp : [T, MAX_CACHE_BLOCK_SIZE, BLOCK_N, std::max(K, N)] // @@ -976,7 +1091,7 @@ at::Tensor fused_experts_cpu( if (moe_comp_method == CPUQuantMethod::INT8_W8A8) { buffer_size_nbytes += std::max(M * K, M * topk * N) + M * topk * sizeof(float); } - if (moe_comp_method == CPUQuantMethod::FP8_W8A16) { + if (moe_comp_method == CPUQuantMethod::FP8_W8A16 || moe_comp_method == CPUQuantMethod::MXFP4) { buffer_size_nbytes += M * topk * 2 * N * 2 + num_threads * MAX_CACHE_BLOCK_SIZE * BLOCK_N * std::max(K, N) * 2; } if (moe_comp_method == CPUQuantMethod::INT4_W4A8) { @@ -1029,9 +1144,11 @@ at::Tensor fused_experts_cpu( float* __restrict__ C_tmp = (float*)((void*)(A_tmp + num_threads * BLOCK_M * K)); scalar_t* __restrict__ intermediate_cache0 = (scalar_t*)((void*)(C_tmp + num_threads * 2 * BLOCK_M * BLOCK_N)); scalar_t* __restrict__ B_tmp = (scalar_t*)((void*)(intermediate_cache0 + M * topk * 2 * N)); + bool with_bias = w1_bias.has_value(); + auto act_func = alpha.has_value() && limit.has_value() ? CPUActMethod::swiglu : CPUActMethod::silu_and_mul; CHECK_MOE_SCALES_FP8(1, 2); - fused_experts_fp8_kernel_impl( + fused_experts_fp_kernel_impl( out_hidden_states.data_ptr(), intermediate_cache0, intermediate_cache1, @@ -1042,6 +1159,8 @@ at::Tensor fused_experts_cpu( hidden_states.data_ptr(), packed_w1.data_ptr(), packed_w2.data_ptr(), + with_bias ? w1_bias.value().data_ptr() : nullptr, + with_bias ? w2_bias.value().data_ptr() : nullptr, w1s.data_ptr(), w2s.data_ptr(), block_size_N, @@ -1055,7 +1174,56 @@ at::Tensor fused_experts_cpu( K, E, topk, - num_tokens_post_pad); + num_tokens_post_pad, + alpha.has_value() ? float(alpha.value()) : 0, + limit.has_value() ? float(limit.value()) : 0, + act_func, + with_bias); + } else if (moe_comp_method == CPUQuantMethod::MXFP4) { + scalar_t* __restrict__ A_tmp = (scalar_t*)((void*)(intermediate_cache2 + M * topk * K)); + float* __restrict__ C_tmp = (float*)((void*)(A_tmp + num_threads * BLOCK_M * K)); + scalar_t* __restrict__ intermediate_cache0 = (scalar_t*)((void*)(C_tmp + num_threads * 2 * BLOCK_M * BLOCK_N)); + scalar_t* __restrict__ B_tmp = (scalar_t*)((void*)(intermediate_cache0 + M * topk * 2 * N)); + bool with_bias = w1_bias.has_value(); + auto act_func = alpha.has_value() && limit.has_value() ? CPUActMethod::swiglu : CPUActMethod::silu_and_mul; + + // mxfp4 supports only group size of 32 (2^5) + constexpr int64_t group_size = 32; + auto w1s = w1_scale.value(); + auto w2s = w2_scale.value(); + TORCH_CHECK(w1s.numel(), E * 2 * N * K >> 5); + TORCH_CHECK(w2s.numel(), E * K * N >> 5); + fused_experts_fp_kernel_impl( + out_hidden_states.data_ptr(), + intermediate_cache0, + intermediate_cache1, + intermediate_cache2, + A_tmp, + B_tmp, + C_tmp, + hidden_states.data_ptr(), + packed_w1.data_ptr(), + packed_w2.data_ptr(), + with_bias ? w1_bias.value().data_ptr() : nullptr, + with_bias ? w2_bias.value().data_ptr() : nullptr, + w1s.data_ptr(), + w2s.data_ptr(), + /*block_size_N*/ 1, + /*block_size_K*/ group_size, + topk_weights_.data_ptr(), + sorted_ids, + expert_ids, + offsets, + M, + N, + K, + E, + topk, + num_tokens_post_pad, + alpha.has_value() ? float(alpha.value()) : 0, + limit.has_value() ? float(limit.value()) : 0, + act_func, + with_bias); } else if (moe_comp_method == CPUQuantMethod::INT4_W4A8) { uint8_t* __restrict__ A_tmp = (uint8_t*)((void*)(intermediate_cache2 + M * topk * K)); float* __restrict__ C_tmp = (float*)((void*)(A_tmp + num_threads * BLOCK_M * K)); @@ -1101,6 +1269,8 @@ at::Tensor fused_experts_cpu( } else { scalar_t* __restrict__ A_tmp = intermediate_cache2 + M * topk * K; float* __restrict__ C_tmp = (float*)((void*)(A_tmp + num_threads * BLOCK_M * K)); + bool with_bias = w1_bias.has_value(); + auto act_func = alpha.has_value() && limit.has_value() ? CPUActMethod::swiglu : CPUActMethod::silu_and_mul; fused_experts_kernel_impl( out_hidden_states.data_ptr(), @@ -1111,6 +1281,8 @@ at::Tensor fused_experts_cpu( hidden_states.data_ptr(), packed_w1.data_ptr(), packed_w2.data_ptr(), + with_bias ? w1_bias.value().data_ptr() : nullptr, + with_bias ? w2_bias.value().data_ptr() : nullptr, topk_weights_.data_ptr(), sorted_ids, expert_ids, @@ -1120,7 +1292,11 @@ at::Tensor fused_experts_cpu( K, E, topk, - num_tokens_post_pad); + num_tokens_post_pad, + alpha.has_value() ? float(alpha.value()) : 0, + limit.has_value() ? float(limit.value()) : 0, + act_func, + with_bias); } }); return out_hidden_states; @@ -1183,7 +1359,11 @@ at::Tensor shared_expert_cpu( CHECK_EQ(packed_w2.size(1), packed_N); // check scales - check_moe_scales(use_int8_w8a8, use_fp8_w8a16, w1_scale, w2_scale, block_size); + if (use_int8_w8a8) { + check_moe_scales(w1_scale, w2_scale, block_size); + } else if (use_fp8_w8a16) { + check_moe_scales(w1_scale, w2_scale, block_size); + } at::Tensor out_hidden_states = inplace ? hidden_states : at::empty_like(hidden_states); diff --git a/sgl-kernel/csrc/cpu/moe.h b/sgl-kernel/csrc/cpu/moe.h index f9d8afca9..71f829c26 100644 --- a/sgl-kernel/csrc/cpu/moe.h +++ b/sgl-kernel/csrc/cpu/moe.h @@ -171,3 +171,109 @@ inline void silu_and_mul_stub( out_vec.store(out + d); } } + +template +inline void copy_mul_stub(scalar_t* __restrict__ out, const float* __restrict__ input, float weight, int64_t size) { + using bVec = at::vec::Vectorized; + using fVec = at::vec::Vectorized; + constexpr int kVecSize = bVec::size(); + const fVec weight_vec = fVec(weight); + int64_t d; +#pragma GCC unroll 4 + for (d = 0; d <= size - kVecSize; d += kVecSize) { + fVec data0 = fVec::loadu(input + d) * weight_vec; + fVec data1 = fVec::loadu(input + d + fVec::size()) * weight_vec; + bVec out_vec = convert_from_float_ext(data0, data1); + out_vec.store(out + d); + } + for (; d < size; ++d) { + out[d] = static_cast(input[d] * weight); + } +} + +// input = input + input2 +inline void add_bias_stub(float* __restrict__ input, const float* __restrict__ input2, int64_t size) { + using fVec = at::vec::Vectorized; + constexpr int kVecSize = fVec::size(); + int64_t d; +#pragma GCC unroll 4 + for (d = 0; d <= size - kVecSize; d += kVecSize) { + fVec x_fvec = fVec::loadu(input + d); + fVec y_fvec = fVec::loadu(input2 + d); + x_fvec = x_fvec + y_fvec; + x_fvec.store(input + d); + } + for (; d < size; ++d) { + input[d] = input[d] + input2[d]; + } +} + +template +inline void copy_mul_stub(scalar_t* __restrict__ out, const scalar_t* __restrict__ input, float weight, int64_t size) { + using bVec = at::vec::Vectorized; + using fVec = at::vec::Vectorized; + constexpr int kVecSize = bVec::size(); + const fVec weight_vec = fVec(weight); + int64_t d; +#pragma GCC unroll 4 + for (d = 0; d <= size - kVecSize; d += kVecSize) { + bVec x = bVec::loadu(input + d); + fVec x0, x1; + std::tie(x0, x1) = at::vec::convert_to_float(x); + x0 = x0 * weight_vec; + x1 = x1 * weight_vec; + bVec out_vec = convert_from_float_ext(x0, x1); + out_vec.store(out + d); + } + for (; d < size; ++d) { + out[d] = static_cast(input[d] * weight); + } +} + +template +inline void clamp_sigmoid_and_mul_stub( + scalar_t* __restrict__ out, + const scalar_t* __restrict__ input, + int64_t size, + const float alpha, + const float limit) { + using bVec = at::vec::Vectorized; + using fVec = at::vec::Vectorized; + const fVec one = fVec(1.f); + const fVec zero = fVec(0.f); + const fVec limit_v = fVec(limit); + const fVec nlimit_v = fVec(-limit); + const fVec alpha_v = fVec(alpha); + + // no remainder +#pragma GCC unroll 4 + for (int64_t d = 0; d < size; d += bVec::size()) { + bVec x = bVec::loadu(input + d); + fVec x0_, y0_; + std::tie(x0_, y0_) = at::vec::convert_to_float(x); + float tmp_buffer[fVec::size() * 2]; // 32 + float tmp_glu[fVec::size()]; // 16 + float tmp_linear[fVec::size()]; // 16 + x0_.store(tmp_buffer); + y0_.store(tmp_buffer + fVec::size()); + // interleaved: x[2i] = glu, x[2i+1] = linear + for (int j = 0; j < fVec::size(); ++j) { + // x0 [0,2,..30] + tmp_glu[j] = tmp_buffer[j * 2]; + // y0 [1,3,...31] + tmp_linear[j] = tmp_buffer[j * 2 + 1]; + } + fVec x0 = fVec::loadu(tmp_glu); + fVec y0 = fVec::loadu(tmp_linear); + + // clamp + x0 = at::vec::minimum(x0, limit_v); + y0 = at::vec::minimum(limit_v, at::vec::maximum(nlimit_v, y0)); + // x * sigmoid(x * alpha) + x0 = x0 / (one + (x0 * alpha_v).neg().exp_u20()); + // (y + 1) * x + y0 = y0 + one; + x0 = x0 * y0; + convert_from_float_and_store(out + d / 2, x0); + } +} diff --git a/sgl-kernel/csrc/cpu/moe_fp8.cpp b/sgl-kernel/csrc/cpu/moe_fp8.cpp index ecd4e2adb..53a643d47 100644 --- a/sgl-kernel/csrc/cpu/moe_fp8.cpp +++ b/sgl-kernel/csrc/cpu/moe_fp8.cpp @@ -2,8 +2,8 @@ #include "gemm.h" #include "moe.h" -template -void fused_experts_fp8_kernel_impl( +template +void fused_experts_fp_kernel_impl( scalar_t* __restrict__ output, scalar_t* __restrict__ ic0, scalar_t* __restrict__ ic1, @@ -12,10 +12,12 @@ void fused_experts_fp8_kernel_impl( scalar_t* __restrict__ B_tmp, float* __restrict__ C_tmp, const scalar_t* __restrict__ input, - const at::Float8_e4m3fn* __restrict__ packed_w1, - const at::Float8_e4m3fn* __restrict__ packed_w2, - const float* __restrict__ w1s, - const float* __restrict__ w2s, + const packed_t* __restrict__ packed_w1, + const packed_t* __restrict__ packed_w2, + const float* __restrict__ w1_bias, + const float* __restrict__ w2_bias, + const param_t* __restrict__ w1s, + const param_t* __restrict__ w2s, int64_t block_size_N, int64_t block_size_K, const float* __restrict__ topk_weights, @@ -27,7 +29,11 @@ void fused_experts_fp8_kernel_impl( int64_t K, int64_t E, int64_t topk, - int64_t num_tokens_post_pad) { + int64_t num_tokens_post_pad, + float alpha, + float limit, + CPUActMethod act_func, + bool with_bias) { constexpr int64_t BLOCK_M = block_size_m(); constexpr int64_t BLOCK_N = block_size_n(); @@ -37,12 +43,20 @@ void fused_experts_fp8_kernel_impl( int64_t scale_size_N = div_up(2 * N, block_size_N); int64_t scale_size_K = div_up(K, block_size_K); int64_t blocks_n_per_group = block_size_N / BLOCK_N; + std::function scale_offset_per_block; + if constexpr (is_mxfp4) { + scale_offset_per_block = [&](int64_t a) { return a * BLOCK_N; }; + } else { + scale_offset_per_block = [&](int64_t a) { return a / blocks_n_per_group; }; + } - const int64_t stride_e = 2 * N * K; - const int64_t stride_n = K; + const int64_t packed_K = get_row_size(K); + + const int64_t stride_e = 2 * N * packed_K; + const int64_t stride_n = packed_K; int64_t avg_M = std::max(int64_t(1), M * topk / E); - const bool use_brgemm = can_use_brgemm(avg_M); + const bool use_brgemm = can_use_brgemm(avg_M); int64_t B_tmp_size_per_thread = MAX_CACHE_BLOCK_SIZE * BLOCK_N * std::max(K, N); @@ -52,14 +66,15 @@ void fused_experts_fp8_kernel_impl( int tid = get_thread_num(); scalar_t* __restrict__ A = A_tmp + tid * BLOCK_M * K; - loop_2d(mb0, mb1, nb0, nb1, BLOCK_N * K, [&](int64_t mb, int64_t nb, int64_t nb_offset) { + loop_2d(mb0, mb1, nb0, nb1, BLOCK_N * K, [&](int64_t mb, int64_t nb, int64_t nb_offset) { int64_t n_size = std::min(2 * N - nb * BLOCK_N, BLOCK_N); // B shape [K, n_size] in vnni format int32_t expert_id = expert_ids[mb]; - const at::Float8_e4m3fn* __restrict__ B = packed_w1 + expert_id * stride_e + nb * BLOCK_N * stride_n; - const float* __restrict__ Bs = - w1s + expert_id * scale_size_N * scale_size_K + (nb / blocks_n_per_group) * scale_size_K; + const packed_t* __restrict__ B = packed_w1 + expert_id * stride_e + nb * BLOCK_N * stride_n; + const param_t* __restrict__ Bs = + w1s + expert_id * scale_size_N * scale_size_K + scale_offset_per_block(nb) * scale_size_K; + const float* __restrict__ B_bias = with_bias ? w1_bias + expert_id * 2 * N + nb * BLOCK_N : nullptr; // do unpacking for the first row or a new expert int32_t pre_expert_id = mb == 0 ? -1 : expert_ids[mb - 1]; @@ -83,6 +98,7 @@ void fused_experts_fp8_kernel_impl( /* C */ ic0 + offset * 2 * N + nb * BLOCK_N, /* Btmp */ B_tmp + tid * B_tmp_size_per_thread + nb_offset * BLOCK_N * K, /* Ctmp */ C_tmp + tid * 2 * BLOCK_M * BLOCK_N, + /* Bbias */ B_bias, /* scale */ Bs, /* M */ m_size, /* N */ n_size, @@ -101,12 +117,20 @@ void fused_experts_fp8_kernel_impl( }); // stage 1.5: intermediate_cache1 = silu(intermediate_cache0) - at::parallel_for(0, M * topk, 0, [&](int64_t begin, int64_t end) { - for (int64_t m = begin; m < end; ++m) { - silu_and_mul_stub(ic1 + m * N, ic0 + m * 2 * N, ic0 + m * 2 * N + N, N); - } - }); - + if (act_func == CPUActMethod::silu_and_mul) { + at::parallel_for(0, M * topk, 0, [&](int64_t begin, int64_t end) { + for (int64_t m = begin; m < end; ++m) { + silu_and_mul_stub(ic1 + m * N, ic0 + m * 2 * N, ic0 + m * 2 * N + N, N); + } + }); + } else if (act_func == CPUActMethod::swiglu) { + at::parallel_for(0, M * topk, 0, [&](int64_t begin, int64_t end) { + for (int64_t m = begin; m < end; ++m) { + clamp_sigmoid_and_mul_stub(ic1 + m * N, ic0 + m * 2 * N, N, alpha, limit); + clamp_sigmoid_and_mul_stub(ic1 + m * N + N / 2, ic0 + m * 2 * N + N, N, alpha, limit); + } + }); + } // stage 2: intermediate_cache2 = intermediate_cache1 @ w2 // w2 : [E, K, N] as [E, OC, IC] const int64_t OC = K; // rename K as OC @@ -115,15 +139,16 @@ void fused_experts_fp8_kernel_impl( const int64_t NB2 = div_up(OC, BLOCK_N); scale_size_N = div_up(K, block_size_N); scale_size_K = div_up(N, block_size_K); - const int64_t stride_e2 = OC * IC; - const int64_t stride_oc = IC; + const int64_t packed_IC = get_row_size(IC); + const int64_t stride_e2 = OC * packed_IC; + const int64_t stride_oc = packed_IC; // parallel on [MB2, NB2] parallel_2d(MB2, NB2, [&](int64_t mb0, int64_t mb1, int64_t nb0, int64_t nb1) { int tid = get_thread_num(); alignas(64) scalar_t C[BLOCK_M * BLOCK_K]; - loop_2d(mb0, mb1, nb0, nb1, BLOCK_N * IC, [&](int64_t mb, int64_t nb, int64_t nb_offset) { + loop_2d(mb0, mb1, nb0, nb1, BLOCK_N * IC, [&](int64_t mb, int64_t nb, int64_t nb_offset) { int64_t m_size = offsets[mb + 1] - offsets[mb]; int64_t n_size = std::min(OC - nb * BLOCK_N, BLOCK_N); @@ -134,9 +159,10 @@ void fused_experts_fp8_kernel_impl( // B shape [IC, n_size] in vnni format int32_t expert_id = expert_ids[mb]; - const at::Float8_e4m3fn* __restrict__ B = packed_w2 + expert_id * stride_e2 + nb * BLOCK_N * stride_oc; - const float* __restrict__ Bs = - w2s + expert_id * scale_size_N * scale_size_K + (nb / blocks_n_per_group) * scale_size_K; + const packed_t* __restrict__ B = packed_w2 + expert_id * stride_e2 + nb * BLOCK_N * stride_oc; + const param_t* __restrict__ Bs = + w2s + expert_id * scale_size_N * scale_size_K + scale_offset_per_block(nb) * scale_size_K; + const float* __restrict__ B_bias = with_bias ? w2_bias + expert_id * OC + nb * BLOCK_N : nullptr; // do unpacking for the first row or a new expert int32_t pre_expert_id = mb == 0 ? -1 : expert_ids[mb - 1]; @@ -148,6 +174,7 @@ void fused_experts_fp8_kernel_impl( /* C */ C, /* Btmp */ B_tmp + tid * B_tmp_size_per_thread + nb_offset * BLOCK_N * IC, /* Ctmp */ C_tmp + tid * 2 * BLOCK_M * BLOCK_N, + /* Bbias */ B_bias, /* scale */ Bs, /* M */ m_size, /* N */ n_size, @@ -182,35 +209,43 @@ void fused_experts_fp8_kernel_impl( }); } -#define INSTANTIATE_MOE_FP8_TEMPLATE(TYPE) \ - template void fused_experts_fp8_kernel_impl( \ - TYPE* __restrict__ output, \ - TYPE* __restrict__ ic0, \ - TYPE* __restrict__ ic1, \ - TYPE* __restrict__ ic2, \ - TYPE* __restrict__ A_tmp, \ - TYPE* __restrict__ B_tmp, \ - float* __restrict__ C_tmp, \ - const TYPE* __restrict__ input, \ - const at::Float8_e4m3fn* __restrict__ packed_w1, \ - const at::Float8_e4m3fn* __restrict__ packed_w2, \ - const float* __restrict__ w1s, \ - const float* __restrict__ w2s, \ - int64_t block_size_N, \ - int64_t block_size_K, \ - const float* __restrict__ topk_weights, \ - const int32_t* __restrict__ sorted_ids, \ - const int32_t* __restrict__ expert_ids, \ - const int32_t* __restrict__ offsets, \ - int64_t M, \ - int64_t N, \ - int64_t K, \ - int64_t E, \ - int64_t topk, \ - int64_t num_tokens_post_pad) +#define INSTANTIATE_MOE_FP_TEMPLATE(TYPE1, TYPE2, TYPE3, IS_MXFP4) \ + template void fused_experts_fp_kernel_impl( \ + TYPE1* __restrict__ output, \ + TYPE1* __restrict__ ic0, \ + TYPE1* __restrict__ ic1, \ + TYPE1* __restrict__ ic2, \ + TYPE1* __restrict__ A_tmp, \ + TYPE1* __restrict__ B_tmp, \ + float* __restrict__ C_tmp, \ + const TYPE1* __restrict__ input, \ + const TYPE2* __restrict__ packed_w1, \ + const TYPE2* __restrict__ packed_w2, \ + const float* __restrict__ w1_bias, \ + const float* __restrict__ w2_bias, \ + const TYPE3* __restrict__ w1s, \ + const TYPE3* __restrict__ w2s, \ + int64_t block_size_N, \ + int64_t block_size_K, \ + const float* __restrict__ topk_weights, \ + const int32_t* __restrict__ sorted_ids, \ + const int32_t* __restrict__ expert_ids, \ + const int32_t* __restrict__ offsets, \ + int64_t M, \ + int64_t N, \ + int64_t K, \ + int64_t E, \ + int64_t topk, \ + int64_t num_tokens_post_pad, \ + float alpha, \ + float limit, \ + CPUActMethod act_func, \ + bool with_bias) -INSTANTIATE_MOE_FP8_TEMPLATE(at::BFloat16); -INSTANTIATE_MOE_FP8_TEMPLATE(at::Half); +INSTANTIATE_MOE_FP_TEMPLATE(at::BFloat16, at::Float8_e4m3fn, float, false); +INSTANTIATE_MOE_FP_TEMPLATE(at::Half, at::Float8_e4m3fn, float, false); +INSTANTIATE_MOE_FP_TEMPLATE(at::BFloat16, uint8_t, uint8_t, true); +INSTANTIATE_MOE_FP_TEMPLATE(at::Half, uint8_t, uint8_t, true); template void shared_expert_fp8_kernel_impl( @@ -261,6 +296,7 @@ void shared_expert_fp8_kernel_impl( /* C */ ic0 + mb * BLOCK_M * 2 * N + nb * BLOCK_N, /* Btmp */ B_tmp + tid * B_tmp_size_per_thread + nb_offset * BLOCK_N * K, /* Ctmp */ C_tmp + tid * 2 * BLOCK_M * BLOCK_N, + /* Bbias */ nullptr, /* scale */ w1s + (nb / blocks_n_per_group) * scale_size_K, /* M */ m_size, /* N */ n_size, @@ -312,6 +348,7 @@ void shared_expert_fp8_kernel_impl( /* C */ C, /* Btmp */ B_tmp + tid * B_tmp_size_per_thread + nb_offset * BLOCK_N * IC, /* Ctmp */ C_tmp + tid * 2 * BLOCK_M * BLOCK_N, + /* Bbias */ nullptr, /* scale */ w2s + (nb / blocks_n_per_group) * scale_size_K, /* M */ m_size, /* N */ n_size, diff --git a/sgl-kernel/csrc/cpu/qkv_proj.cpp b/sgl-kernel/csrc/cpu/qkv_proj.cpp index 624c0f0f8..f929efb92 100644 --- a/sgl-kernel/csrc/cpu/qkv_proj.cpp +++ b/sgl-kernel/csrc/cpu/qkv_proj.cpp @@ -210,6 +210,7 @@ void segment_gemm_kernel_impl( /* C */ C + mb_start * ldc + local_nb_start, /* Btmp*/ Btmp + tid * BLOCK_N * K, /* Ctmp*/ Ctmp, + /*Bbias*/ nullptr, /* Bs */ Bs + (new_nb / blocks_n_per_group) * scale_size_K, /* M */ mb_size, /* N */ nb_size, diff --git a/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp b/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp index 5ed291986..3ec8fcdc2 100644 --- a/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp +++ b/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp @@ -100,7 +100,9 @@ void decode_attention_cpu( double sm_scale, double logit_cap, bool is_cross_attn, - std::optional encoder_lens); + int64_t slidling_window_size, + std::optional encoder_lens, + std::optional sinks); void extend_attention_cpu( at::Tensor& q_extend, @@ -118,7 +120,9 @@ void extend_attention_cpu( double sm_scale, double logit_cap, bool is_cross_attn, - std::optional encoder_lens); + int64_t sliding_window_size, + std::optional encoder_lens, + std::optional sinks); // flash attention at::Tensor flash_attn_varlen_func( @@ -215,6 +219,8 @@ at::Tensor fused_linear_sigmoid_mul( // bmm void bmm_cpu(at::Tensor& out, at::Tensor& mat1, at::Tensor& mat2, bool is_vnni, const std::optional& scale); +#if !defined(SGLANG_CPU_ARM64_SKIP_X86_ONLY_OPS) + // fused moe at::Tensor fused_experts_cpu( at::Tensor& hidden_states, @@ -229,9 +235,11 @@ at::Tensor fused_experts_cpu( const std::optional& w1_zero, const std::optional& w2_zero, const std::optional> block_size, + const std::optional& w1_bias, + const std::optional& w2_bias, + const std::optional& alpha, + const std::optional& limit, bool is_vnni); - -#if !defined(SGLANG_CPU_ARM64_SKIP_X86_ONLY_OPS) at::Tensor shared_expert_cpu( at::Tensor& hidden_states, at::Tensor& w1, @@ -313,6 +321,23 @@ at::Tensor causal_conv1d_update_cpu( const std::optional& conv_state_indices, int64_t pad_slot_id, bool is_vnni); +#else + +// fused moe +at::Tensor fused_experts_cpu( + at::Tensor& hidden_states, + at::Tensor& w1, + at::Tensor& w2, + at::Tensor& topk_weights, + at::Tensor& topk_ids, + bool inplace, + int64_t moe_comp_method, + const std::optional& w1_scale, + const std::optional& w2_scale, + const std::optional& w1_zero, + const std::optional& w2_zero, + const std::optional> block_size, + bool is_vnni); #endif // conv3d fast path for patch embedding @@ -478,15 +503,16 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) { "decode_attention_cpu(Tensor query, Tensor k_cache, Tensor v_cahce, Tensor(a!) output, Tensor? key, Tensor? " "value, " "Tensor loc, Tensor attn_logits, Tensor req_to_token, Tensor req_pool_indices, Tensor seq_lens, float sm_scale, " - "float logit_cap, bool is_cross_attn, Tensor? encoder_lens) -> ()"); + "float logit_cap, bool is_cross_attn, int sliding_window_size, Tensor? encoder_lens, Tensor? sinks) -> ()"); m.impl("decode_attention_cpu", torch::kCPU, &decode_attention_cpu); // extend m.def( "extend_attention_cpu(Tensor q_extend, Tensor? k_extend, Tensor? v_extend, Tensor(a!) o_extend, Tensor k_buffer, " "Tensor v_buffer, Tensor req_to_token, Tensor req_pool_indices, Tensor seq_lens, Tensor extend_seq_lens, Tensor " - "extend_start_loc, int max_len_extend, float sm_scale, float logit_cap, bool is_cross_attn, Tensor? " - "encoder_lens) -> ()"); + "extend_start_loc, int max_len_extend, float sm_scale, float logit_cap, bool is_cross_attn, int " + "sliding_window_size, Tensor? " + "encoder_lens, Tensor? sinks) -> ()"); m.impl("extend_attention_cpu", torch::kCPU, &extend_attention_cpu); // flash attn @@ -561,14 +587,14 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) { m.def("bmm_cpu(Tensor(a!) out, Tensor mat1, Tensor mat2, bool is_vnni, Tensor? scale) -> ()"); m.impl("bmm_cpu", torch::kCPU, &bmm_cpu); +#if !defined(SGLANG_CPU_ARM64_SKIP_X86_ONLY_OPS) // moe m.def( "fused_experts_cpu(Tensor hidden_states, Tensor w1, Tensor w2, Tensor topk_weights, Tensor topk_ids, bool " "inplace, int moe_comp_method, Tensor? w1_scale, Tensor? w2_scale, " - "Tensor? w1_zero, Tensor? w2_zero, int[]? block_size, bool is_vnni) -> Tensor"); + "Tensor? w1_zero, Tensor? w2_zero, int[]? block_size, Tensor? w1_bias, Tensor? w2_bias, float? alpha, float? " + "limit, bool is_vnni) -> Tensor"); m.impl("fused_experts_cpu", torch::kCPU, &fused_experts_cpu); - -#if !defined(SGLANG_CPU_ARM64_SKIP_X86_ONLY_OPS) // weight absorption m.def( "qkv_proj_with_rope(Tensor hidden_states, Tensor q_a_proj_weight, Tensor q_b_proj_weight, Tensor " @@ -607,6 +633,13 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) { "causal_conv1d_update_cpu(Tensor x, Tensor(a!) conv_states, Tensor weight, Tensor? bias, bool silu_activation," "Tensor? cache_seqlens, Tensor? conv_state_indices, int pad_slot_id, bool is_vnni) -> Tensor"); m.impl("causal_conv1d_update_cpu", torch::kCPU, &causal_conv1d_update_cpu); +#else + // moe + m.def( + "fused_experts_cpu(Tensor hidden_states, Tensor w1, Tensor w2, Tensor topk_weights, Tensor topk_ids, bool " + "inplace, int moe_comp_method, Tensor? w1_scale, Tensor? w2_scale, " + "Tensor? w1_zero, Tensor? w2_zero, int[]? block_size, bool is_vnni) -> Tensor"); + m.impl("fused_experts_cpu", torch::kCPU, &fused_experts_cpu); #endif // conv3d fast path for patch embedding diff --git a/sgl-kernel/csrc/cpu/vec.h b/sgl-kernel/csrc/cpu/vec.h index a37bc6ba2..f1de4f9ca 100644 --- a/sgl-kernel/csrc/cpu/vec.h +++ b/sgl-kernel/csrc/cpu/vec.h @@ -16,6 +16,15 @@ inline Vectorized convert_from_float_ext(const Vectorized& a, c return at::vec::convert_from_float(a, b); } +template +inline void convert_from_float_and_store(scalar_t* out, const Vectorized& a) { + float out_buffer[at::vec::Vectorized::size()]; + a.store(out_buffer); + for (int i = 0; i < 16; i++) { + out[i] = (scalar_t)out_buffer[i]; + } +} + // allow f16, bf16 template , int> = 1> inline std::tuple, Vectorized> load_float_vec2(const scalar_t* __restrict__ data) { @@ -45,6 +54,11 @@ convert_from_float_ext(const Vectorized& a, const Vectorize return (__m512i)(_mm512_cvtne2ps_pbh(__m512(b), __m512(a))); } +template <> +inline void convert_from_float_and_store(at::BFloat16* out, const Vectorized& a) { + _mm256_storeu_si256((__m256i*)out, (__m256i)(_mm512_cvtneps_pbh(__m512(a)))); +} + #define CVT_BF16_TO_FP32(a) _mm512_castsi512_ps(_mm512_slli_epi32(_mm512_cvtepu16_epi32(a), 16)) #define CVT_FP16_TO_FP32(a) _mm512_cvtph_ps(a) diff --git a/test/registered/cpu/test_decode.py b/test/registered/cpu/test_decode.py index a29d9b76c..494fe73f7 100644 --- a/test/registered/cpu/test_decode.py +++ b/test/registered/cpu/test_decode.py @@ -12,6 +12,91 @@ torch.manual_seed(1234) class TestDecodeAttention(CustomTestCase): + def _scaled_dot_product_attention(self, Q, K, V, S, scaling, sliding_window): + # sliding_window <= 0 means no sliding window + # Q: [n_tokens_q, n_heads, q_mult, d_head] + # K: [n_tokens_kv, n_heads, d_head] + # V: [n_tokens_kv, n_heads, d_head] + n_tokens_q, n_heads, q_mult, d_head = Q.shape + n_tokens_kv = K.shape[0] + assert K.shape == (n_tokens_kv, n_heads, d_head) + assert V.shape == (n_tokens_kv, n_heads, d_head) + K = K[:, :, None, :].expand(-1, -1, q_mult, -1) + V = V[:, :, None, :].expand(-1, -1, q_mult, -1) + S = S.reshape(n_heads, q_mult, 1, 1).expand(-1, -1, n_tokens_q, -1) + if n_tokens_q == n_tokens_kv: # Prefill + mask = torch.triu( + Q.new_full((n_tokens_q, n_tokens_kv), -float("inf")), diagonal=1 + ) + else: # Decode + mask = Q.new_zeros((n_tokens_q, n_tokens_kv)) + if sliding_window is not None and sliding_window > 0: + mask += torch.tril( + mask.new_full((n_tokens_q, n_tokens_kv), -float("inf")), + diagonal=n_tokens_kv - n_tokens_q - sliding_window, + ) + QK = torch.einsum("qhmd,khmd->hmqk", Q, K) + QK *= scaling + QK += mask[None, None, :, :] + QK = torch.cat([QK, S], dim=-1) + W = torch.softmax(QK, dim=-1) + W = W[..., :-1] + attn = torch.einsum("hmqk,khmd->qhmd", W, V) + return attn.reshape(n_tokens_q, -1) + + def _run_sdpa_forward_decode_sink( + self, + query: torch.Tensor, + output: torch.Tensor, + k_cache: torch.Tensor, + v_cache: torch.Tensor, + req_to_token: torch.Tensor, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + num_kv_heads: int, + q_mult: int, + scaling=None, + sliding_window=None, + attention_sinks=None, + enable_gqa=False, + causal=False, + ): + # [num_tokens, num_heads, head_size] -> [num_heads, num_tokens, head_size] + query = query.movedim(0, query.dim() - 2) + start_q, start_kv = 0, 0 + for seq_idx in range(seq_lens.shape[0]): + # TODO: this loop process a sequence per iter, this is inefficient. + # Need optimize the performance later. + seq_len_q = 1 + seq_len_kv = seq_lens[seq_idx] + end_q = start_q + seq_len_q + end_kv = start_kv + seq_len_kv + per_req_query = query[:, start_q:end_q, :] + # get key and value from cache. per_req_tokens contains the kv cache + # index for each token in the sequence. + req_pool_idx = req_pool_indices[seq_idx] + per_req_tokens = req_to_token[req_pool_idx, :seq_len_kv] + + per_req_query = per_req_query.permute(1, 0, 2).reshape( + seq_len_q, num_kv_heads, q_mult, per_req_query.shape[-1] + ) + per_req_key = k_cache[per_req_tokens].movedim(0, query.dim() - 2) + per_req_value = v_cache[per_req_tokens].movedim(0, query.dim() - 2) + per_req_key = per_req_key.permute(1, 0, 2) + per_req_value = per_req_value.permute(1, 0, 2) + + per_req_out = self._scaled_dot_product_attention( + per_req_query, + per_req_key, + per_req_value, + attention_sinks, + scaling=scaling, + sliding_window=sliding_window, + ).reshape(seq_len_q, -1, per_req_value.shape[-1]) + output[start_q:end_q, :, :] = per_req_out + start_q, start_kv = end_q, end_kv + return output + def _run_sdpa_forward_decode( self, query: torch.Tensor, @@ -71,11 +156,11 @@ class TestDecodeAttention(CustomTestCase): return output def _test_grouped_decode_attention_once( - self, B, H_Q, H_KV, D, D_V, is_cross_attn, dtype, device + self, B, H_Q, H_KV, D, D_V, sliding_window, sink, is_cross_attn, dtype, device ): # This represents the number of tokens already in the sequence seq_len = 1024 - encoder_len = 10 + encoder_len = 0 if sink else 10 total_tokens = B * (seq_len + encoder_len) sm_scale = 1.0 / (D**0.5) logit_cap = 0.0 @@ -84,6 +169,7 @@ class TestDecodeAttention(CustomTestCase): # q represents the new token being generated, one per batch q = torch.randn(B, H_Q, D, dtype=dtype, device=device) + sinks = torch.rand(H_Q, dtype=dtype, device=device) * 10 # k_buffer and v_buffer represent all previous tokens k_buffer = torch.randn(total_tokens, H_KV, D, dtype=dtype, device=device) @@ -137,22 +223,41 @@ class TestDecodeAttention(CustomTestCase): sm_scale, logit_cap, is_cross_attn, + sliding_window if sliding_window is not None else 0, encoder_lens, + sinks if sink else None, ) - self._run_sdpa_forward_decode( - q, - o_grouped, - k_buffer, - v_buffer, - req_to_token, - b_req_idx, - b_seq_len, - scaling=sm_scale, - enable_gqa=enable_gqa, - encoder_lens=encoder_lens, - is_cross_attn=is_cross_attn, - ) + if sink: + self._run_sdpa_forward_decode_sink( + q, + o_grouped, + k_buffer, + v_buffer, + req_to_token, + b_req_idx, + b_seq_len, + num_kv_heads=H_KV, + q_mult=H_Q // H_KV if enable_gqa else 1, + scaling=sm_scale, + sliding_window=sliding_window if sliding_window is not None else None, + attention_sinks=sinks, + enable_gqa=enable_gqa, + ) + else: + self._run_sdpa_forward_decode( + q, + o_grouped, + k_buffer, + v_buffer, + req_to_token, + b_req_idx, + b_seq_len, + scaling=sm_scale, + enable_gqa=enable_gqa, + encoder_lens=encoder_lens, + is_cross_attn=is_cross_attn, + ) cos_sim = torch.nn.functional.cosine_similarity( o.flatten(), o_grouped.flatten(), dim=0 ) @@ -176,11 +281,26 @@ class TestDecodeAttention(CustomTestCase): for B, H_Q, H_KV, D, D_V in configs: for dtype in [torch.bfloat16, torch.float16]: + for sink in [True, False]: + if D != D_V and sink: + continue + for sliding_window in [None, 10]: + if sliding_window is not None and not sink: + continue + self._test_grouped_decode_attention_once( + B, + H_Q, + H_KV, + D, + D_V, + sliding_window, + sink, + False, + dtype=dtype, + device=device, + ) self._test_grouped_decode_attention_once( - B, H_Q, H_KV, D, D_V, False, dtype=dtype, device=device - ) - self._test_grouped_decode_attention_once( - B, H_Q, H_KV, D, D_V, True, dtype=dtype, device=device + B, H_Q, H_KV, D, D_V, None, False, True, dtype=dtype, device=device ) def test_grouped_decode_attention(self): diff --git a/test/registered/cpu/test_extend.py b/test/registered/cpu/test_extend.py index ab55fe383..994ead071 100644 --- a/test/registered/cpu/test_extend.py +++ b/test/registered/cpu/test_extend.py @@ -12,6 +12,37 @@ torch.manual_seed(1234) class TestExtendAttention(CustomTestCase): + def _scaled_dot_product_attention(self, Q, K, V, S, scaling, sliding_window): + # sliding_window <= 0 means no sliding window + # Q: [n_tokens_q, n_heads, q_mult, d_head] + # K: [n_tokens_kv, n_heads, d_head] + # V: [n_tokens_kv, n_heads, d_head] + n_tokens_q, n_heads, q_mult, d_head = Q.shape + n_tokens_kv = K.shape[0] + assert K.shape == (n_tokens_kv, n_heads, d_head) + assert V.shape == (n_tokens_kv, n_heads, d_head) + K = K[:, :, None, :].expand(-1, -1, q_mult, -1) + V = V[:, :, None, :].expand(-1, -1, q_mult, -1) + S = S.reshape(n_heads, q_mult, 1, 1).expand(-1, -1, n_tokens_q, -1) + if n_tokens_q == n_tokens_kv: # Prefill + mask = torch.triu( + Q.new_full((n_tokens_q, n_tokens_kv), -float("inf")), diagonal=1 + ) + else: # Decode + mask = Q.new_zeros((n_tokens_q, n_tokens_kv)) + if sliding_window is not None and sliding_window > 0: + mask += torch.tril( + mask.new_full((n_tokens_q, n_tokens_kv), -float("inf")), + diagonal=n_tokens_kv - n_tokens_q - sliding_window, + ) + QK = torch.einsum("qhmd,khmd->hmqk", Q, K) + QK *= scaling + QK += mask[None, None, :, :] + QK = torch.cat([QK, S], dim=-1) + W = torch.softmax(QK, dim=-1) + W = W[..., :-1] + attn = torch.einsum("hmqk,khmd->qhmd", W, V) + return attn.reshape(n_tokens_q, -1) def _run_sdpa_forward_extend( self, @@ -86,6 +117,70 @@ class TestExtendAttention(CustomTestCase): start_q, start_kv = end_q, end_kv return output + def _run_sdpa_forward_extend_sink( + self, + query: torch.Tensor, + output: torch.Tensor, + k_cache: torch.Tensor, + v_cache: torch.Tensor, + req_to_token: torch.Tensor, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + extend_prefix_lens: torch.Tensor, + extend_seq_lens: torch.Tensor, + num_kv_heads: int, + q_mult: int, + scaling=None, + sliding_window=None, + attention_sinks=None, + enable_gqa=False, + causal=False, + ): + assert seq_lens.shape[0] == extend_prefix_lens.shape[0] + assert seq_lens.shape[0] == extend_seq_lens.shape[0] + # [num_tokens, num_heads, head_size] -> [num_heads, num_tokens, head_size] + query = query.movedim(0, query.dim() - 2) + start_q, start_kv = 0, 0 + for seq_idx in range(seq_lens.shape[0]): + # TODO: this loop process a sequence per iter, this is inefficient. + # Need optimize the performance later. + extend_seq_len_q = extend_seq_lens[seq_idx] + prefill_seq_len_q = extend_prefix_lens[seq_idx] + seq_len_kv = seq_lens[seq_idx] + end_q = start_q + extend_seq_len_q + end_kv = start_kv + seq_len_kv + per_req_query = query[:, start_q:end_q, :] + per_req_query_redudant = torch.empty( + (per_req_query.shape[0], seq_len_kv, per_req_query.shape[2]), + dtype=per_req_query.dtype, + device=per_req_query.device, + ) + per_req_query_redudant[:, prefill_seq_len_q:, :] = per_req_query + + # get key and value from cache. per_req_tokens contains the kv cache + # index for each token in the sequence. + req_pool_idx = req_pool_indices[seq_idx] + per_req_tokens = req_to_token[req_pool_idx, :seq_len_kv] + per_req_key = k_cache[per_req_tokens].movedim(0, query.dim() - 2) + per_req_value = v_cache[per_req_tokens].movedim(0, query.dim() - 2) + + per_req_query_redudant = per_req_query_redudant.permute(1, 0, 2).reshape( + seq_len_kv, num_kv_heads, q_mult, per_req_query_redudant.shape[-1] + ) + per_req_key = per_req_key.permute(1, 0, 2) + per_req_value = per_req_value.permute(1, 0, 2) + per_req_out_redudant = self._scaled_dot_product_attention( + per_req_query_redudant, + per_req_key, + per_req_value, + attention_sinks, + scaling=scaling, + sliding_window=sliding_window, + ).reshape(seq_len_kv, -1, per_req_value.shape[-1]) + output[start_q:end_q, :, :] = per_req_out_redudant[prefill_seq_len_q:, :, :] + start_q, start_kv = end_q, end_kv + return output + def _test_extend_attention_once( self, B, @@ -94,6 +189,8 @@ class TestExtendAttention(CustomTestCase): H_KV, D, DV, + sliding_window=None, + has_sink=False, mla=False, is_cross_attn=False, *, @@ -108,9 +205,13 @@ class TestExtendAttention(CustomTestCase): b_seq_len_prefix = torch.as_tensor(b_seq_len_prefix, dtype=torch.int32) encoder_lens = torch.randint(1, N_CTX // 2, (B,), dtype=torch.int64) + scale = 20 if mla: b_seq_len_prefix.zero_() encoder_lens.zero_() + if has_sink: + encoder_lens.zero_() + scale = 1 if b_seq_len_extend is None: b_seq_len_extend = torch.randint(1, N_CTX // 2, (B,), dtype=torch.int32) @@ -143,6 +244,7 @@ class TestExtendAttention(CustomTestCase): k_extend = torch.empty((extend_token_num, H_KV, D), dtype=dtype) v_extend = torch.empty((extend_token_num, H_KV, DV), dtype=dtype) q_extend = torch.empty((extend_token_num, H_Q, D), dtype=dtype) + sinks = torch.rand(H_Q, dtype=dtype) for i in range(B): extend_start_in_buffer = ( @@ -158,7 +260,7 @@ class TestExtendAttention(CustomTestCase): extend_start_in_buffer:extend_end_in_buffer ] q_extend[extend_start:extend_end] = ( - torch.randn((b_seq_len_extend[i], H_Q, D), dtype=dtype) * 20 + torch.randn((b_seq_len_extend[i], H_Q, D), dtype=dtype) * scale ) # q_extend, k_extend, v_extend, k_buffer and v_buffer supports non-contiguous tensors @@ -182,22 +284,40 @@ class TestExtendAttention(CustomTestCase): enable_gqa = H_Q != H_KV o_ref = torch.empty((extend_token_num, H_Q, DV), dtype=dtype) - self._run_sdpa_forward_extend( - q_extend, - o_ref, - k_buffer, - v_buffer, - req_to_tokens, - b_req_idx, - b_seq_len, - b_seq_len_prefix, - b_seq_len_extend, - scaling=sm_scale, - enable_gqa=enable_gqa, - causal=not is_cross_attn, - is_cross_attn=is_cross_attn, - encoder_lens=encoder_lens, - ) + if has_sink: + self._run_sdpa_forward_extend_sink( + q_extend, + o_ref, + k_buffer, + v_buffer, + req_to_tokens, + b_req_idx, + b_seq_len, + b_seq_len_prefix, + b_seq_len_extend, + H_KV, + H_Q // H_KV if enable_gqa else 1, + scaling=sm_scale, + sliding_window=sliding_window, + attention_sinks=sinks, + ) + else: + self._run_sdpa_forward_extend( + q_extend, + o_ref, + k_buffer, + v_buffer, + req_to_tokens, + b_req_idx, + b_seq_len, + b_seq_len_prefix, + b_seq_len_extend, + scaling=sm_scale, + enable_gqa=enable_gqa, + causal=not is_cross_attn, + is_cross_attn=is_cross_attn, + encoder_lens=encoder_lens, + ) o_extend = torch.empty((extend_token_num, H_Q, DV), dtype=dtype) torch.ops.sgl_kernel.extend_attention_cpu( @@ -216,7 +336,9 @@ class TestExtendAttention(CustomTestCase): sm_scale, logit_cap, is_cross_attn, + sliding_window if sliding_window is not None else 0, encoder_lens, + sinks if has_sink else None, ) torch.testing.assert_close(o_ref, o_extend, atol=1e-2, rtol=1e-2) @@ -227,16 +349,29 @@ class TestExtendAttention(CustomTestCase): if is_mla and is_cross_attn: continue self._test_extend_attention_once( - 1, 123, 1, 1, 128, 96, is_mla, is_cross_attn + 1, 123, 1, 1, 128, 96, None, False, is_mla, is_cross_attn ) self._test_extend_attention_once( - 1, 123, 16, 1, 128, 96, is_mla, is_cross_attn + 1, 123, 16, 1, 128, 96, None, False, is_mla, is_cross_attn ) self._test_extend_attention_once( - 4, 1230, 16, 4, 128, 96, is_mla, is_cross_attn + 4, 1230, 16, 4, 128, 96, None, False, is_mla, is_cross_attn ) self._test_extend_attention_once( - 1, 9000, 16, 1, 32, 32, is_mla, is_cross_attn + 1, 9000, 16, 1, 32, 32, None, False, is_mla, is_cross_attn + ) + for has_sink in [True, False]: + for sliding_window in [None, 10, 128]: + if not has_sink and sliding_window is not None: + continue + self._test_extend_attention_once( + 1, 123, 16, 4, 64, 64, sliding_window, has_sink, False, False + ) + self._test_extend_attention_once( + 1, 20, 16, 1, 64, 64, sliding_window, has_sink, False, False + ) + self._test_extend_attention_once( + 1, 20, 1, 1, 64, 64, sliding_window, has_sink, False, False ) def test_extend_attention_large_seq_causal_mask(self): diff --git a/test/registered/cpu/test_gemm.py b/test/registered/cpu/test_gemm.py index 3bd0bb961..8e8083e18 100644 --- a/test/registered/cpu/test_gemm.py +++ b/test/registered/cpu/test_gemm.py @@ -1,12 +1,13 @@ -import itertools import unittest # TODO: use interface in cpu.py import torch import torch.nn as nn from utils import ( + MXFP4QuantizeUtil, convert_weight, native_w8a8_per_token_matmul, + parametrize, per_token_quant_int8, precision, unpack_and_dequant_awq, @@ -31,30 +32,15 @@ class Mod(nn.Module): class TestGemm(CustomTestCase): - M = [1, 101] - N = [16, 32 * 13] - K = [32 * 16] - has_bias = [False, True] - dim = [2, 3, 4, 5] - - M_int8 = [2, 128] - N_int8 = [32 * 12] - K_int8 = [32 * 17] - - M_fp8 = [1, 11] - N_fp8 = [128, 224] - K_fp8 = [512, 576] - - M_awq = [1, 32] - N_awq = [4096] - K_awq = [4096] - - M_gptq = [1, 32] - N_gptq = [4096] - K_gptq = [4096] - - def _bf16_gemm(self, M, N, K, has_bias, dim): + @parametrize( + M=[1, 101], + N=[16, 32 * 13], + K=[32 * 16], + has_bias=[False, True], + dim=[2, 3, 4, 5], + ) + def test_bf16_gemm(self, M, N, K, has_bias, dim): mat1 = torch.randn(M, K, dtype=torch.bfloat16) mat2 = torch.randn(N, K, dtype=torch.bfloat16) if dim == 3: @@ -84,24 +70,14 @@ class TestGemm(CustomTestCase): torch.testing.assert_close(ref, out, atol=atol, rtol=rtol) torch.testing.assert_close(ref, out2, atol=atol, rtol=rtol) - def test_bf16_gemm(self): - for params in itertools.product( - self.M, - self.N, - self.K, - self.has_bias, - self.dim, - ): - with self.subTest( - M=params[0], - N=params[1], - K=params[2], - has_bias=params[3], - dim=params[4], - ): - self._bf16_gemm(*params) - - def _bf16_gemm_with_small_oc(self, M, N, K, has_bias, use_post_sigmul): + @parametrize( + M=[1, 8, 32, 1024], + N=[12, 1], + K=[32 * 16], + has_bias=[False, True], + use_post_sigmul=[False, True], + ) + def bf16_gemm_with_small_oc(self, M, N, K, has_bias, use_post_sigmul): use_post_sigmul = use_post_sigmul and N == 1 mat_mul = ( None if not use_post_sigmul else torch.randn(M, 2 * K, dtype=torch.bfloat16) @@ -132,20 +108,8 @@ class TestGemm(CustomTestCase): atol = rtol = precision[ref.dtype] torch.testing.assert_close(ref, out, atol=atol, rtol=rtol) - def test_bf16_gemm_with_small_oc(self): - for params in itertools.product( - [1, 8, 32, 1024], [12, 1], self.K, self.has_bias, [False, True] - ): - with self.subTest( - M=params[0], - N=params[1], - K=params[2], - has_bias=params[3], - use_post_sigmul=params[4], - ): - self._bf16_gemm_with_small_oc(*params) - - def _int8_gemm(self, M, N, K, has_bias): + @parametrize(M=[2, 128], N=[32 * 12], K=[32 * 17], has_bias=[False, True]) + def test_int8_gemm(self, M, N, K, has_bias): dtype = torch.bfloat16 A = torch.randn((M, K), dtype=dtype) / 10 Aq, As = per_token_quant_int8(A) @@ -175,35 +139,21 @@ class TestGemm(CustomTestCase): ) torch.testing.assert_close(ref_out, fused_out, atol=atol, rtol=rtol) - def test_int8_gemm(self): - for params in itertools.product( - self.M_int8, - self.N_int8, - self.K_int8, - self.has_bias, - ): - with self.subTest( - M=params[0], - N=params[1], - K=params[2], - has_bias=params[3], - ): - self._int8_gemm(*params) - - def _fp8_gemm(self, M, N, K, has_bias): + @parametrize(M=[1, 11], N=[128, 224], K=[512, 576], has_bias=[False, True]) + def test_fp8_gemm(self, M, N, K, has_bias): prepack = True chunk = False scale_block_size_N = 64 scale_block_size_K = 128 assert scale_block_size_N <= N assert scale_block_size_K <= K - A_dtype = torch.bfloat16 + dtype = torch.bfloat16 model = Mod(K, N, has_bias).eval() if chunk: - data = torch.randn(M, K + 6, dtype=A_dtype).narrow(1, 0, K) + data = torch.randn(M, K + 6, dtype=dtype).narrow(1, 0, K) else: - data = torch.randn(M, K, dtype=A_dtype) + data = torch.randn(M, K, dtype=dtype) weight = model.linear.weight # (N, K) @@ -211,18 +161,18 @@ class TestGemm(CustomTestCase): bias = model.linear.bias fp8_weight, scales, dq_weight = convert_weight( - weight, [scale_block_size_N, scale_block_size_K], A_dtype + weight, [scale_block_size_N, scale_block_size_K], dtype ) if has_bias: - ref = torch.matmul(data.to(A_dtype), dq_weight.T) + bias.to(A_dtype) + ref = torch.matmul(data.to(dtype), dq_weight.T) + bias.to(dtype) else: - ref = torch.matmul(data.to(A_dtype), dq_weight.T) + ref = torch.matmul(data.to(dtype), dq_weight.T) if prepack: fp8_weight = torch.ops.sgl_kernel.convert_weight_packed(fp8_weight) - opt = torch.ops.sgl_kernel.fp8_scaled_mm_cpu( + out = torch.ops.sgl_kernel.fp8_scaled_mm_cpu( data, fp8_weight, scales, @@ -232,24 +182,41 @@ class TestGemm(CustomTestCase): prepack, ) atol = rtol = precision[ref.dtype] - torch.testing.assert_close(ref, opt, atol=atol, rtol=rtol) + torch.testing.assert_close(ref, out, atol=atol, rtol=rtol) - def test_fp8_gemm(self): - for params in itertools.product( - self.M_fp8, - self.N_fp8, - self.K_fp8, - self.has_bias, - ): - with self.subTest( - M=params[0], - N=params[1], - K=params[2], - has_bias=params[3], - ): - self._fp8_gemm(*params) + @parametrize(M=[1, 11], N=[128, 224], K=[512, 576], has_bias=[False, True]) + def test_mxfp4_gemm(self, M, N, K, has_bias): + prepack = True + dtype = torch.bfloat16 - def _int4_awq_gemm(self, M, N, K, group_size, has_bias): + A = torch.randn((M, K), dtype=dtype) / 10 + + # we randomly generate Bq and Bs, then dequantize it to BFloat16 as reference + Bq = torch.randint(0, 256, (N, K // 2), dtype=torch.uint8) + Bs = torch.randint(126, 127, (N, K // 32), dtype=torch.uint8) + + Bdq = MXFP4QuantizeUtil.dequantize(Bq, dtype, Bs) + + B_packed = torch.ops.sgl_kernel.convert_weight_packed(Bq) + Bs_packed = torch.ops.sgl_kernel.convert_scale_packed(Bs) + + bias = torch.randn(N) if has_bias else None + + ref = torch.matmul(A.float(), Bdq.float().t()).bfloat16() + if bias is not None: + ref.add_(bias.view(1, -1)) + + out = torch.ops.sgl_kernel.mxfp4_scaled_mm_cpu( + A, B_packed, Bs_packed, bias, prepack + ) + + atol = rtol = precision[ref.dtype] + torch.testing.assert_close(ref, out, atol=atol, rtol=rtol) + + @parametrize( + M=[1, 32], N=[4096], K=[4096], group_size=[128], has_bias=[False, True] + ) + def test_int4_awq_gemm(self, M, N, K, group_size, has_bias): awq_weight = torch.randint(-128, 128, (K, N // 8)).to(torch.int) awq_zero = torch.randint(0, 10, (K // group_size, N // 8)).to(torch.int) awq_scales = torch.rand(int(K // group_size), N).to(torch.bfloat16) @@ -281,20 +248,10 @@ class TestGemm(CustomTestCase): atol = rtol = precision[ref_res.dtype] torch.testing.assert_close(ref_res, target_res, atol=atol, rtol=rtol) - def test_int4_awq_gemm(self): - for params in itertools.product( - self.M_awq, self.N_awq, self.K_awq, [128], self.has_bias - ): - with self.subTest( - M=params[0], - N=params[1], - K=params[2], - group_size=params[3], - has_bias=params[4], - ): - self._int4_awq_gemm(*params) - - def _int4_gptq_gemm(self, M, N, K, group_size, has_bias): + @parametrize( + M=[1, 32], N=[4096], K=[4096], group_size=[128], has_bias=[False, True] + ) + def test_int4_gptq_gemm(self, M, N, K, group_size, has_bias): torch.manual_seed(127) gptq_weight = torch.randint(-128, 128, (K // 8, N)).to(torch.int) gptq_zero = torch.randint(0, 10, (K // group_size, N // 8)).to(torch.int) @@ -326,19 +283,6 @@ class TestGemm(CustomTestCase): atol = rtol = precision[ref_res.dtype] torch.testing.assert_close(ref_res, target_res, atol=atol, rtol=rtol) - def test_int4_gptq_gemm(self): - for params in itertools.product( - self.M_gptq, self.N_gptq, self.K_gptq, [128], self.has_bias - ): - with self.subTest( - M=params[0], - N=params[1], - K=params[2], - group_size=params[3], - has_bias=params[4], - ): - self._int4_gptq_gemm(*params) - if __name__ == "__main__": unittest.main() diff --git a/test/registered/cpu/test_mla.py b/test/registered/cpu/test_mla.py index da11a281f..5d87bd749 100644 --- a/test/registered/cpu/test_mla.py +++ b/test/registered/cpu/test_mla.py @@ -119,6 +119,8 @@ class TestMLA(CustomTestCase): sm_scale, logit_cap, False, + 0, + None, None, ) diff --git a/test/registered/cpu/test_moe.py b/test/registered/cpu/test_moe.py index 7ad54d891..076aa7347 100644 --- a/test/registered/cpu/test_moe.py +++ b/test/registered/cpu/test_moe.py @@ -1,4 +1,3 @@ -import itertools import math import unittest @@ -9,18 +8,21 @@ from sglang.srt.layers.amx_utils import CPUQuantMethod kernel = torch.ops.sgl_kernel -torch.manual_seed(128) +torch.manual_seed(1234) from utils import ( BLOCK_K, BLOCK_N, + MXFP4QuantizeUtil, factor_for_scale, fp8_max, fp8_min, native_fp8_fused_moe, + parametrize, precision, scaled_weight, torch_naive_fused_moe, + torch_naive_fused_moe_gptoss, torch_w8a8_per_column_fused_moe, unpack_and_dequant_awq, ) @@ -60,37 +62,18 @@ def fused_moe(a, w1, w2, score, topk, renormalize, prepack): None, None, None, + None, + None, + None, + None, prepack, ) class TestFusedExperts(CustomTestCase): - M = [2, 114] - N = [32] - K = [32] - E = [4] - topk = [2] - renormalize = [False, True] - M_int8 = [1, 39] - N_int8 = [128] - K_int8 = [256] - E_int8 = [8] - topk_int8 = [3] - - M_fp8 = [2, 121] - N_fp8 = [352, 512] - K_fp8 = [256, 320] - E_fp8 = [8] - topk_fp8 = [4] - - M_int4 = [1, 6] - N_int4 = [512] - K_int4 = [256] - E_int4 = [8] - topk_int4 = [4] - - def _bf16_moe(self, m, n, k, e, topk, renormalize): + @parametrize(m=[2, 114], n=[32], k=[32], e=[4], topk=[2], renormalize=[False, True]) + def test_bf16_moe(self, m, n, k, e, topk, renormalize): dtype = torch.bfloat16 prepack = True @@ -105,26 +88,51 @@ class TestFusedExperts(CustomTestCase): atol = rtol = precision[torch_output.dtype] torch.testing.assert_close(torch_output, fused_output, atol=atol, rtol=rtol) - def test_bf16_moe(self): - for params in itertools.product( - self.M, - self.N, - self.K, - self.E, - self.topk, - self.renormalize, - ): - with self.subTest( - m=params[0], - n=params[1], - k=params[2], - e=params[3], - topk=params[4], - renormalize=params[5], - ): - self._bf16_moe(*params) + @parametrize( + m=[1, 32], n=[128, 64], k=[128, 64], e=[4], topk=[2], renormalize=[False] + ) + def test_bf16_moe_bias(self, m, n, k, e, topk, renormalize): + dtype = torch.bfloat16 - def _int8_moe(self, M, N, K, E, topk): + a = torch.randn((m, k), device="cpu", dtype=dtype) / 10 + w1 = torch.randn((e, 2 * n, k), device="cpu", dtype=dtype) / 10 + w1_b = torch.randn((e, 2 * n), device="cpu", dtype=torch.float) / 10 + w2 = torch.randn((e, k, n), device="cpu", dtype=dtype) / 10 + w2_b = torch.randn((e, k), device="cpu", dtype=torch.float) / 10 + score = torch.randn((m, e), device="cpu", dtype=dtype) + score = torch.softmax(score, dim=-1, dtype=torch.float32) + topk_weight, topk_ids = torch.topk(score, topk) + alpha = 1.702 + limit = 7.0 + torch_output = torch_naive_fused_moe_gptoss( + a, w1, w2, w1_b, w2_b, topk_weight, topk_ids, renormalize, alpha, limit, e + ) + packed_w1 = kernel.convert_weight_packed(w1) + packed_w2 = kernel.convert_weight_packed(w2) + fused_output = torch.ops.sgl_kernel.fused_experts_cpu( + a, + packed_w1, + packed_w2, + topk_weight, + topk_ids.to(torch.int), + False, # inplace # See [Note] inplace should be False in fused_experts. + CPUQuantMethod.UNQUANT, + None, # w1_scale + None, # w2_scale + None, # w1_zp + None, # w2_zp + None, # block_size + w1_b, + w2_b, + alpha, + limit, + True, # is_vnni + ) + atol = rtol = precision[torch_output.dtype] + torch.testing.assert_close(torch_output, fused_output, atol=atol, rtol=rtol) + + @parametrize(M=[1, 39], N=[128], K=[256], E=[8], topk=[3]) + def test_int8_moe(self, M, N, K, E, topk): dtype = torch.bfloat16 prepack = True @@ -173,6 +181,10 @@ class TestFusedExperts(CustomTestCase): None, None, None, + None, + None, + None, + None, prepack, ) @@ -182,24 +194,8 @@ class TestFusedExperts(CustomTestCase): atol = rtol = 0.02 torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol) - def test_int8_moe(self): - for params in itertools.product( - self.M_int8, - self.N_int8, - self.K_int8, - self.E_int8, - self.topk_int8, - ): - with self.subTest( - M=params[0], - N=params[1], - K=params[2], - E=params[3], - topk=params[4], - ): - self._int8_moe(*params) - - def _fp8_moe(self, M, N, K, E, topk): + @parametrize(M=[2, 121], N=[352, 512], K=[256, 320], E=[8], topk=[4]) + def test_fp8_moe(self, M, N, K, E, topk): dtype = torch.bfloat16 a = torch.randn(M, K, dtype=dtype) / math.sqrt(K) @@ -245,30 +241,132 @@ class TestFusedExperts(CustomTestCase): None, None, [BLOCK_N, BLOCK_K], + None, + None, + None, + None, True, ) atol = rtol = precision[dtype] torch.testing.assert_close(ref_out.bfloat16(), out, atol=atol, rtol=rtol) - def test_fp8_moe(self): - for params in itertools.product( - self.M_fp8, - self.N_fp8, - self.K_fp8, - self.E_fp8, - self.topk_fp8, - ): - with self.subTest( - M=params[0], - N=params[1], - K=params[2], - E=params[3], - topk=params[4], - ): - self._fp8_moe(*params) + @parametrize(M=[2, 121], N=[352, 512], K=[256, 320], E=[8], topk=[4]) + def test_mxfp4_moe(self, M, N, K, E, topk): + dtype = torch.bfloat16 - def _int4_moe(self, M, N, K, E, topk, group_size=128): + a = torch.randn(M, K, dtype=dtype) / 10 + + w1_bf16 = torch.randn((E, 2 * N, K), dtype=dtype) / 10 + w1q, w1s = MXFP4QuantizeUtil.quantize(w1_bf16) + w1s = w1s.reshape(E, 2 * N, K // 32) + w1dq = MXFP4QuantizeUtil.dequantize(w1q, dtype, w1s) + + w2_bf16 = torch.randn((E, K, N), dtype=dtype) / 10 + w2q, w2s = MXFP4QuantizeUtil.quantize(w2_bf16) + w2s = w2s.reshape(E, K, N // 32) + w2dq = MXFP4QuantizeUtil.dequantize(w2q, dtype, w2s) + + score = torch.randn((M, E), dtype=dtype) + score = torch.softmax(score, dim=-1, dtype=torch.float32) + topk_weight, topk_ids = torch.topk(score, topk) + + w1 = kernel.convert_weight_packed(w1q) + w2 = kernel.convert_weight_packed(w2q) + w1s = kernel.convert_scale_packed(w1s) + w2s = kernel.convert_scale_packed(w2s) + + ref_out = native_fp8_fused_moe( + a, w1dq.float(), w2dq.float(), topk_weight, topk_ids, topk + ) + out = kernel.fused_experts_cpu( + a, + w1, + w2, + topk_weight, + topk_ids.to(torch.int32), + False, + CPUQuantMethod.MXFP4, + w1s, + w2s, + None, + None, + None, + None, + None, + None, + None, + True, + ) + + atol = rtol = precision[dtype] + torch.testing.assert_close(ref_out.bfloat16(), out, atol=atol, rtol=rtol) + + @parametrize( + m=[1, 32], n=[128, 64], k=[128, 64], e=[4], topk=[2], renormalize=[False] + ) + def test_mxfp4_moe_bias(self, m, n, k, e, topk, renormalize): + dtype = torch.bfloat16 + + a = torch.randn((m, k), device="cpu", dtype=dtype) / 10 + w1_bf16 = torch.randn((e, 2 * n, k), device="cpu", dtype=dtype) / 10 + w1q, w1s = MXFP4QuantizeUtil.quantize(w1_bf16) + w1s = w1s.reshape(e, 2 * n, k // 32) + w1dq = MXFP4QuantizeUtil.dequantize(w1q, dtype, w1s) + w1_b = torch.randn((e, 2 * n), device="cpu", dtype=torch.float32) / 10 + w2_bf16 = torch.randn((e, k, n), device="cpu", dtype=dtype) / 10 + w2q, w2s = MXFP4QuantizeUtil.quantize(w2_bf16) + w2s = w2s.reshape(e, k, n // 32) + w2dq = MXFP4QuantizeUtil.dequantize(w2q, dtype, w2s) + w2_b = torch.randn((e, k), device="cpu", dtype=torch.float32) / 10 + score = torch.randn((m, e), device="cpu", dtype=dtype) + score = torch.softmax(score, dim=-1, dtype=torch.float32) + topk_weight, topk_ids = torch.topk(score, topk) + alpha = 1.702 + limit = 7.0 + torch_output = torch_naive_fused_moe_gptoss( + a, + w1dq, + w2dq, + w1_b, + w2_b, + topk_weight, + topk_ids, + renormalize, + alpha, + limit, + e, + ) + + w1 = kernel.convert_weight_packed(w1q) + w2 = kernel.convert_weight_packed(w2q) + w1s = kernel.convert_scale_packed(w1s) + w2s = kernel.convert_scale_packed(w2s) + + fused_output = torch.ops.sgl_kernel.fused_experts_cpu( + a, + w1, + w2, + topk_weight, + topk_ids.to(torch.int32), + False, # inplace # See [Note] inplace should be False in fused_experts. + CPUQuantMethod.MXFP4, # use_mxfp4 + w1s, # w1_scale + w2s, # w2_scale + None, # w1_zp + None, # w2_zp + None, # block_size + w1_b, + w2_b, + alpha, + limit, + True, # is_vnni + ) + atol = rtol = precision[torch_output.dtype] + torch.testing.assert_close(torch_output, fused_output, atol=atol, rtol=rtol) + + @parametrize(M=[1, 6], N=[512], K=[256], E=[8], topk=[4]) + def test_int4_moe(self, M, N, K, E, topk, group_size=128): dtype = torch.bfloat16 a = torch.rand(M, K, dtype=dtype) / math.sqrt(K) @@ -327,29 +425,16 @@ class TestFusedExperts(CustomTestCase): awq_w13_zero_pack, awq_w2_zero_pack, None, + None, + None, + None, + None, True, ) atol = rtol = precision[dtype] torch.testing.assert_close(ref_out.bfloat16(), out, atol=atol, rtol=rtol) - def test_int4_moe(self): - for params in itertools.product( - self.M_int4, - self.N_int4, - self.K_int4, - self.E_int4, - self.topk_int4, - ): - with self.subTest( - M=params[0], - N=params[1], - K=params[2], - E=params[3], - topk=params[4], - ): - self._int4_moe(*params) - if __name__ == "__main__": unittest.main() diff --git a/test/registered/cpu/utils.py b/test/registered/cpu/utils.py index 57e7e74c1..307cf8fb0 100644 --- a/test/registered/cpu/utils.py +++ b/test/registered/cpu/utils.py @@ -221,6 +221,107 @@ def torch_naive_fused_moe(a, w1, w2, score, topk, renormalize): ).sum(dim=1) +def moe_gptoss_act(x, alpha: float = 1.702, limit: float = 7.0): + x_glu, x_linear = x[..., ::2], x[..., 1::2] + # Clamp the input values + x_glu = x_glu.clamp(min=None, max=limit) + x_linear = x_linear.clamp(min=-limit, max=limit) + out_glu = x_glu * torch.sigmoid(alpha * x_glu) + # Note we add an extra bias of 1 to the linear layer + return out_glu * (x_linear + 1.0) + + +def torch_naive_gptoss_fused_moe( + x, + w1, + w2, + w1_bias, + w2_bias, + topk_weights, + topk_ids, + activation_alpha, + swiglu_limit, + len_experts, +) -> torch.Tensor: + + # Ref code from https://huggingface.co/deepseek-ai/DeepSeek-V2/blob/e0828e3cc0a03408724b80c3cc92c8e072db8d01/modeling_deepseek.py#L589 + cnts = topk_ids.new_zeros((topk_ids.shape[0], len_experts)) + cnts.scatter_(1, topk_ids.to(torch.int64), 1) + tokens_per_expert = cnts.sum(dim=0) + idxs = topk_ids.view(-1).argsort() + + sorted_tokens = x[idxs // topk_ids.shape[1]] + tokens_per_expert = tokens_per_expert.cpu().numpy() + + outputs = [] + start_idx = 0 + for i, num_tokens in enumerate(tokens_per_expert): + end_idx = start_idx + num_tokens + if num_tokens == 0: + continue + tokens_for_this_expert = sorted_tokens[start_idx:end_idx] + + layer_w13_weight = w1[i] + layer_w13_weight_bias = w1_bias[i] + layer_w2_weight_bias = w2_bias[i] + layer_w2_weight = w2[i] + + gate_up = F.linear( + tokens_for_this_expert, + layer_w13_weight, + bias=layer_w13_weight_bias.to(torch.bfloat16), + ) + gate_up = moe_gptoss_act(gate_up, activation_alpha, swiglu_limit) + expert_out = F.linear( + gate_up, layer_w2_weight, bias=layer_w2_weight_bias.to(torch.bfloat16) + ) + outputs.append(expert_out) + start_idx = end_idx + + outs = torch.cat(outputs, dim=0) if len(outputs) else sorted_tokens.new_empty(0) + new_x = torch.empty_like(outs) + + new_x[idxs] = outs + final_out = ( + new_x.view(*topk_ids.shape, -1) + .type(topk_weights.dtype) + .mul_(topk_weights.unsqueeze(dim=-1)) + .sum(dim=1) + .type(new_x.dtype) + ) + return final_out + + +def torch_naive_fused_moe_gptoss( + a, + w1, + w2, + w1_bias, + w2_bias, + topk_weight, + topk_ids, + renormalize, + activation_alpha, + swiglu_limit, + len_experts, +): + if renormalize: + topk_weight = topk_weight / topk_weight.sum(dim=-1, keepdim=True) + + return torch_naive_gptoss_fused_moe( + a, + w1, + w2, + w1_bias, + w2_bias, + topk_weight, + topk_ids, + activation_alpha, + swiglu_limit, + len_experts, + ) + + def torch_w8a8_per_column_fused_moe(a, w1, w2, w1_s, w2_s, topk_weight, topk_ids, topk): """This function performs fused moe with per-column int8 quantization using native torch.""" @@ -294,6 +395,121 @@ def native_fp8_fused_moe(a, w1, w2, topk_weight, topk_ids, topk): ) +# https://github.com/NVIDIA/TensorRT-Model-Optimizer/blob/main/modelopt/torch/quantization/qtensor/mxfp4_tensor.py +class MXFP4QuantizeUtil: + E2M1_max = 6.0 + + E2M1_values = [0, 0.5, 1, 1.5, 2, 3, 4, 6] + E2M1_bounds = torch.tensor([0.25, 0.75, 1.25, 1.75, 2.5, 3.5, 5]) + + block_size = 32 + + @classmethod + def quantize(cls, input: torch.Tensor) -> tuple: + """Converting a tensor to a quantized format based on MXFP4 quantization. Only E4M3 is supported. + Args: + input (torch.Tensor): The input tensor to be quantized. + """ + + def cast_fp4(x): + sign = torch.sign(x) + sign_bit = (2 - sign) // 2 + ord_ = torch.sum( + (x.abs().unsqueeze(-1) - cls.E2M1_bounds.to(x.device)) > 0, dim=-1 + ) + fp4_val = (sign_bit * 0b1000 + ord_).to(torch.uint8) + return fp4_val + + def fuse_uint4_to_uint8(x): + # If the last dimension is odd, pad with zeros + # If this behavior is not desired, please modify the code accordingly + left_side = x[..., 0::2] # Even indices (0, 2, 4...) + right_side = x[..., 1::2] # Odd indices (1, 3, 5...) + new_data = ( + right_side.clone() << 4 + ) # Put odd indices (higher addresses) in high bits + new_data[ + ..., : left_side.shape[-1] + ] += left_side # Put even indices in low bits + return new_data + + original_shape = input.shape + original_dtype = input.dtype + input = input.view(-1, cls.block_size) + # get scales + input_amax = input.abs().max(dim=-1, keepdim=True).values + descale = input_amax / cls.E2M1_max + min_value = torch.tensor(-127.0, device=descale.device) + e8m0_scale = torch.ceil(torch.maximum(torch.log2(descale), min_value)) + + input = (input / torch.exp2(e8m0_scale)).view(original_shape) + input_q = cast_fp4(input) + input_q = fuse_uint4_to_uint8(input_q) + e8m0_scale = (e8m0_scale + 127).to(torch.uint8) + return input_q, e8m0_scale + + @classmethod + def dequantize(cls, quantized_data, dtype: torch.dtype, scale): + """Dequantze MXFP4 packed tensor to a target dtype.""" + + def unfuse_uint8_to_uint4(x): + """Unfuse uint8 values back to uint4 values. + This is the inverse operation of fuse_uint4_to_uint8. + """ + # Extract the lower 4 bits (even indices) + left_side = x & 0x0F + + # Extract the upper 4 bits (odd indices) + right_side = (x >> 4) & 0x0F + + # Create a new tensor with alternating values + shape = list(x.shape) + shape[-1] = shape[-1] * 2 + result = torch.zeros(shape, dtype=torch.uint8, device=x.device) + + # Fill in the values - even indices get low bits, odd indices get high bits + result[..., 0::2] = left_side # Even indices from low bits + result[..., 1::2] = right_side # Odd indices from high bits + + return result + + e8m0_scale = scale + + # Unfuse the uint8 values back to uint4 + x_unfused = unfuse_uint8_to_uint4(quantized_data) + # print("@@@ x_unfused: ", x_unfused) + # Extract sign and magnitude + sign = 1 - 2 * ((x_unfused & 0b1000) >> 3).to( + torch.float32 + ) # Extract sign bit and convert to +1/-1 + magnitude = x_unfused & 0b0111 # Extract magnitude bits + magnitude = magnitude.to(torch.long) + + # Create a tensor with the E2M1 values + values = torch.tensor(cls.E2M1_values, device=quantized_data.device) + + # Use gather to index the values tensor properly + # We need to reshape magnitude to match the dimensions we want to gather along + original_shape = magnitude.shape + x_float = values[magnitude.reshape(-1)].reshape(original_shape) + + # Apply sign and scale + x_float = sign.float() * x_float + + # Reshape to apply block-wise scaling + x_float = x_float.reshape(-1, cls.block_size) + + # Apply the E8M0 scale + scale_factor = torch.exp2(e8m0_scale.float() - 127) + scale_factor = scale_factor.reshape(-1, 1) # Reshape for proper broadcasting + + # Apply scaling and reshape back to original shape + x_float = x_float * scale_factor + + # Reshape back to the original shape + return x_float.reshape(original_shape).to(dtype) + + def make_non_contiguous(x: torch.Tensor) -> torch.Tensor: """ Make a tensor non-contiguous by slicing it via last dimension. diff --git a/test/srt/cpu/test_decode.py b/test/srt/cpu/test_decode.py index b04f3d51c..d4bba2620 100644 --- a/test/srt/cpu/test_decode.py +++ b/test/srt/cpu/test_decode.py @@ -9,6 +9,91 @@ torch.manual_seed(1234) class TestDecodeAttention(CustomTestCase): + def _scaled_dot_product_attention(self, Q, K, V, S, scaling, sliding_window): + # sliding_window <= 0 means no sliding window + # Q: [n_tokens_q, n_heads, q_mult, d_head] + # K: [n_tokens_kv, n_heads, d_head] + # V: [n_tokens_kv, n_heads, d_head] + n_tokens_q, n_heads, q_mult, d_head = Q.shape + n_tokens_kv = K.shape[0] + assert K.shape == (n_tokens_kv, n_heads, d_head) + assert V.shape == (n_tokens_kv, n_heads, d_head) + K = K[:, :, None, :].expand(-1, -1, q_mult, -1) + V = V[:, :, None, :].expand(-1, -1, q_mult, -1) + S = S.reshape(n_heads, q_mult, 1, 1).expand(-1, -1, n_tokens_q, -1) + if n_tokens_q == n_tokens_kv: # Prefill + mask = torch.triu( + Q.new_full((n_tokens_q, n_tokens_kv), -float("inf")), diagonal=1 + ) + else: # Decode + mask = Q.new_zeros((n_tokens_q, n_tokens_kv)) + if sliding_window is not None and sliding_window > 0: + mask += torch.tril( + mask.new_full((n_tokens_q, n_tokens_kv), -float("inf")), + diagonal=n_tokens_kv - n_tokens_q - sliding_window, + ) + QK = torch.einsum("qhmd,khmd->hmqk", Q, K) + QK *= scaling + QK += mask[None, None, :, :] + QK = torch.cat([QK, S], dim=-1) + W = torch.softmax(QK, dim=-1) + W = W[..., :-1] + attn = torch.einsum("hmqk,khmd->qhmd", W, V) + return attn.reshape(n_tokens_q, -1) + + def _run_sdpa_forward_decode_sink( + self, + query: torch.Tensor, + output: torch.Tensor, + k_cache: torch.Tensor, + v_cache: torch.Tensor, + req_to_token: torch.Tensor, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + num_kv_heads: int, + q_mult: int, + scaling=None, + sliding_window=None, + attention_sinks=None, + enable_gqa=False, + causal=False, + ): + # [num_tokens, num_heads, head_size] -> [num_heads, num_tokens, head_size] + query = query.movedim(0, query.dim() - 2) + start_q, start_kv = 0, 0 + for seq_idx in range(seq_lens.shape[0]): + # TODO: this loop process a sequence per iter, this is inefficient. + # Need optimize the performance later. + seq_len_q = 1 + seq_len_kv = seq_lens[seq_idx] + end_q = start_q + seq_len_q + end_kv = start_kv + seq_len_kv + per_req_query = query[:, start_q:end_q, :] + # get key and value from cache. per_req_tokens contains the kv cache + # index for each token in the sequence. + req_pool_idx = req_pool_indices[seq_idx] + per_req_tokens = req_to_token[req_pool_idx, :seq_len_kv] + + per_req_query = per_req_query.permute(1, 0, 2).reshape( + seq_len_q, num_kv_heads, q_mult, per_req_query.shape[-1] + ) + per_req_key = k_cache[per_req_tokens].movedim(0, query.dim() - 2) + per_req_value = v_cache[per_req_tokens].movedim(0, query.dim() - 2) + per_req_key = per_req_key.permute(1, 0, 2) + per_req_value = per_req_value.permute(1, 0, 2) + + per_req_out = self._scaled_dot_product_attention( + per_req_query, + per_req_key, + per_req_value, + attention_sinks, + scaling=scaling, + sliding_window=sliding_window, + ).reshape(seq_len_q, -1, per_req_value.shape[-1]) + output[start_q:end_q, :, :] = per_req_out + start_q, start_kv = end_q, end_kv + return output + def _run_sdpa_forward_decode( self, query: torch.Tensor, @@ -68,11 +153,11 @@ class TestDecodeAttention(CustomTestCase): return output def _test_grouped_decode_attention_once( - self, B, H_Q, H_KV, D, D_V, is_cross_attn, dtype, device + self, B, H_Q, H_KV, D, D_V, sliding_window, sink, is_cross_attn, dtype, device ): # This represents the number of tokens already in the sequence seq_len = 1024 - encoder_len = 10 + encoder_len = 0 if sink else 10 total_tokens = B * (seq_len + encoder_len) sm_scale = 1.0 / (D**0.5) logit_cap = 0.0 @@ -81,6 +166,7 @@ class TestDecodeAttention(CustomTestCase): # q represents the new token being generated, one per batch q = torch.randn(B, H_Q, D, dtype=dtype, device=device) + sinks = torch.rand(H_Q, dtype=dtype, device=device) * 10 # k_buffer and v_buffer represent all previous tokens k_buffer = torch.randn(total_tokens, H_KV, D, dtype=dtype, device=device) @@ -134,22 +220,41 @@ class TestDecodeAttention(CustomTestCase): sm_scale, logit_cap, is_cross_attn, + sliding_window if sliding_window is not None else 0, encoder_lens, + sinks if sink else None, ) - self._run_sdpa_forward_decode( - q, - o_grouped, - k_buffer, - v_buffer, - req_to_token, - b_req_idx, - b_seq_len, - scaling=sm_scale, - enable_gqa=enable_gqa, - encoder_lens=encoder_lens, - is_cross_attn=is_cross_attn, - ) + if sink: + self._run_sdpa_forward_decode_sink( + q, + o_grouped, + k_buffer, + v_buffer, + req_to_token, + b_req_idx, + b_seq_len, + num_kv_heads=H_KV, + q_mult=H_Q // H_KV if enable_gqa else 1, + scaling=sm_scale, + sliding_window=sliding_window if sliding_window is not None else None, + attention_sinks=sinks, + enable_gqa=enable_gqa, + ) + else: + self._run_sdpa_forward_decode( + q, + o_grouped, + k_buffer, + v_buffer, + req_to_token, + b_req_idx, + b_seq_len, + scaling=sm_scale, + enable_gqa=enable_gqa, + encoder_lens=encoder_lens, + is_cross_attn=is_cross_attn, + ) cos_sim = torch.nn.functional.cosine_similarity( o.flatten(), o_grouped.flatten(), dim=0 ) @@ -173,11 +278,26 @@ class TestDecodeAttention(CustomTestCase): for B, H_Q, H_KV, D, D_V in configs: for dtype in [torch.bfloat16, torch.float16]: + for sink in [True, False]: + if D != D_V and sink: + continue + for sliding_window in [None, 10]: + if sliding_window is not None and not sink: + continue + self._test_grouped_decode_attention_once( + B, + H_Q, + H_KV, + D, + D_V, + sliding_window, + sink, + False, + dtype=dtype, + device=device, + ) self._test_grouped_decode_attention_once( - B, H_Q, H_KV, D, D_V, False, dtype=dtype, device=device - ) - self._test_grouped_decode_attention_once( - B, H_Q, H_KV, D, D_V, True, dtype=dtype, device=device + B, H_Q, H_KV, D, D_V, None, False, True, dtype=dtype, device=device ) def test_grouped_decode_attention(self): diff --git a/test/srt/cpu/test_extend.py b/test/srt/cpu/test_extend.py index 58a22c37b..31d8cc601 100644 --- a/test/srt/cpu/test_extend.py +++ b/test/srt/cpu/test_extend.py @@ -9,6 +9,37 @@ torch.manual_seed(1234) class TestExtendAttention(CustomTestCase): + def _scaled_dot_product_attention(self, Q, K, V, S, scaling, sliding_window): + # sliding_window <= 0 means no sliding window + # Q: [n_tokens_q, n_heads, q_mult, d_head] + # K: [n_tokens_kv, n_heads, d_head] + # V: [n_tokens_kv, n_heads, d_head] + n_tokens_q, n_heads, q_mult, d_head = Q.shape + n_tokens_kv = K.shape[0] + assert K.shape == (n_tokens_kv, n_heads, d_head) + assert V.shape == (n_tokens_kv, n_heads, d_head) + K = K[:, :, None, :].expand(-1, -1, q_mult, -1) + V = V[:, :, None, :].expand(-1, -1, q_mult, -1) + S = S.reshape(n_heads, q_mult, 1, 1).expand(-1, -1, n_tokens_q, -1) + if n_tokens_q == n_tokens_kv: # Prefill + mask = torch.triu( + Q.new_full((n_tokens_q, n_tokens_kv), -float("inf")), diagonal=1 + ) + else: # Decode + mask = Q.new_zeros((n_tokens_q, n_tokens_kv)) + if sliding_window is not None and sliding_window > 0: + mask += torch.tril( + mask.new_full((n_tokens_q, n_tokens_kv), -float("inf")), + diagonal=n_tokens_kv - n_tokens_q - sliding_window, + ) + QK = torch.einsum("qhmd,khmd->hmqk", Q, K) + QK *= scaling + QK += mask[None, None, :, :] + QK = torch.cat([QK, S], dim=-1) + W = torch.softmax(QK, dim=-1) + W = W[..., :-1] + attn = torch.einsum("hmqk,khmd->qhmd", W, V) + return attn.reshape(n_tokens_q, -1) def _run_sdpa_forward_extend( self, @@ -83,6 +114,70 @@ class TestExtendAttention(CustomTestCase): start_q, start_kv = end_q, end_kv return output + def _run_sdpa_forward_extend_sink( + self, + query: torch.Tensor, + output: torch.Tensor, + k_cache: torch.Tensor, + v_cache: torch.Tensor, + req_to_token: torch.Tensor, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + extend_prefix_lens: torch.Tensor, + extend_seq_lens: torch.Tensor, + num_kv_heads: int, + q_mult: int, + scaling=None, + sliding_window=None, + attention_sinks=None, + enable_gqa=False, + causal=False, + ): + assert seq_lens.shape[0] == extend_prefix_lens.shape[0] + assert seq_lens.shape[0] == extend_seq_lens.shape[0] + # [num_tokens, num_heads, head_size] -> [num_heads, num_tokens, head_size] + query = query.movedim(0, query.dim() - 2) + start_q, start_kv = 0, 0 + for seq_idx in range(seq_lens.shape[0]): + # TODO: this loop process a sequence per iter, this is inefficient. + # Need optimize the performance later. + extend_seq_len_q = extend_seq_lens[seq_idx] + prefill_seq_len_q = extend_prefix_lens[seq_idx] + seq_len_kv = seq_lens[seq_idx] + end_q = start_q + extend_seq_len_q + end_kv = start_kv + seq_len_kv + per_req_query = query[:, start_q:end_q, :] + per_req_query_redudant = torch.empty( + (per_req_query.shape[0], seq_len_kv, per_req_query.shape[2]), + dtype=per_req_query.dtype, + device=per_req_query.device, + ) + per_req_query_redudant[:, prefill_seq_len_q:, :] = per_req_query + + # get key and value from cache. per_req_tokens contains the kv cache + # index for each token in the sequence. + req_pool_idx = req_pool_indices[seq_idx] + per_req_tokens = req_to_token[req_pool_idx, :seq_len_kv] + per_req_key = k_cache[per_req_tokens].movedim(0, query.dim() - 2) + per_req_value = v_cache[per_req_tokens].movedim(0, query.dim() - 2) + + per_req_query_redudant = per_req_query_redudant.permute(1, 0, 2).reshape( + seq_len_kv, num_kv_heads, q_mult, per_req_query_redudant.shape[-1] + ) + per_req_key = per_req_key.permute(1, 0, 2) + per_req_value = per_req_value.permute(1, 0, 2) + per_req_out_redudant = self._scaled_dot_product_attention( + per_req_query_redudant, + per_req_key, + per_req_value, + attention_sinks, + scaling=scaling, + sliding_window=sliding_window, + ).reshape(seq_len_kv, -1, per_req_value.shape[-1]) + output[start_q:end_q, :, :] = per_req_out_redudant[prefill_seq_len_q:, :, :] + start_q, start_kv = end_q, end_kv + return output + def _test_extend_attention_once( self, B, @@ -91,6 +186,8 @@ class TestExtendAttention(CustomTestCase): H_KV, D, DV, + sliding_window=None, + has_sink=False, mla=False, is_cross_attn=False, *, @@ -105,9 +202,13 @@ class TestExtendAttention(CustomTestCase): b_seq_len_prefix = torch.as_tensor(b_seq_len_prefix, dtype=torch.int32) encoder_lens = torch.randint(1, N_CTX // 2, (B,), dtype=torch.int64) + scale = 20 if mla: b_seq_len_prefix.zero_() encoder_lens.zero_() + if has_sink: + encoder_lens.zero_() + scale = 1 if b_seq_len_extend is None: b_seq_len_extend = torch.randint(1, N_CTX // 2, (B,), dtype=torch.int32) @@ -140,6 +241,7 @@ class TestExtendAttention(CustomTestCase): k_extend = torch.empty((extend_token_num, H_KV, D), dtype=dtype) v_extend = torch.empty((extend_token_num, H_KV, DV), dtype=dtype) q_extend = torch.empty((extend_token_num, H_Q, D), dtype=dtype) + sinks = torch.rand(H_Q, dtype=dtype) for i in range(B): extend_start_in_buffer = ( @@ -155,7 +257,7 @@ class TestExtendAttention(CustomTestCase): extend_start_in_buffer:extend_end_in_buffer ] q_extend[extend_start:extend_end] = ( - torch.randn((b_seq_len_extend[i], H_Q, D), dtype=dtype) * 20 + torch.randn((b_seq_len_extend[i], H_Q, D), dtype=dtype) * scale ) # q_extend, k_extend, v_extend, k_buffer and v_buffer supports non-contiguous tensors @@ -179,22 +281,40 @@ class TestExtendAttention(CustomTestCase): enable_gqa = H_Q != H_KV o_ref = torch.empty((extend_token_num, H_Q, DV), dtype=dtype) - self._run_sdpa_forward_extend( - q_extend, - o_ref, - k_buffer, - v_buffer, - req_to_tokens, - b_req_idx, - b_seq_len, - b_seq_len_prefix, - b_seq_len_extend, - scaling=sm_scale, - enable_gqa=enable_gqa, - causal=not is_cross_attn, - is_cross_attn=is_cross_attn, - encoder_lens=encoder_lens, - ) + if has_sink: + self._run_sdpa_forward_extend_sink( + q_extend, + o_ref, + k_buffer, + v_buffer, + req_to_tokens, + b_req_idx, + b_seq_len, + b_seq_len_prefix, + b_seq_len_extend, + H_KV, + H_Q // H_KV if enable_gqa else 1, + scaling=sm_scale, + sliding_window=sliding_window, + attention_sinks=sinks, + ) + else: + self._run_sdpa_forward_extend( + q_extend, + o_ref, + k_buffer, + v_buffer, + req_to_tokens, + b_req_idx, + b_seq_len, + b_seq_len_prefix, + b_seq_len_extend, + scaling=sm_scale, + enable_gqa=enable_gqa, + causal=not is_cross_attn, + is_cross_attn=is_cross_attn, + encoder_lens=encoder_lens, + ) o_extend = torch.empty((extend_token_num, H_Q, DV), dtype=dtype) torch.ops.sgl_kernel.extend_attention_cpu( @@ -213,7 +333,9 @@ class TestExtendAttention(CustomTestCase): sm_scale, logit_cap, is_cross_attn, + sliding_window if sliding_window is not None else 0, encoder_lens, + sinks if has_sink else None, ) torch.testing.assert_close(o_ref, o_extend, atol=1e-2, rtol=1e-2) @@ -224,16 +346,29 @@ class TestExtendAttention(CustomTestCase): if is_mla and is_cross_attn: continue self._test_extend_attention_once( - 1, 123, 1, 1, 128, 96, is_mla, is_cross_attn + 1, 123, 1, 1, 128, 96, None, False, is_mla, is_cross_attn ) self._test_extend_attention_once( - 1, 123, 16, 1, 128, 96, is_mla, is_cross_attn + 1, 123, 16, 1, 128, 96, None, False, is_mla, is_cross_attn ) self._test_extend_attention_once( - 4, 1230, 16, 4, 128, 96, is_mla, is_cross_attn + 4, 1230, 16, 4, 128, 96, None, False, is_mla, is_cross_attn ) self._test_extend_attention_once( - 1, 9000, 16, 1, 32, 32, is_mla, is_cross_attn + 1, 9000, 16, 1, 32, 32, None, False, is_mla, is_cross_attn + ) + for has_sink in [True, False]: + for sliding_window in [None, 10, 128]: + if not has_sink and sliding_window is not None: + continue + self._test_extend_attention_once( + 1, 123, 16, 4, 64, 64, sliding_window, has_sink, False, False + ) + self._test_extend_attention_once( + 1, 20, 16, 1, 64, 64, sliding_window, has_sink, False, False + ) + self._test_extend_attention_once( + 1, 20, 1, 1, 64, 64, sliding_window, has_sink, False, False ) def test_extend_attention_large_seq_causal_mask(self): diff --git a/test/srt/cpu/test_mla.py b/test/srt/cpu/test_mla.py index 3db205581..2047720f3 100644 --- a/test/srt/cpu/test_mla.py +++ b/test/srt/cpu/test_mla.py @@ -116,6 +116,8 @@ class TestMLA(CustomTestCase): sm_scale, logit_cap, False, + 0, + None, None, )