InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
moe_wna16_gemm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_MOE_WNA16_GEMM_H_
2#define INFINI_OPS_BASE_MOE_WNA16_GEMM_H_
3
4#include <cassert>
5#include <cstdint>
6#include <limits>
7#include <optional>
8
9#include "operator.h"
10
11namespace infini::ops {
12
13// Aligned with vLLM `_moe_C::moe_wna16_gemm` at commit
14// ffc4f08c8ee130d4ea6347c1bf31ffd4f8af28ab.
15class MoeWna16Gemm : public Operator<MoeWna16Gemm> {
16 public:
17 MoeWna16Gemm(const Tensor input, const Tensor b_qweight,
18 const Tensor b_scales, std::optional<Tensor> b_qzeros,
19 std::optional<Tensor> topk_weights,
20 const Tensor sorted_token_ids, const Tensor expert_ids,
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},
30 : 0},
31 sorted_token_ids_size_{sorted_token_ids.numel()},
32 expert_ids_size_{expert_ids.numel()},
33 top_k_{top_k},
34 block_size_m_{block_size_m},
35 block_size_n_{block_size_n},
36 block_size_k_{block_size_k},
37 bit_{bit},
38 dtype_{input.dtype()},
39 has_qzeros_{b_qzeros.has_value()},
40 has_topk_weights_{topk_weights.has_value()},
41 device_type_{input.device().type()},
42 device_index_{input.device().index()} {
43 Validate(input, b_qweight, b_scales, b_qzeros, topk_weights,
44 sorted_token_ids, expert_ids, num_tokens_post_pad, output);
45 }
46
47 virtual void operator()(
48 const Tensor input, const Tensor b_qweight, const Tensor b_scales,
49 std::optional<Tensor> b_qzeros, std::optional<Tensor> topk_weights,
50 const Tensor sorted_token_ids, const Tensor expert_ids,
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;
54
55 protected:
57 const Tensor input, const Tensor b_qweight, const Tensor b_scales,
58 std::optional<Tensor> b_qzeros, std::optional<Tensor> topk_weights,
59 const Tensor sorted_token_ids, const Tensor expert_ids,
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 {
64 assert(top_k == top_k_ && block_size_m == block_size_m_ &&
65 block_size_n == block_size_n_ && block_size_k == block_size_k_ &&
66 bit == bit_ &&
67 "`MoeWna16Gemm` attributes changed after descriptor creation");
68
69 const auto same_device = [&](const Tensor tensor) {
70 return tensor.device().type() == device_type_ &&
71 tensor.device().index() == device_index_;
72 };
73 const auto values_per_byte = static_cast<Tensor::Size>(8 / bit_);
74 const auto matches =
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 &&
81 b_scales.size(0) == num_experts_ && b_scales.size(1) == n_ &&
82 b_scales.size(2) == num_groups_ && b_scales.dtype() == dtype_ &&
83 b_scales.IsContiguous() && same_device(b_scales) &&
84 sorted_token_ids.ndim() == 1 &&
85 sorted_token_ids.numel() == sorted_token_ids_size_ &&
86 sorted_token_ids.dtype() == DataType::kInt32 &&
87 sorted_token_ids.IsContiguous() && same_device(sorted_token_ids) &&
88 expert_ids.ndim() == 1 && expert_ids.numel() == expert_ids_size_ &&
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) &&
98 b_qzeros.has_value() == has_qzeros_ &&
99 topk_weights.has_value() == has_topk_weights_;
100 assert(matches && "`MoeWna16Gemm` call metadata must match descriptor");
101
102 if (b_qzeros) {
103 assert(b_qzeros->ndim() == 3 && b_qzeros->size(0) == num_experts_ &&
104 b_qzeros->size(1) == n_ / values_per_byte &&
105 b_qzeros->size(2) == num_groups_ &&
106 b_qzeros->dtype() == DataType::kUInt8 &&
107 b_qzeros->IsContiguous() && same_device(*b_qzeros) &&
108 "`MoeWna16Gemm` zero-point metadata must match descriptor");
109 }
110
111 if (topk_weights) {
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");
117 }
118 }
119
120 Tensor::Size num_experts_{0};
121
122 Tensor::Size m_{0};
123
124 Tensor::Size n_{0};
125
126 Tensor::Size k_{0};
127
128 Tensor::Size num_groups_{0};
129
130 Tensor::Size group_size_{0};
131
132 Tensor::Size sorted_token_ids_size_{0};
133
134 Tensor::Size expert_ids_size_{0};
135
136 int64_t top_k_{0};
137
138 int64_t block_size_m_{0};
139
140 int64_t block_size_n_{0};
141
142 int64_t block_size_k_{0};
143
144 int64_t bit_{0};
145
146 DataType dtype_;
147
148 bool has_qzeros_{false};
149
150 bool has_topk_weights_{false};
151
152 Device::Type device_type_;
153
155
156 private:
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,
160 const Tensor sorted_token_ids, const Tensor expert_ids,
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");
164 assert((bit_ == 4 || bit_ == 8) &&
165 "`MoeWna16Gemm` requires 4-bit or 8-bit weights");
166 assert(m_ > 0 && n_ > 0 && k_ > 0 && num_experts_ > 0 && top_k_ > 0 &&
167 "`MoeWna16Gemm` requires positive dimensions");
168 assert(num_groups_ > 0 && group_size_ > 0 &&
169 "`MoeWna16Gemm` requires an integral quantization group size");
170
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 &&
174 group_size_ % 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 &&
177 block_size_n_ <= 1024 && block_size_k_ > 0 && block_size_k_ <= k_ &&
178 block_size_k_ % group_size_ == 0 && k_ % block_size_k_ == 0 &&
179 block_size_n_ % values_per_word == 0 &&
180 (block_size_k_ / values_per_word) % 4 == 0 &&
181 "`MoeWna16Gemm` received unsupported block sizes");
182 const auto groups_per_block = block_size_k_ / group_size_;
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");
186
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());
191 assert(num_experts_ <= kMaxU16 && group_size_ <= kMaxU16 &&
192 top_k_ <= static_cast<int64_t>(kMaxU16) &&
193 block_size_m_ <= static_cast<int64_t>(kMaxU16) &&
194 block_size_n_ <= static_cast<int64_t>(kMaxU16) &&
195 block_size_k_ <= static_cast<int64_t>(kMaxU16) && m_ <= kMaxU32 &&
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");
201
202 auto effective_sorted_size = sorted_token_ids_size_;
203 if (m_ <= block_size_m_) {
204 const auto limit = m_ * block_size_m_ * top_k_;
205 if (effective_sorted_size > limit) {
206 effective_sorted_size = limit;
207 }
208 }
209 const auto num_token_blocks =
210 (effective_sorted_size + block_size_m_ - 1) / block_size_m_;
211 assert(sorted_token_ids_size_ > 0 && expert_ids_size_ >= 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");
215 assert((n_ + block_size_n_ - 1) / block_size_n_ <= 65535 &&
216 (k_ + block_size_k_ - 1) / block_size_k_ <= 65535 &&
217 "`MoeWna16Gemm` grid dimensions exceed CUDA limits");
218
219 ValidateCallMetadata(input, b_qweight, b_scales, b_qzeros, topk_weights,
220 sorted_token_ids, expert_ids, num_tokens_post_pad,
222 bit_, output);
223 }
224};
225
226} // namespace infini::ops
227
228#endif // INFINI_OPS_BASE_MOE_WNA16_GEMM_H_
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