22 batch_size_{ndim_ == 4 ? out.size(-4) : 1},
23 seq_len_{out.size(-3)},
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();
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");
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 "
53 assert(sin_table.dtype() == cos_table.dtype() &&
54 "`RotaryEmbeddingInfinilm` requires sin_table and cos_table same "
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");
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, "
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");