use super::*;
#[derive(Clone,Debug)]
pub struct TensorParallelTransformerInput<B:Backend,const D:usize=2> {
pub tokens:Tensor<B,D,Int>,
pub positions:Option<Tensor<B,D,Int>>,
pub token_types:Option<Tensor<B,D,Int>>,
pub embedding_dtypes:Option<(FloatDType,FloatDType)>,
}
impl<B:Backend,const D:usize> TensorParallelTransformerInput<B,D> {
pub fn new(tokens:Tensor<B,D,Int>) -> Self {Self {tokens,positions:None,token_types:None,embedding_dtypes:None}}
pub fn with_positions(mut self,positions:Tensor<B,D,Int>) -> Self {self.positions = Some(positions);self}
pub fn with_token_types(mut self,token_types:Tensor<B,D,Int>) -> Self {self.token_types = Some(token_types);self}
pub fn with_embedding_dtypes(mut self,compute:FloatDType,output:FloatDType) -> Self {
self.embedding_dtypes = Some((compute,output));self
}
}
impl<B:Backend> TensorParallelTransformerInput<B,1> {
pub(super) fn into_batched(self,packed:&PackedSequenceLayout) -> TensorParallelTransformerInput<B> {
let count = self.tokens.dims()[0];assert_eq!(count,packed.tokens(),"parallel model packed input/document lengths differ");
let convert = |ids:Tensor<B,1,Int>| {assert_eq!(ids.dims(),[count],"parallel packed table metadata length differs");ids.reshape([1,count])};
TensorParallelTransformerInput {tokens:self.tokens.reshape([1,count]),positions:self.positions.map(convert),
token_types:self.token_types.map(convert),embedding_dtypes:self.embedding_dtypes}
}
}
impl<B:Backend> TensorParallelTransformerInput<B> {
pub fn embed_inference<C:BroadcastTensorCollective<B>>(self,embeddings:&TensorParallelTransformerEmbeddings<B>,
communicator:C,layout:&VocabParallelLossLayout) -> Result<Tensor<B,3>,C::Error> {
match self.embedding_dtypes {
Some((compute,output))=>embeddings.forward_inference_with_compute_dtype(self.tokens,self.positions,self.token_types,communicator,layout,compute,output),
None=>embeddings.forward_inference(self.tokens,self.positions,self.token_types,communicator,layout),
}
}
}
impl<B:Backend,S:CheckpointStrategy> TensorParallelTransformerInput<Autodiff<B,S>> {
pub fn embed<C:BroadcastTensorCollective<B>>(self,embeddings:&TensorParallelTransformerEmbeddings<Autodiff<B,S>>,
communicator:C,layout:&VocabParallelLossLayout) -> Result<Tensor<Autodiff<B,S>,3>,C::Error> {
self.embed_with_dropout(embeddings,communicator,layout,|dropout,input|dropout.forward(input))
}
pub fn embed_with_dropout<C,F>(self,embeddings:&TensorParallelTransformerEmbeddings<Autodiff<B,S>>,communicator:C,
layout:&VocabParallelLossLayout,dropout:F) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C:BroadcastTensorCollective<B>,F:FnOnce(&Dropout,Tensor<Autodiff<B,S>,3>)->Tensor<Autodiff<B,S>,3> {
match self.embedding_dtypes {
Some((compute,output))=>embeddings.forward_with_compute_dtype_and_dropout(self.tokens,self.positions,self.token_types,communicator,layout,compute,output,dropout),
None=>embeddings.forward_with_dropout(self.tokens,self.positions,self.token_types,communicator,layout,dropout),
}
}
}