1#ifndef INFINI_OPS_BASE_ROTARY_EMBEDDING_H_
2#define INFINI_OPS_BASE_ROTARY_EMBEDDING_H_
14 std::optional<Tensor> key,
const Tensor cos_sin_cache,
15 int64_t head_size,
bool is_neox, int64_t rope_dim_offset = 0,
33 rot_dim_{static_cast<int64_t>(cos_sin_cache.size(1))},
57 "`RotaryEmbedding` requires 1D or 2D `positions`");
59 "`RotaryEmbedding` requires int64 `positions`");
60 assert(positions.stride(-1) == 1 &&
62 "`RotaryEmbedding` requires contiguous `positions`");
65 "`RotaryEmbedding` received unsupported `query` rank");
67 "`RotaryEmbedding` requires query hidden size divisible by "
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");
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");
87 assert(query.size(0) == positions.size(0) &&
88 (!key.has_value() || key->size(0) == positions.size(0)) &&
89 "`RotaryEmbedding` requires matching token counts");
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");
99 assert(query.size(-1) ==
static_cast<Tensor::Size
>(
head_size_) &&
100 "`RotaryEmbedding` query head dimension does not match "
104 if (key.has_value()) {
108 key->dtype() ==
query_type_ && key->stride(-1) == 1 &&
109 "`RotaryEmbedding` key layout or dtype is incompatible with "
112 "`RotaryEmbedding` requires query heads divisible by key heads");
114 assert(key->size(-1) ==
static_cast<Tensor::Size
>(
head_size_) &&
115 "`RotaryEmbedding` key head dimension does not match "
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;
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