ClikaRT::nn::RotaryEmbedding
class
Header: ClikaRT/nn/rotary_embedding.h
Inherits: ClikaRT::nn::Module
Rotary position embedding module: the angle tables built once at make, the rotation applied per call. The module form of ops::rotary_embedding.
forward(input, position_ids) rotates the leading dim entries of every head vector of input ([.., S, D] with D >= dim) by the angle of the token's position: row position_ids[.., s] of the tables, or position s when position_ids is absent. The output keeps the input's shape and dtype. forward_qk rotates a query / key pair with one angle gather.
auto rope = ClikaRT::nn::RotaryEmbedding::make(/*dim=*/128, /*max_positions=*/4096);
auto [q_rot, k_rot] = rope->forward_qk(q, k, position_ids);
Static member functions
make()
static std::shared_ptr<RotaryEmbedding> make(
std::int64_t dim,
std::int64_t max_positions,
RotaryEmbeddingOptions options = {}
)
Defaults: every option at its default (RotaryEmbeddingOptions).
Declared in ClikaRT/nn/rotary_embedding.h, line 130
Member functions
~RotaryEmbedding()
~RotaryEmbedding() override
Declared in ClikaRT/nn/rotary_embedding.h, line 135
to_impl(StreamOrDevice)
virtual Result<void> to_impl(StreamOrDevice where) override
Move the tables to a placement (Device / Stream); the rotation is rebuilt there on the next forward.
Declared in ClikaRT/nn/rotary_embedding.h, line 139
to_impl(DataType)
A dtype move leaves the tables at Float32 (the rotation's output follows its input's dtype), so a model-wide cast walks through this module without effect and returns success.
Declared in ClikaRT/nn/rotary_embedding.h, line 143
forward()
Tensor forward(Tensor input, OptionalTensor position_ids = {}) const
forward_impl, unwrapped: raises ClikaRT::Error on failure.
Declared in ClikaRT/nn/rotary_embedding.h, line 153
operator()()
Tensor operator()(Tensor input, OptionalTensor position_ids = {}) const
Same as forward.
Declared in ClikaRT/nn/rotary_embedding.h, line 157
forward_qk()
std::array<Tensor 2> forward_qk(
Tensor query,
Tensor key,
OptionalTensor position_ids = {}
) const
forward_qk_impl, unwrapped: raises ClikaRT::Error on failure.
Declared in ClikaRT/nn/rotary_embedding.h, line 167
cos()
Tensor cos() const
The cosine table, [max_positions, dim / 2] Float32 (the cos buffer).
Declared in ClikaRT/nn/rotary_embedding.h, line 173
sin()
Tensor sin() const
The sine table, [max_positions, dim / 2] Float32 (the sin buffer).
Declared in ClikaRT/nn/rotary_embedding.h, line 175
dim()
std::int64_t dim() const noexcept
The rotated head prefix.
Declared in ClikaRT/nn/rotary_embedding.h, line 177
max_positions()
std::int64_t max_positions() const noexcept
The number of positions the tables cover.
Declared in ClikaRT/nn/rotary_embedding.h, line 179
options()
const RotaryEmbeddingOptions& options() const noexcept
The settings this module was made with.
Declared in ClikaRT/nn/rotary_embedding.h, line 181
Protected member functions
initialize_impl()
virtual Result<void> initialize_impl() override
The pack hook: prepares the rotation (idempotent, thread-safe).
Declared in ClikaRT/nn/rotary_embedding.h, line 190