InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
get_cutlass_moe_mm_data.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_GET_CUTLASS_MOE_MM_DATA_H_
2#define INFINI_OPS_BASE_GET_CUTLASS_MOE_MM_DATA_H_
3
4#include <cassert>
5#include <cstdint>
6#include <functional>
7#include <limits>
8#include <optional>
9
10#include "operator.h"
11
12namespace infini::ops {
13
14// Aligned with vLLM `get_cutlass_moe_mm_data`.
15class GetCutlassMoeMmData : public Operator<GetCutlassMoeMmData> {
16 public:
17 GetCutlassMoeMmData(const Tensor topk_ids, const int64_t num_experts,
18 const int64_t n, const int64_t k, Tensor expert_offsets,
19 Tensor problem_sizes1, Tensor problem_sizes2,
20 Tensor input_permutation, Tensor output_permutation,
21 std::optional<Tensor> blockscale_offsets = std::nullopt)
22 : GetCutlassMoeMmData{topk_ids,
23 num_experts,
24 n,
25 k,
26 true,
27 expert_offsets,
28 problem_sizes1,
29 problem_sizes2,
30 input_permutation,
31 output_permutation,
32 blockscale_offsets} {}
33
34 GetCutlassMoeMmData(const Tensor topk_ids, const int64_t num_experts,
35 const int64_t n, const int64_t k, const bool is_gated,
36 Tensor expert_offsets, Tensor problem_sizes1,
37 Tensor problem_sizes2, Tensor input_permutation,
38 Tensor output_permutation,
39 std::optional<Tensor> blockscale_offsets = std::nullopt)
40 : topk_ids_metadata_{topk_ids},
41 expert_offsets_metadata_{expert_offsets},
42 problem_sizes1_metadata_{problem_sizes1},
43 problem_sizes2_metadata_{problem_sizes2},
44 input_permutation_metadata_{input_permutation},
45 output_permutation_metadata_{output_permutation},
46 blockscale_offsets_metadata_{blockscale_offsets},
47 num_experts_{num_experts},
48 n_{n},
49 k_{k},
50 is_gated_{is_gated},
51 numel_{topk_ids.numel()},
52 topk_{topk_ids.ndim() == 2 ? topk_ids.size(1) : 0},
53 device_index_{topk_ids.device().index()} {
54 Validate(topk_ids, expert_offsets, problem_sizes1, problem_sizes2,
55 input_permutation, output_permutation, blockscale_offsets);
56 }
57
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,
61 Tensor problem_sizes2, Tensor input_permutation,
62 Tensor output_permutation,
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,
66 blockscale_offsets);
67 }
68
69 virtual void operator()(
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,
72 Tensor problem_sizes1, Tensor problem_sizes2, Tensor input_permutation,
73 Tensor output_permutation,
74 std::optional<Tensor> blockscale_offsets = std::nullopt) const = 0;
75
76 protected:
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 {
83 assert(num_experts == num_experts_ && n == n_ && k == k_ &&
84 is_gated == is_gated_ &&
85 "`GetCutlassMoeMmData` attributes changed after descriptor "
86 "creation");
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");
91 }
92
94
96
98
100
102
104
105 std::optional<Tensor> blockscale_offsets_metadata_;
106
107 int64_t num_experts_{0};
108
109 int64_t n_{0};
110
111 int64_t k_{0};
112
113 bool is_gated_{true};
114
115 Tensor::Size numel_{0};
116
117 Tensor::Size topk_{0};
118
120
121 private:
122 bool CallMetadataMatches(
123 const Tensor topk_ids, const Tensor expert_offsets,
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 =
129 blockscale_offsets_metadata_.has_value() ==
130 blockscale_offsets.has_value() &&
132 same_metadata(*blockscale_offsets_metadata_, *blockscale_offsets));
133
134 return same_metadata(topk_ids_metadata_, topk_ids) &&
135 same_metadata(expert_offsets_metadata_, expert_offsets) &&
136 same_metadata(problem_sizes1_metadata_, problem_sizes1) &&
137 same_metadata(problem_sizes2_metadata_, problem_sizes2) &&
138 same_metadata(input_permutation_metadata_, input_permutation) &&
139 same_metadata(output_permutation_metadata_, output_permutation) &&
140 same_blockscale_metadata;
141 }
142
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`");
153 assert(numel_ > 0 && topk_ > 0 &&
154 "`GetCutlassMoeMmData` requires non-empty `topk_ids`");
155 assert(num_experts_ > 0 &&
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");
163 assert(numel_ <=
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");
173 }
174
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();
178 };
179 assert(
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");
187
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`]");
210
211 const auto is_int32_contiguous = [](const Tensor tensor) {
212 return tensor.dtype() == DataType::kInt32 && tensor.IsContiguous();
213 };
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");
221 }
222};
223
224} // namespace infini::ops
225
226#endif // INFINI_OPS_BASE_GET_CUTLASS_MOE_MM_DATA_H_
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