use super::*;
impl<B:Backend> TensorParallelTransformerModel<B> {
pub fn forward_hidden_inference_with<C,F>(&self,input:TensorParallelTransformerInput<B>,input_group:C,
input_layout:&VocabParallelLossLayout,layer:F) -> Result<Tensor<B,3>,C::Error>
where C:BroadcastTensorCollective<B>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<B>,Tensor<B,3>)->Result<Tensor<B,3>,C::Error> {
let hidden = input.embed_inference(&self.embeddings,input_group,input_layout)?;
self.backbone.forward_inference_with(hidden,layer).map(|hidden|self.finish(hidden))
}
pub fn forward_packed_hidden_inference_with<C,F>(&self,input:TensorParallelTransformerInput<B,1>,packed:&PackedSequenceLayout,
input_group:C,input_layout:&VocabParallelLossLayout,layer:F) -> Result<Tensor<B,2>,C::Error>
where C:BroadcastTensorCollective<B>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<B>,Tensor<B,2>)->Result<Tensor<B,2>,C::Error> {
let hidden = input.into_batched(packed).embed_inference(&self.embeddings,input_group,input_layout)?;
let width = hidden.dims()[2];let hidden = hidden.reshape([packed.tokens(),width]);
self.backbone.forward_packed_inference_with(hidden,layer).map(|hidden|self.finish(hidden))
}
pub fn forward_cached_hidden_inference_with<C,F>(&self,input:TensorParallelTransformerInput<B>,cache:&mut TransformerKvCache<B>,
input_group:C,input_layout:&VocabParallelLossLayout,layer:F) -> Result<Tensor<B,3>,C::Error>
where C:BroadcastTensorCollective<B>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<B>,Tensor<B,3>,&mut ProjectedKvCache<B>)->Result<Tensor<B,3>,C::Error> {
let hidden = input.embed_inference(&self.embeddings,input_group,input_layout)?;
self.backbone.forward_cached_inference_with(hidden,cache,layer).map(|hidden|self.finish(hidden))
}
}
impl<B:Backend,S:CheckpointStrategy> TensorParallelTransformerModel<Autodiff<B,S>> {
pub fn forward_hidden_with<C,F>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>>,input_group:C,
input_layout:&VocabParallelLossLayout,layer:F) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C:BroadcastTensorCollective<B>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,3>)->Result<Tensor<Autodiff<B,S>,3>,C::Error> {
self.forward_hidden_with_dropout(input,input_group,input_layout,layer,|dropout,input|dropout.forward(input))
}
pub fn forward_hidden_with_dropout<C,F,I>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>>,input_group:C,
input_layout:&VocabParallelLossLayout,layer:F,input_dropout:I) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C:BroadcastTensorCollective<B>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,3>)->Result<Tensor<Autodiff<B,S>,3>,C::Error>,
I:FnOnce(&Dropout,Tensor<Autodiff<B,S>,3>)->Tensor<Autodiff<B,S>,3> {
let hidden = input.embed_with_dropout(&self.embeddings,input_group,input_layout,input_dropout)?;
self.backbone.forward_with(hidden,layer).map(|hidden|self.finish(hidden))
}
pub fn forward_packed_hidden_with<C,F>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>,1>,packed:&PackedSequenceLayout,
input_group:C,input_layout:&VocabParallelLossLayout,layer:F) -> Result<Tensor<Autodiff<B,S>,2>,C::Error>
where C:BroadcastTensorCollective<B>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,2>)->Result<Tensor<Autodiff<B,S>,2>,C::Error> {
self.forward_packed_hidden_with_dropout(input,packed,input_group,input_layout,layer,|dropout,input|dropout.forward(input))
}
pub fn forward_packed_hidden_with_dropout<C,F,I>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>,1>,packed:&PackedSequenceLayout,
input_group:C,input_layout:&VocabParallelLossLayout,layer:F,input_dropout:I) -> Result<Tensor<Autodiff<B,S>,2>,C::Error>
where C:BroadcastTensorCollective<B>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,2>)->Result<Tensor<Autodiff<B,S>,2>,C::Error>,
I:FnOnce(&Dropout,Tensor<Autodiff<B,S>,3>)->Tensor<Autodiff<B,S>,3> {
let hidden = input.into_batched(packed).embed_with_dropout(&self.embeddings,input_group,input_layout,input_dropout)?;
let width = hidden.dims()[2];let hidden = hidden.reshape([packed.tokens(),width]);
self.backbone.forward_packed_with(hidden,layer).map(|hidden|self.finish(hidden))
}
pub fn forward_cached_hidden_with<C,F>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>>,cache:&mut TransformerKvCache<Autodiff<B,S>>,
input_group:C,input_layout:&VocabParallelLossLayout,layer:F) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C:BroadcastTensorCollective<B>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,3>,&mut ProjectedKvCache<Autodiff<B,S>>)
->Result<Tensor<Autodiff<B,S>,3>,C::Error> {
let hidden = input.embed(&self.embeddings,input_group,input_layout)?;
self.backbone.forward_cached_with(hidden,cache,layer).map(|hidden|self.finish(hidden))
}
}