use super::*;
use crate::{activation::Activation,attention::GroupedQueryAttention,
transformer::{AdaptedProjection,AdaptedGroupedQueryAttention,AdaptedFeedForward,DenseFeedForward,DenseTransformerNorm}};
#[derive(Module,Debug)]
pub enum FullyShardedAdaptedProjection<B:Backend> {
Dense(FullyShardedLinear<B>),
LoRA(FullyShardedLoRALinear<B>),
}
#[derive(Module,Debug)]
pub enum FullyShardedActivation<B:Backend> {
Stateless(Activation<B>),
PRelu(FullyShardedPRelu<B>),
SwiGlu(FullyShardedSwiGlu<B>),
}
#[derive(Module,Debug)]
pub struct FullyShardedPRelu<B:Backend> {
pub alpha:ShardedParameter<B>,
pub alpha_value:f64,
}
#[derive(Module,Debug)]
pub struct FullyShardedSwiGlu<B:Backend> {
pub inner:FullyShardedLinear<B>,
pub outer:FullyShardedLinear<B>,
}
#[derive(Module,Debug)]
pub enum FullyShardedTransformerNorm<B:Backend> {
Layer(FullyShardedLayerNorm<B>),
Rms(FullyShardedRmsNorm<B>),
}
#[derive(Module,Debug)]
pub struct FullyShardedGroupedQueryAttention<B:Backend> {
pub query:FullyShardedAdaptedProjection<B>,
pub key:FullyShardedAdaptedProjection<B>,
pub value:FullyShardedAdaptedProjection<B>,
pub output:FullyShardedAdaptedProjection<B>,
pub dropout:crate::Dropout,
pub query_heads:usize,
pub kv_heads:usize,
pub head_dimension:usize,
}
#[derive(Module,Debug)]
pub struct FullyShardedFeedForward<B:Backend> {
pub up:FullyShardedAdaptedProjection<B>,
pub gate:Option<FullyShardedAdaptedProjection<B>>,
pub down:FullyShardedAdaptedProjection<B>,
pub activation:FullyShardedActivation<B>,
pub dropout:crate::Dropout,
}
impl<B:Backend> ShardingContext<B> {
pub fn adapted_projection(&mut self,projection:AdaptedProjection<B>) -> FullyShardedAdaptedProjection<B> {
match projection {
AdaptedProjection::Dense(layer)=>FullyShardedAdaptedProjection::Dense(self.linear(layer)),
AdaptedProjection::LoRA(layer)=>FullyShardedAdaptedProjection::LoRA(self.lora(layer)),
}
}
pub fn activation(&mut self,activation:Activation<B>) -> FullyShardedActivation<B> {
match activation {
Activation::PRelu(layer)=>FullyShardedActivation::PRelu(FullyShardedPRelu {alpha:self.parameter(layer.alpha),alpha_value:layer.alpha_value}),
Activation::SwiGlu(layer)=>FullyShardedActivation::SwiGlu(FullyShardedSwiGlu {inner:self.linear(layer.linear_inner),outer:self.linear(layer.linear_outer)}),
other @ (Activation::Gelu(_) | Activation::Relu(_) | Activation::LeakyRelu(_) | Activation::Selu(_)
| Activation::Sigmoid(_) | Activation::Tanh(_) | Activation::HardSigmoid(_) | Activation::HardSwish(_)
| Activation::Softplus(_) | Activation::Softsign(_) | Activation::Elu(_) | Activation::Celu(_)
| Activation::ThresholdedRelu(_) | Activation::HardShrink(_) | Activation::SoftShrink(_) | Activation::Shrink(_)
| Activation::Silu(_))=>FullyShardedActivation::Stateless(other),
}
}
pub fn normalization(&mut self,norm:DenseTransformerNorm<B>) -> FullyShardedTransformerNorm<B> {
match norm {
DenseTransformerNorm::Layer(layer)=>FullyShardedTransformerNorm::Layer(self.layer_norm(layer)),
DenseTransformerNorm::Rms(layer)=>FullyShardedTransformerNorm::Rms(self.rms_norm(layer)),
}
}
pub fn grouped_attention(&mut self,attention:GroupedQueryAttention<B>) -> FullyShardedGroupedQueryAttention<B> {
self.adapted_attention(AdaptedGroupedQueryAttention {
query:AdaptedProjection::Dense(attention.query),key:AdaptedProjection::Dense(attention.key),
value:AdaptedProjection::Dense(attention.value),output:AdaptedProjection::Dense(attention.output),
dropout:attention.dropout,query_heads:attention.query_heads,kv_heads:attention.kv_heads,head_dimension:attention.head_dimension,
})
}
pub fn adapted_attention(&mut self,attention:AdaptedGroupedQueryAttention<B>) -> FullyShardedGroupedQueryAttention<B> {
FullyShardedGroupedQueryAttention {
query:self.adapted_projection(attention.query),key:self.adapted_projection(attention.key),
value:self.adapted_projection(attention.value),output:self.adapted_projection(attention.output),
dropout:attention.dropout,query_heads:attention.query_heads,kv_heads:attention.kv_heads,head_dimension:attention.head_dimension,
}
}
pub fn feed_forward(&mut self,feed:DenseFeedForward<B>) -> FullyShardedFeedForward<B> {
self.adapted_feed_forward(AdaptedFeedForward {up:AdaptedProjection::Dense(feed.up),
gate:feed.gate.map(AdaptedProjection::Dense),down:AdaptedProjection::Dense(feed.down),activation:feed.activation,dropout:feed.dropout})
}
pub fn adapted_feed_forward(&mut self,feed:AdaptedFeedForward<B>) -> FullyShardedFeedForward<B> {
FullyShardedFeedForward {up:self.adapted_projection(feed.up),gate:feed.gate.map(|gate|self.adapted_projection(gate)),
down:self.adapted_projection(feed.down),activation:self.activation(feed.activation),dropout:feed.dropout}
}
}
macro_rules! gathered_components {
($backend:ty,[$($generics:tt)*],$gather:ident) => {
impl<$($generics)*> FullyShardedAdaptedProjection<$backend> {
pub fn $gather<C:BroadcastTensorCollective<B>>(&self,communicator:C) -> Result<AdaptedProjection<$backend>,C::Error> {
match self {
Self::Dense(layer)=>layer.$gather(communicator).map(AdaptedProjection::Dense),
Self::LoRA(layer)=>layer.$gather(communicator).map(AdaptedProjection::LoRA),
}
}
}
impl<$($generics)*> FullyShardedActivation<$backend> {
pub fn $gather<C:BroadcastTensorCollective<B>>(&self,communicator:C) -> Result<Activation<$backend>,C::Error> {
Ok(match self {
Self::Stateless(layer)=>{
assert!(!matches!(layer,Activation::PRelu(_)|Activation::SwiGlu(_)),"parameterized activation must use local shard storage");
layer.clone()
}
Self::PRelu(layer)=>Activation::PRelu(crate::activation::PRelu {
alpha:Param::initialized(layer.alpha.local.id,layer.alpha.$gather::<C,1>(communicator)?),alpha_value:layer.alpha_value,
}),
Self::SwiGlu(layer)=>Activation::SwiGlu(crate::activation::SwiGlu {
linear_inner:layer.inner.$gather(communicator.clone())?,linear_outer:layer.outer.$gather(communicator)?,
}),
})
}
}
impl<$($generics)*> FullyShardedTransformerNorm<$backend> {
pub fn $gather<C:BroadcastTensorCollective<B>>(&self,communicator:C) -> Result<DenseTransformerNorm<$backend>,C::Error> {
match self {
Self::Layer(layer)=>layer.$gather(communicator).map(DenseTransformerNorm::Layer),
Self::Rms(layer)=>layer.$gather(communicator).map(DenseTransformerNorm::Rms),
}
}
}
impl<$($generics)*> FullyShardedGroupedQueryAttention<$backend> {
pub fn $gather<C:BroadcastTensorCollective<B>>(&self,communicator:C) -> Result<AdaptedGroupedQueryAttention<$backend>,C::Error> {
Ok(AdaptedGroupedQueryAttention {query:self.query.$gather(communicator.clone())?,key:self.key.$gather(communicator.clone())?,
value:self.value.$gather(communicator.clone())?,output:self.output.$gather(communicator)?,dropout:self.dropout.clone(),
query_heads:self.query_heads,kv_heads:self.kv_heads,head_dimension:self.head_dimension})
}
}
impl<$($generics)*> FullyShardedFeedForward<$backend> {
pub fn $gather<C:BroadcastTensorCollective<B>>(&self,communicator:C) -> Result<AdaptedFeedForward<$backend>,C::Error> {
Ok(AdaptedFeedForward {up:self.up.$gather(communicator.clone())?,
gate:self.gate.as_ref().map(|gate|gate.$gather(communicator.clone())).transpose()?,
down:self.down.$gather(communicator.clone())?,activation:self.activation.$gather(communicator)?,dropout:self.dropout.clone()})
}
}
};
}
gathered_components!(B,[B:Backend],gather_inference);
gathered_components!(Autodiff<B,S>,[B:Backend,S:CheckpointStrategy],gather);
impl<B:Backend,S:CheckpointStrategy> FullyShardedFeedForward<Autodiff<B,S>> {
pub fn forward<C:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<Autodiff<B,S>,D>,communicator:C)
-> Result<Tensor<Autodiff<B,S>,D>,C::Error> {
Ok(self.gather(communicator)?.forward(input))
}
}
impl<B:Backend> FullyShardedFeedForward<B> {
pub fn forward_inference<C:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,communicator:C)
-> Result<Tensor<B,D>,C::Error> {
Ok(self.gather_inference(communicator)?.forward(input))
}
}