1#ifndef INFINI_OPS_BASE_GET_CUTLASS_MOE_MM_DATA_H_
2#define INFINI_OPS_BASE_GET_CUTLASS_MOE_MM_DATA_H_
18 const int64_t n,
const int64_t k,
Tensor expert_offsets,
21 std::optional<Tensor> blockscale_offsets = std::nullopt)
32 blockscale_offsets} {}
35 const int64_t n,
const int64_t k,
const bool is_gated,
39 std::optional<Tensor> blockscale_offsets = std::nullopt)
52 topk_{topk_ids.ndim() == 2 ? topk_ids.size(1) : 0},
54 Validate(topk_ids, expert_offsets, problem_sizes1, problem_sizes2,
55 input_permutation, output_permutation, blockscale_offsets);
59 const Tensor topk_ids,
const int64_t num_experts,
const int64_t n,
60 const int64_t k,
Tensor expert_offsets,
Tensor problem_sizes1,
63 std::optional<Tensor> blockscale_offsets = std::nullopt)
const {
64 (*this)(topk_ids, num_experts, n, k,
true, expert_offsets, problem_sizes1,
65 problem_sizes2, input_permutation, output_permutation,
70 const Tensor topk_ids,
const int64_t num_experts,
const int64_t n,
71 const int64_t k,
const bool is_gated,
Tensor expert_offsets,
74 std::optional<Tensor> blockscale_offsets = std::nullopt)
const = 0;
78 const Tensor topk_ids,
const int64_t num_experts,
const int64_t n,
79 const int64_t k,
const bool is_gated,
const Tensor expert_offsets,
80 const Tensor problem_sizes1,
const Tensor problem_sizes2,
81 const Tensor input_permutation,
const Tensor output_permutation,
82 const std::optional<Tensor> blockscale_offsets)
const {
85 "`GetCutlassMoeMmData` attributes changed after descriptor "
87 assert(CallMetadataMatches(topk_ids, expert_offsets, problem_sizes1,
88 problem_sizes2, input_permutation,
89 output_permutation, blockscale_offsets) &&
90 "`GetCutlassMoeMmData` tensor metadata differs from its descriptor");
122 bool CallMetadataMatches(
124 const Tensor problem_sizes1,
const Tensor problem_sizes2,
125 const Tensor input_permutation,
const Tensor output_permutation,
126 const std::optional<Tensor> blockscale_offsets)
const {
127 const std::equal_to<Tensor> same_metadata;
128 const auto same_blockscale_metadata =
130 blockscale_offsets.has_value() &&
140 same_blockscale_metadata;
143 void Validate(
const Tensor topk_ids,
const Tensor expert_offsets,
144 const Tensor problem_sizes1,
const Tensor problem_sizes2,
145 const Tensor input_permutation,
const Tensor output_permutation,
146 const std::optional<Tensor> blockscale_offsets)
const {
147 assert(topk_ids.ndim() == 2 &&
148 "`GetCutlassMoeMmData` requires 2D `topk_ids`");
149 assert(topk_ids.dtype() == DataType::kInt32 &&
150 "`GetCutlassMoeMmData` requires int32 `topk_ids`");
151 assert(topk_ids.IsContiguous() &&
152 "`GetCutlassMoeMmData` requires contiguous `topk_ids`");
154 "`GetCutlassMoeMmData` requires non-empty `topk_ids`");
156 num_experts_ <= std::numeric_limits<int32_t>::max() / 3 &&
157 "`GetCutlassMoeMmData` requires indexable int32 `num_experts`");
158 const auto max_n =
is_gated_ ? std::numeric_limits<int32_t>::max() / 2
159 : std::numeric_limits<int32_t>::max();
160 assert(
n_ > 0 && n_ <= max_n && k_ > 0 &&
161 k_ <= std::numeric_limits<int32_t>::max() &&
162 "`GetCutlassMoeMmData` requires positive int32 GEMM dimensions");
164 static_cast<Tensor::Size
>(std::numeric_limits<int32_t>::max()) &&
165 "`GetCutlassMoeMmData` requires int32-addressable routing indices");
166 if (blockscale_offsets) {
167 constexpr int64_t kBlockscalePadding = 127;
168 const int64_t max_blockscale_padding =
num_experts_ * kBlockscalePadding;
169 assert(max_blockscale_padding <= std::numeric_limits<int32_t>::max() &&
170 static_cast<int64_t
>(
numel_) <=
171 std::numeric_limits<int32_t>::max() - max_blockscale_padding &&
172 "`GetCutlassMoeMmData` blockscale offsets must fit int32");
175 const auto same_device_as_topk_ids = [&](
const Tensor tensor) {
176 return tensor.device().type() == topk_ids.device().type() &&
177 tensor.device().index() == topk_ids.device().index();
180 same_device_as_topk_ids(expert_offsets) &&
181 same_device_as_topk_ids(problem_sizes1) &&
182 same_device_as_topk_ids(problem_sizes2) &&
183 same_device_as_topk_ids(input_permutation) &&
184 same_device_as_topk_ids(output_permutation) &&
185 (!blockscale_offsets || same_device_as_topk_ids(*blockscale_offsets)) &&
186 "`GetCutlassMoeMmData` requires all tensors on the same device");
188 const auto num_experts =
static_cast<Tensor::Size
>(
num_experts_);
189 assert(expert_offsets.ndim() == 1 &&
190 expert_offsets.numel() == num_experts + 1 &&
191 "`GetCutlassMoeMmData` requires `expert_offsets` shape "
192 "[`num_experts + 1`]");
193 assert(problem_sizes1.ndim() == 2 &&
194 problem_sizes1.size(0) == num_experts &&
195 problem_sizes1.size(1) == 3 && problem_sizes2.ndim() == 2 &&
196 problem_sizes2.shape() == problem_sizes1.shape() &&
197 "`GetCutlassMoeMmData` requires problem sizes shape "
198 "[`num_experts`, 3]");
199 assert(input_permutation.ndim() == 1 &&
200 input_permutation.numel() ==
numel_ &&
201 output_permutation.ndim() == 1 &&
202 output_permutation.numel() ==
numel_ &&
203 "`GetCutlassMoeMmData` requires permutation shape "
204 "[`topk_ids.numel()`]");
205 assert((!blockscale_offsets ||
206 (blockscale_offsets->ndim() == 1 &&
207 blockscale_offsets->numel() == num_experts + 1)) &&
208 "`GetCutlassMoeMmData` requires `blockscale_offsets` shape "
209 "[`num_experts + 1`]");
211 const auto is_int32_contiguous = [](
const Tensor tensor) {
212 return tensor.dtype() == DataType::kInt32 && tensor.IsContiguous();
214 assert(is_int32_contiguous(expert_offsets) &&
215 is_int32_contiguous(problem_sizes1) &&
216 is_int32_contiguous(problem_sizes2) &&
217 is_int32_contiguous(input_permutation) &&
218 is_int32_contiguous(output_permutation) &&
219 (!blockscale_offsets || is_int32_contiguous(*blockscale_offsets)) &&
220 "`GetCutlassMoeMmData` requires contiguous int32 outputs");
Definition get_cutlass_moe_mm_data.h:15
Tensor output_permutation_metadata_
Definition get_cutlass_moe_mm_data.h:103
GetCutlassMoeMmData(const Tensor topk_ids, const int64_t num_experts, const int64_t n, const int64_t k, const bool is_gated, Tensor expert_offsets, Tensor problem_sizes1, Tensor problem_sizes2, Tensor input_permutation, Tensor output_permutation, std::optional< Tensor > blockscale_offsets=std::nullopt)
Definition get_cutlass_moe_mm_data.h:34
int64_t num_experts_
Definition get_cutlass_moe_mm_data.h:107
Tensor::Size topk_
Definition get_cutlass_moe_mm_data.h:117
Tensor expert_offsets_metadata_
Definition get_cutlass_moe_mm_data.h:95
int64_t n_
Definition get_cutlass_moe_mm_data.h:109
Tensor::Size numel_
Definition get_cutlass_moe_mm_data.h:115
Tensor problem_sizes1_metadata_
Definition get_cutlass_moe_mm_data.h:97
int device_index_
Definition get_cutlass_moe_mm_data.h:119
std::optional< Tensor > blockscale_offsets_metadata_
Definition get_cutlass_moe_mm_data.h:105
bool is_gated_
Definition get_cutlass_moe_mm_data.h:113
Tensor input_permutation_metadata_
Definition get_cutlass_moe_mm_data.h:101
virtual void operator()(const Tensor topk_ids, const int64_t num_experts, const int64_t n, const int64_t k, const bool is_gated, Tensor expert_offsets, Tensor problem_sizes1, Tensor problem_sizes2, Tensor input_permutation, Tensor output_permutation, std::optional< Tensor > blockscale_offsets=std::nullopt) const =0
int64_t k_
Definition get_cutlass_moe_mm_data.h:111
Tensor topk_ids_metadata_
Definition get_cutlass_moe_mm_data.h:93
void ValidateCallMetadata(const Tensor topk_ids, const int64_t num_experts, const int64_t n, const int64_t k, const bool is_gated, const Tensor expert_offsets, const Tensor problem_sizes1, const Tensor problem_sizes2, const Tensor input_permutation, const Tensor output_permutation, const std::optional< Tensor > blockscale_offsets) const
Definition get_cutlass_moe_mm_data.h:77
void operator()(const Tensor topk_ids, const int64_t num_experts, const int64_t n, const int64_t k, Tensor expert_offsets, Tensor problem_sizes1, Tensor problem_sizes2, Tensor input_permutation, Tensor output_permutation, std::optional< Tensor > blockscale_offsets=std::nullopt) const
Definition get_cutlass_moe_mm_data.h:58
Tensor problem_sizes2_metadata_
Definition get_cutlass_moe_mm_data.h:99
GetCutlassMoeMmData(const Tensor topk_ids, const int64_t num_experts, const int64_t n, const int64_t k, Tensor expert_offsets, Tensor problem_sizes1, Tensor problem_sizes2, Tensor input_permutation, Tensor output_permutation, std::optional< Tensor > blockscale_offsets=std::nullopt)
Definition get_cutlass_moe_mm_data.h:17
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8