use super::*;
use ruda_model::tensor::{Bool,IntegerTensorCollective};
use crate::transformer::TransformerProjection;
use crate::{attention::{DenseAttentionMask,DenseAttentionOptions,PackedSequenceLayout,PackedAttentionOptions,PackedDocumentAttentionMask},
cache::TransformerKvCache,transformer::{AwqTransformerHead,AwqTransformerModel}};
use crate::loss::{CausalCrossEntropyConfig,CausalLoss};
use ruda_autodiff::collective::{CollectiveScope,ScopedCollectiveError};
#[derive(Module,Debug)]
pub struct FullyShardedAwqTransformerHead<B:Backend,P:Module<B>=FullyShardedAwqProjection<B>> {
pub projection:P,
pub normalization:Option<FullyShardedTransformerNorm<B>>,
pub dropout:crate::Dropout,
}
#[derive(Module,Debug)]
pub struct FullyShardedAwqTransformerModel<B:Backend,P:Module<B>=FullyShardedAwqProjection<B>> {
pub embeddings:FullyShardedTransformerEmbeddings<B>,
pub backbone:FullyShardedAwqTransformerStack<B,P>,
pub normalization:Option<FullyShardedTransformerNorm<B>>,
pub head:FullyShardedAwqTransformerHead<B,P>,
}
impl<B:Backend> ShardingContext<B> {
pub fn awq_transformer_head<P:ShardTransformerProjection<B>>(&mut self,head:AwqTransformerHead<B,P>) -> FullyShardedAwqTransformerHead<B,P::Sharded> {
FullyShardedAwqTransformerHead {projection:self.awq_projection(head.projection),
normalization:head.normalization.map(|norm|self.normalization(norm)),dropout:head.dropout}
}
pub fn awq_transformer_model<P:ShardTransformerProjection<B>>(&mut self,model:AwqTransformerModel<B,P>) -> FullyShardedAwqTransformerModel<B,P::Sharded> {
FullyShardedAwqTransformerModel {embeddings:self.transformer_embeddings(model.embeddings),
backbone:self.awq_transformer_stack(model.backbone),normalization:model.normalization.map(|norm|self.normalization(norm)),
head:self.awq_transformer_head(model.head)}
}
}
impl<B:Backend,P:Module<B>> FullyShardedAwqTransformerModel<B,P> {
pub fn from_full<Q:ShardTransformerProjection<B,Sharded=P>>(model:AwqTransformerModel<B,Q>,rank:usize,world:usize) -> Self {ShardingContext::new(rank,world).awq_transformer_model(model)}
pub fn new_kv_cache(&self,initial_capacity:usize) -> TransformerKvCache<B> {self.backbone.new_kv_cache(initial_capacity)}
}
macro_rules! awq_model_execution {
($backend:ty,[$($generics:tt)*],$gather:ident,$embed:ident,$forward:ident,$packed:ident,$hidden:ident,$packed_hidden:ident) => {
impl<$($generics)*,P:GatherTransformerProjection<$backend,B>> FullyShardedAwqTransformerHead<$backend,P>
where P::Gathered:TransformerProjection<$backend> {
pub fn $gather<C:IntegerTensorCollective<B>>(&self,communicator:C) -> Result<AwqTransformerHead<$backend,P::Gathered>,C::Error> {
Ok(AwqTransformerHead::from_projection(self.projection.gather_projection(communicator.clone())?,
self.normalization.as_ref().map(|norm|norm.$gather(communicator)).transpose()?,self.dropout.clone()))
}
pub fn $forward<C:IntegerTensorCollective<B>,const D:usize>(&self,hidden:Tensor<$backend,D>,communicator:C)
-> Result<Tensor<$backend,D>,FullyShardedAwqError<C::Error,<P::Gathered as TransformerProjection<$backend>>::Error>> {
self.$gather(communicator).map_err(FullyShardedAwqError::Collective)?.forward(hidden).map_err(FullyShardedAwqError::Projection)
}
}
impl<$($generics)*,P:GatherTransformerProjection<$backend,B>> FullyShardedAwqTransformerModel<$backend,P>
where P::Gathered:TransformerProjection<$backend> {
pub fn $forward<C,F>(&self,input:FullyShardedTransformerInput<$backend>,masks:DenseAttentionMask<$backend>,
options:DenseAttentionOptions,communicator:C,positions:F)
-> Result<Tensor<$backend,3>,FullyShardedAwqError<C::Error,<P::Gathered as TransformerProjection<$backend>>::Error>>
where C:IntegerTensorCollective<B>,F:FnMut(usize,Tensor<$backend,4>,Tensor<$backend,4>)->(Tensor<$backend,4>,Tensor<$backend,4>) {
let hidden=self.$hidden(input,masks,options,communicator.clone(),positions)?;
self.head.$forward(hidden,communicator)
}
pub fn $hidden<C,F>(&self,input:FullyShardedTransformerInput<$backend>,masks:DenseAttentionMask<$backend>,
options:DenseAttentionOptions,communicator:C,positions:F)
-> Result<Tensor<$backend,3>,FullyShardedAwqError<C::Error,<P::Gathered as TransformerProjection<$backend>>::Error>>
where C:IntegerTensorCollective<B>,F:FnMut(usize,Tensor<$backend,4>,Tensor<$backend,4>)->(Tensor<$backend,4>,Tensor<$backend,4>) {
let hidden=self.embeddings.$embed(input,communicator.clone()).map_err(FullyShardedAwqError::Collective)?;
let hidden=self.backbone.$forward(hidden,masks,options,communicator.clone(),positions)?;
let hidden=if let Some(norm)=&self.normalization {norm.$gather(communicator.clone()).map_err(FullyShardedAwqError::Collective)?.forward(hidden)} else {hidden};
Ok(hidden)
}
pub fn $packed<C,F>(&self,input:FullyShardedTransformerInput<$backend,1>,layout:&PackedSequenceLayout,
masks:Option<&[PackedDocumentAttentionMask<$backend>]>,options:PackedAttentionOptions,communicator:C,positions:F)
-> Result<Tensor<$backend,2>,FullyShardedAwqError<C::Error,<P::Gathered as TransformerProjection<$backend>>::Error>>
where C:IntegerTensorCollective<B>,F:FnMut(usize,Tensor<$backend,3>,Tensor<$backend,3>)->(Tensor<$backend,3>,Tensor<$backend,3>) {
let hidden=self.$packed_hidden(input,layout,masks,options,communicator.clone(),positions)?;
self.head.$forward(hidden,communicator)
}
pub fn $packed_hidden<C,F>(&self,input:FullyShardedTransformerInput<$backend,1>,layout:&PackedSequenceLayout,
masks:Option<&[PackedDocumentAttentionMask<$backend>]>,options:PackedAttentionOptions,communicator:C,positions:F)
-> Result<Tensor<$backend,2>,FullyShardedAwqError<C::Error,<P::Gathered as TransformerProjection<$backend>>::Error>>
where C:IntegerTensorCollective<B>,F:FnMut(usize,Tensor<$backend,3>,Tensor<$backend,3>)->(Tensor<$backend,3>,Tensor<$backend,3>) {
let input=super::model::packed_input(input,layout);
let hidden=self.embeddings.$embed(input,communicator.clone()).map_err(FullyShardedAwqError::Collective)?
.reshape([layout.tokens(),self.embeddings.hidden_width()]);
let hidden=self.backbone.$packed(hidden,layout,masks,options,communicator.clone(),positions)?;
let hidden=if let Some(norm)=&self.normalization {norm.$gather(communicator.clone()).map_err(FullyShardedAwqError::Collective)?.forward(hidden)} else {hidden};
Ok(hidden)
}
}
};
}
awq_model_execution!(B,[B:Backend],gather_inference,forward_inference,forward_inference,forward_packed_inference,forward_hidden_inference,forward_packed_hidden_inference);
awq_model_execution!(Autodiff<B,S>,[B:Backend,S:CheckpointStrategy],gather,forward,forward,forward_packed,forward_hidden,forward_packed_hidden);
impl<B:Backend,P:GatherTransformerProjection<B,B>> FullyShardedAwqTransformerModel<B,P>
where P::Gathered:TransformerProjection<B> {
pub fn forward_cached_inference<C,F>(&self,input:FullyShardedTransformerInput<B>,visible:Option<Tensor<B,2,Bool>>,
cache:&mut TransformerKvCache<B>,masks:DenseAttentionMask<B>,options:DenseAttentionOptions,communicator:C,mut positions:F)
-> Result<Tensor<B,3>,FullyShardedAwqError<C::Error,<P::Gathered as TransformerProjection<B>>::Error>>
where C:IntegerTensorCollective<B>,F:FnMut(usize,Tensor<B,4>,Tensor<B,4>,usize)->(Tensor<B,4>,Tensor<B,4>) {
let mut hidden=self.embeddings.forward_inference(input,communicator.clone()).map_err(FullyShardedAwqError::Collective)?;
cache.validate_layers(self.backbone.blocks.len());let rows=(hidden.dims()[0],hidden.dims()[1]);
let next=cache.position().checked_add(rows.1).expect("cached sharded model position overflows");
for (index,block) in self.backbone.blocks.iter().enumerate() {
hidden=block.forward_cached_inference(hidden,visible.clone(),&mut cache.layers_mut()[index],masks.clone(),options,communicator.clone(),
|query,key,position|positions(index,query,key,position))?;
assert_eq!((hidden.dims()[0],hidden.dims()[1]),rows,"cached model layer changed actual rows");
}
cache.finish_chunk(next);
let hidden=if let Some(norm)=&self.normalization {norm.gather_inference(communicator.clone()).map_err(FullyShardedAwqError::Collective)?.forward(hidden)} else {hidden};
self.head.forward_inference(hidden,communicator)
}
}
#[derive(Debug)]
pub enum FullyShardedAwqTrainingError<C:core::fmt::Debug,Q:core::fmt::Debug> {
Model(FullyShardedAwqError<C,Q>),
Loss(ScopedCollectiveError<C>),
}
impl<C:core::fmt::Debug,Q:core::fmt::Debug> core::fmt::Display for FullyShardedAwqTrainingError<C,Q> {
fn fmt(&self,f:&mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {Self::Model(error)=>write!(f,"AWQ training model: {error}"),Self::Loss(error)=>write!(f,"AWQ training loss: {error}")}
}
}
impl<C:core::fmt::Debug,Q:core::fmt::Debug> core::error::Error for FullyShardedAwqTrainingError<C,Q> {}
macro_rules! awq_head_loss {
($backend:ty,[$($generics:tt)*],$gather:ident,$causal:ident,$packed:ident) => {
impl<$($generics)*,P:GatherTransformerProjection<$backend,B>> FullyShardedAwqTransformerHead<$backend,P>
where P::Gathered:TransformerProjection<$backend> {
pub fn $causal<C:IntegerTensorCollective<B>>(&self,hidden:Tensor<$backend,3>,labels:Tensor<$backend,2,Int>,
criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C)
-> Result<CausalLoss<$backend>,FullyShardedAwqError<C::Error,<P::Gathered as TransformerProjection<$backend>>::Error>> {
let head=self.$gather(communicator).map_err(FullyShardedAwqError::Collective)?;
criterion.try_forward_hidden_with_smoothing(hidden,labels,|rows|head.forward(rows),label_smoothing).map_err(FullyShardedAwqError::Projection)
}
pub fn $packed<C:IntegerTensorCollective<B>>(&self,hidden:Tensor<$backend,2>,labels:Tensor<$backend,1,Int>,layout:&PackedSequenceLayout,
criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C)
-> Result<CausalLoss<$backend>,FullyShardedAwqError<C::Error,<P::Gathered as TransformerProjection<$backend>>::Error>> {
let head=self.$gather(communicator).map_err(FullyShardedAwqError::Collective)?;
criterion.try_forward_packed_hidden_with_smoothing(hidden,labels,layout,|rows|head.forward(rows),label_smoothing).map_err(FullyShardedAwqError::Projection)
}
}
};
}
awq_head_loss!(B,[B:Backend],gather_inference,forward_causal_loss_inference,forward_packed_causal_loss_inference);
awq_head_loss!(Autodiff<B,S>,[B:Backend,S:CheckpointStrategy],gather,forward_causal_loss,forward_packed_causal_loss);
impl<B:Backend,S:CheckpointStrategy,P:GatherTransformerProjection<Autodiff<B,S>,B>> FullyShardedAwqTransformerModel<Autodiff<B,S>,P>
where P::Gathered:TransformerProjection<Autodiff<B,S>> {
pub fn forward_causal_with_positions<C,F>(&self,input:FullyShardedTransformerInput<Autodiff<B,S>>,labels:Tensor<Autodiff<B,S>,2,Int>,
masks:DenseAttentionMask<Autodiff<B,S>>,options:DenseAttentionOptions,criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C,positions:F)
-> Result<FullyShardedLoss<B,S>,FullyShardedAwqTrainingError<C::Error,<P::Gathered as TransformerProjection<Autodiff<B,S>>>::Error>>
where C:IntegerTensorCollective<B>,F:FnMut(usize,Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>)->(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>) {
let scope=CollectiveScope::<B,S>::new();let transport=scope.bind(communicator.clone());
let hidden=self.forward_hidden(input,masks,options,transport.clone(),positions).map_err(FullyShardedAwqTrainingError::Model)?;
let loss=self.head.forward_causal_loss(hidden,labels,criterion,label_smoothing,transport).map_err(FullyShardedAwqTrainingError::Model)?;
complete_fully_sharded_loss(&scope,loss.loss_sum,loss.valid_tokens,communicator).map_err(FullyShardedAwqTrainingError::Loss)
}
pub fn forward_packed_causal_with_positions<C,F>(&self,input:FullyShardedTransformerInput<Autodiff<B,S>,1>,labels:Tensor<Autodiff<B,S>,1,Int>,
layout:&PackedSequenceLayout,masks:Option<&[PackedDocumentAttentionMask<Autodiff<B,S>>]>,options:PackedAttentionOptions,
criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C,positions:F)
-> Result<FullyShardedLoss<B,S>,FullyShardedAwqTrainingError<C::Error,<P::Gathered as TransformerProjection<Autodiff<B,S>>>::Error>>
where C:IntegerTensorCollective<B>,F:FnMut(usize,Tensor<Autodiff<B,S>,3>,Tensor<Autodiff<B,S>,3>)->(Tensor<Autodiff<B,S>,3>,Tensor<Autodiff<B,S>,3>) {
let scope=CollectiveScope::<B,S>::new();let transport=scope.bind(communicator.clone());
let hidden=self.forward_packed_hidden(input,layout,masks,options,transport.clone(),positions).map_err(FullyShardedAwqTrainingError::Model)?;
let loss=self.head.forward_packed_causal_loss(hidden,labels,layout,criterion,label_smoothing,transport).map_err(FullyShardedAwqTrainingError::Model)?;
complete_fully_sharded_loss(&scope,loss.loss_sum,loss.valid_tokens,communicator).map_err(FullyShardedAwqTrainingError::Loss)
}
}