InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
rotary_embedding_infinilm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_ROTARY_EMBEDDING_INFINILM_H_
2#define INFINI_OPS_BASE_ROTARY_EMBEDDING_INFINILM_H_
3
4#include <cassert>
5#include <cstddef>
6
7#include "data_type.h"
8#include "operator.h"
9#include "tensor.h"
10
11namespace infini::ops {
12
15class [[deprecated("Use `RotaryEmbedding` instead.")]] RotaryEmbeddingInfinilm
16 : public Operator<RotaryEmbeddingInfinilm> {
17 public:
18 RotaryEmbeddingInfinilm(const Tensor input, const Tensor pos_ids,
19 const Tensor sin_table, const Tensor cos_table,
20 bool is_neox, Tensor out)
21 : ndim_{out.ndim()},
22 batch_size_{ndim_ == 4 ? out.size(-4) : 1},
23 seq_len_{out.size(-3)},
24 nhead_{out.size(-2)},
25 table_dim_{sin_table.size(1)},
26 has_batch_dim_{ndim_ == 4},
27 pos_has_batch_dim_{pos_ids.ndim() == 2},
28 input_strides_{input.strides()},
29 out_strides_{out.strides()},
30 pos_strides_{pos_ids.strides()} {
31 const auto head_dim = out.size(-1);
32 const auto table_len = sin_table.size(0);
33 const auto angle_dtype = sin_table.dtype();
34 const auto pos_dtype = pos_ids.dtype();
35
36 assert(input.shape() == out.shape() &&
37 "`RotaryEmbeddingInfinilm` requires `input` and `out` same shape");
38 assert(input.dtype() == out.dtype() &&
39 "`RotaryEmbeddingInfinilm` requires `input` and `out` same dtype");
40 assert((ndim_ == 3 || ndim_ == 4) &&
41 "`RotaryEmbeddingInfinilm` requires 3D or 4D tensor");
42 assert(head_dim % 2 == 0 &&
43 "`RotaryEmbeddingInfinilm` requires head dimension to be even");
44 assert(
45 head_dim == table_dim_ * 2 &&
46 "`RotaryEmbeddingInfinilm` requires table dim to be half of head dim");
47 assert(pos_ids.ndim() == 1 || pos_ids.ndim() == 2);
48 assert((pos_dtype == DataType::kInt32 || pos_dtype == DataType::kInt64) &&
49 "`RotaryEmbeddingInfinilm` requires int32 or int64 position ids");
50 assert(sin_table.shape() == cos_table.shape() &&
51 "`RotaryEmbeddingInfinilm` requires sin_table and cos_table same "
52 "shape");
53 assert(sin_table.dtype() == cos_table.dtype() &&
54 "`RotaryEmbeddingInfinilm` requires sin_table and cos_table same "
55 "dtype");
56 assert((angle_dtype == DataType::kFloat16 ||
57 angle_dtype == DataType::kBFloat16 ||
58 angle_dtype == DataType::kFloat32) &&
59 "`RotaryEmbeddingInfinilm` requires float sin/cos tables");
60 assert(sin_table.ndim() == 2 && cos_table.ndim() == 2 &&
61 "`RotaryEmbeddingInfinilm` requires 2D sin/cos tables");
62 assert(
63 table_len >= seq_len_ &&
64 "`RotaryEmbeddingInfinilm` requires table length >= sequence length");
65 assert((pos_has_batch_dim_ ? (pos_ids.size(0) == batch_size_ &&
66 pos_ids.size(1) == seq_len_)
67 : (pos_ids.size(0) == seq_len_)) &&
68 "`RotaryEmbeddingInfinilm` requires pos_ids shape [seq] or [batch, "
69 "seq]");
70 assert(out_strides_[ndim_ - 1] == 1 && input_strides_[ndim_ - 1] == 1 &&
71 "`RotaryEmbeddingInfinilm` requires contiguous head dimension");
72 assert(sin_table.strides()[1] == 1 && cos_table.strides()[1] == 1 &&
73 "`RotaryEmbeddingInfinilm` requires contiguous table dimension");
74 }
75
76 virtual void operator()(const Tensor input, const Tensor pos_ids,
77 const Tensor sin_table, const Tensor cos_table,
78 bool is_neox, Tensor out) const = 0;
79
80 protected:
81 Tensor::Size ndim_{0};
82
83 Tensor::Size batch_size_{0};
84
85 Tensor::Size seq_len_{0};
86
87 Tensor::Size nhead_{0};
88
89 Tensor::Size table_dim_{0};
90
91 bool has_batch_dim_{false};
92
93 bool pos_has_batch_dim_{false};
94
95 Tensor::Strides input_strides_;
96
97 Tensor::Strides out_strides_;
98
99 Tensor::Strides pos_strides_;
100};
101
102} // namespace infini::ops
103
104#endif
Definition generated/include/operator.h:282
Definition rotary_embedding_infinilm.h:16
Tensor::Strides out_strides_
Definition rotary_embedding_infinilm.h:97
Tensor::Strides input_strides_
Definition rotary_embedding_infinilm.h:95
Tensor::Strides pos_strides_
Definition rotary_embedding_infinilm.h:99
virtual void operator()(const Tensor input, const Tensor pos_ids, const Tensor sin_table, const Tensor cos_table, bool is_neox, Tensor out) const =0
RotaryEmbeddingInfinilm(const Tensor input, const Tensor pos_ids, const Tensor sin_table, const Tensor cos_table, bool is_neox, Tensor out)
Definition rotary_embedding_infinilm.h:18
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8