use super::{PackedAttentionOptions, PackedCausalAlignment};
use crate::{Dropout, DropoutConfig, Linear, LinearConfig};
use ruda_model::{
config::Config, module::Module,
tensor::{Bool, DType, FloatDType, Int, Tensor, backend::Backend},
};
#[cfg(not(feature = "std"))]
#[allow(unused_imports)]
use num_traits::Float as _;
pub type DenseAttentionOptions = PackedAttentionOptions;
#[derive(Clone, Debug)]
pub struct DenseAttentionMask<B: Backend> {
pub query_valid: Option<Tensor<B, 2, Bool>>,
pub key_valid: Option<Tensor<B, 2, Bool>>,
pub allowed: Option<Tensor<B, 4, Bool>>,
pub bias: Option<Tensor<B, 4>>,
}
impl<B: Backend> Default for DenseAttentionMask<B> {
fn default() -> Self {
Self { query_valid: None, key_valid: None, allowed: None, bias: None }
}
}
fn check_broadcast(shape: [usize; 4], target: [usize; 4]) {
assert!(shape.iter().zip(target).all(|(&size, expected)| size == 1 || size == expected),
"attention mask/bias does not broadcast to the actual score geometry");
}
fn connected_zero<B: Backend>(tensor: Tensor<B, 4>, dtype: DType) -> Tensor<B, 1> {
let excluded = Tensor::<B, 4, Bool>::zeros(tensor.dims(), &tensor.device()).bool_not();
tensor.cast(dtype).mask_fill(excluded, 0).sum()
}
pub fn dense_scaled_dot_product_attention<B: Backend>(
query: Tensor<B, 4>, key: Tensor<B, 4>, value: Tensor<B, 4>,
masks: DenseAttentionMask<B>, options: DenseAttentionOptions, dropout: Option<&Dropout>,
) -> Tensor<B, 4> {
let [batch, heads, queries, features] = query.dims();
let [key_batch, kv_heads, keys, key_features] = key.dims();
let [value_batch, value_heads, values, value_features] = value.dims();
assert!(heads > 0 && kv_heads > 0 && heads.is_multiple_of(kv_heads) && features > 0 && value_features > 0,
"invalid grouped attention head/feature geometry");
assert_eq!((batch, keys, kv_heads, features), (key_batch, values, value_heads, key_features),
"key/value payloads or QK feature geometry differ");
assert_eq!(batch, value_batch, "query/value batches differ");
let device = query.device();
let storage = query.dtype();
assert!(matches!(storage, DType::F32 | DType::F16 | DType::BF16 | DType::F64)
&& key.dtype() == storage && value.dtype() == storage, "attention requires matching floating input storage");
assert!(device == key.device() && device == value.device(), "attention inputs must share a device");
let compute = if storage == DType::F64 { DType::F64 } else { DType::F32 };
let score_shape = [batch, heads, queries, keys];
if let Some(mask) = &masks.query_valid {
assert_eq!(mask.dims(), [batch, queries], "query visibility differs from actual query tokens");
assert_eq!(mask.device(), device, "query visibility must share the device");
}
if let Some(mask) = &masks.key_valid {
assert_eq!(mask.dims(), [batch, keys], "key visibility differs from actual key tokens");
assert_eq!(mask.device(), device, "key visibility must share the device");
}
if let Some(mask) = &masks.allowed {
check_broadcast(mask.dims(), score_shape);
assert_eq!(mask.device(), device, "allowed edges must share the device");
}
if let Some(bias) = &masks.bias {
check_broadcast(bias.dims(), score_shape);
assert_eq!(bias.device(), device, "attention bias must share the device");
assert!(bias.dtype() == storage || bias.dtype() == compute, "attention bias storage cannot be silently narrowed");
}
let scale = options.scale.unwrap_or(1.0 / (features as f64).sqrt());
assert!(scale.is_finite() && (compute == DType::F64 || (scale as f32).is_finite()),
"attention scale must be finite in the selected compute precision");
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 attention dropout probability"); }
if batch == 0 || queries == 0 || keys == 0
|| dropout.is_some_and(|dropout|dropout.prob == 1.0 && B::ad_enabled(&device)) {
let mut zero = connected_zero(query, compute) + connected_zero(key, compute) + connected_zero(value, compute);
if let Some(bias) = masks.bias { zero = zero + connected_zero(bias, compute); }
return (Tensor::<B, 4>::zeros([batch, heads, queries, value_features], (&device, compute))
+ zero.reshape([1, 1, 1, 1])).cast(storage);
}
let groups = heads / kv_heads;
let mut query = query.cast(compute);
let mut key = key.cast(compute);
let mut value = value.cast(compute);
if let Some(mask) = &masks.query_valid {
query = query.mask_fill(mask.clone().bool_not().reshape([batch,1,queries,1])
.expand([batch,heads,queries,features]),0);
}
if let Some(mask) = &masks.key_valid {
let excluded = mask.clone().bool_not().reshape([batch,1,keys,1]);
key = key.mask_fill(excluded.clone().expand([batch,kv_heads,keys,features]),0);
value = value.mask_fill(excluded.expand([batch,kv_heads,keys,value_features]),0);
}
let key = key.reshape([batch, kv_heads, 1, keys, features])
.repeat_dim(2, groups).reshape([batch, heads, keys, features]);
let value = value.reshape([batch, kv_heads, 1, keys, value_features])
.repeat_dim(2, groups).reshape([batch, heads, keys, value_features]);
let mut scores = query.matmul(key.swap_dims(2, 3)).mul_scalar(scale);
if let Some(bias) = masks.bias { scores = scores + bias.cast(compute).expand(score_shape); }
let mut excluded = Tensor::<B, 4, Bool>::zeros(score_shape, &device);
if let Some(mask) = masks.allowed { excluded = excluded.bool_or(mask.expand(score_shape).bool_not()); }
if let Some(mask) = masks.query_valid {
excluded = excluded.bool_or(mask.reshape([batch, 1, queries, 1]).expand(score_shape).bool_not());
}
if let Some(mask) = masks.key_valid {
excluded = excluded.bool_or(mask.reshape([batch, 1, 1, keys]).expand(score_shape).bool_not());
}
if options.causal || options.window.is_some() {
let query_length = i64::try_from(queries).expect("query positions exceed integer range");
let key_length = i64::try_from(keys).expect("key positions exceed integer range");
let offset = match options.alignment {
PackedCausalAlignment::UpperLeft => 0,
PackedCausalAlignment::LowerRight => key_length - query_length,
};
let rows = Tensor::<B, 1, Int>::arange(0..query_length, (&device, DType::I64))
.add_scalar(offset).reshape([1, 1, queries, 1]).expand(score_shape);
let columns = Tensor::<B, 1, Int>::arange(0..key_length, (&device, DType::I64))
.reshape([1, 1, 1, keys]).expand(score_shape);
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) < queries.max(keys) {
excluded = excluded.bool_or(columns.clone().lower(rows.clone().sub_scalar(left as i64)));
}
if right >= 0 && (right as usize) < queries.max(keys) {
excluded = excluded.bool_or(columns.greater(rows.add_scalar(right as i64)));
}
}
}
scores = scores.mask_fill(excluded, f64::NEG_INFINITY);
let fully_excluded = scores.clone().equal_elem(f64::NEG_INFINITY).all_dim(3);
let maximum = scores.clone().max_dim(3).mask_fill(fully_excluded.clone(), 0);
let exponentials = (scores - maximum).exp();
let denominator = exponentials.clone().sum_dim(3).clamp_min(1);
let mut weights = exponentials / denominator;
if let Some(dropout) = dropout {
weights = dropout.forward(weights);
}
weights.matmul(value).mask_fill(fully_excluded.expand([batch,heads,queries,value_features]),0).cast(storage)
}
#[derive(Config, Debug)]
pub struct GroupedQueryAttentionConfig {
pub d_model: usize,
pub query_heads: usize,
pub kv_heads: usize,
pub head_dimension: usize,
#[config(default = false)]
pub bias: bool,
#[config(default = 0.0)]
pub dropout: f64,
}
#[derive(Module, Debug)]
pub struct GroupedQueryAttention<B: Backend> {
pub query: Linear<B>,
pub key: Linear<B>,
pub value: Linear<B>,
pub output: Linear<B>,
pub dropout: Dropout,
pub query_heads: usize,
pub kv_heads: usize,
pub head_dimension: usize,
}
impl GroupedQueryAttentionConfig {
pub fn init<B: Backend>(&self, device: &B::Device) -> GroupedQueryAttention<B> {
assert!(self.d_model > 0 && self.query_heads > 0 && self.kv_heads > 0
&& self.head_dimension > 0 && self.query_heads.is_multiple_of(self.kv_heads), "invalid grouped projection geometry");
let queries = self.query_heads.checked_mul(self.head_dimension).expect("query projection width overflow");
let keys = self.kv_heads.checked_mul(self.head_dimension).expect("KV projection width overflow");
GroupedQueryAttention {
query: LinearConfig::new(self.d_model, queries).with_bias(self.bias).init(device),
key: LinearConfig::new(self.d_model, keys).with_bias(self.bias).init(device),
value: LinearConfig::new(self.d_model, keys).with_bias(self.bias).init(device),
output: LinearConfig::new(queries, self.d_model).with_bias(self.bias).init(device),
dropout: DropoutConfig::new(self.dropout).init(),
query_heads: self.query_heads, kv_heads: self.kv_heads, head_dimension: self.head_dimension,
}
}
}
impl<B: Backend> GroupedQueryAttention<B> {
pub fn from_projections(query: Linear<B>,key: Linear<B>,value: Linear<B>,output: Linear<B>,
query_heads: usize,kv_heads: usize,head_dimension: usize,dropout: Dropout) -> Self {
assert!(query_heads > 0 && kv_heads > 0 && head_dimension > 0
&& query_heads.is_multiple_of(kv_heads),"invalid grouped projection head geometry");
let query_width = query_heads.checked_mul(head_dimension).expect("query width overflow");
let key_width = kv_heads.checked_mul(head_dimension).expect("KV width overflow");
let query_shape = query.weight.val().dims();
let key_shape = key.weight.val().dims();
let value_shape = value.weight.val().dims();
assert_eq!(query_shape[1],query_width,"actual query weight differs from head geometry");
assert_eq!(key_shape[1],key_width,"actual key weight differs from head geometry");
assert_eq!(value_shape,key_shape,"actual key/value input and head widths differ");
assert_eq!(output.weight.val().dims(),[query_width,query_shape[0]],"actual attention output/residual widths differ");
assert!(dropout.prob.is_finite() && (0.0..=1.0).contains(&dropout.prob),"invalid attention dropout");
Self {query,key,value,output,dropout,query_heads,kv_heads,head_dimension}
}
pub fn project(&self, query: Tensor<B, 3>, key: Tensor<B, 3>, value: Tensor<B, 3>)
-> (Tensor<B, 4>, Tensor<B, 4>, Tensor<B, 4>) {
let [batch, queries, _] = query.dims();
let [key_batch, keys, _] = key.dims();
let [value_batch, values, _] = value.dims();
assert_eq!((batch, keys), (key_batch, values), "grouped projection batches/key lengths differ");
assert_eq!(batch, value_batch, "grouped value batch differs");
let query = self.query.forward(query).reshape([batch, queries, self.query_heads, self.head_dimension]).swap_dims(1, 2);
let key = self.key.forward(key).reshape([batch, keys, self.kv_heads, self.head_dimension]).swap_dims(1, 2);
let value = self.value.forward(value).reshape([batch, keys, self.kv_heads, self.head_dimension]).swap_dims(1, 2);
(query, key, value)
}
pub fn project_with_compute_dtype(&self,query: Tensor<B,3>,key: Tensor<B,3>,value: Tensor<B,3>,
dtype: FloatDType) -> (Tensor<B,4>,Tensor<B,4>,Tensor<B,4>) {
let [batch,queries,_] = query.dims();
let [key_batch,keys,_] = key.dims();
let [value_batch,values,_] = value.dims();
assert_eq!((batch,keys),(key_batch,values),"grouped projection batches/key lengths differ");
assert_eq!(batch,value_batch,"grouped value batch differs");
let project = |layer: &Linear<B>,input: Tensor<B,3>|ruda_model::tensor::module::linear(
input.cast(dtype),layer.weight.val().cast(dtype),layer.bias.as_ref().map(|bias|bias.val().cast(dtype)));
let query = project(&self.query,query).reshape([batch,queries,self.query_heads,self.head_dimension]).swap_dims(1,2);
let key = project(&self.key,key).reshape([batch,keys,self.kv_heads,self.head_dimension]).swap_dims(1,2);
let value = project(&self.value,value).reshape([batch,keys,self.kv_heads,self.head_dimension]).swap_dims(1,2);
(query,key,value)
}
pub fn forward_projected(&self, query: Tensor<B, 4>, key: Tensor<B, 4>, value: Tensor<B, 4>,
masks: DenseAttentionMask<B>, options: DenseAttentionOptions) -> Tensor<B, 3> {
let [batch, heads, queries, width] = query.dims();
assert_eq!((heads, width), (self.query_heads, self.head_dimension), "projected query heads differ from the output weight");
assert_eq!((key.dims()[1], key.dims()[3]), (self.kv_heads, self.head_dimension), "projected key geometry differs");
assert_eq!((value.dims()[1], value.dims()[3]), (self.kv_heads, self.head_dimension), "projected value geometry differs");
let context = dense_scaled_dot_product_attention(query, key, value, masks, options, Some(&self.dropout));
self.output.forward(context.swap_dims(1, 2).reshape([batch, queries, self.query_heads * self.head_dimension]))
}
pub fn forward(&self, query: Tensor<B, 3>, key: Tensor<B, 3>, value: Tensor<B, 3>,
masks: DenseAttentionMask<B>, options: DenseAttentionOptions) -> Tensor<B, 3> {
let (query, key, value) = self.project(query, key, value);
self.forward_projected(query, key, value, masks, options)
}
pub fn forward_with_compute_dtype(&self,query: Tensor<B,3>,key: Tensor<B,3>,value: Tensor<B,3>,
mut masks: DenseAttentionMask<B>,options: DenseAttentionOptions,dtype: FloatDType) -> Tensor<B,3> {
let storage = query.dtype();
let (query,key,value) = self.project_with_compute_dtype(query,key,value,dtype);
if let Some(bias) = masks.bias.take() { masks.bias = Some(bias.cast(dtype)); }
let [batch,heads,queries,width] = query.dims();
let context = dense_scaled_dot_product_attention(query,key,value,masks,options,Some(&self.dropout))
.swap_dims(1,2).reshape([batch,queries,heads*width]);
ruda_model::tensor::module::linear(context,self.output.weight.val().cast(dtype),
self.output.bias.as_ref().map(|bias|bias.val().cast(dtype))).cast(storage)
}
}