InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
topk_softmax.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_TOPK_SOFTMAX_H_
2#define INFINI_OPS_BASE_TOPK_SOFTMAX_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 `_moe_C::topk_softmax`.
15class TopkSoftmax : public Operator<TopkSoftmax> {
16 public:
17 TopkSoftmax(const Tensor gating_output, std::optional<Tensor> bias,
18 std::optional<Tensor> is_padding, const bool renormalize,
19 Tensor topk_weights, Tensor topk_indices,
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},
24 input_dtype_{gating_output.dtype()},
25 index_dtype_{topk_indices.dtype()},
26 renormalize_{renormalize},
27 device_index_{gating_output.device().index()},
28 gating_output_metadata_{gating_output},
29 bias_metadata_{bias},
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);
36 }
37
38 virtual void operator()(const Tensor gating_output,
39 std::optional<Tensor> bias,
40 std::optional<Tensor> is_padding,
41 const bool renormalize, Tensor topk_weights,
42 Tensor topk_indices,
43 Tensor token_expert_indices) const = 0;
44
45 protected:
46 void ValidateCallMetadata(const Tensor gating_output,
47 std::optional<Tensor> bias,
48 std::optional<Tensor> is_padding,
49 const bool renormalize, Tensor topk_weights,
50 Tensor topk_indices,
51 Tensor token_expert_indices) const {
52 assert(renormalize == renormalize_ &&
53 "`TopkSoftmax` attributes changed after descriptor creation");
54
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));
60 };
61 const auto matches =
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");
69 }
70
71 Tensor::Size num_tokens_{0};
72
73 Tensor::Size num_experts_{0};
74
75 Tensor::Size topk_{0};
76
77 DataType input_dtype_;
78
79 DataType index_dtype_;
80
81 bool renormalize_{false};
82
84
85 private:
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`");
91 assert((input_dtype_ == DataType::kFloat32 ||
92 input_dtype_ == DataType::kFloat16 ||
93 input_dtype_ == DataType::kBFloat16) &&
94 "`TopkSoftmax` supports float32, float16, and bfloat16 input");
95 assert(gating_output.IsContiguous() &&
96 "`TopkSoftmax` requires contiguous `gating_output`");
97 assert(num_experts_ > 0 && topk_ > 0 && topk_ <= num_experts_ &&
98 "`TopkSoftmax` requires `topk` in `[1, num_experts]`");
99 assert(num_tokens_ <=
100 static_cast<Tensor::Size>(std::numeric_limits<int32_t>::max()) &&
101 num_experts_ <=
102 static_cast<Tensor::Size>(std::numeric_limits<int32_t>::max()) &&
103 topk_ <=
104 static_cast<Tensor::Size>(std::numeric_limits<int32_t>::max()) &&
105 "`TopkSoftmax` dimensions must fit int32 indexing");
106 assert(num_tokens_ <=
107 static_cast<Tensor::Size>(std::numeric_limits<int32_t>::max()) /
108 topk_ &&
109 "`TopkSoftmax` output indices must fit int32 indexing");
110
111 const Tensor::Shape output_shape{num_tokens_, topk_};
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`");
118 assert((index_dtype_ == DataType::kInt32 ||
119 index_dtype_ == DataType::kUInt32 ||
120 index_dtype_ == DataType::kInt64) &&
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");
127
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();
131 };
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");
135
136 if (bias) {
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");
143 }
144
145 if (is_padding) {
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");
153 }
154 }
155
156 Tensor gating_output_metadata_;
157
158 std::optional<Tensor> bias_metadata_;
159
160 std::optional<Tensor> is_padding_metadata_;
161
162 Tensor topk_weights_metadata_;
163
164 Tensor topk_indices_metadata_;
165
166 Tensor token_expert_indices_metadata_;
167};
168
169} // namespace infini::ops
170
171#endif // INFINI_OPS_BASE_TOPK_SOFTMAX_H_
Definition generated/include/operator.h:282
Definition topk_softmax.h:15
Tensor::Size num_experts_
Definition topk_softmax.h:73
Tensor::Size num_tokens_
Definition topk_softmax.h:71
bool renormalize_
Definition topk_softmax.h:81
DataType input_dtype_
Definition topk_softmax.h:77
DataType index_dtype_
Definition topk_softmax.h:79
int device_index_
Definition topk_softmax.h:83
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
Tensor::Size topk_
Definition topk_softmax.h:75
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8