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