From 3f89ad2ff896a39544479e0b628eee87ab6ec6d0 Mon Sep 17 00:00:00 2001 From: Jiacheng Huang Date: Sat, 8 Aug 2026 20:46:49 +0800 Subject: [PATCH 1/2] feat(nvidia): link vLLM `moe_wna16_marlin_gemm` --- src/base/moe_wna16_marlin_gemm.h | 301 ++++++++++++++++++ .../nvidia/ops/moe_wna16_marlin_gemm/vllm.cc | 8 + .../nvidia/ops/moe_wna16_marlin_gemm/vllm.h | 111 +++++++ .../ops/moe_wna16_marlin_gemm/vllm.yaml | 11 + src/linked/torch/ops/moe_wna16_marlin_gemm.h | 73 +++++ tests/test_moe_wna16_marlin_gemm.py | 237 ++++++++++++++ 6 files changed, 741 insertions(+) create mode 100644 src/base/moe_wna16_marlin_gemm.h create mode 100644 src/linked/torch/nvidia/ops/moe_wna16_marlin_gemm/vllm.cc create mode 100644 src/linked/torch/nvidia/ops/moe_wna16_marlin_gemm/vllm.h create mode 100644 src/linked/torch/nvidia/ops/moe_wna16_marlin_gemm/vllm.yaml create mode 100644 src/linked/torch/ops/moe_wna16_marlin_gemm.h create mode 100644 tests/test_moe_wna16_marlin_gemm.py diff --git a/src/base/moe_wna16_marlin_gemm.h b/src/base/moe_wna16_marlin_gemm.h new file mode 100644 index 000000000..6a790ffbe --- /dev/null +++ b/src/base/moe_wna16_marlin_gemm.h @@ -0,0 +1,301 @@ +#ifndef INFINI_OPS_BASE_MOE_WNA16_MARLIN_GEMM_H_ +#define INFINI_OPS_BASE_MOE_WNA16_MARLIN_GEMM_H_ + +#include +#include +#include +#include +#include + +#include "operator.h" + +namespace infini::ops { + +// Aligned with vLLM `_moe_C::moe_wna16_marlin_gemm` at commit +// bcc0a3cbefe55f99da4821f9d89106e3d71e4867. +class MoeWna16MarlinGemm : public Operator { + public: + MoeWna16MarlinGemm( + const Tensor a, const Tensor b_q_weight, const Tensor b_scales, + std::optional global_scale, std::optional b_zeros_or_none, + std::optional g_idx_or_none, std::optional perm_or_none, + const Tensor workspace, const Tensor sorted_token_ids, + const Tensor expert_ids, const Tensor num_tokens_past_padded, + const Tensor topk_weights, const int64_t moe_block_size, + const int64_t top_k, const bool mul_topk_weights, const bool is_ep, + const int64_t b_q_type_id, const int64_t size_m, const int64_t size_n, + const int64_t size_k, const bool is_full_k, const bool use_atomic_add, + const bool use_fp32_reduce, const bool is_zp_float, Tensor out) + : a_metadata_{a}, + b_q_weight_metadata_{b_q_weight}, + b_scales_metadata_{b_scales}, + global_scale_metadata_{global_scale}, + b_zeros_or_none_metadata_{b_zeros_or_none}, + g_idx_or_none_metadata_{g_idx_or_none}, + perm_or_none_metadata_{perm_or_none}, + workspace_metadata_{workspace}, + sorted_token_ids_metadata_{sorted_token_ids}, + expert_ids_metadata_{expert_ids}, + num_tokens_past_padded_metadata_{num_tokens_past_padded}, + topk_weights_metadata_{topk_weights}, + out_metadata_{out}, + moe_block_size_{moe_block_size}, + top_k_{top_k}, + mul_topk_weights_{mul_topk_weights}, + is_ep_{is_ep}, + b_q_type_id_{b_q_type_id}, + size_m_{size_m}, + size_n_{size_n}, + size_k_{size_k}, + is_full_k_{is_full_k}, + use_atomic_add_{use_atomic_add}, + use_fp32_reduce_{use_fp32_reduce}, + is_zp_float_{is_zp_float}, + device_index_{a.device().index()} { + Validate(a, b_q_weight, b_scales, global_scale, b_zeros_or_none, + g_idx_or_none, perm_or_none, workspace, sorted_token_ids, + expert_ids, num_tokens_past_padded, topk_weights, out); + } + + virtual void operator()( + const Tensor a, const Tensor b_q_weight, const Tensor b_scales, + std::optional global_scale, std::optional b_zeros_or_none, + std::optional g_idx_or_none, std::optional perm_or_none, + const Tensor workspace, const Tensor sorted_token_ids, + const Tensor expert_ids, const Tensor num_tokens_past_padded, + const Tensor topk_weights, const int64_t moe_block_size, + const int64_t top_k, const bool mul_topk_weights, const bool is_ep, + const int64_t b_q_type_id, const int64_t size_m, const int64_t size_n, + const int64_t size_k, const bool is_full_k, const bool use_atomic_add, + const bool use_fp32_reduce, const bool is_zp_float, Tensor out) const = 0; + + protected: + void ValidateCallMetadata( + const Tensor a, const Tensor b_q_weight, const Tensor b_scales, + std::optional global_scale, std::optional b_zeros_or_none, + std::optional g_idx_or_none, std::optional perm_or_none, + const Tensor workspace, const Tensor sorted_token_ids, + const Tensor expert_ids, const Tensor num_tokens_past_padded, + const Tensor topk_weights, const int64_t moe_block_size, + const int64_t top_k, const bool mul_topk_weights, const bool is_ep, + const int64_t b_q_type_id, const int64_t size_m, const int64_t size_n, + const int64_t size_k, const bool is_full_k, const bool use_atomic_add, + const bool use_fp32_reduce, const bool is_zp_float, Tensor out) const { + assert(moe_block_size == moe_block_size_ && top_k == top_k_ && + mul_topk_weights == mul_topk_weights_ && is_ep == is_ep_ && + b_q_type_id == b_q_type_id_ && size_m == size_m_ && + size_n == size_n_ && size_k == size_k_ && is_full_k == is_full_k_ && + use_atomic_add == use_atomic_add_ && + use_fp32_reduce == use_fp32_reduce_ && is_zp_float == is_zp_float_ && + "`MoeWna16MarlinGemm` attributes changed after descriptor " + "creation"); + + const std::equal_to same_metadata; + const auto optional_matches = [&](const std::optional& expected, + const std::optional& actual) { + return expected.has_value() == actual.has_value() && + (!expected || same_metadata(*expected, *actual)); + }; + const auto matches = + same_metadata(a_metadata_, a) && + same_metadata(b_q_weight_metadata_, b_q_weight) && + same_metadata(b_scales_metadata_, b_scales) && + optional_matches(global_scale_metadata_, global_scale) && + optional_matches(b_zeros_or_none_metadata_, b_zeros_or_none) && + optional_matches(g_idx_or_none_metadata_, g_idx_or_none) && + optional_matches(perm_or_none_metadata_, perm_or_none) && + same_metadata(workspace_metadata_, workspace) && + same_metadata(sorted_token_ids_metadata_, sorted_token_ids) && + same_metadata(expert_ids_metadata_, expert_ids) && + same_metadata(num_tokens_past_padded_metadata_, + num_tokens_past_padded) && + same_metadata(topk_weights_metadata_, topk_weights) && + same_metadata(out_metadata_, out); + assert(matches && + "`MoeWna16MarlinGemm` tensor metadata must match descriptor"); + } + + private: + void Validate(const Tensor a, const Tensor b_q_weight, const Tensor b_scales, + std::optional global_scale, + std::optional b_zeros_or_none, + std::optional g_idx_or_none, + std::optional perm_or_none, const Tensor workspace, + const Tensor sorted_token_ids, const Tensor expert_ids, + const Tensor num_tokens_past_padded, const Tensor topk_weights, + const Tensor out) const { + assert(a.ndim() == 2 && a.size(0) == size_m_ && a.size(1) == size_k_ && + "`MoeWna16MarlinGemm` `a` shape must match `size_m` and " + "`size_k`"); + assert( + (a.dtype() == DataType::kFloat16 || a.dtype() == DataType::kBFloat16) && + a.IsContiguous() && + "`MoeWna16MarlinGemm` requires contiguous float16 or bfloat16 " + "`a`"); + assert(size_m_ > 0 && size_n_ > 0 && size_k_ > 0 && top_k_ > 0 && + size_k_ % 16 == 0 && size_n_ % 64 == 0 && + "`MoeWna16MarlinGemm` received unsupported dimensions"); + assert((moe_block_size_ == 8 || + (moe_block_size_ >= 16 && moe_block_size_ <= 64 && + moe_block_size_ % 16 == 0)) && + "`MoeWna16MarlinGemm` received an unsupported `moe_block_size`"); + assert(size_m_ <= std::numeric_limits::max() / top_k_ && + "`MoeWna16MarlinGemm` output dimensions overflow"); + + constexpr int64_t kUint4B8 = 1125899907892224; + constexpr int64_t kUint8B128 = 1125899923621888; + constexpr int64_t kUint4 = 1125899906843648; + constexpr int64_t kUint8 = 1125899906844672; + constexpr int64_t kFloat8E4M3Fn = 2814749767172868; + const auto supported_qtype = + b_q_type_id_ == kUint4B8 || b_q_type_id_ == kUint8B128 || + b_q_type_id_ == kUint4 || b_q_type_id_ == kUint8 || + b_q_type_id_ == kFloat8E4M3Fn; + const auto has_zero_points = b_zeros_or_none.has_value(); + assert(supported_qtype && + has_zero_points == + (b_q_type_id_ == kUint4 || b_q_type_id_ == kUint8) && + !global_scale.has_value() && + (!is_zp_float_ || + (has_zero_points && a.dtype() == DataType::kFloat16)) && + "`MoeWna16MarlinGemm` received an unsupported quantization " + "configuration"); + + const auto pack_factor = + (b_q_type_id_ == kUint8B128 || b_q_type_id_ == kUint8 || + b_q_type_id_ == kFloat8E4M3Fn) + ? 4 + : 8; + assert(size_n_ <= std::numeric_limits::max() / 16 && + b_q_weight.ndim() == 3 && b_q_weight.size(1) == size_k_ / 16 && + b_q_weight.size(2) == size_n_ * 16 / pack_factor && + b_q_weight.dtype() == DataType::kInt32 && + b_q_weight.IsContiguous() && + "`MoeWna16MarlinGemm` received invalid packed weights"); + assert(b_scales.ndim() == 3 && b_scales.size(0) == b_q_weight.size(0) && + b_scales.size(1) > 0 && b_scales.size(2) == size_n_ && + size_k_ % b_scales.size(1) == 0 && b_scales.dtype() == a.dtype() && + b_scales.IsContiguous() && + "`MoeWna16MarlinGemm` received invalid weight scales"); + if (b_zeros_or_none) { + assert(b_zeros_or_none->ndim() == 3 && + b_zeros_or_none->size(0) == b_q_weight.size(0) && + b_zeros_or_none->size(1) == b_scales.size(1) && + b_zeros_or_none->size(2) == + (is_zp_float_ ? size_n_ : size_n_ / pack_factor) && + (is_zp_float_ ? b_zeros_or_none->dtype() == a.dtype() + : b_zeros_or_none->dtype() == DataType::kInt32) && + "`MoeWna16MarlinGemm` received invalid zero points"); + } + + const auto same_device = [&](const Tensor tensor) { + return tensor.device().type() == a.device().type() && + tensor.device().index() == a.device().index(); + }; + const auto valid_optional = [&](const std::optional& tensor) { + return !tensor || (tensor->IsContiguous() && same_device(*tensor)); + }; + assert(same_device(b_q_weight) && same_device(b_scales) && + valid_optional(global_scale) && valid_optional(b_zeros_or_none) && + valid_optional(g_idx_or_none) && valid_optional(perm_or_none) && + same_device(workspace) && same_device(sorted_token_ids) && + same_device(expert_ids) && same_device(num_tokens_past_padded) && + same_device(topk_weights) && same_device(out) && + "`MoeWna16MarlinGemm` requires all tensors on the input device"); + + assert(g_idx_or_none.has_value() == perm_or_none.has_value() && + "`MoeWna16MarlinGemm` requires `g_idx_or_none` and `perm_or_none` " + "together"); + if (g_idx_or_none) { + assert(g_idx_or_none->ndim() > 0 && perm_or_none->ndim() > 0 && + g_idx_or_none->size(-1) == perm_or_none->size(-1) && + (g_idx_or_none->size(-1) == 0 || + g_idx_or_none->size(-1) == size_k_) && + g_idx_or_none->dtype() == DataType::kInt32 && + perm_or_none->dtype() == DataType::kInt32 && + (!is_full_k_ || b_scales.size(1) > 1) && + "`MoeWna16MarlinGemm` received invalid activation-order " + "metadata"); + } + + assert(workspace.ndim() == 1 && workspace.numel() > 0 && + workspace.dtype() == DataType::kInt32 && workspace.IsContiguous() && + "`MoeWna16MarlinGemm` requires a non-empty int32 workspace"); + assert(sorted_token_ids.ndim() == 1 && + sorted_token_ids.dtype() == DataType::kInt32 && + sorted_token_ids.IsContiguous() && expert_ids.ndim() == 1 && + expert_ids.dtype() == DataType::kInt32 && + expert_ids.IsContiguous() && num_tokens_past_padded.numel() == 1 && + num_tokens_past_padded.dtype() == DataType::kInt32 && + num_tokens_past_padded.IsContiguous() && + "`MoeWna16MarlinGemm` received invalid routing metadata"); + assert(topk_weights.shape() == + Tensor::Shape({static_cast(size_m_), + static_cast(top_k_)}) && + topk_weights.dtype() == DataType::kFloat32 && + topk_weights.IsContiguous() && + "`MoeWna16MarlinGemm` received invalid top-k weights"); + assert(out.shape() == + Tensor::Shape({static_cast(size_m_ * top_k_), + static_cast(size_n_)}) && + out.dtype() == a.dtype() && out.IsContiguous() && + "`MoeWna16MarlinGemm` output metadata is invalid"); + } + + Tensor a_metadata_; + + Tensor b_q_weight_metadata_; + + Tensor b_scales_metadata_; + + std::optional global_scale_metadata_; + + std::optional b_zeros_or_none_metadata_; + + std::optional g_idx_or_none_metadata_; + + std::optional perm_or_none_metadata_; + + Tensor workspace_metadata_; + + Tensor sorted_token_ids_metadata_; + + Tensor expert_ids_metadata_; + + Tensor num_tokens_past_padded_metadata_; + + Tensor topk_weights_metadata_; + + Tensor out_metadata_; + + int64_t moe_block_size_{0}; + + int64_t top_k_{0}; + + bool mul_topk_weights_{false}; + + bool is_ep_{false}; + + int64_t b_q_type_id_{0}; + + int64_t size_m_{0}; + + int64_t size_n_{0}; + + int64_t size_k_{0}; + + bool is_full_k_{false}; + + bool use_atomic_add_{false}; + + bool use_fp32_reduce_{false}; + + bool is_zp_float_{false}; + + protected: + int device_index_{0}; +}; + +} // namespace infini::ops + +#endif // INFINI_OPS_BASE_MOE_WNA16_MARLIN_GEMM_H_ diff --git a/src/linked/torch/nvidia/ops/moe_wna16_marlin_gemm/vllm.cc b/src/linked/torch/nvidia/ops/moe_wna16_marlin_gemm/vllm.cc new file mode 100644 index 000000000..e6331d401 --- /dev/null +++ b/src/linked/torch/nvidia/ops/moe_wna16_marlin_gemm/vllm.cc @@ -0,0 +1,8 @@ +#include "linked/torch/nvidia/ops/moe_wna16_marlin_gemm/vllm.h" + +namespace infini::ops::linked::torch { + +template class TorchMoeWna16MarlinGemm< + ::infini::ops::linked::torch::nvidia::VllmMoeWna16MarlinGemm>; + +} // namespace infini::ops::linked::torch diff --git a/src/linked/torch/nvidia/ops/moe_wna16_marlin_gemm/vllm.h b/src/linked/torch/nvidia/ops/moe_wna16_marlin_gemm/vllm.h new file mode 100644 index 000000000..bffd2774f --- /dev/null +++ b/src/linked/torch/nvidia/ops/moe_wna16_marlin_gemm/vllm.h @@ -0,0 +1,111 @@ +#ifndef INFINI_OPS_LINKED_TORCH_NVIDIA_OPS_MOE_WNA16_MARLIN_GEMM_VLLM_H_ +#define INFINI_OPS_LINKED_TORCH_NVIDIA_OPS_MOE_WNA16_MARLIN_GEMM_VLLM_H_ + +#include +#include +#include + +#include +#include + +#include "linked/torch/nvidia/c10.h" +#include "linked/torch/ops/moe_wna16_marlin_gemm.h" + +namespace infini::ops::linked::torch::nvidia { + +struct VllmMoeWna16MarlinGemm : C10 { + static void Validate(const int64_t b_q_type_id, const bool has_global_scale) { + constexpr int64_t kFloat4E2M1F = 562949953487106; + TORCH_CHECK( + b_q_type_id != kFloat4E2M1F, + "Linked `moe_wna16_marlin_gemm` does not support `float4_e2m1f` " + "because InfiniRT cannot represent its float8 scales."); + TORCH_CHECK( + !has_global_scale, + "Linked `moe_wna16_marlin_gemm` does not support `global_scale`."); + } + + static void Call(at::Tensor a, at::Tensor out, at::Tensor b_q_weight, + at::Tensor b_scales, std::optional global_scale, + std::optional b_zeros_or_none, + std::optional g_idx_or_none, + std::optional perm_or_none, at::Tensor workspace, + at::Tensor sorted_token_ids, at::Tensor expert_ids, + at::Tensor num_tokens_past_padded, at::Tensor topk_weights, + int64_t moe_block_size, int64_t top_k, bool mul_topk_weights, + bool is_ep, int64_t b_q_type_id, int64_t size_m, + int64_t size_n, int64_t size_k, bool is_full_k, + bool use_atomic_add, bool use_fp32_reduce, + bool is_zp_float) { + static const auto op = c10::Dispatcher::singleton().findSchemaOrThrow( + "_moe_C::moe_wna16_marlin_gemm", ""); + c10::Stack stack; + stack.reserve(25); + stack.emplace_back(std::move(a)); + stack.emplace_back(out); + stack.emplace_back(std::move(b_q_weight)); + stack.emplace_back(std::move(b_scales)); + stack.emplace_back(global_scale ? c10::IValue(std::move(*global_scale)) + : c10::IValue()); + stack.emplace_back(b_zeros_or_none + ? c10::IValue(std::move(*b_zeros_or_none)) + : c10::IValue()); + stack.emplace_back(g_idx_or_none ? c10::IValue(std::move(*g_idx_or_none)) + : c10::IValue()); + stack.emplace_back(perm_or_none ? c10::IValue(std::move(*perm_or_none)) + : c10::IValue()); + stack.emplace_back(std::move(workspace)); + stack.emplace_back(std::move(sorted_token_ids)); + stack.emplace_back(std::move(expert_ids)); + stack.emplace_back(std::move(num_tokens_past_padded)); + stack.emplace_back(std::move(topk_weights)); + stack.emplace_back(moe_block_size); + stack.emplace_back(top_k); + stack.emplace_back(mul_topk_weights); + stack.emplace_back(is_ep); + stack.emplace_back(b_q_type_id); + stack.emplace_back(size_m); + stack.emplace_back(size_n); + stack.emplace_back(size_k); + stack.emplace_back(is_full_k); + stack.emplace_back(use_atomic_add); + stack.emplace_back(use_fp32_reduce); + stack.emplace_back(is_zp_float); + op.callBoxed(&stack); + + TORCH_CHECK(stack.size() == 1, + "Linked `moe_wna16_marlin_gemm` returned an unexpected " + "number of values."); + auto result = std::move(stack.front()).toTensor(); + TORCH_CHECK(result.unsafeGetTensorImpl() == out.unsafeGetTensorImpl(), + "Linked `moe_wna16_marlin_gemm` did not return the provided " + "output tensor."); + } +}; + +} // namespace infini::ops::linked::torch::nvidia + +namespace infini::ops::linked::torch { + +extern template class TorchMoeWna16MarlinGemm< + ::infini::ops::linked::torch::nvidia::VllmMoeWna16MarlinGemm>; + +} // namespace infini::ops::linked::torch + +namespace infini::ops { + +template <> +class Operator + : public linked::torch::TorchMoeWna16MarlinGemm< + linked::torch::nvidia::VllmMoeWna16MarlinGemm> { + public: + using linked::torch::TorchMoeWna16MarlinGemm< + linked::torch::nvidia::VllmMoeWna16MarlinGemm>::TorchMoeWna16MarlinGemm; + + using linked::torch::TorchMoeWna16MarlinGemm< + linked::torch::nvidia::VllmMoeWna16MarlinGemm>::operator(); +}; + +} // namespace infini::ops + +#endif // INFINI_OPS_LINKED_TORCH_NVIDIA_OPS_MOE_WNA16_MARLIN_GEMM_VLLM_H_ diff --git a/src/linked/torch/nvidia/ops/moe_wna16_marlin_gemm/vllm.yaml b/src/linked/torch/nvidia/ops/moe_wna16_marlin_gemm/vllm.yaml new file mode 100644 index 000000000..19ec54e44 --- /dev/null +++ b/src/linked/torch/nvidia/ops/moe_wna16_marlin_gemm/vllm.yaml @@ -0,0 +1,11 @@ +library: vllm_moe +operator_schema: >- + _moe_C::moe_wna16_marlin_gemm(Tensor! a, Tensor? c_or_none, + Tensor! b_q_weight, Tensor! b_scales, Tensor? global_scale, + Tensor? b_zeros_or_none, Tensor? g_idx_or_none, Tensor? perm_or_none, + Tensor! workspace, Tensor sorted_token_ids, Tensor! expert_ids, + Tensor! num_tokens_past_padded, Tensor! topk_weights, int moe_block_size, + int top_k, bool mul_topk_weights, bool is_ep, int b_q_type_id, int size_m, + int size_n, int size_k, bool is_full_k, bool use_atomic_add, + bool use_fp32_reduce, bool is_zp_float) -> Tensor +dispatch_key: CUDA diff --git a/src/linked/torch/ops/moe_wna16_marlin_gemm.h b/src/linked/torch/ops/moe_wna16_marlin_gemm.h new file mode 100644 index 000000000..d7f951ffc --- /dev/null +++ b/src/linked/torch/ops/moe_wna16_marlin_gemm.h @@ -0,0 +1,73 @@ +#ifndef INFINI_OPS_LINKED_TORCH_OPS_MOE_WNA16_MARLIN_GEMM_H_ +#define INFINI_OPS_LINKED_TORCH_OPS_MOE_WNA16_MARLIN_GEMM_H_ + +#include +#include + +#include "base/moe_wna16_marlin_gemm.h" +#include "torch/tensor_.h" + +namespace infini::ops::linked::torch { + +template +class TorchMoeWna16MarlinGemm : public ::infini::ops::MoeWna16MarlinGemm { + public: + using ::infini::ops::MoeWna16MarlinGemm::MoeWna16MarlinGemm; + + using ::infini::ops::MoeWna16MarlinGemm::operator(); + + void operator()(const Tensor a, const Tensor b_q_weight, + const Tensor b_scales, std::optional global_scale, + std::optional b_zeros_or_none, + std::optional g_idx_or_none, + std::optional perm_or_none, const Tensor workspace, + const Tensor sorted_token_ids, const Tensor expert_ids, + const Tensor num_tokens_past_padded, + const Tensor topk_weights, const int64_t moe_block_size, + const int64_t top_k, const bool mul_topk_weights, + const bool is_ep, const int64_t b_q_type_id, + const int64_t size_m, const int64_t size_n, + const int64_t size_k, const bool is_full_k, + const bool use_atomic_add, const bool use_fp32_reduce, + const bool is_zp_float, Tensor out) const override { + ValidateCallMetadata(a, b_q_weight, b_scales, global_scale, b_zeros_or_none, + g_idx_or_none, perm_or_none, workspace, + sorted_token_ids, expert_ids, num_tokens_past_padded, + topk_weights, moe_block_size, top_k, mul_topk_weights, + is_ep, b_q_type_id, size_m, size_n, size_k, is_full_k, + use_atomic_add, use_fp32_reduce, is_zp_float, out); + + Backend::Validate(b_q_type_id, global_scale.has_value()); + + const typename Backend::StreamGuard stream_guard{ + Backend::GetStreamFromExternal(stream_, device_index_)}; + Backend::Call(ToAten(a), ToAten(out), ToAten(b_q_weight), ToAten(b_scales), + ToOptionalAten(global_scale), ToOptionalAten(b_zeros_or_none), + ToOptionalAten(g_idx_or_none), ToOptionalAten(perm_or_none), + ToAten(workspace), ToAten(sorted_token_ids), + ToAten(expert_ids), ToAten(num_tokens_past_padded), + ToAten(topk_weights), moe_block_size, top_k, mul_topk_weights, + is_ep, b_q_type_id, size_m, size_n, size_k, is_full_k, + use_atomic_add, use_fp32_reduce, is_zp_float); + } + + private: + at::Tensor ToAten(const Tensor tensor) const { + return ToAtenTensor(const_cast(tensor.data()), + tensor.shape(), tensor.strides(), + tensor.dtype(), device_index_); + } + + std::optional ToOptionalAten( + const std::optional& tensor) const { + if (!tensor) { + return std::nullopt; + } + + return ToAten(*tensor); + } +}; + +} // namespace infini::ops::linked::torch + +#endif // INFINI_OPS_LINKED_TORCH_OPS_MOE_WNA16_MARLIN_GEMM_H_ diff --git a/tests/test_moe_wna16_marlin_gemm.py b/tests/test_moe_wna16_marlin_gemm.py new file mode 100644 index 000000000..1d297ddd6 --- /dev/null +++ b/tests/test_moe_wna16_marlin_gemm.py @@ -0,0 +1,237 @@ +import infini.ops +import pytest +import torch + +from tests.utils import get_stream + + +if not hasattr(infini.ops, "MoeWna16MarlinGemm"): + pytest.skip( + "`MoeWna16MarlinGemm` is not available on this platform", + allow_module_level=True, + ) + + +@pytest.mark.parametrize( + "dtype, b_q_type_id, num_bits, rtol, atol", + ( + (torch.float16, 1125899907892224, 4, 2e-2, 2e-2), + (torch.bfloat16, 1125899907892224, 4, 5e-2, 5e-2), + (torch.float16, 1125899923621888, 8, 2e-2, 2e-2), + ), +) +def test_moe_wna16_marlin_gemm( + dtype, + b_q_type_id, + num_bits, + rtol, + atol, + device, + implementation_index, +): + if device != "cuda": + pytest.skip("`moe_wna16_marlin_gemm` requires the NVIDIA backend") + + provider_case = _make_case(device, dtype, b_q_type_id, num_bits) + case = _make_case(device, dtype, b_q_type_id, num_bits) + expected = provider_case["out"] + provider_result = _call_provider(provider_case, expected) + + assert provider_result.data_ptr() == expected.data_ptr() + result = _call_infini(case, implementation_index, get_stream(case["a"].device)) + + assert result is None + torch.testing.assert_close(case["out"], expected, rtol=rtol, atol=atol) + + +def test_moe_wna16_marlin_gemm_non_default_stream(device, implementation_index): + if device != "cuda": + pytest.skip("non-default CUDA streams require the NVIDIA backend") + + provider_case = _make_case(device, torch.float16, 1125899907892224, 4) + case = _make_case(device, torch.float16, 1125899907892224, 4) + expected = provider_case["out"] + _call_provider(provider_case, expected) + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + + _call_infini(case, implementation_index, stream.cuda_stream) + + stream.synchronize() + torch.testing.assert_close(case["out"], expected, rtol=2e-2, atol=2e-2) + + +def test_moe_wna16_marlin_gemm_optional_act_order(device, implementation_index): + if device != "cuda": + pytest.skip("activation-order metadata requires the NVIDIA backend") + + provider_case = _make_case( + device, torch.float16, 1125899907892224, 4, has_act_order=True + ) + case = _make_case(device, torch.float16, 1125899907892224, 4, has_act_order=True) + expected = provider_case["out"] + _call_provider(provider_case, expected) + + _call_infini(case, implementation_index, get_stream(case["a"].device)) + + torch.testing.assert_close(case["out"], expected, rtol=2e-2, atol=2e-2) + + +def test_moe_wna16_marlin_gemm_zero_points(device, implementation_index): + if device != "cuda": + pytest.skip("zero-point metadata requires the NVIDIA backend") + + provider_case = _make_case( + device, torch.float16, 1125899906843648, 4, has_zero_points=True + ) + case = _make_case(device, torch.float16, 1125899906843648, 4, has_zero_points=True) + expected = provider_case["out"] + provider_result = _call_provider(provider_case, expected) + + assert provider_result.data_ptr() == expected.data_ptr() + _call_infini(case, implementation_index, get_stream(case["a"].device)) + + torch.testing.assert_close(case["out"], expected, rtol=2e-2, atol=2e-2) + + +def _make_case( + device, + dtype, + b_q_type_id, + num_bits, + *, + has_act_order=False, + has_zero_points=False, +): + torch.manual_seed(0) + size_m, size_n, size_k = 1, 128, 256 + top_k = 2 + num_experts = 4 + moe_block_size = 16 + route_count = size_m * top_k + padding = route_count + sorted_token_ids = torch.tensor( + (0,) + (padding,) * 15 + (1,) + (padding,) * 45, + dtype=torch.int32, + device=device, + ) + pack_factor = 32 // num_bits + num_groups = 8 if has_act_order else 1 + scale = 0.0 if has_act_order else 0.02 + g_idx_or_none = None + perm_or_none = None + if has_act_order: + g_idx_or_none = ( + torch.arange(num_groups, dtype=torch.int32, device=device) + .repeat_interleave(size_k // num_groups) + .repeat(num_experts, 1) + ) + perm_or_none = torch.arange(size_k, dtype=torch.int32, device=device).repeat( + num_experts, 1 + ) + + b_zeros_or_none = None + if has_zero_points: + b_zeros_or_none = torch.zeros( + (num_experts, num_groups, size_n // pack_factor), + dtype=torch.int32, + device=device, + ) + + return { + "a": torch.randn((size_m, size_k), dtype=dtype, device=device), + "b_q_weight": torch.zeros( + (num_experts, size_k // 16, size_n * 16 // pack_factor), + dtype=torch.int32, + device=device, + ), + "b_scales": torch.full( + (num_experts, num_groups, size_n), scale, dtype=dtype, device=device + ), + "global_scale": None, + "b_zeros_or_none": b_zeros_or_none, + "g_idx_or_none": g_idx_or_none, + "perm_or_none": perm_or_none, + "workspace": torch.zeros(432, dtype=torch.int32, device=device), + "sorted_token_ids": sorted_token_ids, + "expert_ids": torch.tensor((0, 1, 0, 0), dtype=torch.int32, device=device), + "num_tokens_past_padded": torch.tensor((32,), dtype=torch.int32, device=device), + "topk_weights": torch.tensor( + ((0.75, 0.25),), dtype=torch.float32, device=device + ), + "moe_block_size": moe_block_size, + "top_k": top_k, + "mul_topk_weights": True, + "is_ep": False, + "b_q_type_id": b_q_type_id, + "size_m": size_m, + "size_n": size_n, + "size_k": size_k, + "is_full_k": True, + "use_atomic_add": dtype == torch.float16, + "use_fp32_reduce": True, + "is_zp_float": False, + "out": torch.zeros((route_count, size_n), dtype=dtype, device=device), + } + + +def _call_provider(case, out): + return torch.ops._moe_C.moe_wna16_marlin_gemm( + case["a"], + out, + case["b_q_weight"], + case["b_scales"], + case["global_scale"], + case["b_zeros_or_none"], + case["g_idx_or_none"], + case["perm_or_none"], + case["workspace"], + case["sorted_token_ids"], + case["expert_ids"], + case["num_tokens_past_padded"], + case["topk_weights"], + case["moe_block_size"], + case["top_k"], + case["mul_topk_weights"], + case["is_ep"], + case["b_q_type_id"], + case["size_m"], + case["size_n"], + case["size_k"], + case["is_full_k"], + case["use_atomic_add"], + case["use_fp32_reduce"], + case["is_zp_float"], + ) + + +def _call_infini(case, implementation_index, stream): + return infini.ops.moe_wna16_marlin_gemm( + case["a"], + case["b_q_weight"], + case["b_scales"], + case["global_scale"], + case["b_zeros_or_none"], + case["g_idx_or_none"], + case["perm_or_none"], + case["workspace"], + case["sorted_token_ids"], + case["expert_ids"], + case["num_tokens_past_padded"], + case["topk_weights"], + case["moe_block_size"], + case["top_k"], + case["mul_topk_weights"], + case["is_ep"], + case["b_q_type_id"], + case["size_m"], + case["size_n"], + case["size_k"], + case["is_full_k"], + case["use_atomic_add"], + case["use_fp32_reduce"], + case["is_zp_float"], + case["out"], + stream=stream, + implementation_index=implementation_index, + ) From 510d120aef7f5c30683b32e668a011fa8b52365c Mon Sep 17 00:00:00 2001 From: Jiacheng Huang Date: Sat, 8 Aug 2026 21:42:31 +0800 Subject: [PATCH 2/2] fix(linked): support multiple operator schemas --- docs/linked-operators.md | 8 +- scripts/resolve_linked_ops.py | 82 +++++++++++++++------ tests/test_resolve_linked_ops.py | 121 ++++++++++++++++++++++++++++++- 3 files changed, 183 insertions(+), 28 deletions(-) diff --git a/docs/linked-operators.md b/docs/linked-operators.md index c06d71bc8..abf15de6e 100644 --- a/docs/linked-operators.md +++ b/docs/linked-operators.md @@ -38,7 +38,7 @@ required_symbols: - silu_and_mul(at::Tensor&, at::Tensor&) ``` -A Dispatcher implementation instead declares its exact schema and required +A Dispatcher implementation instead declares its exact schema and a required dispatch key: ```yaml @@ -49,6 +49,12 @@ operator_schema: >- dispatch_key: CUDA ``` +When one implementation requires multiple Dispatcher operators from the same +library and dispatch key, `operator_schema` may instead be a non-empty YAML list +of unique, non-empty schema strings. The resolver validates every schema while +force-loading the library once. A single schema is emitted as a scalar string, +including when the input uses a one-item list. + Each implementation uses exactly one contract form. The resolver rejects a partial Dispatcher contract or a binding that mixes both forms. diff --git a/scripts/resolve_linked_ops.py b/scripts/resolve_linked_ops.py index 03f63be9e..36de60ed8 100644 --- a/scripts/resolve_linked_ops.py +++ b/scripts/resolve_linked_ops.py @@ -78,7 +78,7 @@ class BindingConfig: source: pathlib.Path library: str required_symbols: tuple[str, ...] - operator_schema: str | None + operator_schema: str | tuple[str, ...] | None dispatch_key: str | None @@ -114,6 +114,34 @@ def _require_string(data, key, path): return value.strip() +def _require_operator_schema(data, path): + value = data["operator_schema"] + if isinstance(value, str): + if not value.strip(): + raise ResolutionError( + f"{path}: operator_schema must be a non-empty string or list of " + "non-empty strings" + ) + + return value.strip() + + if ( + not isinstance(value, list) + or not value + or any(not isinstance(schema, str) or not schema.strip() for schema in value) + ): + raise ResolutionError( + f"{path}: operator_schema must be a non-empty string or list of " + "non-empty strings" + ) + + schemas = tuple(schema.strip() for schema in value) + if len(schemas) != len(set(schemas)): + raise ResolutionError(f"{path}: operator_schema contains duplicates") + + return schemas[0] if len(schemas) == 1 else schemas + + def _require_relative_glob(data, key, path): value = _require_string(data, key, path) normalized = pathlib.PurePosixPath(value.replace("\\", "/")) @@ -194,7 +222,7 @@ def _load_bindings(platform_dir, device, transport, selected_ops): operator_schema = None else: symbols = () - operator_schema = _require_string(data, "operator_schema", path) + operator_schema = _require_operator_schema(data, path) if dispatch_key is None: raise ResolutionError(f"{path}: operator_schema requires dispatch_key") dispatch_key = _require_string(data, "dispatch_key", path) @@ -441,26 +469,30 @@ def _verify_dispatcher_contracts(contracts): " torch.ops.load_library(library_path)\n" " loaded.add(library_path)\n" "for contract in contracts:\n" - " expected = torch._C.parse_schema(contract['schema'])\n" - " actual = torch._C._dispatch_find_schema_or_throw(\n" - " expected.name, expected.overload_name\n" - " ).schema()\n" - " if str(actual) != str(expected):\n" - " sys.exit(\n" - " f\"{contract['binding_path']}: expected {expected}, \"\n" - " f'found {actual}'\n" - " )\n" - " qualified_name = expected.name\n" - " if expected.overload_name:\n" - " qualified_name += f'.{expected.overload_name}'\n" - " dispatch_key = contract['dispatch_key']\n" - " if not torch._C._dispatch_has_kernel_for_dispatch_key(\n" - " qualified_name, dispatch_key\n" - " ):\n" - " sys.exit(\n" - " f\"{contract['binding_path']}: {qualified_name} has no \"\n" - " f'{dispatch_key} kernel'\n" - " )\n" + " schemas = contract['schema']\n" + " if isinstance(schemas, str):\n" + " schemas = [schemas]\n" + " for schema in schemas:\n" + " expected = torch._C.parse_schema(schema)\n" + " actual = torch._C._dispatch_find_schema_or_throw(\n" + " expected.name, expected.overload_name\n" + " ).schema()\n" + " if str(actual) != str(expected):\n" + " sys.exit(\n" + " f\"{contract['binding_path']}: expected {expected}, \"\n" + " f'found {actual}'\n" + " )\n" + " qualified_name = expected.name\n" + " if expected.overload_name:\n" + " qualified_name += f'.{expected.overload_name}'\n" + " dispatch_key = contract['dispatch_key']\n" + " if not torch._C._dispatch_has_kernel_for_dispatch_key(\n" + " qualified_name, dispatch_key\n" + " ):\n" + " sys.exit(\n" + " f\"{contract['binding_path']}: {qualified_name} has no \"\n" + " f'{dispatch_key} kernel'\n" + " )\n" ) try: subprocess.run( @@ -653,7 +685,11 @@ def resolve_linked_ops( if binding.required_symbols: operator["required_symbols"] = list(binding.required_symbols) else: - operator["operator_schema"] = binding.operator_schema + operator["operator_schema"] = ( + list(binding.operator_schema) + if isinstance(binding.operator_schema, tuple) + else binding.operator_schema + ) operator["dispatch_key"] = binding.dispatch_key operators.append(operator) diff --git a/tests/test_resolve_linked_ops.py b/tests/test_resolve_linked_ops.py index 72a7b45d0..c824d2e89 100644 --- a/tests/test_resolve_linked_ops.py +++ b/tests/test_resolve_linked_ops.py @@ -96,8 +96,9 @@ def test_resolve_collects_selected_implementation_and_library(monkeypatch, tmp_p assert "ignored" not in manifest +@pytest.mark.parametrize("schema_form", ("scalar", "list")) def test_resolve_validates_dispatcher_contract_and_force_loads_library( - monkeypatch, tmp_path + monkeypatch, tmp_path, schema_form ): module = _load_resolver_module() source_root = tmp_path / "linked" @@ -111,8 +112,13 @@ def test_resolve_validates_dispatcher_contract_and_force_loads_library( "_C::gptq_marlin_repack(Tensor b_q_weight, Tensor perm, " "SymInt size_k, SymInt size_n, int num_bits, bool is_a_8bit) -> Tensor" ) + if schema_form == "scalar": + schema_yaml = f"operator_schema: {schema}\n" + else: + schema_yaml = f"operator_schema:\n - {schema}\n" + (op_dir / "vllm.yaml").write_text( - f"library: vllm\noperator_schema: {schema}\ndispatch_key: CUDA\n" + f"library: vllm\n{schema_yaml}dispatch_key: CUDA\n" ) (op_dir / "vllm.h").write_text("// declaration\n") (op_dir / "vllm.cc").write_text("// definition\n") @@ -167,6 +173,67 @@ def test_resolve_validates_dispatcher_contract_and_force_loads_library( assert str(library_path).replace("\\", "/") in force_load_block +def test_resolve_validates_multiple_dispatcher_schemas(monkeypatch, tmp_path): + module = _load_resolver_module() + source_root = tmp_path / "linked" + platform = source_root / "torch" / "nvidia" + op_dir = platform / "ops" / "fused_marlin_moe" + op_dir.mkdir(parents=True) + (platform / "vllm.yaml").write_text( + "python_distribution_package: vllm\nlibrary_glob: vllm/_C*.so\n" + ) + schemas = ( + "_moe_C::moe_align_block_size(Tensor topk_ids) -> ()", + "_moe_C::moe_wna16_marlin_gemm(Tensor input) -> Tensor", + ) + (op_dir / "vllm.yaml").write_text( + "library: vllm\n" + "operator_schema:\n" + f" - {schemas[0]}\n" + f" - {schemas[1]}\n" + "dispatch_key: CUDA\n" + ) + (op_dir / "vllm.h").write_text("// declaration\n") + (op_dir / "vllm.cc").write_text("// definition\n") + library_path = tmp_path / "site-packages" / "vllm" / "_moe_C.abi3.so" + library_path.parent.mkdir(parents=True) + library_path.touch() + + monkeypatch.setattr( + module, "_locate_distribution_library", lambda config: library_path + ) + contracts = [] + monkeypatch.setattr( + module, + "_verify_dispatcher_contracts", + lambda resolved: contracts.extend(resolved), + ) + + output_dir = tmp_path / "generated" + payload = module.resolve_linked_ops( + ["nvidia"], + ["fused_marlin_moe"], + source_root=source_root, + output_dir=output_dir, + ) + + assert len(contracts) == 1 + contract, resolved_path = contracts[0] + assert contract.operator_schema == schemas + assert contract.dispatch_key == "CUDA" + assert resolved_path == library_path + assert payload["operators"][0]["operator_schema"] == list(schemas) + assert json.loads((output_dir / "resolved.json").read_text()) == payload + assert len(payload["libraries"]) == 1 + assert payload["libraries"][0]["force_load"] + + manifest = (output_dir / "manifest.cmake").read_text() + force_load_block = manifest.split( + "set(INFINI_OPS_LINKED_FORCE_LOAD_LIBRARIES", maxsplit=1 + )[1].split(")", maxsplit=1)[0] + assert force_load_block.count(str(library_path).replace("\\", "/")) == 1 + + def test_resolve_supports_multiple_implementations_for_one_operator( monkeypatch, tmp_path ): @@ -261,6 +328,51 @@ def test_resolve_requires_one_complete_binding_contract(tmp_path, binding, messa ) +@pytest.mark.parametrize( + ("operator_schema", "message"), + ( + ( + "operator_schema: []\n", + "operator_schema must be a non-empty string or list of non-empty strings", + ), + ( + "operator_schema: 1\n", + "operator_schema must be a non-empty string or list of non-empty strings", + ), + ( + "operator_schema: ''\n", + "operator_schema must be a non-empty string or list of non-empty strings", + ), + ( + "operator_schema:\n - _C::op() -> Tensor\n - 1\n", + "operator_schema must be a non-empty string or list of non-empty strings", + ), + ( + "operator_schema:\n - ''\n", + "operator_schema must be a non-empty string or list of non-empty strings", + ), + ( + "operator_schema:\n - _C::op() -> Tensor\n - ' _C::op() -> Tensor '\n", + "operator_schema contains duplicates", + ), + ), +) +def test_resolve_rejects_invalid_operator_schema(tmp_path, operator_schema, message): + module = _load_resolver_module() + source_root = tmp_path / "linked" + _, op_dir = _write_linked_config(source_root) + (op_dir / "vllm.yaml").write_text( + f"library: vllm\n{operator_schema}dispatch_key: CUDA\n" + ) + + with pytest.raises(module.ResolutionError, match=message): + module.resolve_linked_ops( + ["metax"], + source_root=source_root, + output_dir=tmp_path / "generated", + ) + + @pytest.mark.parametrize( ("library_extra", "binding_extra", "unknown_key"), ( @@ -502,6 +614,7 @@ def test_dispatcher_contract_validation_uses_one_isolated_process( ): module = _load_resolver_module() schema = "_C::op(Tensor input) -> Tensor" + schemas = (schema, "_C::other(Tensor input) -> Tensor") config = module.BindingConfig( device="nvidia", name="op", @@ -511,7 +624,7 @@ def test_dispatcher_contract_validation_uses_one_isolated_process( source=tmp_path / "vllm.cc", library="vllm", required_symbols=(), - operator_schema=schema, + operator_schema=schemas, dispatch_key="CUDA", ) other_schema = "other::op(Tensor input) -> Tensor" @@ -547,7 +660,7 @@ def fake_run(command, **kwargs): { "binding_path": str(config.path), "library_path": str(library_path), - "schema": schema, + "schema": list(schemas), "dispatch_key": "CUDA", }, {