use alloc::vec;
use alloc::vec::Vec;
use ruda_model::{config::Config, module::Module,
tensor::{Bool, DType, Int, Tensor, backend::Backend}};
use crate::{Linear, LinearConfig, LayerNorm, LayerNormConfig};
use super::{SparseRotaryEmbedding, indexer_kl_loss, sparse_gather_entries, sparse_stable_topk};
use super::sparse_ops::{positions, work_dtype};
use super::CompressedAttentionProjection;
#[cfg(not(feature = "std"))]
#[allow(unused_imports)]
use num_traits::Float as _;
#[derive(Config, Debug)]
pub struct LightningIndexerConfig {
pub width: usize,
#[config(default = 4)]
pub num_heads: usize,
#[config(default = 16)]
pub head_dim: usize,
#[config(default = "None")]
pub query_dim: Option<usize>,
#[config(default = 32)]
pub topk: usize,
#[config(default = 32)]
pub query_chunk_size: usize,
#[config(default = 128)]
pub key_chunk_size: usize,
#[config(default = false)]
pub external_keys: bool,
#[config(default = true)]
pub detach_inputs: bool,
#[config(default = 0)]
pub rope_dim: usize,
#[config(default = 10000.0)]
pub rope_base: f64,
#[config(default = 1e-6)]
pub epsilon: f64,
}
#[derive(Module, Debug)]
pub struct LightningIndexer<B: Backend, P: Module<B> = Linear<B>> {
pub query: P,
pub head_weight: P,
pub key: Option<P>,
pub key_norm: Option<LayerNorm<B>>,
pub rotary: SparseRotaryEmbedding,
pub width: usize,
pub query_dim: usize,
pub num_heads: usize,
pub head_dim: usize,
pub topk: usize,
pub query_chunk_size: usize,
pub key_chunk_size: usize,
pub detach_inputs: bool,
}
pub type DSAIndexer<B, P = Linear<B>> = LightningIndexer<B, P>;
#[derive(Clone, Debug)]
pub struct IndexerMask<B: Backend> {
pub allowed: Option<Tensor<B, 3, Bool>>,
pub key_valid: Option<Tensor<B, 2, Bool>>,
pub query_valid: Option<Tensor<B, 2, Bool>>,
pub key_end_positions: Option<Tensor<B, 1, Int>>,
pub query_positions: Option<Tensor<B, 1, Int>>,
}
impl<B: Backend> Default for IndexerMask<B> {
fn default() -> Self {
Self { allowed: None, key_valid: None, query_valid: None, key_end_positions: None, query_positions: None }
}
}
#[derive(Clone, Debug)]
pub struct IndexerOutput<B: Backend> {
pub indices: Tensor<B, 3, Int>,
pub scores: Tensor<B, 3>,
pub valid: Tensor<B, 3, Bool>,
}
impl LightningIndexerConfig {
pub fn init<B: Backend>(&self, device: &B::Device) -> LightningIndexer<B> {
let query_dim = self.query_dim.unwrap_or(self.width);
assert!(self.width > 0 && query_dim > 0 && self.num_heads > 0 && self.head_dim > 0
&& self.topk > 0 && self.query_chunk_size > 0 && self.key_chunk_size > 0 && self.epsilon.is_finite() && self.epsilon > 0.0,
"invalid indexer dimensions/chunk budget/epsilon");
let query_output = self.num_heads.checked_mul(self.head_dim).expect("indexer query width overflow");
let query = LinearConfig::new(query_dim, query_output).with_bias(false).init(device);
let head_weight = LinearConfig::new(self.width, self.num_heads).with_bias(false).init(device);
let (key, key_norm) = if self.external_keys { (None, None) } else {
(Some(LinearConfig::new(self.width, self.head_dim).with_bias(false).init(device)),
Some(LayerNormConfig::new(self.head_dim).with_epsilon(self.epsilon).init(device)))
};
LightningIndexer::from_parts(query, head_weight, key, key_norm,
SparseRotaryEmbedding::new(self.rope_dim, self.rope_base), self.topk,
self.query_chunk_size, self.key_chunk_size, self.detach_inputs)
}
}
impl<B: Backend, P: CompressedAttentionProjection<B>> LightningIndexer<B, P> {
pub fn from_parts(query: P, head_weight: P, key: Option<P>,
key_norm: Option<LayerNorm<B>>, rotary: SparseRotaryEmbedding, topk: usize,
query_chunk_size: usize, key_chunk_size: usize, detach_inputs: bool) -> Self {
let [query_dim, query_output] = query.dimensions();
let [width, num_heads] = head_weight.dimensions();
assert!(width > 0 && query_dim > 0 && num_heads > 0 && query_output > 0
&& query_output.is_multiple_of(num_heads) && topk > 0 && query_chunk_size > 0 && key_chunk_size > 0,
"invalid loaded indexer geometry/budget");
let head_dim = query_output / num_heads;
assert!(rotary.rope_dim <= head_dim, "indexer rotary width exceeds a head");
assert!(!query.has_bias() && !head_weight.has_bias(), "indexer projections must be bias-free");
let device = query.device();
assert_eq!(head_weight.device(), device, "indexer projection devices differ");
match (&key, &key_norm) {
(Some(key), Some(norm)) => {
assert_eq!(key.dimensions(), [width, head_dim], "indexer token-key geometry differs");
assert!(!key.has_bias(), "indexer token-key projection must be bias-free");
assert_eq!(norm.gamma.val().dims(), [head_dim], "indexer key normalization width differs");
assert!(norm.epsilon().is_finite() && norm.epsilon() > 0.0, "invalid indexer key normalization epsilon");
assert!(key.device() == device && norm.gamma.val().device() == device
&& norm.beta.as_ref().is_none_or(|beta| beta.val().device() == device),
"indexer key parameters must share a device");
}
(None, None) => {}
_ => panic!("indexer token-key projection and normalization must be supplied together"),
}
Self { query, head_weight, key, key_norm, rotary, width, query_dim, num_heads, head_dim,
topk, query_chunk_size, key_chunk_size, detach_inputs }
}
pub fn project_keys(&self, input: Tensor<B, 3>, pos: Option<Tensor<B, 1, Int>>) -> Tensor<B, 3> {
assert_eq!(input.dims()[2], self.width, "indexer token feature width differs");
assert_eq!(input.device(), self.query.device(), "indexer token feature device differs");
work_dtype(input.dtype());
let key = self.key.as_ref().expect("external-key indexer has no token-key projection");
let norm = self.key_norm.as_ref().expect("indexer token-key normalization missing");
let input = if self.detach_inputs { input.detach() } else { input };
let pos = pos.unwrap_or_else(|| positions::<B>(input.dims()[1], 0, &input.device()));
self.rotary.forward_shared(norm.forward(key.forward(input)), pos)
}
fn queries(&self, input: Tensor<B, 3>, latent: Option<Tensor<B, 3>>, pos: Option<Tensor<B, 1, Int>>,
selection: bool) -> (Tensor<B, 4>, Tensor<B, 3>) {
let [batch, tokens, width] = input.dims();
assert_eq!(width, self.width, "indexer feature width differs");
assert_eq!(input.device(), self.query.device(), "indexer feature/parameter device differs");
let latent = latent.unwrap_or_else(|| input.clone());
assert_eq!(latent.dims(), [batch, tokens, self.query_dim], "indexer query latent geometry differs");
assert_eq!(latent.device(), input.device(), "indexer query latent device differs");
work_dtype(latent.dtype());
let compute = work_dtype(input.dtype());
let pos = pos.unwrap_or_else(|| positions::<B>(tokens, 0, &input.device()));
let input = if selection || self.detach_inputs { input.detach() } else { input };
let latent = if selection || self.detach_inputs { latent.detach() } else { latent };
let q = self.query.forward_in_compute(latent.cast(compute), compute, 0..self.num_heads * self.head_dim, selection)
.reshape([batch, tokens, self.num_heads, self.head_dim]);
let scale = 1.0 / ((self.head_dim as f64) * (self.num_heads as f64)).sqrt();
let weights = self.head_weight.forward_in_compute(input.cast(compute), compute, 0..self.num_heads, selection).mul_scalar(scale);
(self.rotary.forward(q, pos, false), weights)
}
fn score(&self, query: Tensor<B, 4>, weights: Tensor<B, 3>, keys: Tensor<B, 3>) -> Tensor<B, 3> {
let [batch, queries, heads, width] = query.dims();
let [key_batch, count, key_width] = keys.dims();
assert_eq!((key_batch, key_width), (batch, width), "indexer key geometry differs");
assert_eq!(keys.device(), query.device(), "indexer key device differs");
work_dtype(keys.dtype());
let compute = query.dtype();
if count == 0 || queries == 0 || batch == 0 {
return Tensor::zeros([batch, queries, count], (&query.device(), compute))
+ (query.sum() + weights.sum() + keys.cast(compute).sum()).mul_scalar(0).reshape([1, 1, 1]);
}
let dots = query.swap_dims(1, 2).matmul(keys.cast(compute).reshape([batch, 1, count, width]).swap_dims(2, 3));
(ruda_model::tensor::activation::relu(dots) * weights.swap_dims(1, 2).reshape([batch, heads, queries, 1]))
.sum_dim(1).reshape([batch, queries, count])
}
pub fn scores(&self, input: Tensor<B, 3>, keys: Tensor<B, 3>, latent: Option<Tensor<B, 3>>,
pos: Option<Tensor<B, 1, Int>>) -> Tensor<B, 3> {
let (query, weights) = self.queries(input, latent, pos, false);
self.score(query, weights, keys)
}
pub fn selected_scores(&self, input: Tensor<B, 3>, keys: Tensor<B, 3>, indices: Tensor<B, 3, Int>,
latent: Option<Tensor<B, 3>>, pos: Option<Tensor<B, 1, Int>>) -> Tensor<B, 3> {
let (query, weights) = self.queries(input, latent, pos, false);
let [batch, tokens, heads, width] = query.dims();
assert_eq!(indices.dims()[..2], [batch, tokens], "indexer selected query geometry differs");
assert_eq!(keys.dims()[2], width, "indexer selected key width differs");
let count = indices.dims()[2];
let entries = sparse_gather_entries(keys.cast(query.dtype()), indices);
if count == 0 || tokens == 0 || batch == 0 {
return Tensor::zeros([batch, tokens, count], (&query.device(), query.dtype()))
+ (query.sum() + weights.sum() + entries.sum()).mul_scalar(0).reshape([1, 1, 1]);
}
(ruda_model::tensor::activation::relu(query.matmul(entries.swap_dims(2, 3)))
* weights.reshape([batch, tokens, heads, 1])).sum_dim(2).reshape([batch, tokens, count])
}
fn validate_mask(&self, input: &Tensor<B, 3>, keys: &Tensor<B, 3>, mask: &IndexerMask<B>) {
let [batch, queries, width] = input.dims();
assert_eq!(width, self.width, "indexer feature width differs");
let count = keys.dims()[1];
assert_eq!(keys.dims(), [batch, count, self.head_dim], "indexer prepared key geometry differs");
assert_eq!(keys.device(), input.device(), "indexer prepared key device differs");
if let Some(valid) = &mask.allowed {
assert_eq!(valid.dims(), [batch, queries, count], "indexer allowed geometry differs");
assert_eq!(valid.device(), input.device(), "indexer allowed device differs");
}
for (valid, shape) in [(&mask.query_valid, [batch, queries]), (&mask.key_valid, [batch, count])] {
if let Some(valid) = valid {
assert_eq!(valid.dims(), shape, "indexer token validity geometry differs");
assert_eq!(valid.device(), input.device(), "indexer token validity device differs");
}
}
for (pos, length) in [(&mask.query_positions, queries), (&mask.key_end_positions, count)] {
if let Some(pos) = pos {
assert_eq!(pos.dims(), [length], "indexer position metadata geometry differs");
assert_eq!(pos.device(), input.device(), "indexer position metadata device differs");
assert_eq!(pos.dtype(), DType::I64, "indexer absolute positions must use I64");
}
}
}
pub fn select(&self, input: Tensor<B, 3>, keys: Tensor<B, 3>, latent: Option<Tensor<B, 3>>,
mask: IndexerMask<B>) -> Tensor<B, 3, Int> {
self.validate_mask(&input, &keys, &mask);
let [batch, tokens, _] = input.dims();
let count = keys.dims()[1];
let device = input.device();
let pos = mask.query_positions.clone().unwrap_or_else(|| positions::<B>(tokens, 0, &device));
let (query, weights) = self.queries(input, latent, Some(pos.clone()), true);
let keys = keys.detach();
if batch == 0 || tokens == 0 || count == 0 || self.topk == 0 {
return Tensor::empty([batch, tokens, self.topk.min(count)], (&device, DType::I64));
}
let mut chunks = Vec::new();
for begin in (0..tokens).step_by(self.query_chunk_size) {
let end = tokens.min(begin.saturating_add(self.query_chunk_size));
let length = end - begin;
let mut best = Tensor::<B, 3>::empty([batch, length, 0], (&device, query.dtype()));
let mut ids = Tensor::<B, 3, Int>::empty([batch, length, 0], (&device, DType::I64));
for start in (0..count).step_by(self.key_chunk_size) {
let stop = count.min(start.saturating_add(self.key_chunk_size));
let score = self.score(query.clone().slice_dim(1, begin..end), weights.clone().slice_dim(1, begin..end),
keys.clone().slice_dim(1, start..stop));
let shape = [batch, length, stop - start];
let mut valid = Tensor::<B, 3, Bool>::zeros(shape, &device).bool_not();
if let Some(allowed) = &mask.allowed { valid = valid.bool_and(allowed.clone().slice([0..batch, begin..end, start..stop])); }
if let Some(key_valid) = &mask.key_valid {
valid = valid.bool_and(key_valid.clone().slice_dim(1, start..stop).reshape([batch, 1, stop - start]).expand(shape));
}
if let Some(query_valid) = &mask.query_valid {
valid = valid.bool_and(query_valid.clone().slice_dim(1, begin..end).reshape([batch, length, 1]).expand(shape));
}
if let Some(ends) = &mask.key_end_positions {
let ends = ends.clone().slice_dim(0, start..stop).reshape([1, 1, stop - start]).expand(shape);
let rows = pos.clone().slice_dim(0, begin..end).reshape([1, length, 1]).expand(shape);
valid = valid.bool_and(ends.lower_equal(rows));
}
let next_ids = positions::<B>(stop - start, start, &device).reshape([1, 1, stop - start]).expand(shape);
let combined_valid = Tensor::cat(vec![ids.clone().greater_equal_elem(0), valid], 2);
let selected = sparse_stable_topk(Tensor::cat(vec![best, score], 2), self.topk,
Some(Tensor::cat(vec![ids, next_ids], 2)), Some(combined_valid));
(best, ids) = selected;
}
chunks.push(ids);
}
Tensor::cat(chunks, 1)
}
pub fn forward_prepared(&self, input: Tensor<B, 3>, keys: Tensor<B, 3>, latent: Option<Tensor<B, 3>>,
mut mask: IndexerMask<B>, causal: bool) -> IndexerOutput<B> {
if causal && mask.key_end_positions.is_none() {
mask.key_end_positions = Some(positions::<B>(keys.dims()[1], 0, &input.device()));
}
let pos = mask.query_positions.clone();
let indices = self.select(input.clone(), keys.clone(), latent.clone(), mask);
let scores = self.selected_scores(input, keys, indices.clone(), latent, pos);
let valid = indices.clone().greater_equal_elem(0);
IndexerOutput { indices, scores, valid }
}
pub fn forward(&self, input: Tensor<B, 3>, key_states: Option<Tensor<B, 3>>,
key_positions: Option<Tensor<B, 1, Int>>, latent: Option<Tensor<B, 3>>,
mask: IndexerMask<B>, causal: bool) -> IndexerOutput<B> {
let keys = self.project_keys(key_states.unwrap_or_else(|| input.clone()), key_positions);
self.forward_prepared(input, keys, latent, mask, causal)
}
pub fn distillation_loss(&self, input: Tensor<B, 3>, keys: Tensor<B, 3>, teacher: Tensor<B, 3>,
allowed: Option<Tensor<B, 3, Bool>>, latent: Option<Tensor<B, 3>>, pos: Option<Tensor<B, 1, Int>>) -> Tensor<B, 1> {
indexer_kl_loss(self.scores(input, keys, latent, pos), teacher, allowed)
}
pub fn selected_distillation_loss(&self, input: Tensor<B, 3>, keys: Tensor<B, 3>, indices: Tensor<B, 3, Int>,
teacher: Tensor<B, 3>, latent: Option<Tensor<B, 3>>, pos: Option<Tensor<B, 1, Int>>) -> Tensor<B, 1> {
let valid = indices.clone().greater_equal_elem(0);
indexer_kl_loss(self.selected_scores(input, keys, indices, latent, pos), teacher, Some(valid))
}
}