InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
moe_sum.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_MOE_SUM_H_
2#define INFINI_OPS_BASE_MOE_SUM_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_sum`.
14class MoeSum : public Operator<MoeSum> {
15 public:
16 MoeSum(const Tensor input, Tensor output)
17 : MoeSum{input, std::nullopt, std::nullopt, output} {}
18
19 MoeSum(const Tensor input, std::optional<Tensor> topk_ids,
20 std::optional<Tensor> expert_map, Tensor output)
21 : num_tokens_{input.ndim() == 3 ? input.size(0) : 0},
22 topk_{input.ndim() == 3 ? input.size(1) : 0},
23 hidden_size_{input.ndim() == 3 ? input.size(2) : 0},
24 input_strides_{input.strides()},
25 output_strides_{output.strides()},
26 dtype_{input.dtype()},
27 device_type_{input.device().type()},
28 has_topk_ids_{topk_ids.has_value()},
29 topk_ids_dtype_{topk_ids ? topk_ids->dtype() : DataType::kInt32},
31 topk_ids && topk_ids->ndim() == 2 ? topk_ids->stride(0) : 0},
33 topk_ids && topk_ids->ndim() == 2 ? topk_ids->stride(1) : 0},
34 has_expert_map_{expert_map.has_value()},
35 expert_map_size_{expert_map ? expert_map->numel() : 0},
37 expert_map && expert_map->ndim() == 1 ? expert_map->stride(0) : 0},
38 device_index_{input.device().index()} {
39 assert(input.ndim() == 3 && output.ndim() == 2 &&
40 "`MoeSum` requires `[num_tokens, topk, hidden_size]` input and "
41 "`[num_tokens, hidden_size]` output");
42 assert(output.size(0) == num_tokens_ && output.size(1) == hidden_size_ &&
43 "`MoeSum` output shape is incompatible with the input");
44 assert(topk_ > 0 && "`MoeSum` requires at least one top-k slot");
45 assert((dtype_ == DataType::kFloat32 || dtype_ == DataType::kFloat16 ||
46 dtype_ == DataType::kBFloat16) &&
47 "`MoeSum` supports float32, float16, and bfloat16 inputs");
48 assert(output.dtype() == dtype_ &&
49 "`MoeSum` input and output dtypes must match");
50 assert(output.IsContiguous() && "`MoeSum` requires contiguous output");
51
52 constexpr auto kMaxSignedIndex =
53 static_cast<Tensor::Size>(std::numeric_limits<int64_t>::max());
54 assert(num_tokens_ <= kMaxSignedIndex && topk_ <= kMaxSignedIndex &&
55 hidden_size_ <= kMaxSignedIndex &&
56 "`MoeSum` dimensions must fit signed index arithmetic");
57 assert(num_tokens_ <= std::numeric_limits<int>::max() &&
58 "`MoeSum` token count exceeds the CUDA grid limit");
59 assert(
60 (hidden_size_ == 0 || num_tokens_ <= kMaxSignedIndex / hidden_size_) &&
61 "`MoeSum` output size must fit signed index arithmetic");
62
63 const auto offsets_fit = [](const auto& shape, const auto& strides) {
64 uint64_t max_offset = 0;
65 constexpr auto kMaxOffset =
66 static_cast<uint64_t>(std::numeric_limits<int64_t>::max());
67
68 for (Tensor::Size dim = 0; dim < shape.size(); ++dim) {
69 if (static_cast<uint64_t>(shape[dim]) > kMaxOffset) {
70 return false;
71 }
72 if (strides[dim] < 0) {
73 return false;
74 }
75 if (shape[dim] == 0) {
76 continue;
77 }
78
79 const auto extent = static_cast<uint64_t>(shape[dim] - 1);
80 const auto stride = static_cast<uint64_t>(strides[dim]);
81 if (extent != 0 && stride > kMaxOffset / extent) {
82 return false;
83 }
84
85 const auto term = extent * stride;
86 if (term > kMaxOffset - max_offset) {
87 return false;
88 }
89 max_offset += term;
90 }
91
92 return true;
93 };
94 assert(offsets_fit(input.shape(), input.strides()) &&
95 "`MoeSum` input requires non-negative strides with signed offsets");
96
97 const auto same_device_as_input = [&](const Tensor tensor) {
98 return tensor.device().type() == input.device().type() &&
99 tensor.device().index() == input.device().index();
100 };
101 assert(same_device_as_input(output) &&
102 "`MoeSum` input and output must be on the same device");
103 assert((!expert_map || topk_ids) &&
104 "`MoeSum` expert_map requires topk_ids");
105
106 if (topk_ids) {
107 assert(topk_ids->ndim() == 2 && topk_ids->size(0) == num_tokens_ &&
108 topk_ids->size(1) == topk_ &&
109 "`MoeSum` topk_ids must have shape `[num_tokens, topk]`");
110 assert((topk_ids_dtype_ == DataType::kInt32 ||
111 topk_ids_dtype_ == DataType::kInt64) &&
112 "`MoeSum` topk_ids must have int32 or int64 dtype");
113 assert(same_device_as_input(*topk_ids) &&
114 "`MoeSum` topk_ids must be on the input device");
115 assert(offsets_fit(topk_ids->shape(), topk_ids->strides()) &&
116 "`MoeSum` topk_ids requires non-negative strides with signed "
117 "offsets");
118 }
119
120 if (expert_map) {
121 assert(expert_map->ndim() == 1 &&
122 expert_map->dtype() == DataType::kInt32 &&
123 "`MoeSum` expert_map must be a 1D int32 tensor");
124 assert(same_device_as_input(*expert_map) &&
125 "`MoeSum` expert_map must be on the input device");
126 assert(offsets_fit(expert_map->shape(), expert_map->strides()) &&
127 "`MoeSum` expert_map requires non-negative strides with signed "
128 "offsets");
129 }
130 }
131
132 void operator()(const Tensor input, Tensor output) const {
133 (*this)(input, std::nullopt, std::nullopt, output);
134 }
135
136 virtual void operator()(const Tensor input, std::optional<Tensor> topk_ids,
137 std::optional<Tensor> expert_map,
138 Tensor output) const = 0;
139
140 protected:
141 void ValidateCallMetadata(const Tensor input, std::optional<Tensor> topk_ids,
142 std::optional<Tensor> expert_map,
143 const Tensor output) const {
144 const auto same_device_as_descriptor = [&](const Tensor tensor) {
145 return tensor.device().type() == device_type_ &&
146 tensor.device().index() == device_index_;
147 };
148 auto matches =
149 input.ndim() == 3 && input.size(0) == num_tokens_ &&
150 input.size(1) == topk_ && input.size(2) == hidden_size_ &&
151 input.strides() == input_strides_ && input.dtype() == dtype_ &&
152 same_device_as_descriptor(input) && output.ndim() == 2 &&
153 output.size(0) == num_tokens_ && output.size(1) == hidden_size_ &&
154 output.strides() == output_strides_ && output.dtype() == dtype_ &&
155 same_device_as_descriptor(output) &&
156 topk_ids.has_value() == has_topk_ids_ &&
157 expert_map.has_value() == has_expert_map_;
158
159 if (matches && topk_ids) {
160 matches = topk_ids->ndim() == 2 && topk_ids->size(0) == num_tokens_ &&
161 topk_ids->size(1) == topk_ &&
162 topk_ids->stride(0) == topk_ids_token_stride_ &&
163 topk_ids->stride(1) == topk_ids_slot_stride_ &&
164 topk_ids->dtype() == topk_ids_dtype_ &&
165 same_device_as_descriptor(*topk_ids);
166 }
167
168 if (matches && expert_map) {
169 matches = expert_map->ndim() == 1 &&
170 expert_map->numel() == expert_map_size_ &&
171 expert_map->stride(0) == expert_map_stride_ &&
172 expert_map->dtype() == DataType::kInt32 &&
173 same_device_as_descriptor(*expert_map);
174 }
175
176 assert(matches && "`MoeSum` call metadata must match descriptor");
177 }
178
179 Tensor::Size num_tokens_{0};
180
181 Tensor::Size topk_{0};
182
183 Tensor::Size hidden_size_{0};
184
185 Tensor::Strides input_strides_;
186
187 Tensor::Strides output_strides_;
188
189 DataType dtype_;
190
191 Device::Type device_type_;
192
193 bool has_topk_ids_{false};
194
196
197 Tensor::Stride topk_ids_token_stride_{0};
198
199 Tensor::Stride topk_ids_slot_stride_{0};
200
201 bool has_expert_map_{false};
202
203 Tensor::Size expert_map_size_{0};
204
205 Tensor::Stride expert_map_stride_{0};
206
208};
209
210} // namespace infini::ops
211
212#endif // INFINI_OPS_BASE_MOE_SUM_H_
Definition moe_sum.h:14
Tensor::Stride expert_map_stride_
Definition moe_sum.h:205
bool has_topk_ids_
Definition moe_sum.h:193
Tensor::Size hidden_size_
Definition moe_sum.h:183
void operator()(const Tensor input, Tensor output) const
Definition moe_sum.h:132
void ValidateCallMetadata(const Tensor input, std::optional< Tensor > topk_ids, std::optional< Tensor > expert_map, const Tensor output) const
Definition moe_sum.h:141
MoeSum(const Tensor input, Tensor output)
Definition moe_sum.h:16
DataType dtype_
Definition moe_sum.h:189
bool has_expert_map_
Definition moe_sum.h:201
DataType topk_ids_dtype_
Definition moe_sum.h:195
Tensor::Strides input_strides_
Definition moe_sum.h:185
virtual void operator()(const Tensor input, std::optional< Tensor > topk_ids, std::optional< Tensor > expert_map, Tensor output) const =0
Tensor::Stride topk_ids_slot_stride_
Definition moe_sum.h:199
Tensor::Size num_tokens_
Definition moe_sum.h:179
Device::Type device_type_
Definition moe_sum.h:191
Tensor::Stride topk_ids_token_stride_
Definition moe_sum.h:197
int device_index_
Definition moe_sum.h:207
Tensor::Size topk_
Definition moe_sum.h:181
Tensor::Size expert_map_size_
Definition moe_sum.h:203
Tensor::Strides output_strides_
Definition moe_sum.h:187
MoeSum(const Tensor input, std::optional< Tensor > topk_ids, std::optional< Tensor > expert_map, Tensor output)
Definition moe_sum.h:19
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
template void uint64_t
Definition operator_call_instantiations.h:101
infini::rt::TensorView Tensor
Definition tensor.h:8