use ruda_autodiff::{Autodiff,checkpoint::strategy::CheckpointStrategy,tensor_parallel as region};
use ruda_model::{module::Module,tensor::{Tensor,backend::Backend,module::linear}};
use crate::transformer::DenseFeedForward;
use region::BroadcastTensorCollective;
#[derive(Module,Debug)]
pub struct TensorParallelFeedForward<B: Backend> {
pub local: DenseFeedForward<B>,
}
impl<B: Backend> TensorParallelFeedForward<B> {
pub fn from_shard(local: DenseFeedForward<B>) -> Self {
let [width,inner] = local.up.weight.val().dims();
assert!(width > 0 && inner > 0,"parallel feed-forward widths must be positive");
assert_eq!(local.down.weight.val().dims(),[inner,width],"parallel feed-forward output rows differ");
if let Some(gate) = &local.gate {assert_eq!(gate.weight.val().dims(),[width,inner],"parallel gate/value columns differ");}
Self {local}
}
fn partial<const D: usize>(&self,input: Tensor<B,D>) -> Tensor<B,D> {
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(),"parallel gate activation changed local intermediate geometry");
activated*up
} else {self.local.activation.forward(up)};
linear(self.local.dropout.forward(value),self.local.down.weight.val(),None)
}
fn bias<const D: usize>(&self,output: Tensor<B,D>) -> Tensor<B,D> {
if let Some(bias) = &self.local.down.bias {
let mut shape = [1;D];shape[D-1] = bias.val().dims()[0];
output+bias.val().reshape(shape)
} else {output}
}
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,"parallel FFN requires a feature axis");
let output = communicator.all_reduce_sum(self.partial(input).into_primitive().tensor())?;
Ok(self.bias(Tensor::from_primitive(ruda_model::tensor::TensorPrimitive::Float(output))))
}
}
impl<B: Backend,S: CheckpointStrategy> TensorParallelFeedForward<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_activation(input,communicator,|module,input|Ok(module.forward(input)))
}
pub fn forward_with_activation<C,F,const D: usize>(&self,input: Tensor<Autodiff<B,S>,D>,communicator: C,activation: F)
-> Result<Tensor<Autodiff<B,S>,D>,C::Error>
where C: BroadcastTensorCollective<B>,
F: FnOnce(&crate::activation::Activation<Autodiff<B,S>>,Tensor<Autodiff<B,S>,D>)->Result<Tensor<Autodiff<B,S>,D>,C::Error> {
assert!(D > 0,"parallel FFN requires a feature axis");
let input = region::copy_to_region(input,communicator.clone())?;
let up = self.local.up.forward(input.clone());
let value = if let Some(gate) = &self.local.gate {
let activated = activation(&self.local.activation,gate.forward(input))?;
assert_eq!(activated.dims(),up.dims(),"parallel activation changed actual gate geometry");
activated*up
} else {activation(&self.local.activation,up)?};
let partial = linear(self.local.dropout.forward(value),self.local.down.weight.val(),None);
Ok(self.bias(region::reduce_from_region(partial,communicator)?))
}
pub fn forward_with_replicated_activation<C,K,const D: usize>(&self,input: Tensor<Autodiff<B,S>,D>,communicator: C,activation_group: K)
-> Result<Tensor<Autodiff<B,S>,D>,C::Error>
where C: BroadcastTensorCollective<B>,K: BroadcastTensorCollective<B,Error=C::Error> {
self.forward_with_activation(input,communicator,|module,input|
super::copy_replicated_module_to_region::<B,S,K,_>(module.clone(),activation_group).map(|module|module.forward(input)))
}
}