InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
topk_sigmoid.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_TOPK_SIGMOID_H_
2#define INFINI_OPS_BASE_TOPK_SIGMOID_H_
3
4#include <cassert>
5#include <cmath>
6#include <cstdint>
7#include <functional>
8#include <limits>
9#include <optional>
10
11#include "operator.h"
12
13namespace infini::ops {
14
15// Aligned with vLLM `_moe_C::topk_sigmoid`.
16class TopkSigmoid : public Operator<TopkSigmoid> {
17 public:
18 TopkSigmoid(const Tensor gating_output,
19 std::optional<Tensor> e_score_correction_bias,
20 std::optional<Tensor> is_padding, const bool renormalize,
21 const double routed_scaling_factor, Tensor topk_weights,
22 Tensor topk_ids, Tensor token_expert_indices)
23 : num_tokens_{gating_output.ndim() == 2 ? gating_output.size(0) : 0},
24 num_experts_{gating_output.ndim() == 2 ? gating_output.size(1) : 0},
25 topk_{topk_weights.ndim() == 2 ? topk_weights.size(1) : 0},
26 input_dtype_{gating_output.dtype()},
27 index_dtype_{topk_ids.dtype()},
28 renormalize_{renormalize},
29 routed_scaling_factor_{routed_scaling_factor},
30 device_index_{gating_output.device().index()},
31 gating_output_metadata_{gating_output},
32 e_score_correction_bias_metadata_{e_score_correction_bias},
33 is_padding_metadata_{is_padding},
34 topk_weights_metadata_{topk_weights},
35 topk_ids_metadata_{topk_ids},
36 token_expert_indices_metadata_{token_expert_indices} {
37 Validate(gating_output, e_score_correction_bias, is_padding, topk_weights,
38 topk_ids, token_expert_indices);
39 }
40
41 virtual void operator()(const Tensor gating_output,
42 std::optional<Tensor> e_score_correction_bias,
43 std::optional<Tensor> is_padding,
44 const bool renormalize,
45 const double routed_scaling_factor,
46 Tensor topk_weights, Tensor topk_ids,
47 Tensor token_expert_indices) const = 0;
48
49 protected:
50 void ValidateCallMetadata(const Tensor gating_output,
51 std::optional<Tensor> e_score_correction_bias,
52 std::optional<Tensor> is_padding,
53 const bool renormalize,
54 const double routed_scaling_factor,
55 Tensor topk_weights, Tensor topk_ids,
56 Tensor token_expert_indices) const {
57 assert(renormalize == renormalize_ &&
58 routed_scaling_factor == routed_scaling_factor_ &&
59 "`TopkSigmoid` attributes changed after descriptor creation");
60
61 const std::equal_to<Tensor> same_metadata;
62 const auto optional_matches = [&](const std::optional<Tensor>& expected,
63 const std::optional<Tensor>& actual) {
64 return expected.has_value() == actual.has_value() &&
65 (!expected || same_metadata(*expected, *actual));
66 };
67 const auto matches =
68 same_metadata(gating_output_metadata_, gating_output) &&
69 optional_matches(e_score_correction_bias_metadata_,
70 e_score_correction_bias) &&
71 optional_matches(is_padding_metadata_, is_padding) &&
72 same_metadata(topk_weights_metadata_, topk_weights) &&
73 same_metadata(topk_ids_metadata_, topk_ids) &&
74 same_metadata(token_expert_indices_metadata_, token_expert_indices);
75 assert(matches && "`TopkSigmoid` call metadata must match descriptor");
76 }
77
78 Tensor::Size num_tokens_{0};
79
80 Tensor::Size num_experts_{0};
81
82 Tensor::Size topk_{0};
83
84 DataType input_dtype_;
85
86 DataType index_dtype_;
87
88 bool renormalize_{false};
89
91
93
94 private:
95 void Validate(const Tensor gating_output,
96 std::optional<Tensor> e_score_correction_bias,
97 std::optional<Tensor> is_padding, Tensor topk_weights,
98 Tensor topk_ids, Tensor token_expert_indices) const {
99 assert(gating_output.ndim() == 2 &&
100 "`TopkSigmoid` requires 2D `gating_output`");
101 assert((input_dtype_ == DataType::kFloat32 ||
102 input_dtype_ == DataType::kFloat16 ||
103 input_dtype_ == DataType::kBFloat16) &&
104 "`TopkSigmoid` supports float32, float16, and bfloat16 input");
105 assert(gating_output.IsContiguous() &&
106 "`TopkSigmoid` requires contiguous `gating_output`");
107 assert(num_experts_ > 0 && topk_ > 0 && topk_ <= num_experts_ &&
108 "`TopkSigmoid` requires `topk` in `[1, num_experts]`");
109 assert(std::isfinite(routed_scaling_factor_) &&
110 "`TopkSigmoid` requires a finite `routed_scaling_factor`");
111 assert(num_tokens_ <=
112 static_cast<Tensor::Size>(std::numeric_limits<int32_t>::max()) &&
113 num_experts_ <=
114 static_cast<Tensor::Size>(std::numeric_limits<int32_t>::max()) &&
115 topk_ <=
116 static_cast<Tensor::Size>(std::numeric_limits<int32_t>::max()) &&
117 "`TopkSigmoid` dimensions must fit int32 indexing");
118 assert(num_tokens_ <=
119 static_cast<Tensor::Size>(std::numeric_limits<int32_t>::max()) /
120 topk_ &&
121 "`TopkSigmoid` output indices must fit int32 indexing");
122
123 const Tensor::Shape output_shape{num_tokens_, topk_};
124 assert(topk_weights.shape() == output_shape &&
125 topk_ids.shape() == output_shape &&
126 token_expert_indices.shape() == output_shape &&
127 "`TopkSigmoid` outputs must have shape `[num_tokens, topk]`");
128 assert(topk_weights.dtype() == DataType::kFloat32 &&
129 "`TopkSigmoid` requires float32 `topk_weights`");
130 assert((index_dtype_ == DataType::kInt32 ||
131 index_dtype_ == DataType::kUInt32 ||
132 index_dtype_ == DataType::kInt64) &&
133 "`TopkSigmoid` requires int32, uint32, or int64 `topk_ids`");
134 assert(token_expert_indices.dtype() == DataType::kInt32 &&
135 "`TopkSigmoid` requires int32 `token_expert_indices`");
136 assert(topk_weights.IsContiguous() && topk_ids.IsContiguous() &&
137 token_expert_indices.IsContiguous() &&
138 "`TopkSigmoid` requires contiguous outputs");
139
140 const auto same_device = [&](const Tensor tensor) {
141 return tensor.device().type() == gating_output.device().type() &&
142 tensor.device().index() == gating_output.device().index();
143 };
144 assert(same_device(topk_weights) && same_device(topk_ids) &&
145 same_device(token_expert_indices) &&
146 "`TopkSigmoid` requires all tensors on the same device");
147
148 if (e_score_correction_bias) {
149 assert(e_score_correction_bias->ndim() == 1 &&
150 e_score_correction_bias->numel() == num_experts_ &&
151 "`TopkSigmoid` requires `e_score_correction_bias` shape "
152 "`[num_experts]`");
153 assert(e_score_correction_bias->dtype() == DataType::kFloat32 &&
154 e_score_correction_bias->IsContiguous() &&
155 "`TopkSigmoid` requires contiguous float32 "
156 "`e_score_correction_bias`");
157 assert(same_device(*e_score_correction_bias) &&
158 "`TopkSigmoid` requires `e_score_correction_bias` on the input "
159 "device");
160 }
161
162 if (is_padding) {
163 assert(is_padding->ndim() == 1 && is_padding->numel() == num_tokens_ &&
164 "`TopkSigmoid` requires `is_padding` shape `[num_tokens]`");
165 assert(is_padding->dtype() == DataType::kBool &&
166 is_padding->IsContiguous() &&
167 "`TopkSigmoid` requires contiguous bool `is_padding`");
168 assert(same_device(*is_padding) &&
169 "`TopkSigmoid` requires `is_padding` on the input device");
170 }
171 }
172
173 Tensor gating_output_metadata_;
174
175 std::optional<Tensor> e_score_correction_bias_metadata_;
176
177 std::optional<Tensor> is_padding_metadata_;
178
179 Tensor topk_weights_metadata_;
180
181 Tensor topk_ids_metadata_;
182
183 Tensor token_expert_indices_metadata_;
184};
185
186} // namespace infini::ops
187
188#endif // INFINI_OPS_BASE_TOPK_SIGMOID_H_
Definition generated/include/operator.h:282
Definition topk_sigmoid.h:16
DataType input_dtype_
Definition topk_sigmoid.h:84
bool renormalize_
Definition topk_sigmoid.h:88
Tensor::Size topk_
Definition topk_sigmoid.h:82
virtual void operator()(const Tensor gating_output, std::optional< Tensor > e_score_correction_bias, std::optional< Tensor > is_padding, const bool renormalize, const double routed_scaling_factor, Tensor topk_weights, Tensor topk_ids, Tensor token_expert_indices) const =0
double routed_scaling_factor_
Definition topk_sigmoid.h:90
Tensor::Size num_tokens_
Definition topk_sigmoid.h:78
TopkSigmoid(const Tensor gating_output, std::optional< Tensor > e_score_correction_bias, std::optional< Tensor > is_padding, const bool renormalize, const double routed_scaling_factor, Tensor topk_weights, Tensor topk_ids, Tensor token_expert_indices)
Definition topk_sigmoid.h:18
Tensor::Size num_experts_
Definition topk_sigmoid.h:80
int device_index_
Definition topk_sigmoid.h:92
void ValidateCallMetadata(const Tensor gating_output, std::optional< Tensor > e_score_correction_bias, std::optional< Tensor > is_padding, const bool renormalize, const double routed_scaling_factor, Tensor topk_weights, Tensor topk_ids, Tensor token_expert_indices) const
Definition topk_sigmoid.h:50
DataType index_dtype_
Definition topk_sigmoid.h:86
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8