use ruda_autodiff::{Autodiff,checkpoint::strategy::CheckpointStrategy,tensor_parallel as region};
use ruda_model::{module::Module,tensor::{FloatDType,Tensor,backend::Backend,module::linear}};
use crate::{Linear,attention::{GroupedQueryAttention,DenseAttentionMask,DenseAttentionOptions,dense_scaled_dot_product_attention}};
use region::BroadcastTensorCollective;
mod cached;
mod packed;
mod masks;
pub use masks::*;
#[derive(Clone,Debug)]
pub struct AttentionParallelGroups<C,K=C> {
pub heads: C,
pub kv_replicas: Option<K>,
}
impl<C> AttentionParallelGroups<C,C> {
pub fn sharded(heads: C) -> Self {Self {heads,kv_replicas:None}}
}
impl<C,K> AttentionParallelGroups<C,K> {
pub fn replicated_kv(heads: C,kv_replicas: K) -> Self {Self {heads,kv_replicas:Some(kv_replicas)}}
}
#[derive(Module,Debug)]
pub struct TensorParallelGroupedQueryAttention<B: Backend> {
pub local: GroupedQueryAttention<B>,
}
fn projection<B: Backend,const D: usize>(input: Tensor<B,D>,weight: Tensor<B,2>,bias: Option<Tensor<B,1>>,
compute: Option<FloatDType>) -> Tensor<B,D> {
if let Some(dtype) = compute {linear(input.cast(dtype),weight.cast(dtype),bias.map(|bias|bias.cast(dtype)))}
else {linear(input,weight,bias)}
}
impl<B: Backend> TensorParallelGroupedQueryAttention<B> {
pub fn from_shard(local: GroupedQueryAttention<B>) -> Self {
let query = local.query.weight.val().dims();
let key = local.key.weight.val().dims();
assert!(local.query_heads > 0 && local.kv_heads > 0 && local.head_dimension > 0
&& local.query_heads.is_multiple_of(local.kv_heads),"invalid local parallel attention head geometry");
assert_eq!(query[1],local.query_heads.checked_mul(local.head_dimension).expect("parallel query width overflow"),"query columns differ from local heads");
assert_eq!(key[1],local.kv_heads.checked_mul(local.head_dimension).expect("parallel KV width overflow"),"KV columns differ from local heads");
assert_eq!(local.value.weight.val().dims(),key,"parallel key/value projection geometry differs");
assert_eq!(local.output.weight.val().dims(),[query[1],query[0]],"parallel output rows/residual width differ");
for layer in [&local.query,&local.key,&local.value,&local.output] {
if let Some(bias) = &layer.bias {assert_eq!(bias.val().dims(),[layer.weight.val().dims()[1]],"parallel projection bias differs from actual output width");}
}
Self {local}
}
fn partial(&self,query: Tensor<B,4>,key: Tensor<B,4>,value: Tensor<B,4>,mut masks: DenseAttentionMask<B>,
options: DenseAttentionOptions,compute: Option<FloatDType>) -> Tensor<B,3> {
let [batch,heads,tokens,width] = query.dims();
assert_eq!((heads,width),(self.local.query_heads,self.local.head_dimension),"parallel projected query heads differ");
assert_eq!((key.dims()[1],key.dims()[3]),(self.local.kv_heads,width),"parallel projected key heads differ");
assert_eq!((value.dims()[1],value.dims()[3]),(self.local.kv_heads,width),"parallel projected value heads differ");
if let Some(dtype) = compute {masks.bias = masks.bias.map(|bias|bias.cast(dtype));}
let context = dense_scaled_dot_product_attention(query,key,value,masks,options,Some(&self.local.dropout))
.swap_dims(1,2).reshape([batch,tokens,heads*width]);
projection(context,self.local.output.weight.val(),None,compute)
}
fn bias<const D: usize>(&self,output: Tensor<B,D>,compute: Option<FloatDType>) -> Tensor<B,D> {
if let Some(bias) = &self.local.output.bias {
let mut shape = [1;D];
shape[D-1] = bias.val().dims()[0];
let bias = if let Some(dtype) = compute {bias.val().cast(dtype)} else {bias.val()};
output+bias.reshape(shape)
} else {output}
}
pub fn forward_projected_inference<C: BroadcastTensorCollective<B>>(&self,query: Tensor<B,4>,key: Tensor<B,4>,value: Tensor<B,4>,
masks: DenseAttentionMask<B>,options: DenseAttentionOptions,communicator: C) -> Result<Tensor<B,3>,C::Error> {
let partial = self.partial(query,key,value,masks,options,None);
let output = communicator.all_reduce_sum(partial.into_primitive().tensor())?;
Ok(self.bias(Tensor::from_primitive(ruda_model::tensor::TensorPrimitive::Float(output)),None))
}
pub fn forward_inference<C: BroadcastTensorCollective<B>>(&self,query: Tensor<B,3>,key: Tensor<B,3>,value: Tensor<B,3>,
masks: DenseAttentionMask<B>,options: DenseAttentionOptions,communicator: C) -> Result<Tensor<B,3>,C::Error> {
let (query,key,value) = self.local.project(query,key,value);
self.forward_projected_inference(query,key,value,masks,options,communicator)
}
}
impl<B: Backend,S: CheckpointStrategy> TensorParallelGroupedQueryAttention<Autodiff<B,S>> {
fn kv_projection<C,K>(&self,layer: &Linear<Autodiff<B,S>>,input: Tensor<Autodiff<B,S>,3>,
groups: &AttentionParallelGroups<C,K>,compute: Option<FloatDType>) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error> {
let mut weight = layer.weight.val();
let mut bias = layer.bias.as_ref().map(|bias|bias.val());
if let Some(dtype) = compute {
weight = weight.cast(dtype);
bias = bias.map(|bias|bias.cast(dtype));
}
if let Some(replicas) = &groups.kv_replicas {
weight = region::copy_to_region(weight,replicas.clone())?;
bias = bias.map(|bias|region::copy_to_region(bias,replicas.clone())).transpose()?;
}
Ok(projection(input,weight,bias,compute))
}
fn project_copied<C,K>(&self,query: Tensor<Autodiff<B,S>,3>,key: Tensor<Autodiff<B,S>,3>,value: Tensor<Autodiff<B,S>,3>,
groups: &AttentionParallelGroups<C,K>,compute: Option<FloatDType>)
-> Result<(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>),C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error> {
let [batch,queries,_] = query.dims();
let [key_batch,keys,_] = key.dims();
assert_eq!((key_batch,keys),(value.dims()[0],value.dims()[1]),"parallel projected K/V rows differ");
assert_eq!(batch,key_batch,"parallel query/memory batch differs");
let query = projection(query,self.local.query.weight.val(),self.local.query.bias.as_ref().map(|bias|bias.val()),compute);
let key = self.kv_projection(&self.local.key,key,groups,compute)?;
let value = self.kv_projection(&self.local.value,value,groups,compute)?;
Ok((query.reshape([batch,queries,self.local.query_heads,self.local.head_dimension]).swap_dims(1,2),
key.reshape([batch,keys,self.local.kv_heads,self.local.head_dimension]).swap_dims(1,2),
value.reshape([batch,keys,self.local.kv_heads,self.local.head_dimension]).swap_dims(1,2)))
}
pub fn project<C,K>(&self,query: Tensor<Autodiff<B,S>,3>,key: Tensor<Autodiff<B,S>,3>,value: Tensor<Autodiff<B,S>,3>,
groups: &AttentionParallelGroups<C,K>)
-> Result<(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>),C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error> {
self.project_with_compute_dtype_inner(query,key,value,groups,None)
}
fn project_with_compute_dtype_inner<C,K>(&self,query: Tensor<Autodiff<B,S>,3>,key: Tensor<Autodiff<B,S>,3>,value: Tensor<Autodiff<B,S>,3>,
groups: &AttentionParallelGroups<C,K>,compute: Option<FloatDType>)
-> Result<(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>),C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error> {
let query = if let Some(dtype) = compute {query.cast(dtype)} else {query};
let key = if let Some(dtype) = compute {key.cast(dtype)} else {key};
let value = if let Some(dtype) = compute {value.cast(dtype)} else {value};
let query = region::copy_to_region(query,groups.heads.clone())?;
let key = region::copy_to_region(key,groups.heads.clone())?;
let value = region::copy_to_region(value,groups.heads.clone())?;
self.project_copied(query,key,value,groups,compute)
}
pub fn project_with_compute_dtype<C,K>(&self,query: Tensor<Autodiff<B,S>,3>,key: Tensor<Autodiff<B,S>,3>,value: Tensor<Autodiff<B,S>,3>,
groups: &AttentionParallelGroups<C,K>,dtype: FloatDType)
-> Result<(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>),C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error> {
self.project_with_compute_dtype_inner(query,key,value,groups,Some(dtype))
}
pub fn project_self<C,K>(&self,input: Tensor<Autodiff<B,S>,3>,groups: &AttentionParallelGroups<C,K>)
-> Result<(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>),C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error> {
let input = region::copy_to_region(input,groups.heads.clone())?;
self.project_copied(input.clone(),input.clone(),input,groups,None)
}
pub fn forward_projected<C: BroadcastTensorCollective<B>>(&self,query: Tensor<Autodiff<B,S>,4>,key: Tensor<Autodiff<B,S>,4>,value: Tensor<Autodiff<B,S>,4>,
masks: DenseAttentionMask<Autodiff<B,S>>,options: DenseAttentionOptions,communicator: C)
-> Result<Tensor<Autodiff<B,S>,3>,C::Error> {
let partial = self.partial(query,key,value,masks,options,None);
Ok(self.bias(region::reduce_from_region(partial,communicator)?,None))
}
pub fn forward<C,K>(&self,query: Tensor<Autodiff<B,S>,3>,key: Tensor<Autodiff<B,S>,3>,value: Tensor<Autodiff<B,S>,3>,
masks: DenseAttentionMask<Autodiff<B,S>>,options: DenseAttentionOptions,groups: &AttentionParallelGroups<C,K>)
-> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error> {
let (query,key,value) = self.project(query,key,value,groups)?;
self.forward_projected(query,key,value,masks,options,groups.heads.clone())
}
pub fn forward_self<C,K>(&self,input: Tensor<Autodiff<B,S>,3>,masks: DenseAttentionMask<Autodiff<B,S>>,
options: DenseAttentionOptions,groups: &AttentionParallelGroups<C,K>) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error> {
let (query,key,value) = self.project_self(input,groups)?;
self.forward_projected(query,key,value,masks,options,groups.heads.clone())
}
pub fn forward_with_compute_dtype<C,K>(&self,query: Tensor<Autodiff<B,S>,3>,key: Tensor<Autodiff<B,S>,3>,value: Tensor<Autodiff<B,S>,3>,
masks: DenseAttentionMask<Autodiff<B,S>>,options: DenseAttentionOptions,groups: &AttentionParallelGroups<C,K>,dtype: FloatDType)
-> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error> {
let storage = query.dtype();
let (query,key,value) = self.project_with_compute_dtype(query,key,value,groups,dtype)?;
let partial = self.partial(query,key,value,masks,options,Some(dtype));
Ok(self.bias(region::reduce_from_region(partial,groups.heads.clone())?,Some(dtype)).cast(storage))
}
}