diff --git a/vllm/custom-esimd-kernels-vllm/csrc/moe_batch/moe_int4.sycl b/vllm/custom-esimd-kernels-vllm/csrc/moe_batch/moe_int4.sycl index cb77d85..91fd8ae 100644 --- a/vllm/custom-esimd-kernels-vllm/csrc/moe_batch/moe_int4.sycl +++ b/vllm/custom-esimd-kernels-vllm/csrc/moe_batch/moe_int4.sycl @@ -18,6 +18,9 @@ #include #include #include +#include +#include +#include #include "moe_topk.h" #include "../xpu/esimd_kernels/moe_ops.h" // TopK V2 vectorized argmax @@ -32,7 +35,18 @@ static inline sycl::event submit_kernel( const torch::Device& device, const char* desc) { sycl::queue& queue = c10::xpu::getCurrentXPUStream(device.index()).queue(); - return queue.submit(kernel); + sycl::event event = queue.submit(kernel); + static const bool trace_kernels = + std::getenv("LLM_SCALER_MOE_TRACE_KERNELS") != nullptr; + if (trace_kernels) { + auto start = std::chrono::steady_clock::now(); + event.wait(); + auto end = std::chrono::steady_clock::now(); + double ms = std::chrono::duration(end - start).count(); + std::cerr << "[llm-scaler][moe-kernel] device=" << static_cast(device.index()) + << " kernel=\"" << desc << "\" wait_ms=" << ms << std::endl; + } + return event; } // ═══════════════════════════════════════════════════════════════════════════════ @@ -1438,7 +1452,7 @@ template class MoeRouteFillInt4; template class MoeTinyMUpCutlassInt4; -template +template class MoeTinyMDownCutlassInt4; template class MoeTinyMDownCutlassInt4WithSharedFP16; @@ -1667,6 +1681,10 @@ static inline int decode_s4_nibble(uint8_t nibble) { return v >= 8 ? v - 16 : v; } +static inline int decode_u4_nibble(uint8_t nibble) { + return (int)(nibble & 0x0F) - 8; +} + template void moe_tiny_m_up_cutlass_int4_kernel( const fp16* x, @@ -1675,7 +1693,8 @@ void moe_tiny_m_up_cutlass_int4_kernel( const IndexT* topk_idx, fp16* intermediates, const int n_tokens, const int top_k, const int hidden_size, const int intermediate_size, - const torch::Device& device) { + const torch::Device& device, + const bool signed_compact = true) { const int two_inter = 2 * intermediate_size; const int k_bytes = hidden_size / 2; @@ -1712,10 +1731,10 @@ void moe_tiny_m_up_cutlass_int4_kernel( int group = kb >> 6; uint8_t gate_packed = gate_row[kb]; uint8_t up_packed = up_row[kb]; - int gate0 = decode_s4_nibble(gate_packed & 0x0F); - int gate1 = decode_s4_nibble((gate_packed >> 4) & 0x0F); - int up0 = decode_s4_nibble(up_packed & 0x0F); - int up1 = decode_s4_nibble((up_packed >> 4) & 0x0F); + int gate0 = signed_compact ? decode_s4_nibble(gate_packed & 0x0F) : decode_u4_nibble(gate_packed & 0x0F); + int gate1 = signed_compact ? decode_s4_nibble((gate_packed >> 4) & 0x0F) : decode_u4_nibble((gate_packed >> 4) & 0x0F); + int up0 = signed_compact ? decode_s4_nibble(up_packed & 0x0F) : decode_u4_nibble(up_packed & 0x0F); + int up1 = signed_compact ? decode_s4_nibble((up_packed >> 4) & 0x0F) : decode_u4_nibble((up_packed >> 4) & 0x0F); float x0 = (float)x_row[k0]; float x1 = (float)x_row[k0 + 1]; gate_sum += x0 * ((float)gate0 * (float)gate_s[group]); @@ -2153,16 +2172,17 @@ void moe_ws_down_cutlass_int4_with_shared_fp16_kernel( -template +template void moe_tiny_m_down_cutlass_int4_kernel( const fp16* intermediates, const uint8_t* w2, const fp16* w2_scales, - const fp16* topk_weight, + const WeightT* topk_weight, const IndexT* topk_idx, fp16* output, const int n_tokens, const int top_k, const int hidden_size, const int intermediate_size, - const torch::Device& device) { + const torch::Device& device, + const bool signed_compact = true) { const int k_bytes = intermediate_size / 2; const int k_groups = intermediate_size / 128; @@ -2170,7 +2190,7 @@ void moe_tiny_m_down_cutlass_int4_kernel( auto cgf = [&](sycl::handler& cgh) { sycl::local_accessor local_acc(sycl::range<1>(WG_SIZE), cgh); - cgh.parallel_for>( + cgh.parallel_for>( sycl::nd_range<1>(sycl::range<1>(n_tokens * hidden_size * WG_SIZE), sycl::range<1>(WG_SIZE)), [=](sycl::nd_item<1> item) { @@ -2188,7 +2208,7 @@ void moe_tiny_m_down_cutlass_int4_kernel( for (int k = lid; k < intermediate_size; k += WG_SIZE) { uint8_t packed = w_row[k >> 1]; uint8_t nibble = (k & 1) ? ((packed >> 4) & 0x0F) : (packed & 0x0F); - int v = decode_s4_nibble(nibble); + int v = signed_compact ? decode_s4_nibble(nibble) : decode_u4_nibble(nibble); float scale = (float)s_row[k >> 7]; sum += route_weight * (float)in_row[k] * ((float)v * scale); } @@ -2383,14 +2403,15 @@ void moe_tiny_m_down_cutlass_int4_with_shared_fp16_htile_kernel( submit_kernel(cgf, device, "moe tiny m down cutlass int4 with shared fp16 htile"); } -torch::Tensor moe_forward_tiny_cutlass_nmajor_int4( +static torch::Tensor moe_forward_tiny_cutlass_nmajor_int4_impl( torch::Tensor x, torch::Tensor w13_qweight_s4, torch::Tensor w13_scales, torch::Tensor w2_qweight_s4, torch::Tensor w2_scales, torch::Tensor topk_weight, - torch::Tensor topk_idx) { + torch::Tensor topk_idx, + const bool signed_compact) { TORCH_CHECK(x.dim() == 2 && x.size(0) >= 1 && x.size(0) <= 64 && x.is_contiguous()); TORCH_CHECK(x.scalar_type() == torch::kHalf); @@ -2399,7 +2420,7 @@ torch::Tensor moe_forward_tiny_cutlass_nmajor_int4( TORCH_CHECK(w13_scales.scalar_type() == torch::kHalf && w13_scales.is_contiguous()); TORCH_CHECK(w2_scales.scalar_type() == torch::kHalf && w2_scales.is_contiguous()); TORCH_CHECK(topk_weight.dim() == 2 && topk_weight.size(0) == x.size(0) && topk_weight.is_contiguous()); - TORCH_CHECK(topk_weight.scalar_type() == torch::kHalf); + TORCH_CHECK(topk_weight.scalar_type() == torch::kHalf || topk_weight.scalar_type() == torch::kFloat); TORCH_CHECK(topk_idx.dim() == 2 && topk_idx.size(0) == x.size(0) && topk_idx.is_contiguous()); TORCH_CHECK(topk_idx.scalar_type() == torch::kInt || topk_idx.scalar_type() == torch::kLong); @@ -2423,28 +2444,70 @@ torch::Tensor moe_forward_tiny_cutlass_nmajor_int4( (const fp16*)x.data_ptr(), (const uint8_t*)w13_qweight_s4.data_ptr(), (const fp16*)w13_scales.data_ptr(), idx_ptr, (fp16*)intermediates.data_ptr(), n_tokens, top_k, hidden_size, - intermediate_size, x.device()); - moe_tiny_m_down_cutlass_int4_kernel( - (const fp16*)intermediates.data_ptr(), (const uint8_t*)w2_qweight_s4.data_ptr(), - (const fp16*)w2_scales.data_ptr(), (const fp16*)topk_weight.data_ptr(), - idx_ptr, (fp16*)output.data_ptr(), n_tokens, top_k, hidden_size, - intermediate_size, x.device()); + intermediate_size, x.device(), signed_compact); + if (topk_weight.scalar_type() == torch::kFloat) { + moe_tiny_m_down_cutlass_int4_kernel( + (const fp16*)intermediates.data_ptr(), (const uint8_t*)w2_qweight_s4.data_ptr(), + (const fp16*)w2_scales.data_ptr(), (const float*)topk_weight.data_ptr(), + idx_ptr, (fp16*)output.data_ptr(), n_tokens, top_k, hidden_size, + intermediate_size, x.device(), signed_compact); + } else { + moe_tiny_m_down_cutlass_int4_kernel( + (const fp16*)intermediates.data_ptr(), (const uint8_t*)w2_qweight_s4.data_ptr(), + (const fp16*)w2_scales.data_ptr(), (const fp16*)topk_weight.data_ptr(), + idx_ptr, (fp16*)output.data_ptr(), n_tokens, top_k, hidden_size, + intermediate_size, x.device(), signed_compact); + } } else { const int64_t* idx_ptr = topk_idx.data_ptr(); moe_tiny_m_up_cutlass_int4_kernel( (const fp16*)x.data_ptr(), (const uint8_t*)w13_qweight_s4.data_ptr(), (const fp16*)w13_scales.data_ptr(), idx_ptr, (fp16*)intermediates.data_ptr(), n_tokens, top_k, hidden_size, - intermediate_size, x.device()); - moe_tiny_m_down_cutlass_int4_kernel( - (const fp16*)intermediates.data_ptr(), (const uint8_t*)w2_qweight_s4.data_ptr(), - (const fp16*)w2_scales.data_ptr(), (const fp16*)topk_weight.data_ptr(), - idx_ptr, (fp16*)output.data_ptr(), n_tokens, top_k, hidden_size, - intermediate_size, x.device()); + intermediate_size, x.device(), signed_compact); + if (topk_weight.scalar_type() == torch::kFloat) { + moe_tiny_m_down_cutlass_int4_kernel( + (const fp16*)intermediates.data_ptr(), (const uint8_t*)w2_qweight_s4.data_ptr(), + (const fp16*)w2_scales.data_ptr(), (const float*)topk_weight.data_ptr(), + idx_ptr, (fp16*)output.data_ptr(), n_tokens, top_k, hidden_size, + intermediate_size, x.device(), signed_compact); + } else { + moe_tiny_m_down_cutlass_int4_kernel( + (const fp16*)intermediates.data_ptr(), (const uint8_t*)w2_qweight_s4.data_ptr(), + (const fp16*)w2_scales.data_ptr(), (const fp16*)topk_weight.data_ptr(), + idx_ptr, (fp16*)output.data_ptr(), n_tokens, top_k, hidden_size, + intermediate_size, x.device(), signed_compact); + } } return output; } +torch::Tensor moe_forward_tiny_cutlass_nmajor_int4( + torch::Tensor x, + torch::Tensor w13_qweight_s4, + torch::Tensor w13_scales, + torch::Tensor w2_qweight_s4, + torch::Tensor w2_scales, + torch::Tensor topk_weight, + torch::Tensor topk_idx) { + return moe_forward_tiny_cutlass_nmajor_int4_impl( + x, w13_qweight_s4, w13_scales, w2_qweight_s4, w2_scales, + topk_weight, topk_idx, true); +} + +torch::Tensor moe_forward_tiny_cutlass_nmajor_int4_u4( + torch::Tensor x, + torch::Tensor w13_qweight_u4, + torch::Tensor w13_scales, + torch::Tensor w2_qweight_u4, + torch::Tensor w2_scales, + torch::Tensor topk_weight, + torch::Tensor topk_idx) { + return moe_forward_tiny_cutlass_nmajor_int4_impl( + x, w13_qweight_u4, w13_scales, w2_qweight_u4, w2_scales, + topk_weight, topk_idx, false); +} + torch::Tensor moe_tiny_cutlass_nmajor_int4_up( torch::Tensor x, torch::Tensor w13_qweight_s4, @@ -2511,13 +2574,13 @@ torch::Tensor moe_tiny_cutlass_nmajor_int4_down( torch::device(intermediates.device()).dtype(torch::kHalf)); if (topk_idx.scalar_type() == torch::kInt) { - moe_tiny_m_down_cutlass_int4_kernel( + moe_tiny_m_down_cutlass_int4_kernel( (const fp16*)intermediates.data_ptr(), (const uint8_t*)w2_qweight_s4.data_ptr(), (const fp16*)w2_scales.data_ptr(), (const fp16*)topk_weight.data_ptr(), topk_idx.data_ptr(), (fp16*)output.data_ptr(), n_tokens, top_k, hidden_size, intermediate_size, intermediates.device()); } else { - moe_tiny_m_down_cutlass_int4_kernel( + moe_tiny_m_down_cutlass_int4_kernel( (const fp16*)intermediates.data_ptr(), (const uint8_t*)w2_qweight_s4.data_ptr(), (const fp16*)w2_scales.data_ptr(), (const fp16*)topk_weight.data_ptr(), topk_idx.data_ptr(), (fp16*)output.data_ptr(), n_tokens, top_k, @@ -3585,6 +3648,8 @@ TORCH_LIBRARY_FRAGMENT(moe_int4_ops, m) { m.def("moe_route_gather_int4(Tensor route_output, Tensor sorted_rows, Tensor sorted_weights, int n_tokens) -> Tensor"); m.def("moe_forward_tiny_cutlass_nmajor_int4(Tensor x, Tensor w13_qweight_s4, Tensor w13_scales, " "Tensor w2_qweight_s4, Tensor w2_scales, Tensor topk_weight, Tensor topk_idx) -> Tensor"); + m.def("moe_forward_tiny_cutlass_nmajor_int4_u4(Tensor x, Tensor w13_qweight_u4, Tensor w13_scales, " + "Tensor w2_qweight_u4, Tensor w2_scales, Tensor topk_weight, Tensor topk_idx) -> Tensor"); m.def("moe_tiny_cutlass_nmajor_int4_up(Tensor x, Tensor w13_qweight_s4, Tensor w13_scales, Tensor topk_idx) -> Tensor"); m.def("moe_tiny_cutlass_nmajor_int4_down(Tensor intermediates, Tensor w2_qweight_s4, Tensor w2_scales, " "Tensor topk_weight, Tensor topk_idx) -> Tensor"); @@ -3627,6 +3692,7 @@ TORCH_LIBRARY_IMPL(moe_int4_ops, XPU, m) { m.impl("moe_silu_mul_int4", &moe_silu_mul_int4); m.impl("moe_route_gather_int4", &moe_route_gather_int4); m.impl("moe_forward_tiny_cutlass_nmajor_int4", &moe_forward_tiny_cutlass_nmajor_int4); + m.impl("moe_forward_tiny_cutlass_nmajor_int4_u4", &moe_forward_tiny_cutlass_nmajor_int4_u4); m.impl("moe_tiny_cutlass_nmajor_int4_up", &moe_tiny_cutlass_nmajor_int4_up); m.impl("moe_tiny_cutlass_nmajor_int4_down", &moe_tiny_cutlass_nmajor_int4_down); m.impl("moe_forward_tiny_cutlass_nmajor_int4_full_fp16_shared", &moe_forward_tiny_cutlass_nmajor_int4_full_fp16_shared); @@ -3647,6 +3713,7 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("moe_silu_mul_int4", &moe_silu_mul_int4); m.def("moe_route_gather_int4", &moe_route_gather_int4); m.def("moe_forward_tiny_cutlass_nmajor_int4", &moe_forward_tiny_cutlass_nmajor_int4); + m.def("moe_forward_tiny_cutlass_nmajor_int4_u4", &moe_forward_tiny_cutlass_nmajor_int4_u4); m.def("moe_tiny_cutlass_nmajor_int4_up", &moe_tiny_cutlass_nmajor_int4_up); m.def("moe_tiny_cutlass_nmajor_int4_down", &moe_tiny_cutlass_nmajor_int4_down); m.def("moe_forward_tiny_cutlass_nmajor_int4_full_fp16_shared", &moe_forward_tiny_cutlass_nmajor_int4_full_fp16_shared); diff --git a/vllm/custom-esimd-kernels-vllm/python/custom_esimd_kernels_vllm/__init__.py b/vllm/custom-esimd-kernels-vllm/python/custom_esimd_kernels_vllm/__init__.py index ad1bbb8..e17a062 100644 --- a/vllm/custom-esimd-kernels-vllm/python/custom_esimd_kernels_vllm/__init__.py +++ b/vllm/custom-esimd-kernels-vllm/python/custom_esimd_kernels_vllm/__init__.py @@ -1,19 +1,29 @@ +import importlib + import torch + +def _optional_import(name): + try: + return importlib.import_module(f"{__name__}.{name}") + except ImportError: + return None + + # Core ESIMD kernels (4 compiled modules) -from custom_esimd_kernels_vllm import custom_esimd_kernels -from custom_esimd_kernels_vllm import custom_esimd_kernels_lgrf -from custom_esimd_kernels_vllm import custom_esimd_kernels_moe -from custom_esimd_kernels_vllm import custom_esimd_kernels_gemm +custom_esimd_kernels = _optional_import("custom_esimd_kernels") +custom_esimd_kernels_lgrf = _optional_import("custom_esimd_kernels_lgrf") +custom_esimd_kernels_moe = _optional_import("custom_esimd_kernels_moe") +custom_esimd_kernels_gemm = _optional_import("custom_esimd_kernels_gemm") -# Eagle kernels — registers torch.ops.eagle_ops.* -from custom_esimd_kernels_vllm import eagle_ops +# Eagle kernels - registers torch.ops.eagle_ops.* +eagle_ops = _optional_import("eagle_ops") -# MoE Batch kernels — registers torch.ops.moe_ops.* -from custom_esimd_kernels_vllm import moe_ops +# MoE Batch kernels - registers torch.ops.moe_ops.* +moe_ops = _optional_import("moe_ops") -# MoE INT4 Batch kernels — registers torch.ops.moe_int4_ops.* -from custom_esimd_kernels_vllm import moe_int4_ops +# MoE INT4 Batch kernels - registers torch.ops.moe_int4_ops.* +moe_int4_ops = _optional_import("moe_int4_ops") from custom_esimd_kernels_vllm.ops import ( # Core ESIMD ops @@ -70,6 +80,7 @@ from custom_esimd_kernels_vllm.ops import ( moe_forward_full_cutlass_nmajor_int4, moe_forward_full_cutlass_nmajor_int4_with_router, moe_forward_tiny_cutlass_nmajor_int4, + moe_forward_tiny_cutlass_nmajor_int4_u4, moe_forward_tiny_cutlass_nmajor_int4_full_fp16_shared, moe_forward_tiny_cutlass_nmajor_int4_full_fp16_shared_from_logits, moe_tiny_cutlass_nmajor_int4_up, diff --git a/vllm/custom-esimd-kernels-vllm/python/custom_esimd_kernels_vllm/ops.py b/vllm/custom-esimd-kernels-vllm/python/custom_esimd_kernels_vllm/ops.py index 328d283..18c32d7 100644 --- a/vllm/custom-esimd-kernels-vllm/python/custom_esimd_kernels_vllm/ops.py +++ b/vllm/custom-esimd-kernels-vllm/python/custom_esimd_kernels_vllm/ops.py @@ -1062,7 +1062,11 @@ def moe_route_gather_int4( pass output = torch.zeros(n_tokens, route_output.shape[1], dtype=route_output.dtype, device=route_output.device) - output.index_add_(0, sorted_rows, route_output * sorted_weights.unsqueeze(-1)) + output.index_add_( + 0, + sorted_rows, + route_output * sorted_weights.to(route_output.dtype).unsqueeze(-1), + ) return output @@ -1271,6 +1275,31 @@ def moe_forward_tiny_cutlass_nmajor_int4( topk_weights.contiguous(), topk_ids.contiguous()) +def moe_forward_tiny_cutlass_nmajor_int4_u4( + hidden_states: torch.Tensor, + w13_qweight_u4: torch.Tensor, + w13_scales: torch.Tensor, + w2_qweight_u4: torch.Tensor, + w2_scales: torch.Tensor, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, +) -> torch.Tensor: + """Tiny routed MoE using unsigned uint4 N-major weights. + + vLLM stores symmetric W4A16 expert weights as uint4 values with an + implicit zero point of 8. This variant subtracts that zero point inside the + kernel instead of requiring a signed compact-int4 copy. The route weights + may be FP32 or FP16; accepting FP32 avoids a per-layer cast in vLLM decode. + """ + if hidden_states.device.type != "xpu": + raise RuntimeError("tiny CUTLASS N-major INT4 path requires XPU") + return _moe_int4.moe_forward_tiny_cutlass_nmajor_int4_u4( + hidden_states.contiguous(), + w13_qweight_u4.contiguous(), w13_scales.contiguous(), + w2_qweight_u4.contiguous(), w2_scales.contiguous(), + topk_weights.contiguous(), topk_ids.contiguous()) + + def moe_tiny_cutlass_nmajor_int4_up( hidden_states: torch.Tensor, w13_qweight_s4: torch.Tensor, diff --git a/vllm/custom-esimd-kernels-vllm/setup_moe_int4_only.py b/vllm/custom-esimd-kernels-vllm/setup_moe_int4_only.py index 8f992f3..7192240 100644 --- a/vllm/custom-esimd-kernels-vllm/setup_moe_int4_only.py +++ b/vllm/custom-esimd-kernels-vllm/setup_moe_int4_only.py @@ -13,6 +13,7 @@ root = Path(__file__).parent.resolve() import torch torch_include = str(Path(torch.__file__).parent / "include") +venv_lib = str(Path(os.environ.get("VIRTUAL_ENV", "")) / "lib") setup( @@ -31,6 +32,9 @@ setup( root / "csrc" / "xpu" / "esimd_kernels", root / "csrc", ], + library_dirs=[ + venv_lib, + ], extra_compile_args={ "cxx": ["-O3", "-std=c++20"], "sycl": [ @@ -40,12 +44,16 @@ setup( "-fsycl-targets=spir64_gen", "-Xs", "-device bmg", + "-D__DPCPP_SYCL_EXTERNAL_LIBC=__DPCPP_SYCL_EXTERNAL", f"-I{torch_include}", ], }, - extra_link_args=["-Wl,-rpath,$ORIGIN/../../torch/lib"], + extra_link_args=[ + f"-Wl,-rpath,{venv_lib}", + "-Wl,-rpath,$ORIGIN/../../torch/lib", + ], py_limited_api=False, ) ], cmdclass={"build_ext": BuildExtension.with_options(use_ninja=True)}, -) \ No newline at end of file +)