use alloc::{vec, vec::Vec};
use ruda_model::tensor::{Tensor, TensorData, Int, Bool, DType, backend::Backend};
use crate::Dropout;
#[cfg(not(feature = "std"))]
#[allow(unused_imports)]
use num_traits::Float as _;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PackedSequenceLayout {
boundaries: Vec<usize>,
}
impl PackedSequenceLayout {
pub fn new(boundaries: Vec<usize>, tokens: usize) -> Self {
assert!(!boundaries.is_empty() && boundaries[0] == 0 && boundaries.last() == Some(&tokens), "boundaries must cover the actual token payload");
assert!(boundaries.windows(2).all(|pair| pair[0] <= pair[1]), "document boundaries must be nondecreasing");
Self { boundaries }
}
pub fn documents(&self) -> usize { self.boundaries.len() - 1 }
pub fn tokens(&self) -> usize { *self.boundaries.last().unwrap() }
pub fn boundaries(&self) -> &[usize] { &self.boundaries }
pub fn max_length(&self) -> usize {
self.boundaries.windows(2).map(|pair| pair[1] - pair[0]).max().unwrap_or(0)
}
pub fn positions<B: Backend>(&self, device: &B::Device) -> Tensor<B, 1, Int> {
let values: Vec<i64> = self.boundaries.windows(2).flat_map(|pair| 0..pair[1] - pair[0])
.map(|position| i64::try_from(position).expect("document position exceeds integer range")).collect();
Tensor::from_data(TensorData::new(values, [self.tokens()]), (device, DType::I64))
}
pub fn document_starts<B: Backend>(&self, device: &B::Device) -> Tensor<B, 1, Bool> {
let mut starts = vec![false; self.tokens()];
for pair in self.boundaries.windows(2) { if pair[0] < pair[1] { starts[pair[0]] = true; } }
Tensor::from_data(TensorData::new(starts, [self.tokens()]), device)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum PackedCausalAlignment {
UpperLeft,
LowerRight,
}
#[derive(Clone, Copy, Debug)]
pub struct PackedAttentionOptions {
pub scale: Option<f64>,
pub causal: bool,
pub alignment: PackedCausalAlignment,
pub window: Option<(isize, isize)>,
}
impl Default for PackedAttentionOptions {
fn default() -> Self {
Self { scale: None, causal: false, alignment: PackedCausalAlignment::UpperLeft, window: None }
}
}
fn connected_zero<B: Backend>(tensor: Tensor<B, 3>) -> Tensor<B, 1> {
let mask = Tensor::<B, 3, Bool>::zeros(tensor.dims(), &tensor.device()).bool_not();
tensor.cast(DType::F32).mask_fill(mask, 0).sum()
}
pub fn packed_scaled_dot_product_attention<B: Backend>(
query: Tensor<B, 3>, key: Tensor<B, 3>, value: Tensor<B, 3>,
query_layout: &PackedSequenceLayout, key_layout: &PackedSequenceLayout,
options: PackedAttentionOptions, dropout: Option<&Dropout>,
) -> Tensor<B, 3> {
let [queries, heads, features] = query.dims();
let [keys, kv_heads, key_features] = key.dims();
let [values, value_heads, value_features] = value.dims();
assert!(heads > 0 && kv_heads > 0 && heads % kv_heads == 0 && features > 0 && value_features > 0, "invalid attention head/feature geometry");
assert_eq!((keys, kv_heads, features), (values, value_heads, key_features), "key/value or QK features differ");
assert_eq!((query_layout.tokens(), key_layout.tokens()), (queries, keys), "packed metadata does not match payload lengths");
assert_eq!(query_layout.documents(), key_layout.documents(), "query/key document counts differ");
assert!(query.device() == key.device() && query.device() == value.device(), "attention operands must share a device");
let dtype = query.dtype();
assert!(matches!(dtype, DType::F32 | DType::F16 | DType::BF16) && key.dtype() == dtype && value.dtype() == dtype, "attention requires matching FP32/FP16/BF16 storage");
let scale = options.scale.unwrap_or(1.0 / (features as f64).sqrt());
assert!(scale.is_finite() && (scale as f32).is_finite(), "scale must be finite in FP32");
if let Some((left, right)) = options.window { assert!(left >= -1 && right >= -1, "window distances must be nonnegative or -1"); }
if let Some(dropout) = dropout { assert!(dropout.prob.is_finite() && (0.0..=1.0).contains(&dropout.prob), "invalid dropout probability"); }
let device = query.device();
if queries == 0 || dropout.is_some_and(|dropout|dropout.prob == 1.0 && B::ad_enabled(&device)) {
let zero = connected_zero(query) + connected_zero(key) + connected_zero(value);
return (Tensor::<B, 3>::zeros([queries, heads, value_features], (&device, DType::F32)) + zero.reshape([1, 1, 1])).cast(dtype);
}
let query = query.cast(DType::F32);
let key = key.cast(DType::F32);
let value = value.cast(DType::F32);
let mut outputs = Vec::new();
for (qrange, krange) in query_layout.boundaries.windows(2).zip(key_layout.boundaries.windows(2)) {
let qlen = qrange[1] - qrange[0];
let klen = krange[1] - krange[0];
let q = query.clone().slice_dim(0, qrange[0]..qrange[1]).swap_dims(0, 1);
let k = key.clone().slice_dim(0, krange[0]..krange[1]).swap_dims(0, 1);
let v = value.clone().slice_dim(0, krange[0]..krange[1]).swap_dims(0, 1);
if qlen == 0 || klen == 0 {
let zero = connected_zero(q) + connected_zero(k) + connected_zero(v);
outputs.push(Tensor::<B, 3>::zeros([qlen, heads, value_features], (&device, DType::F32)) + zero.reshape([1, 1, 1]));
continue;
}
let groups = heads / kv_heads;
let k = k.reshape([kv_heads, 1, klen, features]).repeat_dim(1, groups).reshape([heads, klen, features]);
let v = v.reshape([kv_heads, 1, klen, value_features]).repeat_dim(1, groups).reshape([heads, klen, value_features]);
let mut scores = q.matmul(k.swap_dims(1, 2)).mul_scalar(scale);
let qlen_i64 = i64::try_from(qlen).expect("query length exceeds integer range");
let klen_i64 = i64::try_from(klen).expect("key length exceeds integer range");
let offset = match options.alignment { PackedCausalAlignment::UpperLeft => 0, PackedCausalAlignment::LowerRight => klen_i64 - qlen_i64 };
let rows = Tensor::<B, 1, Int>::arange(0..qlen_i64, (&device, DType::I64)).add_scalar(offset).reshape([qlen, 1]).repeat_dim(1, klen);
let columns = Tensor::<B, 1, Int>::arange(0..klen_i64, (&device, DType::I64)).reshape([1, klen]).repeat_dim(0, qlen);
let mut excluded = Tensor::<B, 2, Bool>::zeros([qlen, klen], &device);
if options.causal { excluded = excluded.bool_or(columns.clone().greater(rows.clone())); }
if let Some((left, right)) = options.window {
if left >= 0 && (left as usize) < qlen.max(klen) { excluded = excluded.bool_or(columns.clone().lower(rows.clone().sub_scalar(left as i64))); }
if right >= 0 && (right as usize) < qlen.max(klen) { excluded = excluded.bool_or(columns.greater(rows.add_scalar(right as i64))); }
}
let mask = excluded.unsqueeze_dim::<3>(0).repeat_dim(0, heads);
let fully_masked = mask.clone().all_dim(2);
scores = scores.mask_fill(mask, f32::NEG_INFINITY);
let maximum = scores.clone().max_dim(2).mask_fill(fully_masked.clone(), 0);
let exponentials = (scores - maximum).exp();
let denominator = exponentials.clone().sum_dim(2).clamp_min(1);
let mut weights = exponentials / denominator;
if let Some(dropout) = dropout {
weights = dropout.forward(weights);
}
outputs.push(weights.matmul(v).mask_fill(fully_masked.expand([heads,qlen,value_features]),0).swap_dims(0, 1));
}
Tensor::cat(outputs, 0).cast(dtype)
}