InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
moe_align_block_size.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_MOE_ALIGN_BLOCK_SIZE_H_
2#define INFINI_OPS_BASE_MOE_ALIGN_BLOCK_SIZE_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// Align routed token indices into expert-specific blocks following vLLM's
15// low-level `moe_align_block_size` operator.
16class MoeAlignBlockSize : public Operator<MoeAlignBlockSize> {
17 public:
18 MoeAlignBlockSize(const Tensor topk_ids,
19 const std::optional<Tensor> expert_map,
20 const int64_t num_experts, const int64_t block_size,
21 Tensor sorted_token_ids, Tensor experts_ids,
22 Tensor num_tokens_post_pad)
23 : topk_ids_metadata_{topk_ids},
24 expert_map_metadata_{expert_map},
25 sorted_token_ids_metadata_{sorted_token_ids},
26 experts_ids_metadata_{experts_ids},
27 num_tokens_post_pad_metadata_{num_tokens_post_pad},
28 numel_{topk_ids.numel()},
29 num_experts_{num_experts},
30 block_size_{block_size},
31 sorted_token_ids_size_{sorted_token_ids.numel()},
32 experts_ids_size_{experts_ids.numel()} {
33 Validate(topk_ids, expert_map, sorted_token_ids, experts_ids,
34 num_tokens_post_pad);
35 }
36
37 void operator()(const Tensor topk_ids, const std::optional<Tensor> expert_map,
38 const int64_t num_experts, const int64_t block_size,
39 Tensor sorted_token_ids, Tensor experts_ids,
40 Tensor num_tokens_post_pad) const {
41 ValidateInvocation(topk_ids, expert_map, num_experts, block_size,
42 sorted_token_ids, experts_ids, num_tokens_post_pad);
43 Run(topk_ids, expert_map, num_experts, block_size, sorted_token_ids,
44 experts_ids, num_tokens_post_pad);
45 }
46
47 protected:
48 void ValidateInvocation(const Tensor topk_ids,
49 const std::optional<Tensor> maybe_expert_map,
50 const int64_t num_experts, const int64_t block_size,
51 Tensor sorted_token_ids, Tensor experts_ids,
52 Tensor num_tokens_post_pad) const {
53 assert(num_experts == num_experts_ && block_size == block_size_ &&
54 "`MoeAlignBlockSize` attributes changed after descriptor creation");
55
56 assert(CallMetadataMatches(topk_ids, maybe_expert_map, sorted_token_ids,
57 experts_ids, num_tokens_post_pad) &&
58 "`MoeAlignBlockSize` tensor metadata differs from its descriptor");
59 }
60
61 bool CallMetadataMatches(const Tensor topk_ids,
62 const std::optional<Tensor> maybe_expert_map,
63 const Tensor sorted_token_ids,
64 const Tensor experts_ids,
65 const Tensor num_tokens_post_pad) const {
66 const std::equal_to<Tensor> same_metadata;
67 const auto same_expert_map_metadata =
68 expert_map_metadata_.has_value() == maybe_expert_map.has_value() &&
70 same_metadata(*expert_map_metadata_, *maybe_expert_map));
71
72 return same_metadata(topk_ids_metadata_, topk_ids) &&
73 same_expert_map_metadata &&
74 same_metadata(sorted_token_ids_metadata_, sorted_token_ids) &&
75 same_metadata(experts_ids_metadata_, experts_ids) &&
76 same_metadata(num_tokens_post_pad_metadata_, num_tokens_post_pad);
77 }
78
79 void Validate(const Tensor topk_ids,
80 const std::optional<Tensor> maybe_expert_map,
81 Tensor sorted_token_ids, Tensor experts_ids,
82 Tensor num_tokens_post_pad) const {
83 assert(topk_ids.ndim() == 2 &&
84 "`MoeAlignBlockSize` requires 2D `topk_ids`");
85 assert(topk_ids.dtype() == DataType::kInt32 &&
86 "`MoeAlignBlockSize` currently requires int32 `topk_ids`");
87 assert(topk_ids.IsContiguous() &&
88 "`MoeAlignBlockSize` requires contiguous `topk_ids`");
89 assert(num_experts_ > 0 && num_experts_ < 1024 &&
90 "`MoeAlignBlockSize` requires `num_experts` in [1, 1023]");
91 assert(block_size_ > 0 &&
92 "`MoeAlignBlockSize` requires a positive `block_size`");
93 assert(block_size_ <= std::numeric_limits<int32_t>::max() &&
94 numel_ <=
95 static_cast<Tensor::Size>(std::numeric_limits<int32_t>::max()) &&
96 "`MoeAlignBlockSize` requires int32-addressable token indices");
97
98 const auto same_device_as_topk_ids = [&](const Tensor tensor) {
99 return tensor.device().type() == topk_ids.device().type() &&
100 tensor.device().index() == topk_ids.device().index();
101 };
102 assert(same_device_as_topk_ids(sorted_token_ids) &&
103 same_device_as_topk_ids(experts_ids) &&
104 same_device_as_topk_ids(num_tokens_post_pad) &&
105 "`MoeAlignBlockSize` requires all tensors on the same device");
106
107 assert(sorted_token_ids.ndim() == 1 && experts_ids.ndim() == 1 &&
108 num_tokens_post_pad.ndim() == 1 &&
109 num_tokens_post_pad.numel() == 1 &&
110 "`MoeAlignBlockSize` requires 1D output tensors");
111 assert(sorted_token_ids.dtype() == DataType::kInt32 &&
112 experts_ids.dtype() == DataType::kInt32 &&
113 num_tokens_post_pad.dtype() == DataType::kInt32 &&
114 "`MoeAlignBlockSize` requires int32 output tensors");
115 assert(sorted_token_ids.IsContiguous() && experts_ids.IsContiguous() &&
116 num_tokens_post_pad.IsContiguous() &&
117 "`MoeAlignBlockSize` requires contiguous output tensors");
118
119 const auto num_experts = static_cast<Tensor::Size>(num_experts_);
120 const auto block_size = static_cast<Tensor::Size>(block_size_);
121 auto required_sorted_size = numel_ + num_experts * (block_size - 1);
122 if (numel_ < num_experts) {
123 const auto small_input_size = numel_ * block_size;
124 required_sorted_size = small_input_size < required_sorted_size
125 ? small_input_size
126 : required_sorted_size;
127 }
128 const auto required_experts_size =
129 (required_sorted_size + block_size - 1) / block_size;
130 assert(required_sorted_size <=
131 static_cast<Tensor::Size>(std::numeric_limits<int32_t>::max()) &&
132 "`MoeAlignBlockSize` requires int32-addressable padded indices");
133 assert(sorted_token_ids_size_ >= required_sorted_size &&
134 experts_ids_size_ >= required_experts_size &&
135 "`MoeAlignBlockSize` output tensors are too small");
136
137 if (maybe_expert_map) {
138 assert(maybe_expert_map->ndim() == 1 &&
139 maybe_expert_map->numel() ==
140 static_cast<Tensor::Size>(num_experts_) &&
141 "`MoeAlignBlockSize` requires `expert_map` shape "
142 "[`num_experts`]");
143 assert(maybe_expert_map->dtype() == DataType::kInt32 &&
144 "`MoeAlignBlockSize` currently requires int32 `expert_map`");
145 assert(maybe_expert_map->IsContiguous() &&
146 "`MoeAlignBlockSize` requires contiguous `expert_map`");
147 assert(same_device_as_topk_ids(*maybe_expert_map) &&
148 "`MoeAlignBlockSize` requires `expert_map` on the input device");
149 }
150 }
151
152 virtual void Run(const Tensor topk_ids,
153 const std::optional<Tensor> maybe_expert_map,
154 const int64_t num_experts, const int64_t block_size,
155 Tensor sorted_token_ids, Tensor experts_ids,
156 Tensor num_tokens_post_pad) const = 0;
157
159
160 std::optional<Tensor> expert_map_metadata_;
161
163
165
167
168 Tensor::Size numel_{0};
169
170 int64_t num_experts_{0};
171
172 int64_t block_size_{0};
173
174 Tensor::Size sorted_token_ids_size_{0};
175
176 Tensor::Size experts_ids_size_{0};
177};
178
179} // namespace infini::ops
180
181#endif // INFINI_OPS_BASE_MOE_ALIGN_BLOCK_SIZE_H_
Definition moe_align_block_size.h:16
int64_t num_experts_
Definition moe_align_block_size.h:170
Tensor::Size sorted_token_ids_size_
Definition moe_align_block_size.h:174
Tensor::Size numel_
Definition moe_align_block_size.h:168
std::optional< Tensor > expert_map_metadata_
Definition moe_align_block_size.h:160
void Validate(const Tensor topk_ids, const std::optional< Tensor > maybe_expert_map, Tensor sorted_token_ids, Tensor experts_ids, Tensor num_tokens_post_pad) const
Definition moe_align_block_size.h:79
virtual void Run(const Tensor topk_ids, const std::optional< Tensor > maybe_expert_map, const int64_t num_experts, const int64_t block_size, Tensor sorted_token_ids, Tensor experts_ids, Tensor num_tokens_post_pad) const =0
Tensor topk_ids_metadata_
Definition moe_align_block_size.h:158
int64_t block_size_
Definition moe_align_block_size.h:172
Tensor experts_ids_metadata_
Definition moe_align_block_size.h:164
void ValidateInvocation(const Tensor topk_ids, const std::optional< Tensor > maybe_expert_map, const int64_t num_experts, const int64_t block_size, Tensor sorted_token_ids, Tensor experts_ids, Tensor num_tokens_post_pad) const
Definition moe_align_block_size.h:48
MoeAlignBlockSize(const Tensor topk_ids, const std::optional< Tensor > expert_map, const int64_t num_experts, const int64_t block_size, Tensor sorted_token_ids, Tensor experts_ids, Tensor num_tokens_post_pad)
Definition moe_align_block_size.h:18
void operator()(const Tensor topk_ids, const std::optional< Tensor > expert_map, const int64_t num_experts, const int64_t block_size, Tensor sorted_token_ids, Tensor experts_ids, Tensor num_tokens_post_pad) const
Definition moe_align_block_size.h:37
Tensor::Size experts_ids_size_
Definition moe_align_block_size.h:176
Tensor num_tokens_post_pad_metadata_
Definition moe_align_block_size.h:166
bool CallMetadataMatches(const Tensor topk_ids, const std::optional< Tensor > maybe_expert_map, const Tensor sorted_token_ids, const Tensor experts_ids, const Tensor num_tokens_post_pad) const
Definition moe_align_block_size.h:61
Tensor sorted_token_ids_metadata_
Definition moe_align_block_size.h:162
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8