InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
rotary_embedding.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_ROTARY_EMBEDDING_H_
2#define INFINI_OPS_BASE_ROTARY_EMBEDDING_H_
3
4#include <optional>
5
6#include "data_type.h"
7#include "operator.h"
8
9namespace infini::ops {
10
11class RotaryEmbedding : public Operator<RotaryEmbedding> {
12 public:
13 RotaryEmbedding(const Tensor positions, Tensor query,
14 std::optional<Tensor> key, const Tensor cos_sin_cache,
15 int64_t head_size, bool is_neox, int64_t rope_dim_offset = 0,
16 bool inverse = false)
17 : positions_shape_{positions.shape()},
18 query_shape_{query.shape()},
19 key_shape_{key.has_value() ? Tensor::Shape{key->shape()}
20 : Tensor::Shape{}},
21 cos_sin_cache_shape_{cos_sin_cache.shape()},
22 positions_strides_{positions.strides()},
23 query_strides_{query.strides()},
24 key_strides_{key.has_value() ? Tensor::Strides{key->strides()}
25 : Tensor::Strides{}},
26 cos_sin_cache_strides_{cos_sin_cache.strides()},
27 positions_type_{positions.dtype()},
28 query_type_{query.dtype()},
29 cos_sin_cache_type_{cos_sin_cache.dtype()},
30 num_tokens_{positions.numel()},
31 positions_ndim_{positions.ndim()},
32 head_size_{head_size},
33 rot_dim_{static_cast<int64_t>(cos_sin_cache.size(1))},
34 rope_dim_offset_{rope_dim_offset},
35 is_neox_{is_neox},
36 inverse_{inverse},
37 query_hidden_size_{num_tokens_ == 0 ? 0 : query.numel() / num_tokens_},
38 key_hidden_size_{key.has_value() && num_tokens_ != 0
39 ? key->numel() / num_tokens_
40 : 0},
43 ? 0
44 : (key.has_value() ? key_hidden_size_ / head_size_
45 : num_heads_)},
46 query_token_stride_{query.stride(positions_ndim_ - 1)},
47 key_token_stride_{key.has_value() ? key->stride(positions_ndim_ - 1)
48 : 0},
49 query_head_stride_{query.ndim() == positions_ndim_ + 2
50 ? query.stride(-2)
51 : head_size_},
52 key_head_stride_{key.has_value() && key->ndim() == positions_ndim_ + 2
53 ? key->stride(-2)
54 : head_size_},
55 device_index_{query.device().index()} {
56 assert((positions_ndim_ == 1 || positions_ndim_ == 2) &&
57 "`RotaryEmbedding` requires 1D or 2D `positions`");
58 assert(positions_type_ == DataType::kInt64 &&
59 "`RotaryEmbedding` requires int64 `positions`");
60 assert(positions.stride(-1) == 1 &&
61 (positions_ndim_ == 1 || positions.stride(0) == positions.size(1)) &&
62 "`RotaryEmbedding` requires contiguous `positions`");
63 assert((query.ndim() == positions_ndim_ + 1 ||
64 query.ndim() == positions_ndim_ + 2) &&
65 "`RotaryEmbedding` received unsupported `query` rank");
66 assert(head_size_ > 0 && query_hidden_size_ % head_size_ == 0 &&
67 "`RotaryEmbedding` requires query hidden size divisible by "
68 "`head_size`");
69 assert((query_type_ == DataType::kFloat16 ||
70 query_type_ == DataType::kBFloat16 ||
71 query_type_ == DataType::kFloat32) &&
72 "`RotaryEmbedding` supports float16, bfloat16, and float32 query");
73 assert(cos_sin_cache.ndim() == 2 && rot_dim_ % 2 == 0 &&
74 cos_sin_cache.stride(1) == 1 &&
75 "`RotaryEmbedding` requires contiguous 2D `cos_sin_cache` with "
76 "even rotary dimension");
77 assert((cos_sin_cache_type_ == DataType::kFloat16 ||
78 cos_sin_cache_type_ == DataType::kBFloat16 ||
79 cos_sin_cache_type_ == DataType::kFloat32) &&
80 "`RotaryEmbedding` supports float16, bfloat16, and float32 cache");
82 "`RotaryEmbedding` rotary dimensions exceed `head_size`");
83 assert(query.stride(-1) == 1 &&
84 "`RotaryEmbedding` requires contiguous query head dimensions");
85
86 if (positions_ndim_ == 1) {
87 assert(query.size(0) == positions.size(0) &&
88 (!key.has_value() || key->size(0) == positions.size(0)) &&
89 "`RotaryEmbedding` requires matching token counts");
90 } else {
91 assert(query.size(0) == positions.size(0) &&
92 query.size(1) == positions.size(1) &&
93 (!key.has_value() || (key->size(0) == positions.size(0) &&
94 key->size(1) == positions.size(1))) &&
95 "`RotaryEmbedding` requires matching batch and sequence sizes");
96 }
97
98 if (query.ndim() == positions_ndim_ + 2) {
99 assert(query.size(-1) == static_cast<Tensor::Size>(head_size_) &&
100 "`RotaryEmbedding` query head dimension does not match "
101 "`head_size`");
102 }
103
104 if (key.has_value()) {
105 assert((key->ndim() == positions_ndim_ + 1 ||
106 key->ndim() == positions_ndim_ + 2) &&
108 key->dtype() == query_type_ && key->stride(-1) == 1 &&
109 "`RotaryEmbedding` key layout or dtype is incompatible with "
110 "query");
111 assert(num_kv_heads_ > 0 && num_heads_ % num_kv_heads_ == 0 &&
112 "`RotaryEmbedding` requires query heads divisible by key heads");
113 if (key->ndim() == positions_ndim_ + 2) {
114 assert(key->size(-1) == static_cast<Tensor::Size>(head_size_) &&
115 "`RotaryEmbedding` key head dimension does not match "
116 "`head_size`");
117 }
118 }
119 }
120
121 virtual void operator()(const Tensor positions, Tensor query,
122 std::optional<Tensor> key, const Tensor cos_sin_cache,
123 int64_t head_size, bool is_neox,
124 int64_t rope_dim_offset = 0,
125 bool inverse = false) const = 0;
126
127 protected:
128 Tensor::Shape positions_shape_;
129
130 Tensor::Shape query_shape_;
131
132 Tensor::Shape key_shape_;
133
134 Tensor::Shape cos_sin_cache_shape_;
135
136 Tensor::Strides positions_strides_;
137
138 Tensor::Strides query_strides_;
139
140 Tensor::Strides key_strides_;
141
142 Tensor::Strides cos_sin_cache_strides_;
143
145
146 DataType query_type_;
147
149
150 Tensor::Size num_tokens_{0};
151
152 Tensor::Size positions_ndim_{0};
153
154 int64_t head_size_{0};
155
156 int64_t rot_dim_{0};
157
159
160 bool is_neox_{false};
161
162 bool inverse_{false};
163
164 Tensor::Size query_hidden_size_{0};
165
166 Tensor::Size key_hidden_size_{0};
167
168 Tensor::Size num_heads_{0};
169
170 Tensor::Size num_kv_heads_{0};
171
173
175
177
179
181};
182
183} // namespace infini::ops
184
185#endif
Definition generated/include/operator.h:282
Definition rotary_embedding.h:11
int64_t rot_dim_
Definition rotary_embedding.h:156
DataType query_type_
Definition rotary_embedding.h:146
Tensor::Strides query_strides_
Definition rotary_embedding.h:138
int64_t key_token_stride_
Definition rotary_embedding.h:174
int device_index_
Definition rotary_embedding.h:180
Tensor::Shape key_shape_
Definition rotary_embedding.h:132
Tensor::Strides positions_strides_
Definition rotary_embedding.h:136
int64_t key_head_stride_
Definition rotary_embedding.h:178
Tensor::Size num_kv_heads_
Definition rotary_embedding.h:170
Tensor::Size positions_ndim_
Definition rotary_embedding.h:152
bool is_neox_
Definition rotary_embedding.h:160
virtual void operator()(const Tensor positions, Tensor query, std::optional< Tensor > key, const Tensor cos_sin_cache, int64_t head_size, bool is_neox, int64_t rope_dim_offset=0, bool inverse=false) const =0
int64_t rope_dim_offset_
Definition rotary_embedding.h:158
int64_t query_token_stride_
Definition rotary_embedding.h:172
Tensor::Strides key_strides_
Definition rotary_embedding.h:140
int64_t query_head_stride_
Definition rotary_embedding.h:176
int64_t head_size_
Definition rotary_embedding.h:154
Tensor::Strides cos_sin_cache_strides_
Definition rotary_embedding.h:142
Tensor::Size query_hidden_size_
Definition rotary_embedding.h:164
Tensor::Size num_heads_
Definition rotary_embedding.h:168
Tensor::Shape cos_sin_cache_shape_
Definition rotary_embedding.h:134
Tensor::Size key_hidden_size_
Definition rotary_embedding.h:166
Tensor::Shape query_shape_
Definition rotary_embedding.h:130
Tensor::Shape positions_shape_
Definition rotary_embedding.h:128
DataType cos_sin_cache_type_
Definition rotary_embedding.h:148
Tensor::Size num_tokens_
Definition rotary_embedding.h:150
bool inverse_
Definition rotary_embedding.h:162
RotaryEmbedding(const Tensor positions, Tensor query, std::optional< Tensor > key, const Tensor cos_sin_cache, int64_t head_size, bool is_neox, int64_t rope_dim_offset=0, bool inverse=false)
Definition rotary_embedding.h:13
DataType positions_type_
Definition rotary_embedding.h:144
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8