18 std::optional<Tensor> is_padding,
const bool renormalize,
20 Tensor token_expert_indices)
21 :
num_tokens_{gating_output.ndim() == 2 ? gating_output.size(0) : 0},
22 num_experts_{gating_output.ndim() == 2 ? gating_output.size(1) : 0},
23 topk_{topk_weights.ndim() == 2 ? topk_weights.size(1) : 0},
28 gating_output_metadata_{gating_output},
30 is_padding_metadata_{is_padding},
31 topk_weights_metadata_{topk_weights},
32 topk_indices_metadata_{topk_indices},
33 token_expert_indices_metadata_{token_expert_indices} {
34 Validate(gating_output, bias, is_padding, topk_weights, topk_indices,
35 token_expert_indices);
39 std::optional<Tensor> bias,
40 std::optional<Tensor> is_padding,
41 const bool renormalize,
Tensor topk_weights,
43 Tensor token_expert_indices)
const = 0;
47 std::optional<Tensor> bias,
48 std::optional<Tensor> is_padding,
49 const bool renormalize,
Tensor topk_weights,
51 Tensor token_expert_indices)
const {
53 "`TopkSoftmax` attributes changed after descriptor creation");
55 const std::equal_to<Tensor> same_metadata;
56 const auto optional_matches = [&](
const std::optional<Tensor>& expected,
57 const std::optional<Tensor>& actual) {
58 return expected.has_value() == actual.has_value() &&
59 (!expected || same_metadata(*expected, *actual));
62 same_metadata(gating_output_metadata_, gating_output) &&
63 optional_matches(bias_metadata_, bias) &&
64 optional_matches(is_padding_metadata_, is_padding) &&
65 same_metadata(topk_weights_metadata_, topk_weights) &&
66 same_metadata(topk_indices_metadata_, topk_indices) &&
67 same_metadata(token_expert_indices_metadata_, token_expert_indices);
68 assert(matches &&
"`TopkSoftmax` call metadata must match descriptor");
86 void Validate(
const Tensor gating_output, std::optional<Tensor> bias,
87 std::optional<Tensor> is_padding,
Tensor topk_weights,
88 Tensor topk_indices,
Tensor token_expert_indices)
const {
89 assert(gating_output.ndim() == 2 &&
90 "`TopkSoftmax` requires 2D `gating_output`");
94 "`TopkSoftmax` supports float32, float16, and bfloat16 input");
95 assert(gating_output.IsContiguous() &&
96 "`TopkSoftmax` requires contiguous `gating_output`");
98 "`TopkSoftmax` requires `topk` in `[1, num_experts]`");
100 static_cast<Tensor::Size
>(std::numeric_limits<int32_t>::max()) &&
102 static_cast<Tensor::Size
>(std::numeric_limits<int32_t>::max()) &&
104 static_cast<Tensor::Size
>(std::numeric_limits<int32_t>::max()) &&
105 "`TopkSoftmax` dimensions must fit int32 indexing");
107 static_cast<Tensor::Size
>(std::numeric_limits<int32_t>::max()) /
109 "`TopkSoftmax` output indices must fit int32 indexing");
112 assert(topk_weights.shape() == output_shape &&
113 topk_indices.shape() == output_shape &&
114 token_expert_indices.shape() == output_shape &&
115 "`TopkSoftmax` outputs must have shape `[num_tokens, topk]`");
116 assert(topk_weights.dtype() == DataType::kFloat32 &&
117 "`TopkSoftmax` requires float32 `topk_weights`");
121 "`TopkSoftmax` requires int32, uint32, or int64 `topk_indices`");
122 assert(token_expert_indices.dtype() == DataType::kInt32 &&
123 "`TopkSoftmax` requires int32 `token_expert_indices`");
124 assert(topk_weights.IsContiguous() && topk_indices.IsContiguous() &&
125 token_expert_indices.IsContiguous() &&
126 "`TopkSoftmax` requires contiguous outputs");
128 const auto same_device = [&](
const Tensor tensor) {
129 return tensor.device().type() == gating_output.device().type() &&
130 tensor.device().index() == gating_output.device().index();
132 assert(same_device(topk_weights) && same_device(topk_indices) &&
133 same_device(token_expert_indices) &&
134 "`TopkSoftmax` requires all tensors on the same device");
137 assert(bias->ndim() == 1 && bias->numel() ==
num_experts_ &&
138 "`TopkSoftmax` requires `bias` shape `[num_experts]`");
139 assert(bias->dtype() == DataType::kFloat32 && bias->IsContiguous() &&
140 "`TopkSoftmax` requires contiguous float32 `bias`");
141 assert(same_device(*bias) &&
142 "`TopkSoftmax` requires `bias` on the input device");
146 assert(is_padding->ndim() == 1 && is_padding->numel() ==
num_tokens_ &&
147 "`TopkSoftmax` requires `is_padding` shape `[num_tokens]`");
148 assert(is_padding->dtype() == DataType::kBool &&
149 is_padding->IsContiguous() &&
150 "`TopkSoftmax` requires contiguous bool `is_padding`");
151 assert(same_device(*is_padding) &&
152 "`TopkSoftmax` requires `is_padding` on the input device");
156 Tensor gating_output_metadata_;
158 std::optional<Tensor> bias_metadata_;
160 std::optional<Tensor> is_padding_metadata_;
162 Tensor topk_weights_metadata_;
164 Tensor topk_indices_metadata_;
166 Tensor token_expert_indices_metadata_;
virtual void operator()(const Tensor gating_output, std::optional< Tensor > bias, std::optional< Tensor > is_padding, const bool renormalize, Tensor topk_weights, Tensor topk_indices, Tensor token_expert_indices) const =0
void ValidateCallMetadata(const Tensor gating_output, std::optional< Tensor > bias, std::optional< Tensor > is_padding, const bool renormalize, Tensor topk_weights, Tensor topk_indices, Tensor token_expert_indices) const
Definition topk_softmax.h:46
TopkSoftmax(const Tensor gating_output, std::optional< Tensor > bias, std::optional< Tensor > is_padding, const bool renormalize, Tensor topk_weights, Tensor topk_indices, Tensor token_expert_indices)
Definition topk_softmax.h:17