1#ifndef INFINI_OPS_BASE_GROUPED_TOPK_H_
2#define INFINI_OPS_BASE_GROUPED_TOPK_H_
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,
21 :
num_tokens_{scores.ndim() == 2 ? scores.size(0) : 0},
32 scores_metadata_{scores},
34 topk_values_metadata_{topk_values},
35 topk_indices_metadata_{topk_indices} {
36 Validate(scores, bias, topk_values, topk_indices);
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;
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 {
59 "`GroupedTopk` attributes changed after descriptor creation");
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");
93 Tensor topk_indices)
const {
94 assert(scores.ndim() == 2 &&
"`GroupedTopk` requires 2D `scores`");
98 "`GroupedTopk` supports float16, bfloat16, and float32 `scores`");
99 assert(scores.IsContiguous() &&
100 "`GroupedTopk` requires contiguous `scores`");
102 static_cast<Tensor::Size
>(std::numeric_limits<int32_t>::max()) &&
104 static_cast<Tensor::Size
>(std::numeric_limits<int32_t>::max()) &&
105 "`GroupedTopk` dimensions must fit int32 indexing");
107 assert(bias.ndim() == 1 && bias.numel() ==
num_experts_ &&
108 "`GroupedTopk` requires `bias` shape `[num_experts]`");
112 "`GroupedTopk` supports float16, bfloat16, and float32 `bias`");
113 assert(bias.IsContiguous() &&
"`GroupedTopk` requires contiguous `bias`");
116 "`GroupedTopk` requires `num_expert_group` in `[1, 32]`");
118 "`GroupedTopk` requires `topk_group` in `[1, num_expert_group]`");
121 "`GroupedTopk` requires experts divisible by `num_expert_group`");
123 "`GroupedTopk` requires at least two experts per group");
127 "`GroupedTopk` requires `topk` in the selected-group capacity");
129 "`GroupedTopk` requires `scoring_func` 0 (none) or 1 (sigmoid)");
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");
143 const auto same_device = [&](
const Tensor tensor) {
144 return tensor.device().type() == scores.device().type() &&
145 tensor.device().index() == scores.device().index();
147 assert(same_device(bias) && same_device(topk_values) &&
148 same_device(topk_indices) &&
149 "`GroupedTopk` requires all tensors on the same device");
156 Tensor topk_values_metadata_;
158 Tensor topk_indices_metadata_;
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