use super::*;
use tensor_parallel::{inference_gather,inference_scatter,inference_sum,VocabParallelLossLayout};
impl<B:Backend> FullyShardedColumnParallelLinear<B> {
pub fn forward_inference<C:BroadcastTensorCollective<B>,T:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,
data:C,tensor:T,gather_output:bool) -> Result<Tensor<B,D>,HybridParallelError<C::Error,T::Error>> {
assert!(D>0,"native hybrid column projection needs a feature axis");
let output = self.local.forward_inference(input,data).map_err(HybridParallelError::Data)?;
if gather_output {inference_gather(output,tensor,D-1).map_err(HybridParallelError::Tensor)} else {Ok(output)}
}
}
impl<B:Backend> FullyShardedRowParallelLinear<B> {
pub fn forward_inference<C:BroadcastTensorCollective<B>,T:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,
data:C,tensor:T,input_is_parallel:bool) -> Result<Tensor<B,D>,HybridParallelError<C::Error,T::Error>> {
assert!(D>0,"native hybrid row projection needs a feature axis");
let input = if input_is_parallel {input} else {inference_scatter(input,&tensor,D-1)};
let weight = self.local.weight.gather_inference::<C,2>(data.clone()).map_err(HybridParallelError::Data)?;
let output = inference_sum(linear(input,weight,None),tensor).map_err(HybridParallelError::Tensor)?;
Ok(if let Some(bias) = &self.local.bias {
let bias = bias.gather_inference::<C,1>(data).map_err(HybridParallelError::Data)?;
let mut shape = [1;D];shape[D-1] = bias.dims()[0];output+bias.reshape(shape)
} else {output})
}
}
impl<B:Backend> FullyShardedColumnParallelLoRA<B> {
pub fn forward_inference<C:BroadcastTensorCollective<B>,T:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,
adapter_input:Option<Tensor<B,D>>,data:C,tensor:T,gather_output:bool) -> Result<Tensor<B,D>,HybridParallelError<C::Error,T::Error>> {
assert!(D>0,"native hybrid adapter needs a feature axis");
let adapted = adapter_input.unwrap_or_else(||input.clone());assert_eq!(adapted.dims(),input.dims(),"native hybrid adapter input geometry differs");
let base = self.base.forward_inference(input,data.clone(),tensor.clone(),false)?;
let storage = base.dtype();
let a = self.adapter_a.weight.gather_inference::<C,2>(data.clone()).map_err(HybridParallelError::Data)?;
let hidden = linear(adapted.cast(a.dtype()),a,None).cast(self.adapter_b.weight.local.val().dtype());
let update = self.adapter_b.forward_inference(hidden,data).map_err(HybridParallelError::Data)?;
let output = base+update.mul_scalar(self.scale).cast(storage);
if gather_output {inference_gather(output,tensor,D-1).map_err(HybridParallelError::Tensor)} else {Ok(output)}
}
}
impl<B:Backend> FullyShardedRowParallelLoRA<B> {
pub fn forward_inference<C:BroadcastTensorCollective<B>,T:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,
adapter_input:Option<Tensor<B,D>>,data:C,tensor:T,input_is_parallel:bool) -> Result<Tensor<B,D>,HybridParallelError<C::Error,T::Error>> {
assert!(D>0,"native hybrid adapter needs a feature axis");
let adapted = adapter_input.unwrap_or_else(||input.clone());assert_eq!(adapted.dims(),input.dims(),"native hybrid adapter input geometry differs");
let (input,adapted) = if input_is_parallel {(input,adapted)} else {
(inference_scatter(input,&tensor,D-1),inference_scatter(adapted,&tensor,D-1))
};
let base = self.base.forward_inference(input,data.clone(),tensor.clone(),true)?;
let storage = base.dtype();
let hidden = self.adapter_a.forward_inference(adapted.cast(self.adapter_a.weight.local.val().dtype()),data.clone()).map_err(HybridParallelError::Data)?;
let hidden = inference_sum(hidden,tensor).map_err(HybridParallelError::Tensor)?.cast(self.adapter_b.weight.local.val().dtype());
let update = self.adapter_b.forward_inference(hidden,data).map_err(HybridParallelError::Data)?;
Ok(base+update.mul_scalar(self.scale).cast(storage))
}
}
impl<B:Backend> FullyShardedTensorParallelGatedMlp<B> {
pub fn forward_inference_with<C:BroadcastTensorCollective<B>,T:BroadcastTensorCollective<B>,F,const D:usize>(&self,input:Tensor<B,D>,data:C,tensor:T,activation:F)
-> Result<Tensor<B,D>,HybridParallelError<C::Error,T::Error>> where F:FnOnce(Tensor<B,D>)->Tensor<B,D> {
let gate = self.gate.forward_inference(input.clone(),data.clone(),tensor.clone(),false)?;
let up = self.up.forward_inference(input,data.clone(),tensor.clone(),false)?;
let gate = activation(gate);assert_eq!(gate.dims(),up.dims(),"native hybrid gate activation changed geometry");
self.down.forward_inference(gate*up,data,tensor,true)
}
pub fn forward_inference<C:BroadcastTensorCollective<B>,T:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,data:C,tensor:T)
-> Result<Tensor<B,D>,HybridParallelError<C::Error,T::Error>> {
self.forward_inference_with(input,data,tensor,silu)
}
}
impl<B:Backend> FullyShardedVocabParallelEmbedding<B> {
pub fn forward_inference_with_layout<C:BroadcastTensorCollective<B>,T:BroadcastTensorCollective<B>>(&self,tokens:Tensor<B,2,Int>,
data:C,tensor:T,layout:&VocabParallelLossLayout) -> Result<Tensor<B,3>,HybridParallelError<C::Error,T::Error>> {
let weight = self.weight.gather_inference::<C,2>(data).map_err(HybridParallelError::Data)?;
let layer = tensor_parallel::VocabParallelEmbedding {local:crate::Embedding {weight:Param::initialized(self.weight.local.id,weight)},
vocabulary_start:self.vocabulary_start,vocabulary_size:self.vocabulary_size,padding_index:self.padding_index};
layer.forward_inference_with_layout(tokens,tensor,layout).map_err(HybridParallelError::Tensor)
}
}
impl<B:Backend> FullyShardedVocabParallelProjection<B> {
pub fn forward_inference_with_layout<C:BroadcastTensorCollective<B>,T:BroadcastTensorCollective<B>,const D:usize>(&self,input:Tensor<B,D>,
data:C,tensor:T,layout:&VocabParallelLossLayout,gather_output:bool) -> Result<Tensor<B,D>,HybridParallelError<C::Error,T::Error>> {
let weight = self.weight.gather_inference::<C,2>(data.clone()).map_err(HybridParallelError::Data)?;
let bias = self.bias.as_ref().map(|bias|bias.gather_inference::<C,1>(data).map(|value|Param::initialized(bias.local.id,value))).transpose()
.map_err(HybridParallelError::Data)?;
let layer = tensor_parallel::VocabParallelProjection {weight:Param::initialized(self.weight.local.id,weight),bias};
layer.forward_inference_with_layout(input,tensor,layout,gather_output).map_err(HybridParallelError::Tensor)
}
}