InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
grouped_topk.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_GROUPED_TOPK_H_
2#define INFINI_OPS_BASE_GROUPED_TOPK_H_
3
4#include <cassert>
5#include <cstdint>
6#include <functional>
7#include <limits>
8
9#include "operator.h"
10
11namespace infini::ops {
12
13// Aligned with vLLM's low-level `_moe_C::grouped_topk` operator.
14class GroupedTopk : public Operator<GroupedTopk> {
15 public:
16 GroupedTopk(const Tensor scores, const Tensor bias,
17 const int64_t num_expert_group, const int64_t topk_group,
18 const int64_t topk, const bool renormalize,
19 const double routed_scaling_factor, const int64_t scoring_func,
20 Tensor topk_values, Tensor topk_indices)
21 : num_tokens_{scores.ndim() == 2 ? scores.size(0) : 0},
22 num_experts_{scores.ndim() == 2 ? scores.size(1) : 0},
23 num_expert_group_{num_expert_group},
24 topk_group_{topk_group},
25 topk_{topk},
26 renormalize_{renormalize},
27 routed_scaling_factor_{routed_scaling_factor},
28 scoring_func_{scoring_func},
29 scores_dtype_{scores.dtype()},
30 bias_dtype_{bias.dtype()},
31 device_index_{scores.device().index()},
32 scores_metadata_{scores},
33 bias_metadata_{bias},
34 topk_values_metadata_{topk_values},
35 topk_indices_metadata_{topk_indices} {
36 Validate(scores, bias, topk_values, topk_indices);
37 }
38
39 virtual void operator()(const Tensor scores, const Tensor bias,
40 const int64_t num_expert_group,
41 const int64_t topk_group, const int64_t topk,
42 const bool renormalize,
43 const double routed_scaling_factor,
44 const int64_t scoring_func, Tensor topk_values,
45 Tensor topk_indices) const = 0;
46
47 protected:
48 void ValidateCallMetadata(const Tensor scores, const Tensor bias,
49 const int64_t num_expert_group,
50 const int64_t topk_group, const int64_t topk,
51 const bool renormalize,
52 const double routed_scaling_factor,
53 const int64_t scoring_func, Tensor topk_values,
54 Tensor topk_indices) const {
55 assert(num_expert_group == num_expert_group_ && topk_group == topk_group_ &&
56 topk == topk_ && renormalize == renormalize_ &&
57 routed_scaling_factor == routed_scaling_factor_ &&
58 scoring_func == scoring_func_ &&
59 "`GroupedTopk` attributes changed after descriptor creation");
60
61 const std::equal_to<Tensor> same_metadata;
62 const auto matches = same_metadata(scores_metadata_, scores) &&
63 same_metadata(bias_metadata_, bias) &&
64 same_metadata(topk_values_metadata_, topk_values) &&
65 same_metadata(topk_indices_metadata_, topk_indices);
66 assert(matches && "`GroupedTopk` call metadata must match descriptor");
67 }
68
69 Tensor::Size num_tokens_{0};
70
71 Tensor::Size num_experts_{0};
72
74
75 int64_t topk_group_{0};
76
77 int64_t topk_{0};
78
79 bool renormalize_{false};
80
82
83 int64_t scoring_func_{0};
84
85 DataType scores_dtype_;
86
87 DataType bias_dtype_;
88
90
91 private:
92 void Validate(const Tensor scores, const Tensor bias, Tensor topk_values,
93 Tensor topk_indices) const {
94 assert(scores.ndim() == 2 && "`GroupedTopk` requires 2D `scores`");
95 assert((scores_dtype_ == DataType::kFloat16 ||
96 scores_dtype_ == DataType::kBFloat16 ||
97 scores_dtype_ == DataType::kFloat32) &&
98 "`GroupedTopk` supports float16, bfloat16, and float32 `scores`");
99 assert(scores.IsContiguous() &&
100 "`GroupedTopk` requires contiguous `scores`");
101 assert(num_tokens_ <=
102 static_cast<Tensor::Size>(std::numeric_limits<int32_t>::max()) &&
103 num_experts_ <=
104 static_cast<Tensor::Size>(std::numeric_limits<int32_t>::max()) &&
105 "`GroupedTopk` dimensions must fit int32 indexing");
106
107 assert(bias.ndim() == 1 && bias.numel() == num_experts_ &&
108 "`GroupedTopk` requires `bias` shape `[num_experts]`");
109 assert((bias_dtype_ == DataType::kFloat16 ||
110 bias_dtype_ == DataType::kBFloat16 ||
111 bias_dtype_ == DataType::kFloat32) &&
112 "`GroupedTopk` supports float16, bfloat16, and float32 `bias`");
113 assert(bias.IsContiguous() && "`GroupedTopk` requires contiguous `bias`");
114
115 assert(num_expert_group_ > 0 && num_expert_group_ <= 32 &&
116 "`GroupedTopk` requires `num_expert_group` in `[1, 32]`");
117 assert(topk_group_ > 0 && topk_group_ <= num_expert_group_ &&
118 "`GroupedTopk` requires `topk_group` in `[1, num_expert_group]`");
119 assert(num_experts_ > 0 &&
120 num_experts_ % static_cast<Tensor::Size>(num_expert_group_) == 0 &&
121 "`GroupedTopk` requires experts divisible by `num_expert_group`");
122 assert(num_experts_ / static_cast<Tensor::Size>(num_expert_group_) >= 2 &&
123 "`GroupedTopk` requires at least two experts per group");
124 assert(topk_ > 0 && topk_ <= 32 &&
125 topk_ <= topk_group_ * static_cast<int64_t>(num_experts_ /
127 "`GroupedTopk` requires `topk` in the selected-group capacity");
128 assert((scoring_func_ == 0 || scoring_func_ == 1) &&
129 "`GroupedTopk` requires `scoring_func` 0 (none) or 1 (sigmoid)");
130
131 const Tensor::Shape output_shape{num_tokens_,
132 static_cast<Tensor::Size>(topk_)};
133 assert(topk_values.shape() == output_shape &&
134 topk_indices.shape() == output_shape &&
135 "`GroupedTopk` outputs must have shape `[num_tokens, topk]`");
136 assert(topk_values.dtype() == DataType::kFloat32 &&
137 "`GroupedTopk` requires float32 `topk_values`");
138 assert(topk_indices.dtype() == DataType::kInt32 &&
139 "`GroupedTopk` requires int32 `topk_indices`");
140 assert(topk_values.IsContiguous() && topk_indices.IsContiguous() &&
141 "`GroupedTopk` requires contiguous outputs");
142
143 const auto same_device = [&](const Tensor tensor) {
144 return tensor.device().type() == scores.device().type() &&
145 tensor.device().index() == scores.device().index();
146 };
147 assert(same_device(bias) && same_device(topk_values) &&
148 same_device(topk_indices) &&
149 "`GroupedTopk` requires all tensors on the same device");
150 }
151
152 Tensor scores_metadata_;
153
154 Tensor bias_metadata_;
155
156 Tensor topk_values_metadata_;
157
158 Tensor topk_indices_metadata_;
159};
160
161} // namespace infini::ops
162
163#endif // INFINI_OPS_BASE_GROUPED_TOPK_H_
Definition grouped_topk.h:14
int64_t topk_
Definition grouped_topk.h:77
DataType bias_dtype_
Definition grouped_topk.h:87
double routed_scaling_factor_
Definition grouped_topk.h:81
int64_t num_expert_group_
Definition grouped_topk.h:73
Tensor::Size num_experts_
Definition grouped_topk.h:71
DataType scores_dtype_
Definition grouped_topk.h:85
virtual void operator()(const Tensor scores, const Tensor bias, const int64_t num_expert_group, const int64_t topk_group, const int64_t topk, const bool renormalize, const double routed_scaling_factor, const int64_t scoring_func, Tensor topk_values, Tensor topk_indices) const =0
int64_t topk_group_
Definition grouped_topk.h:75
Tensor::Size num_tokens_
Definition grouped_topk.h:69
bool renormalize_
Definition grouped_topk.h:79
int64_t scoring_func_
Definition grouped_topk.h:83
void ValidateCallMetadata(const Tensor scores, const Tensor bias, const int64_t num_expert_group, const int64_t topk_group, const int64_t topk, const bool renormalize, const double routed_scaling_factor, const int64_t scoring_func, Tensor topk_values, Tensor topk_indices) const
Definition grouped_topk.h:48
GroupedTopk(const Tensor scores, const Tensor bias, const int64_t num_expert_group, const int64_t topk_group, const int64_t topk, const bool renormalize, const double routed_scaling_factor, const int64_t scoring_func, Tensor topk_values, Tensor topk_indices)
Definition grouped_topk.h:16
int device_index_
Definition grouped_topk.h:89
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8