use super::{Autodiff,Backend,BroadcastTensorCollective,CheckpointStrategy,Dropout,Module,Tensor,column,row,row_inference,geometry};
use crate::transformer::{AdaptedFeedForward,FeedForwardAdapterTarget};
use crate::transformer::NativeFeedForwardError;
use ruda_model::tensor::NativeSwiGluOps;
#[derive(Module,Debug)]
pub struct TensorParallelAdaptedFeedForward<B: Backend> {
pub local: AdaptedFeedForward<B>,
}
impl<B: Backend> TensorParallelAdaptedFeedForward<B> {
pub fn from_shard(local: AdaptedFeedForward<B>) -> Self {
let [width,inner] = geometry(&local.up);
assert!(width > 0 && inner > 0,"adapted parallel FFN widths must be positive");
assert_eq!(geometry(&local.down),[inner,width],"adapted parallel FFN output rows differ");
if let Some(gate) = &local.gate {assert_eq!(geometry(gate),[width,inner],"adapted parallel gate columns differ");}
Self {local}
}
pub fn forward_inference<C: BroadcastTensorCollective<B>,const D: usize>(&self,input: Tensor<B,D>,communicator: C)
-> Result<Tensor<B,D>,C::Error> {
assert!(D > 0,"adapted native parallel FFN requires a feature axis");
let up = self.local.up.forward(input.clone());
let value = if let Some(gate) = &self.local.gate {
let activated = self.local.activation.forward(gate.forward(input));
assert_eq!(activated.dims(),up.dims(),"adapted native gate changed actual local geometry");
activated*up
} else {self.local.activation.forward(up)};
row_inference(&self.local.down,self.local.dropout.forward(value),&communicator)
}
}
impl<B: Backend,S: CheckpointStrategy> TensorParallelAdaptedFeedForward<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> {
self.forward_with_adapter_dropout(input,communicator,|_,module,input|module.forward(input))
}
pub fn forward_with_adapter_dropout<C,F,const D: usize>(&self,input: Tensor<Autodiff<B,S>,D>,communicator: C,mut dropout: F)
-> Result<Tensor<Autodiff<B,S>,D>,C::Error>
where C: BroadcastTensorCollective<B>,F: FnMut(FeedForwardAdapterTarget,&Dropout,Tensor<Autodiff<B,S>,D>)->Tensor<Autodiff<B,S>,D> {
self.forward_with_transforms(input,communicator,&mut dropout,|module,input|Ok(module.forward(input)))
}
pub fn forward_with_transforms<C,F,A,const D: usize>(&self,input: Tensor<Autodiff<B,S>,D>,communicator: C,mut dropout: F,activation: A)
-> Result<Tensor<Autodiff<B,S>,D>,C::Error>
where C: BroadcastTensorCollective<B>,F: FnMut(FeedForwardAdapterTarget,&Dropout,Tensor<Autodiff<B,S>,D>)->Tensor<Autodiff<B,S>,D>,
A: FnOnce(&crate::activation::Activation<Autodiff<B,S>>,Tensor<Autodiff<B,S>,D>)->Result<Tensor<Autodiff<B,S>,D>,C::Error> {
assert!(D > 0,"adapted parallel FFN requires an actual feature axis");
let up = column(&self.local.up,input.clone(),&communicator,None::<&C>,|module,input|dropout(FeedForwardAdapterTarget::Up,module,input))?;
let value = if let Some(gate) = &self.local.gate {
let gate = column(gate,input,&communicator,None::<&C>,|module,input|dropout(FeedForwardAdapterTarget::Gate,module,input))?;
let activated = activation(&self.local.activation,gate)?;
assert_eq!(activated.dims(),up.dims(),"adapted parallel gate activation changed actual local geometry");
activated*up
} else {activation(&self.local.activation,up)?};
row(&self.local.down,self.local.dropout.forward(value),&communicator,|module,input|dropout(FeedForwardAdapterTarget::Down,module,input))
}
pub fn forward_with_replicated_activation<C,K,F,const D: usize>(&self,input: Tensor<Autodiff<B,S>,D>,communicator: C,activation_group: K,dropout: F)
-> Result<Tensor<Autodiff<B,S>,D>,C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error>,
F: FnMut(FeedForwardAdapterTarget,&Dropout,Tensor<Autodiff<B,S>,D>)->Tensor<Autodiff<B,S>,D> {
self.forward_with_transforms(input,communicator,dropout,|module,input|
super::super::copy_replicated_module_to_region::<B,S,K,_>(module.clone(),activation_group).map(|module|module.forward(input)))
}
}
impl<B: NativeSwiGluOps> TensorParallelAdaptedFeedForward<B> {
pub fn try_forward_native_inference<C: BroadcastTensorCollective<B>, const D: usize>(
&self, input: Tensor<B, D>, communicator: C,
) -> Result<Tensor<B, D>, NativeFeedForwardError<C::Error, B::SwiGluError>> {
assert!(D > 0, "adapted parallel FFN requires a feature axis");
let up = self.local.up.forward(input.clone());
let value = if let Some(gate) = &self.local.gate {
self.local.activation.try_forward_gated_native(gate.forward(input), up)
} else { self.local.activation.try_forward_native(up) }.map_err(NativeFeedForwardError::Activation)?;
row_inference(&self.local.down, self.local.dropout.forward(value), &communicator)
.map_err(NativeFeedForwardError::Execution)
}
}
impl<B: NativeSwiGluOps, S: CheckpointStrategy> TensorParallelAdaptedFeedForward<Autodiff<B, S>> {
pub fn try_forward_native<C: BroadcastTensorCollective<B>, const D: usize>(
&self, input: Tensor<Autodiff<B, S>, D>, communicator: C,
) -> Result<Tensor<Autodiff<B, S>, D>, NativeFeedForwardError<C::Error,
<Autodiff<B, S> as NativeSwiGluOps>::SwiGluError>> {
self.try_forward_native_with_adapter_dropout(input, communicator, |_, module, input| module.forward(input))
}
pub fn try_forward_native_with_adapter_dropout<C, F, const D: usize>(
&self, input: Tensor<Autodiff<B, S>, D>, communicator: C, dropout: F,
) -> Result<Tensor<Autodiff<B, S>, D>, NativeFeedForwardError<C::Error,
<Autodiff<B, S> as NativeSwiGluOps>::SwiGluError>>
where C: BroadcastTensorCollective<B>,
F: FnMut(FeedForwardAdapterTarget, &Dropout, Tensor<Autodiff<B, S>, D>) -> Tensor<Autodiff<B, S>, D> {
self.try_forward_native_with_activation(input, communicator, dropout, |module, input, up| {
match up {
Some(up) => module.try_forward_gated_native(input, up),
None => module.try_forward_native(input),
}.map_err(NativeFeedForwardError::Activation)
})
}
pub fn try_forward_native_with_replicated_activation<C, K, F, const D: usize>(
&self, input: Tensor<Autodiff<B, S>, D>, communicator: C, activation_group: K, dropout: F,
) -> Result<Tensor<Autodiff<B, S>, D>, NativeFeedForwardError<C::Error,
<Autodiff<B, S> as NativeSwiGluOps>::SwiGluError>>
where C: BroadcastTensorCollective<B>, K: BroadcastTensorCollective<B, Error = C::Error>,
F: FnMut(FeedForwardAdapterTarget, &Dropout, Tensor<Autodiff<B, S>, D>) -> Tensor<Autodiff<B, S>, D> {
self.try_forward_native_with_activation(input, communicator, dropout, |module, input, up| {
let module = super::super::copy_replicated_module_to_region::<B, S, K, _>(module.clone(), activation_group)
.map_err(NativeFeedForwardError::Execution)?;
match up {
Some(up) => module.try_forward_gated_native(input, up),
None => module.try_forward_native(input),
}.map_err(NativeFeedForwardError::Activation)
})
}
fn try_forward_native_with_activation<C, F, A, const D: usize>(
&self, input: Tensor<Autodiff<B, S>, D>, communicator: C, mut dropout: F, activation: A,
) -> Result<Tensor<Autodiff<B, S>, D>, NativeFeedForwardError<C::Error,
<Autodiff<B, S> as NativeSwiGluOps>::SwiGluError>>
where C: BroadcastTensorCollective<B>,
F: FnMut(FeedForwardAdapterTarget, &Dropout, Tensor<Autodiff<B, S>, D>) -> Tensor<Autodiff<B, S>, D>,
A: FnOnce(&crate::activation::Activation<Autodiff<B, S>>, Tensor<Autodiff<B, S>, D>,
Option<Tensor<Autodiff<B, S>, D>>) -> Result<Tensor<Autodiff<B, S>, D>,
NativeFeedForwardError<C::Error, <Autodiff<B, S> as NativeSwiGluOps>::SwiGluError>> {
assert!(D > 0, "adapted parallel FFN requires a feature axis");
let up = column(&self.local.up, input.clone(), &communicator, None::<&C>,
|module, input| dropout(FeedForwardAdapterTarget::Up, module, input))
.map_err(NativeFeedForwardError::Execution)?;
let value = if let Some(gate) = &self.local.gate {
let gate = column(gate, input, &communicator, None::<&C>,
|module, input| dropout(FeedForwardAdapterTarget::Gate, module, input))
.map_err(NativeFeedForwardError::Execution)?;
activation(&self.local.activation, gate, Some(up))?
} else { activation(&self.local.activation, up, None)? };
row(&self.local.down, self.local.dropout.forward(value), &communicator,
|module, input| dropout(FeedForwardAdapterTarget::Down, module, input))
.map_err(NativeFeedForwardError::Execution)
}
}