use super::*;
use crate::{Dropout,pool::{pool_sequence,pool_packed_sequences,SequencePooling,SequencePoolOutput},
attention::PackedSequenceLayout,transformer::{TransformerEmbeddings,TransformerHead,AdaptedTransformerHead,
AdaptedProjection,DenseTransformerNorm,SequenceHeadOutput}};
use ruda_model::tensor::Bool;
pub type FullyShardedTransformerInput<B,const D:usize=2> = crate::tensor_parallel::TensorParallelTransformerInput<B,D>;
#[derive(Module,Debug)]
pub struct FullyShardedTransformerEmbeddings<B:Backend> {
pub token:FullyShardedEmbedding<B>,
pub position:Option<FullyShardedEmbedding<B>>,
pub token_type:Option<FullyShardedEmbedding<B>>,
pub normalization:Option<FullyShardedTransformerNorm<B>>,
pub dropout:Dropout,
}
#[derive(Module,Debug)]
pub enum FullyShardedHeadProjection<B:Backend> {
Column(FullyShardedAdaptedProjection<B>),
RowMajor(FullyShardedProjection<B>),
}
#[derive(Module,Debug)]
pub struct FullyShardedTransformerHead<B:Backend> {
pub projection:FullyShardedHeadProjection<B>,
pub normalization:Option<FullyShardedTransformerNorm<B>>,
pub dropout:Dropout,
}
enum GatheredHeadProjection<B:Backend> {
Column(AdaptedProjection<B>),
RowMajor {weight:Tensor<B,2>,bias:Option<Tensor<B,1>>},
}
pub struct GatheredFullyShardedTransformerHead<B:Backend> {
projection:GatheredHeadProjection<B>,
normalization:Option<DenseTransformerNorm<B>>,
dropout:Dropout,
}
impl<B:Backend> ShardingContext<B> {
pub fn transformer_embeddings(&mut self,embeddings:TransformerEmbeddings<B>) -> FullyShardedTransformerEmbeddings<B> {
FullyShardedTransformerEmbeddings {token:self.embedding(embeddings.token),
position:embeddings.position.map(|table|self.embedding(table)),
token_type:embeddings.token_type.map(|table|self.embedding(table)),
normalization:embeddings.normalization.map(|norm|self.normalization(norm)),dropout:embeddings.dropout}
}
pub fn transformer_head(&mut self,head:TransformerHead<B>) -> FullyShardedTransformerHead<B> {
FullyShardedTransformerHead::from_parts(FullyShardedHeadProjection::Column(
FullyShardedAdaptedProjection::Dense(self.linear(head.projection))),
head.normalization.map(|norm|self.normalization(norm)),head.dropout)
}
pub fn adapted_transformer_head(&mut self,head:AdaptedTransformerHead<B>) -> FullyShardedTransformerHead<B> {
FullyShardedTransformerHead::from_parts(FullyShardedHeadProjection::Column(
FullyShardedAdaptedProjection::LoRA(self.lora(head.projection))),
head.normalization.map(|norm|self.normalization(norm)),head.dropout)
}
pub fn tied_transformer_head(&mut self,table:&FullyShardedEmbedding<B>,bias:Option<Param<Tensor<B,1>>>,
normalization:Option<DenseTransformerNorm<B>>,dropout:Dropout) -> FullyShardedTransformerHead<B> {
let projection=self.tied_projection(table,bias);
FullyShardedTransformerHead::from_parts(FullyShardedHeadProjection::RowMajor(projection),
normalization.map(|norm|self.normalization(norm)),dropout)
}
}
pub(super) fn projection_geometry<B:Backend>(projection:&FullyShardedAdaptedProjection<B>) -> (usize,usize) {
let weight=match projection {FullyShardedAdaptedProjection::Dense(layer)=>&layer.weight,
FullyShardedAdaptedProjection::LoRA(layer)=>&layer.base.weight};
assert_eq!(weight.logical_shape.len(),2,"sharded projection must be a matrix");
(weight.logical_shape[0],weight.logical_shape[1])
}
impl<B:Backend> FullyShardedTransformerNorm<B> {
pub fn width(&self) -> usize {
let gamma=match self {Self::Layer(norm)=>&norm.gamma,Self::Rms(norm)=>&norm.gamma};
assert_eq!(gamma.logical_shape.len(),1,"sharded transformer affine must be a vector");gamma.logical_shape[0]
}
}
impl<B:Backend> FullyShardedTransformerEmbeddings<B> {
pub fn hidden_width(&self) -> usize {
assert_eq!(self.token.weight.logical_shape.len(),2,"sharded token table must be a matrix");
self.token.weight.logical_shape[1]
}
}
impl<B:Backend> FullyShardedTransformerHead<B> {
pub fn from_parts(projection:FullyShardedHeadProjection<B>,normalization:Option<FullyShardedTransformerNorm<B>>,dropout:Dropout) -> Self {
let head=Self {projection,normalization,dropout};
if let Some(norm)=&head.normalization {assert_eq!(norm.width(),head.hidden_width(),"sharded head norm/input widths differ");}
assert!(head.dropout.prob.is_finite() && (0.0..=1.0).contains(&head.dropout.prob),"invalid sharded head dropout");head
}
pub fn hidden_width(&self) -> usize {
match &self.projection {FullyShardedHeadProjection::Column(layer)=>projection_geometry(layer).0,
FullyShardedHeadProjection::RowMajor(layer)=>{
assert_eq!(layer.weight.logical_shape.len(),2,"sharded output table must be a matrix");layer.weight.logical_shape[1]
}}
}
pub fn classes(&self) -> usize {
match &self.projection {FullyShardedHeadProjection::Column(layer)=>projection_geometry(layer).1,
FullyShardedHeadProjection::RowMajor(layer)=>layer.weight.logical_shape[0]}
}
}
impl<B:Backend> GatheredFullyShardedTransformerHead<B> {
pub fn forward<const D:usize>(&self,hidden:Tensor<B,D>) -> Tensor<B,D> {
let hidden=if let Some(norm)=&self.normalization {norm.forward(hidden)} else {hidden};
let hidden=self.dropout.forward(hidden);
match &self.projection {GatheredHeadProjection::Column(layer)=>layer.forward(hidden),
GatheredHeadProjection::RowMajor {weight,bias}=>linear(hidden,weight.clone().transpose(),bias.clone())}
}
pub fn forward_pooled(&self,pooled:SequencePoolOutput<B>) -> SequenceHeadOutput<B> {
SequenceHeadOutput {logits:self.forward(pooled.values),valid_rows:pooled.valid_rows,token_counts:pooled.token_counts}
}
pub fn forward_sequence(&self,hidden:Tensor<B,3>,visible:Tensor<B,2,Bool>,pooling:SequencePooling) -> SequenceHeadOutput<B> {
self.forward_pooled(pool_sequence(hidden,visible,pooling))
}
pub fn forward_packed_sequences(&self,hidden:Tensor<B,2>,layout:&PackedSequenceLayout,
visible:Option<Tensor<B,1,Bool>>,pooling:SequencePooling) -> SequenceHeadOutput<B> {
self.forward_pooled(pool_packed_sequences(hidden,layout,visible,pooling))
}
}
macro_rules! gathered_model_parts {
($backend:ty,[$($generics:tt)*],$gather:ident,$forward:ident) => {
impl<$($generics)*> FullyShardedTransformerEmbeddings<$backend> {
pub fn $gather<C:BroadcastTensorCollective<B>>(&self,communicator:C) -> Result<TransformerEmbeddings<$backend>,C::Error> {
let table=|table:&FullyShardedEmbedding<$backend>| {
table.weight.$gather::<C,2>(communicator.clone()).map(|value|crate::Embedding {
weight:Param::initialized(table.weight.local.id,value)})
};
Ok(TransformerEmbeddings::from_tables(table(&self.token)?,
self.position.as_ref().map(&table).transpose()?,self.token_type.as_ref().map(&table).transpose()?,
self.normalization.as_ref().map(|norm|norm.$gather(communicator.clone())).transpose()?,self.dropout.clone()))
}
pub fn $forward<C:BroadcastTensorCollective<B>>(&self,input:FullyShardedTransformerInput<$backend>,communicator:C)
-> Result<Tensor<$backend,3>,C::Error> {
let tables=self.$gather(communicator)?;
Ok(match input.embedding_dtypes {
Some((compute,output))=>tables.forward_with_compute_dtype(input.tokens,input.positions,input.token_types,compute,output),
None=>tables.forward(input.tokens,input.positions,input.token_types),
})
}
}
impl<$($generics)*> FullyShardedTransformerHead<$backend> {
pub fn $gather<C:BroadcastTensorCollective<B>>(&self,communicator:C) -> Result<GatheredFullyShardedTransformerHead<$backend>,C::Error> {
let projection=match &self.projection {
FullyShardedHeadProjection::Column(layer)=>GatheredHeadProjection::Column(layer.$gather(communicator.clone())?),
FullyShardedHeadProjection::RowMajor(layer)=>GatheredHeadProjection::RowMajor {
weight:layer.weight.$gather::<C,2>(communicator.clone())?,
bias:layer.bias.as_ref().map(|bias|bias.$gather::<C,1>(communicator.clone())).transpose()?,
},
};
Ok(GatheredFullyShardedTransformerHead {projection,
normalization:self.normalization.as_ref().map(|norm|norm.$gather(communicator.clone())).transpose()?,dropout:self.dropout.clone()})
}
pub fn $forward<C:BroadcastTensorCollective<B>,const D:usize>(&self,hidden:Tensor<$backend,D>,communicator:C)
-> Result<Tensor<$backend,D>,C::Error> {Ok(self.$gather(communicator)?.forward(hidden))}
}
};
}
gathered_model_parts!(B,[B:Backend],gather_inference,forward_inference);
gathered_model_parts!(Autodiff<B,S>,[B:Backend,S:CheckpointStrategy],gather,forward);