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,
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},
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);
42 std::optional<Tensor> e_score_correction_bias,
43 std::optional<Tensor> is_padding,
44 const bool renormalize,
45 const double routed_scaling_factor,
47 Tensor token_expert_indices)
const = 0;
51 std::optional<Tensor> e_score_correction_bias,
52 std::optional<Tensor> is_padding,
53 const bool renormalize,
54 const double routed_scaling_factor,
56 Tensor token_expert_indices)
const {
59 "`TopkSigmoid` attributes changed after descriptor creation");
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));
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");
95 void Validate(
const Tensor gating_output,
96 std::optional<Tensor> e_score_correction_bias,
97 std::optional<Tensor> is_padding,
Tensor topk_weights,
99 assert(gating_output.ndim() == 2 &&
100 "`TopkSigmoid` requires 2D `gating_output`");
104 "`TopkSigmoid` supports float32, float16, and bfloat16 input");
105 assert(gating_output.IsContiguous() &&
106 "`TopkSigmoid` requires contiguous `gating_output`");
108 "`TopkSigmoid` requires `topk` in `[1, num_experts]`");
110 "`TopkSigmoid` requires a finite `routed_scaling_factor`");
112 static_cast<Tensor::Size
>(std::numeric_limits<int32_t>::max()) &&
114 static_cast<Tensor::Size
>(std::numeric_limits<int32_t>::max()) &&
116 static_cast<Tensor::Size
>(std::numeric_limits<int32_t>::max()) &&
117 "`TopkSigmoid` dimensions must fit int32 indexing");
119 static_cast<Tensor::Size
>(std::numeric_limits<int32_t>::max()) /
121 "`TopkSigmoid` output indices must fit int32 indexing");
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`");
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");
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();
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");
148 if (e_score_correction_bias) {
149 assert(e_score_correction_bias->ndim() == 1 &&
151 "`TopkSigmoid` requires `e_score_correction_bias` shape "
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 "
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");
173 Tensor gating_output_metadata_;
175 std::optional<Tensor> e_score_correction_bias_metadata_;
177 std::optional<Tensor> is_padding_metadata_;
179 Tensor topk_weights_metadata_;
181 Tensor topk_ids_metadata_;
183 Tensor token_expert_indices_metadata_;
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
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
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