Skip to main content

ClikaRT::nn::Attention

class

Header: ClikaRT/nn/attention.h
Inherits: ClikaRT::nn::Module

Attention module: fixed settings, optional bound tensors, one fused call per forward. The module form of ops::attention.

Inputs are head-major with the head dimension innermost: query [B, H_q, S_q, D], key [B, H_kv, S_kv, D], value [B, H_kv, S_kv, Dv] with H_kv dividing H_q (grouped-query attention broadcasts each key / value head over its query group). With num_heads set the inputs are the hidden-folded projection outputs instead, [B, S, H * D], and the result is folded back the same way. Causal attention needs S_q <= S_kv.

auto attn = ClikaRT::nn::Attention::make(
ClikaRT::nn::AttentionOptions().is_causal(true).sliding_window(4096));
auto out = attn->forward(q, k, v); // [B, H_q, S_q, Dv]

Static member functions​

make()​

static std::shared_ptr<Attention> make(AttentionOptions options = {})

Defaults: every option at its default (AttentionOptions).

Declared in ClikaRT/nn/attention.h, line 122

Member functions​

~Attention()​

~Attention() override

Declared in ClikaRT/nn/attention.h, line 126

set_mask()​

void set_mask(Tensor mask)

Bind an attention mask every forward applies (the buffer attn_mask): broadcastable to [B, H_q, S_q, S_kv]; a Bool tensor keeps where true, a float tensor adds to the logits. A mask passed to forward takes precedence for that call. Re-binding replaces the buffer.

Declared in ClikaRT/nn/attention.h, line 134

set_sinks()​

void set_sinks(Tensor sinks)

Bind the learned per-head softmax sink (the parameter sinks, [H_q]): one virtual logit per head folded into every softmax denominator, so a query can attend to nothing. Binds the declared slot when sinks was set at make (the shape must match), or registers the parameter otherwise.

Declared in ClikaRT/nn/attention.h, line 142

set_rope()​

void set_rope(
    Tensor cos,
    Tensor sin,
    ops::RotaryMode mode = ops::RotaryMode::NeoX
)

Bind rotary tables (the buffers rope_cos / rope_sin, each [max_positions, rotary_dim / 2] Float32, the tables RotaryEmbedding::cos() / sin() or ops::generate_rotary_cache produce): every forward then rotates the queries AND the keys by position (position_ids, else 0, 1, 2, ...) before the dot product, so the two share one sequence length: a call with S_q != S_kv is refused while tables are bound. mode is the pair layout.

Declared in ClikaRT/nn/attention.h, line 154

to_impl(StreamOrDevice)​

virtual Result<void> to_impl(StreamOrDevice where) override

Move the bound tensors to a placement (Device / Stream); the attention is rebuilt there on the next forward.

Declared in ClikaRT/nn/attention.h, line 161

to_impl(DataType)​

virtual Result<void> to_impl(DataType dtype) override

Cast the bound sinks vector and a FLOAT mask to dtype; a Bool mask and the rotary tables (always Float32) keep their dtypes.

Declared in ClikaRT/nn/attention.h, line 164

forward()​

Tensor forward(
    Tensor query,
    Tensor key,
    Tensor value,
    OptionalTensor mask = {},
    OptionalTensor position_ids = {}
) const

forward_impl, unwrapped: raises ClikaRT::Error on failure.

Declared in ClikaRT/nn/attention.h, line 183

operator()()​

Tensor operator()(
    Tensor query,
    Tensor key,
    Tensor value,
    OptionalTensor mask = {},
    OptionalTensor position_ids = {}
) const

Same as forward.

Declared in ClikaRT/nn/attention.h, line 190

options()​

const AttentionOptions& options() const noexcept

The settings this module was made with.

Declared in ClikaRT/nn/attention.h, line 198

Protected member functions​

initialize_impl()​

virtual Result<void> initialize_impl() override

The pack hook: prepares the fused attention (idempotent, thread-safe). A declared sinks slot must be bound first.

Declared in ClikaRT/nn/attention.h, line 208