Skip to main content

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)​

virtual Result<void> to_impl(DataType dtype) override

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