19 std::optional<Tensor> b_bias_or_none,
20 const Tensor b_scales, std::optional<Tensor> a_scales,
21 std::optional<Tensor> global_scale,
22 std::optional<Tensor> b_zeros_or_none,
23 std::optional<Tensor> g_idx_or_none,
24 std::optional<Tensor> perm_or_none,
const Tensor workspace,
26 const Tensor num_tokens_past_padded,
27 const Tensor topk_weights,
const int64_t moe_block_size,
28 const int64_t top_k,
const bool mul_topk_weights,
29 const int64_t b_type_id,
const int64_t size_m,
30 const int64_t size_n,
const int64_t size_k,
31 const bool is_full_k,
const bool use_atomic_add,
32 const bool use_fp32_reduce,
const bool is_zp_float,
33 const int64_t thread_k,
const int64_t thread_n,
34 const int64_t blocks_per_sm,
Tensor out)
36 b_q_weight_metadata_{b_q_weight},
37 b_bias_or_none_metadata_{b_bias_or_none},
38 b_scales_metadata_{b_scales},
39 a_scales_metadata_{a_scales},
40 global_scale_metadata_{global_scale},
41 b_zeros_or_none_metadata_{b_zeros_or_none},
42 g_idx_or_none_metadata_{g_idx_or_none},
43 perm_or_none_metadata_{perm_or_none},
44 workspace_metadata_{workspace},
45 sorted_token_ids_metadata_{sorted_token_ids},
46 expert_ids_metadata_{expert_ids},
47 num_tokens_past_padded_metadata_{num_tokens_past_padded},
48 topk_weights_metadata_{topk_weights},
50 moe_block_size_{moe_block_size},
52 mul_topk_weights_{mul_topk_weights},
53 b_type_id_{b_type_id},
57 is_full_k_{is_full_k},
58 use_atomic_add_{use_atomic_add},
59 use_fp32_reduce_{use_fp32_reduce},
60 is_zp_float_{is_zp_float},
63 blocks_per_sm_{blocks_per_sm},
65 Validate(a, b_q_weight, b_bias_or_none, b_scales, a_scales, global_scale,
66 b_zeros_or_none, g_idx_or_none, perm_or_none, workspace,
67 sorted_token_ids, expert_ids, num_tokens_past_padded, topk_weights,
73 std::optional<Tensor> b_bias_or_none,
const Tensor b_scales,
74 std::optional<Tensor> a_scales, std::optional<Tensor> global_scale,
75 std::optional<Tensor> b_zeros_or_none,
76 std::optional<Tensor> g_idx_or_none, std::optional<Tensor> perm_or_none,
78 const Tensor expert_ids,
const Tensor num_tokens_past_padded,
79 const Tensor topk_weights,
const int64_t moe_block_size,
80 const int64_t top_k,
const bool mul_topk_weights,
const int64_t b_type_id,
81 const int64_t size_m,
const int64_t size_n,
const int64_t size_k,
82 const bool is_full_k,
const bool use_atomic_add,
83 const bool use_fp32_reduce,
const bool is_zp_float,
84 const int64_t thread_k,
const int64_t thread_n,
85 const int64_t blocks_per_sm,
Tensor out)
const = 0;
90 std::optional<Tensor> b_bias_or_none,
const Tensor b_scales,
91 std::optional<Tensor> a_scales, std::optional<Tensor> global_scale,
92 std::optional<Tensor> b_zeros_or_none,
93 std::optional<Tensor> g_idx_or_none, std::optional<Tensor> perm_or_none,
95 const Tensor expert_ids,
const Tensor num_tokens_past_padded,
96 const Tensor topk_weights,
const int64_t moe_block_size,
97 const int64_t top_k,
const bool mul_topk_weights,
const int64_t b_type_id,
98 const int64_t size_m,
const int64_t size_n,
const int64_t size_k,
99 const bool is_full_k,
const bool use_atomic_add,
100 const bool use_fp32_reduce,
const bool is_zp_float,
101 const int64_t thread_k,
const int64_t thread_n,
102 const int64_t blocks_per_sm,
Tensor out)
const {
103 assert(moe_block_size == moe_block_size_ && top_k == top_k_ &&
104 mul_topk_weights == mul_topk_weights_ && b_type_id == b_type_id_ &&
105 size_m == size_m_ && size_n == size_n_ && size_k == size_k_ &&
106 is_full_k == is_full_k_ && use_atomic_add == use_atomic_add_ &&
107 use_fp32_reduce == use_fp32_reduce_ && is_zp_float == is_zp_float_ &&
108 thread_k == thread_k_ && thread_n == thread_n_ &&
109 blocks_per_sm == blocks_per_sm_ &&
110 "`MoeWna16MarlinGemm` attributes changed after descriptor "
113 const std::equal_to<Tensor> same_metadata;
114 const auto optional_matches = [&](
const std::optional<Tensor>& expected,
115 const std::optional<Tensor>& actual) {
116 return expected.has_value() == actual.has_value() &&
117 (!expected || same_metadata(*expected, *actual));
120 same_metadata(a_metadata_, a) &&
121 same_metadata(b_q_weight_metadata_, b_q_weight) &&
122 optional_matches(b_bias_or_none_metadata_, b_bias_or_none) &&
123 same_metadata(b_scales_metadata_, b_scales) &&
124 optional_matches(a_scales_metadata_, a_scales) &&
125 optional_matches(global_scale_metadata_, global_scale) &&
126 optional_matches(b_zeros_or_none_metadata_, b_zeros_or_none) &&
127 optional_matches(g_idx_or_none_metadata_, g_idx_or_none) &&
128 optional_matches(perm_or_none_metadata_, perm_or_none) &&
129 same_metadata(workspace_metadata_, workspace) &&
130 same_metadata(sorted_token_ids_metadata_, sorted_token_ids) &&
131 same_metadata(expert_ids_metadata_, expert_ids) &&
132 same_metadata(num_tokens_past_padded_metadata_,
133 num_tokens_past_padded) &&
134 same_metadata(topk_weights_metadata_, topk_weights) &&
135 same_metadata(out_metadata_, out);
137 "`MoeWna16MarlinGemm` tensor metadata must match descriptor");
142 std::optional<Tensor> b_bias_or_none,
const Tensor b_scales,
143 std::optional<Tensor> a_scales,
144 std::optional<Tensor> global_scale,
145 std::optional<Tensor> b_zeros_or_none,
146 std::optional<Tensor> g_idx_or_none,
147 std::optional<Tensor> perm_or_none,
const Tensor workspace,
149 const Tensor num_tokens_past_padded,
const Tensor topk_weights,
151 assert(a.ndim() == 2 && a.size(0) == size_m_ && a.size(1) == size_k_ &&
152 "`MoeWna16MarlinGemm` `a` shape must match `size_m` and "
154 const auto is_a_8bit = a.dtype() == DataType::kInt8;
155 const auto output_dtype = is_a_8bit ? b_scales.dtype() : a.dtype();
157 (a.dtype() == DataType::kFloat16 || a.dtype() == DataType::kBFloat16 ||
160 "`MoeWna16MarlinGemm` requires contiguous float16, bfloat16, or int8 "
162 assert(size_m_ > 0 && size_n_ > 0 && size_k_ > 0 && top_k_ > 0 &&
163 size_k_ % 16 == 0 && size_n_ % 64 == 0 &&
164 "`MoeWna16MarlinGemm` received unsupported dimensions");
165 assert((moe_block_size_ == 8 ||
166 (moe_block_size_ >= 16 && moe_block_size_ <= 64 &&
167 moe_block_size_ % 16 == 0)) &&
168 "`MoeWna16MarlinGemm` received an unsupported `moe_block_size`");
169 assert(size_m_ <= std::numeric_limits<Tensor::Size>::max() / top_k_ &&
170 "`MoeWna16MarlinGemm` output dimensions overflow");
172 constexpr int64_t kUint4B8 = 1125899907892224;
173 constexpr int64_t kUint8B128 = 1125899923621888;
174 constexpr int64_t kUint4 = 1125899906843648;
175 constexpr int64_t kUint8 = 1125899906844672;
176 constexpr int64_t kInt4 = 1125899906908928;
177 constexpr int64_t kInt8 = 1125899906909952;
178 constexpr int64_t kFloat8E4M3Fn = 2814749767172868;
179 constexpr int64_t kFloat4E2M1F = 562949953487106;
180 const auto supported_qtype =
181 b_type_id_ == kUint4B8 || b_type_id_ == kUint8B128 ||
182 b_type_id_ == kUint4 || b_type_id_ == kUint8 || b_type_id_ == kInt4 ||
183 b_type_id_ == kInt8 || b_type_id_ == kFloat8E4M3Fn ||
184 b_type_id_ == kFloat4E2M1F;
185 const auto has_zero_points = b_zeros_or_none.has_value();
186 assert(supported_qtype &&
187 has_zero_points == (b_type_id_ == kUint4 || b_type_id_ == kUint8) &&
189 (has_zero_points && a.dtype() == DataType::kFloat16)) &&
190 "`MoeWna16MarlinGemm` received an unsupported quantization "
193 const auto pack_factor =
194 (b_type_id_ == kUint8B128 || b_type_id_ == kUint8 ||
195 b_type_id_ == kInt8 || b_type_id_ == kFloat8E4M3Fn)
198 assert(size_n_ <= std::numeric_limits<Tensor::Size>::max() / 16 &&
199 b_q_weight.ndim() == 3 && b_q_weight.size(1) == size_k_ / 16 &&
200 b_q_weight.size(2) == size_n_ * 16 / pack_factor &&
201 b_q_weight.dtype() == DataType::kInt32 &&
202 b_q_weight.IsContiguous() &&
203 "`MoeWna16MarlinGemm` received invalid packed weights");
204 assert(b_scales.ndim() == 3 && b_scales.size(0) == b_q_weight.size(0) &&
205 b_scales.size(1) > 0 && b_scales.size(2) == size_n_ &&
206 size_k_ % b_scales.size(1) == 0 &&
207 (output_dtype == DataType::kFloat16 ||
208 output_dtype == DataType::kBFloat16) &&
209 b_scales.dtype() == output_dtype && b_scales.IsContiguous() &&
210 "`MoeWna16MarlinGemm` received invalid weight scales");
212 a_scales.has_value() == is_a_8bit &&
213 "`MoeWna16MarlinGemm` requires activation scales exactly for int8 `a`");
215 assert(a_scales->shape() ==
216 Tensor::Shape({static_cast<Tensor::Size>(size_m_), 1}) &&
217 a_scales->dtype() == DataType::kFloat32 &&
218 a_scales->IsContiguous() &&
219 "`MoeWna16MarlinGemm` received invalid activation scales");
222 if (b_bias_or_none) {
223 assert(b_bias_or_none->ndim() == 2 &&
224 b_bias_or_none->size(0) == b_q_weight.size(0) &&
225 b_bias_or_none->size(1) == size_n_ &&
226 b_bias_or_none->dtype() == output_dtype &&
227 b_bias_or_none->IsContiguous() &&
228 "`MoeWna16MarlinGemm` received invalid bias");
231 if (b_zeros_or_none) {
232 assert(b_zeros_or_none->ndim() == 3 &&
233 b_zeros_or_none->size(0) == b_q_weight.size(0) &&
234 b_zeros_or_none->size(1) == b_scales.size(1) &&
235 b_zeros_or_none->size(2) ==
236 (is_zp_float_ ? size_n_ : size_n_ / pack_factor) &&
237 (is_zp_float_ ? b_zeros_or_none->dtype() == output_dtype
238 : b_zeros_or_none->dtype() == DataType::kInt32) &&
239 "`MoeWna16MarlinGemm` received invalid zero points");
242 const auto same_device = [&](
const Tensor tensor) {
243 return tensor.device().type() == a.device().type() &&
244 tensor.device().index() == a.device().index();
246 const auto valid_optional = [&](
const std::optional<Tensor>& tensor) {
247 return !tensor || (tensor->IsContiguous() && same_device(*tensor));
249 assert(same_device(b_q_weight) && same_device(b_scales) &&
250 valid_optional(b_bias_or_none) && valid_optional(a_scales) &&
251 valid_optional(global_scale) && valid_optional(b_zeros_or_none) &&
252 valid_optional(g_idx_or_none) && valid_optional(perm_or_none) &&
253 same_device(workspace) && same_device(sorted_token_ids) &&
254 same_device(expert_ids) && same_device(num_tokens_past_padded) &&
255 same_device(topk_weights) && same_device(out) &&
256 "`MoeWna16MarlinGemm` requires all tensors on the input device");
258 assert(g_idx_or_none.has_value() == perm_or_none.has_value() &&
259 "`MoeWna16MarlinGemm` requires `g_idx_or_none` and `perm_or_none` "
262 assert(g_idx_or_none->ndim() > 0 && perm_or_none->ndim() > 0 &&
263 g_idx_or_none->size(-1) == perm_or_none->size(-1) &&
264 (g_idx_or_none->size(-1) == 0 ||
265 g_idx_or_none->size(-1) == size_k_) &&
266 g_idx_or_none->dtype() == DataType::kInt32 &&
267 perm_or_none->dtype() == DataType::kInt32 &&
268 (!is_full_k_ || b_scales.size(1) > 1) &&
269 "`MoeWna16MarlinGemm` received invalid activation-order "
273 assert(workspace.ndim() == 1 && workspace.numel() > 0 &&
274 workspace.dtype() == DataType::kInt32 && workspace.IsContiguous() &&
275 "`MoeWna16MarlinGemm` requires a non-empty int32 workspace");
276 assert(sorted_token_ids.ndim() == 1 &&
277 sorted_token_ids.dtype() == DataType::kInt32 &&
278 sorted_token_ids.IsContiguous() && expert_ids.ndim() == 1 &&
279 expert_ids.dtype() == DataType::kInt32 &&
280 expert_ids.IsContiguous() && num_tokens_past_padded.numel() == 1 &&
281 num_tokens_past_padded.dtype() == DataType::kInt32 &&
282 num_tokens_past_padded.IsContiguous() &&
283 "`MoeWna16MarlinGemm` received invalid routing metadata");
284 assert(topk_weights.numel() == size_m_ * top_k_ &&
285 ((!mul_topk_weights_ &&
286 (topk_weights.dtype() == DataType::kFloat16 ||
287 topk_weights.dtype() == DataType::kBFloat16)) ||
288 topk_weights.dtype() == DataType::kFloat32) &&
289 topk_weights.IsContiguous() &&
290 "`MoeWna16MarlinGemm` received invalid top-k weights");
291 assert(out.shape() ==
292 Tensor::Shape({static_cast<Tensor::Size>(size_m_ * top_k_),
293 static_cast<Tensor::Size>(size_n_)}) &&
294 out.dtype() == output_dtype && out.IsContiguous() &&
295 "`MoeWna16MarlinGemm` output metadata is invalid");
300 Tensor b_q_weight_metadata_;
302 std::optional<Tensor> b_bias_or_none_metadata_;
304 Tensor b_scales_metadata_;
306 std::optional<Tensor> a_scales_metadata_;
308 std::optional<Tensor> global_scale_metadata_;
310 std::optional<Tensor> b_zeros_or_none_metadata_;
312 std::optional<Tensor> g_idx_or_none_metadata_;
314 std::optional<Tensor> perm_or_none_metadata_;
316 Tensor workspace_metadata_;
318 Tensor sorted_token_ids_metadata_;
320 Tensor expert_ids_metadata_;
322 Tensor num_tokens_past_padded_metadata_;
324 Tensor topk_weights_metadata_;
328 int64_t moe_block_size_{0};
332 bool mul_topk_weights_{
false};
334 int64_t b_type_id_{0};
342 bool is_full_k_{
false};
344 bool use_atomic_add_{
false};
346 bool use_fp32_reduce_{
false};
348 bool is_zp_float_{
false};
350 int64_t thread_k_{0};
352 int64_t thread_n_{0};
354 int64_t blocks_per_sm_{0};
void ValidateCallMetadata(const Tensor a, const Tensor b_q_weight, std::optional< Tensor > b_bias_or_none, const Tensor b_scales, std::optional< Tensor > a_scales, std::optional< Tensor > global_scale, std::optional< Tensor > b_zeros_or_none, std::optional< Tensor > g_idx_or_none, std::optional< Tensor > 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 int64_t b_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, const int64_t thread_k, const int64_t thread_n, const int64_t blocks_per_sm, Tensor out) const
Definition moe_wna16_marlin_gemm.h:88
virtual void operator()(const Tensor a, const Tensor b_q_weight, std::optional< Tensor > b_bias_or_none, const Tensor b_scales, std::optional< Tensor > a_scales, std::optional< Tensor > global_scale, std::optional< Tensor > b_zeros_or_none, std::optional< Tensor > g_idx_or_none, std::optional< Tensor > 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 int64_t b_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, const int64_t thread_k, const int64_t thread_n, const int64_t blocks_per_sm, Tensor out) const =0
MoeWna16MarlinGemm(const Tensor a, const Tensor b_q_weight, std::optional< Tensor > b_bias_or_none, const Tensor b_scales, std::optional< Tensor > a_scales, std::optional< Tensor > global_scale, std::optional< Tensor > b_zeros_or_none, std::optional< Tensor > g_idx_or_none, std::optional< Tensor > 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 int64_t b_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, const int64_t thread_k, const int64_t thread_n, const int64_t blocks_per_sm, Tensor out)
Definition moe_wna16_marlin_gemm.h:18