1#ifndef INFINI_OPS_BASE_MOE_WNA16_GEMM_H_
2#define INFINI_OPS_BASE_MOE_WNA16_GEMM_H_
18 const Tensor b_scales, std::optional<Tensor> b_qzeros,
19 std::optional<Tensor> topk_weights,
21 const Tensor num_tokens_post_pad,
const int64_t top_k,
22 const int64_t block_size_m,
const int64_t block_size_n,
23 const int64_t block_size_k,
const int64_t bit,
Tensor output)
24 :
num_experts_{b_qweight.ndim() == 3 ? b_qweight.size(0) : 0},
25 m_{input.ndim() == 2 ? input.size(0) : 0},
26 n_{b_qweight.ndim() == 3 ? b_qweight.size(1) : 0},
27 k_{input.ndim() == 2 ? input.size(1) : 0},
28 num_groups_{b_scales.ndim() == 3 ? b_scales.size(2) : 0},
43 Validate(input, b_qweight, b_scales, b_qzeros, topk_weights,
44 sorted_token_ids, expert_ids, num_tokens_post_pad, output);
49 std::optional<Tensor> b_qzeros, std::optional<Tensor> topk_weights,
51 const Tensor num_tokens_post_pad,
const int64_t top_k,
52 const int64_t block_size_m,
const int64_t block_size_n,
53 const int64_t block_size_k,
const int64_t bit,
Tensor output)
const = 0;
58 std::optional<Tensor> b_qzeros, std::optional<Tensor> topk_weights,
60 const Tensor num_tokens_post_pad,
const int64_t top_k,
61 const int64_t block_size_m,
const int64_t block_size_n,
62 const int64_t block_size_k,
const int64_t bit,
63 const Tensor output)
const {
67 "`MoeWna16Gemm` attributes changed after descriptor creation");
69 const auto same_device = [&](
const Tensor tensor) {
73 const auto values_per_byte =
static_cast<Tensor::Size
>(8 /
bit_);
75 input.ndim() == 2 && input.size(0) ==
m_ && input.size(1) ==
k_ &&
76 input.dtype() ==
dtype_ && input.IsContiguous() && same_device(input) &&
77 b_qweight.ndim() == 3 && b_qweight.size(0) ==
num_experts_ &&
78 b_qweight.size(1) ==
n_ && b_qweight.size(2) ==
k_ / values_per_byte &&
79 b_qweight.dtype() == DataType::kUInt8 && b_qweight.IsContiguous() &&
80 same_device(b_qweight) && b_scales.ndim() == 3 &&
83 b_scales.IsContiguous() && same_device(b_scales) &&
84 sorted_token_ids.ndim() == 1 &&
86 sorted_token_ids.dtype() == DataType::kInt32 &&
87 sorted_token_ids.IsContiguous() && same_device(sorted_token_ids) &&
89 expert_ids.dtype() == DataType::kInt32 && expert_ids.IsContiguous() &&
90 same_device(expert_ids) && num_tokens_post_pad.ndim() == 1 &&
91 num_tokens_post_pad.numel() == 1 &&
92 num_tokens_post_pad.dtype() == DataType::kInt32 &&
93 num_tokens_post_pad.IsContiguous() &&
94 same_device(num_tokens_post_pad) && output.ndim() == 3 &&
95 output.size(0) ==
m_ && output.size(1) ==
top_k_ &&
96 output.size(2) ==
n_ && output.dtype() ==
dtype_ &&
97 output.IsContiguous() && same_device(output) &&
100 assert(matches &&
"`MoeWna16Gemm` call metadata must match descriptor");
103 assert(b_qzeros->ndim() == 3 && b_qzeros->size(0) ==
num_experts_ &&
104 b_qzeros->size(1) ==
n_ / values_per_byte &&
106 b_qzeros->dtype() == DataType::kUInt8 &&
107 b_qzeros->IsContiguous() && same_device(*b_qzeros) &&
108 "`MoeWna16Gemm` zero-point metadata must match descriptor");
112 assert(topk_weights->ndim() == 2 && topk_weights->size(0) ==
m_ &&
113 topk_weights->size(1) ==
top_k_ &&
114 topk_weights->dtype() == DataType::kFloat32 &&
115 topk_weights->IsContiguous() && same_device(*topk_weights) &&
116 "`MoeWna16Gemm` top-k weight metadata must match descriptor");
157 void Validate(
const Tensor input,
const Tensor b_qweight,
158 const Tensor b_scales, std::optional<Tensor> b_qzeros,
159 std::optional<Tensor> topk_weights,
161 const Tensor num_tokens_post_pad,
const Tensor output)
const {
162 assert((
dtype_ == DataType::kFloat16 ||
dtype_ == DataType::kBFloat16) &&
163 "`MoeWna16Gemm` supports float16 and bfloat16 inputs");
165 "`MoeWna16Gemm` requires 4-bit or 8-bit weights");
167 "`MoeWna16Gemm` requires positive dimensions");
169 "`MoeWna16Gemm` requires an integral quantization group size");
171 const auto values_per_byte =
static_cast<Tensor::Size
>(8 /
bit_);
172 const auto values_per_word =
static_cast<Tensor::Size
>(32 /
bit_);
173 assert(
k_ % values_per_byte == 0 &&
n_ % values_per_word == 0 &&
175 "`MoeWna16Gemm` packed dimensions are incompatible with `bit`");
176 assert(
block_size_m_ > 0 && block_size_m_ <= 64 && block_size_n_ > 0 &&
181 "`MoeWna16Gemm` received unsupported block sizes");
183 assert((groups_per_block == 1 || groups_per_block == 2 ||
184 groups_per_block == 4 || groups_per_block == 8) &&
185 "`MoeWna16Gemm` requires 1, 2, 4, or 8 groups per K block");
187 constexpr auto kMaxU16 =
188 static_cast<Tensor::Size
>(std::numeric_limits<uint16_t>::max());
189 constexpr auto kMaxU32 =
190 static_cast<Tensor::Size
>(std::numeric_limits<uint32_t>::max());
192 top_k_ <=
static_cast<int64_t
>(kMaxU16) &&
196 n_ <= kMaxU32 &&
k_ <= kMaxU32 &&
197 "`MoeWna16Gemm` dimensions exceed CUDA kernel limits");
198 assert(
m_ <= std::numeric_limits<Tensor::Size>::max() /
top_k_ &&
199 m_ *
top_k_ <= std::numeric_limits<Tensor::Size>::max() /
n_ &&
200 "`MoeWna16Gemm` output dimensions overflow");
205 if (effective_sorted_size > limit) {
206 effective_sorted_size = limit;
209 const auto num_token_blocks =
212 "`MoeWna16Gemm` routing metadata is too small");
213 assert(num_token_blocks <= kMaxU32 &&
214 "`MoeWna16Gemm` token block count exceeds CUDA grid limits");
217 "`MoeWna16Gemm` grid dimensions exceed CUDA limits");
220 sorted_token_ids, expert_ids, num_tokens_post_pad,
Definition moe_wna16_gemm.h:15
int64_t top_k_
Definition moe_wna16_gemm.h:136
virtual void operator()(const Tensor input, const Tensor b_qweight, const Tensor b_scales, std::optional< Tensor > b_qzeros, std::optional< Tensor > topk_weights, const Tensor sorted_token_ids, const Tensor expert_ids, const Tensor num_tokens_post_pad, const int64_t top_k, const int64_t block_size_m, const int64_t block_size_n, const int64_t block_size_k, const int64_t bit, Tensor output) const =0
Tensor::Size group_size_
Definition moe_wna16_gemm.h:130
Tensor::Size num_experts_
Definition moe_wna16_gemm.h:120
Tensor::Size num_groups_
Definition moe_wna16_gemm.h:128
MoeWna16Gemm(const Tensor input, const Tensor b_qweight, const Tensor b_scales, std::optional< Tensor > b_qzeros, std::optional< Tensor > topk_weights, const Tensor sorted_token_ids, const Tensor expert_ids, const Tensor num_tokens_post_pad, const int64_t top_k, const int64_t block_size_m, const int64_t block_size_n, const int64_t block_size_k, const int64_t bit, Tensor output)
Definition moe_wna16_gemm.h:17
int64_t block_size_n_
Definition moe_wna16_gemm.h:140
Tensor::Size k_
Definition moe_wna16_gemm.h:126
int64_t bit_
Definition moe_wna16_gemm.h:144
int device_index_
Definition moe_wna16_gemm.h:154
Tensor::Size expert_ids_size_
Definition moe_wna16_gemm.h:134
Tensor::Size sorted_token_ids_size_
Definition moe_wna16_gemm.h:132
int64_t block_size_k_
Definition moe_wna16_gemm.h:142
Device::Type device_type_
Definition moe_wna16_gemm.h:152
bool has_qzeros_
Definition moe_wna16_gemm.h:148
Tensor::Size n_
Definition moe_wna16_gemm.h:124
void ValidateCallMetadata(const Tensor input, const Tensor b_qweight, const Tensor b_scales, std::optional< Tensor > b_qzeros, std::optional< Tensor > topk_weights, const Tensor sorted_token_ids, const Tensor expert_ids, const Tensor num_tokens_post_pad, const int64_t top_k, const int64_t block_size_m, const int64_t block_size_n, const int64_t block_size_k, const int64_t bit, const Tensor output) const
Definition moe_wna16_gemm.h:56
int64_t block_size_m_
Definition moe_wna16_gemm.h:138
DataType dtype_
Definition moe_wna16_gemm.h:146
Tensor::Size m_
Definition moe_wna16_gemm.h:122
bool has_topk_weights_
Definition moe_wna16_gemm.h:150
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8