use super::*;
impl<B:Backend> ShardedParameter<B> {
pub fn gather_inference<C:BroadcastTensorCollective<B>,const D:usize>(&self,communicator:C)
-> Result<Tensor<B,D>,C::Error> {
self.gather_inference_inner(communicator,None)
}
pub fn gather_inference_with_compute_dtype<C:BroadcastTensorCollective<B>,const D:usize>(&self,communicator:C,dtype:FloatDType)
-> Result<Tensor<B,D>,C::Error> {
self.gather_inference_inner(communicator,Some(dtype.into()))
}
fn gather_inference_inner<C:BroadcastTensorCollective<B>,const D:usize>(&self,communicator:C,dtype:Option<DType>)
-> Result<Tensor<B,D>,C::Error> {
assert_eq!(communicator.rank() as usize,self.rank,"native parameter shard rank differs");
assert_eq!(communicator.world_size() as usize,self.world_size,"native parameter shard world size differs");
let shape:[usize;D] = self.logical_shape.clone().try_into().expect("native logical parameter rank differs");
let elements = self.logical_shape.iter().try_fold(1usize,|total,axis|total.checked_mul(*axis)).expect("native logical parameter size overflow");
let local = self.local.val();
assert_eq!(local.dims(),[elements.div_ceil(self.world_size)],"loaded native local parameter slice length differs");
let local = if let Some(dtype) = dtype {local.cast(dtype)} else {local};
let dtype = local.dtype();let device = local.device();
let padded = local.dims()[0].checked_mul(self.world_size).expect("native parameter gather length overflow");
let output = communicator.all_gather_float(local.into_primitive().tensor())?;
let output = Tensor::<B,1>::from_primitive(TensorPrimitive::Float(output));
assert_eq!(output.dims(),[padded],"native parameter transport gather shape differs");
assert_eq!(output.dtype(),dtype,"native parameter transport changed precision");
assert_eq!(output.device(),device,"native parameter transport changed device");
Ok(output.slice([0..elements]).reshape(shape))
}
}
impl<B:Backend> FullyShardedLinear<B> {
pub fn forward_inference<C:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,communicator:C)
-> Result<Tensor<B,D>,C::Error> {
let weight = self.weight.gather_inference::<C,2>(communicator.clone())?;
let bias = self.bias.as_ref().map(|bias|bias.gather_inference::<C,1>(communicator)).transpose()?;
Ok(linear(input,weight,bias))
}
pub fn forward_inference_with_compute_dtype<C:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,communicator:C,dtype:FloatDType)
-> Result<Tensor<B,D>,C::Error> {
let storage = input.dtype();
let weight = self.weight.gather_inference_with_compute_dtype::<C,2>(communicator.clone(),dtype)?;
let bias = self.bias.as_ref().map(|bias|bias.gather_inference_with_compute_dtype::<C,1>(communicator,dtype)).transpose()?;
Ok(linear(input.cast(dtype),weight,bias).cast(storage))
}
}
impl<B:Backend> FullyShardedEmbedding<B> {
pub fn forward_inference<C:BroadcastTensorCollective<B>>(&self,input:Tensor<B,2,Int>,communicator:C)
-> Result<Tensor<B,3>,C::Error> {
Ok(embedding(self.weight.gather_inference::<C,2>(communicator)?,input))
}
pub fn forward_inference_with_compute_dtype<C:BroadcastTensorCollective<B>>(&self,input:Tensor<B,2,Int>,communicator:C,dtype:FloatDType)
-> Result<Tensor<B,3>,C::Error> {
Ok(embedding(self.weight.gather_inference_with_compute_dtype::<C,2>(communicator,dtype)?,input))
}
}
impl<B:Backend> FullyShardedProjection<B> {
pub fn forward_inference<C:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,communicator:C)
-> Result<Tensor<B,D>,C::Error> {
let weight = self.weight.gather_inference::<C,2>(communicator.clone())?.transpose();
let bias = self.bias.as_ref().map(|bias|bias.gather_inference::<C,1>(communicator)).transpose()?;
Ok(linear(input,weight,bias))
}
pub fn forward_inference_with_compute_dtype<C:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,communicator:C,dtype:FloatDType)
-> Result<Tensor<B,D>,C::Error> {
let storage = input.dtype();
let weight = self.weight.gather_inference_with_compute_dtype::<C,2>(communicator.clone(),dtype)?.transpose();
let bias = self.bias.as_ref().map(|bias|bias.gather_inference_with_compute_dtype::<C,1>(communicator,dtype)).transpose()?;
Ok(linear(input.cast(dtype),weight,bias).cast(storage))
}
}
impl<B:Backend> FullyShardedLoRALinear<B> {
pub fn forward_inference<C:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,communicator:C)
-> Result<Tensor<B,D>,C::Error> {
let adapted = self.dropout.forward(input.clone().cast(self.adapter_a.weight.local.val().dtype()));
let base = self.base.forward_inference(input,communicator.clone())?;
let hidden = self.adapter_a.forward_inference(adapted,communicator.clone())?.cast(self.adapter_b.weight.local.val().dtype());
let update = self.adapter_b.forward_inference(hidden,communicator)?.mul_scalar(self.scale);
let storage = base.dtype();Ok(base+update.cast(storage))
}
}
impl<B:Backend> FullyShardedGatedMLP<B> {
pub fn forward_inference_with<C:BroadcastTensorCollective<B>,F,const D:usize>(&self,input:Tensor<B,D>,communicator:C,activation:F)
-> Result<Tensor<B,D>,C::Error> where F:FnOnce(Tensor<B,D>)->Tensor<B,D> {
let gate = activation(self.gate.forward_inference(input.clone(),communicator.clone())?);
let up = self.up.forward_inference(input,communicator.clone())?;
assert_eq!(gate.dims(),up.dims(),"native sharded gate activation changed geometry");
self.down.forward_inference(gate*up,communicator)
}
pub fn forward_inference<C:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,communicator:C)
-> Result<Tensor<B,D>,C::Error> {
self.forward_inference_with(input,communicator,ruda_model::tensor::activation::silu)
}
}
impl<B:Backend> FullyShardedRmsNorm<B> {
pub fn forward_inference<C:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,communicator:C)
-> Result<Tensor<B,D>,C::Error> {
self.forward_inference_with_compute_dtype(input,communicator,FloatDType::F32)
}
pub fn forward_inference_with_compute_dtype<C:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,communicator:C,dtype:FloatDType)
-> Result<Tensor<B,D>,C::Error> {
assert!(D>0 && self.gamma.logical_shape==[input.dims()[D-1]],"native RMSNorm affine width differs");
let storage = input.dtype();
let gamma = self.gamma.gather_inference_with_compute_dtype::<C,1>(communicator,dtype)?;
let input = input.cast(dtype);let rms = (input.clone().square().mean_dim(D-1)+self.epsilon).sqrt();
Ok(((input/rms)*gamma.unsqueeze::<D>()).cast(storage))
}
}
impl<B:Backend> FullyShardedLayerNorm<B> {
pub fn forward_inference<C:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,communicator:C)
-> Result<Tensor<B,D>,C::Error> {
self.forward_inference_with_compute_dtype(input,communicator,FloatDType::F32)
}
pub fn forward_inference_with_compute_dtype<C:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,communicator:C,dtype:FloatDType)
-> Result<Tensor<B,D>,C::Error> {
assert!(D>0 && self.gamma.logical_shape==[input.dims()[D-1]],"native LayerNorm affine width differs");
let storage = input.dtype();
let gamma = self.gamma.gather_inference_with_compute_dtype::<C,1>(communicator.clone(),dtype)?.into_primitive().tensor();
let beta = self.beta.as_ref().map(|beta| {
assert_eq!(beta.logical_shape,self.gamma.logical_shape,"native LayerNorm affine bias width differs");
beta.gather_inference_with_compute_dtype::<C,1>(communicator,dtype).map(|beta|beta.into_primitive().tensor())
}).transpose()?;
let output = <B as ModuleOps<B>>::layer_norm(input.cast(dtype).into_primitive().tensor(),gamma,beta,self.epsilon);
Ok(Tensor::<B,D>::from_primitive(TensorPrimitive::Float(output)).cast(storage))
}
}