use super::*;
use super::super::AttentionParallelGroups;
use crate::{attention::{PackedAttentionOptions,PackedDocumentAttentionMask},loss::{CausalCrossEntropyConfig,CausalLoss}};
impl<B:Backend> TensorParallelTransformerModel<B> {
pub fn forward_packed_inference<C:BroadcastTensorCollective<B>>(&self,input:TensorParallelTransformerInput<B,1>,packed:&PackedSequenceLayout,
masks:Option<&[PackedDocumentAttentionMask<B>]>,options:PackedAttentionOptions,communicator:C,
input_layout:&VocabParallelLossLayout,output_layout:&VocabParallelLossLayout,gather_output:bool) -> Result<Tensor<B,2>,C::Error> {
self.forward_packed_inference_with_positions(input,packed,masks,options,communicator,input_layout,output_layout,gather_output,|_,query,key|(query,key))
}
pub fn forward_packed_inference_with_positions<C,P>(&self,input:TensorParallelTransformerInput<B,1>,packed:&PackedSequenceLayout,
masks:Option<&[PackedDocumentAttentionMask<B>]>,options:PackedAttentionOptions,communicator:C,
input_layout:&VocabParallelLossLayout,output_layout:&VocabParallelLossLayout,gather_output:bool,mut positions:P)
-> Result<Tensor<B,2>,C::Error>
where C:BroadcastTensorCollective<B>,P:FnMut(usize,Tensor<B,3>,Tensor<B,3>)->(Tensor<B,3>,Tensor<B,3>) {
self.forward_packed_inference_with(input,packed,communicator.clone(),input_layout,|index,block,hidden|
block.forward_packed_inference(hidden,packed,masks,options,communicator.clone(),|query,key|positions(index,query,key)),
communicator.clone(),output_layout,gather_output)
}
}
impl<B:Backend,S:CheckpointStrategy> TensorParallelTransformerModel<Autodiff<B,S>> {
pub fn forward_packed<C,K>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>,1>,packed:&PackedSequenceLayout,
masks:Option<&[PackedDocumentAttentionMask<Autodiff<B,S>>]>,options:PackedAttentionOptions,groups:&AttentionParallelGroups<C,K>,
input_layout:&VocabParallelLossLayout,output_layout:&VocabParallelLossLayout,gather_output:bool)
-> Result<Tensor<Autodiff<B,S>,2>,C::Error>
where C:BroadcastTensorCollective<B>,K:BroadcastTensorCollective<B,Error=C::Error> {
self.forward_packed_with_positions(input,packed,masks,options,groups,input_layout,output_layout,gather_output,|_,query,key|(query,key))
}
pub fn forward_packed_with_positions<C,K,P>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>,1>,packed:&PackedSequenceLayout,
masks:Option<&[PackedDocumentAttentionMask<Autodiff<B,S>>]>,options:PackedAttentionOptions,groups:&AttentionParallelGroups<C,K>,
input_layout:&VocabParallelLossLayout,output_layout:&VocabParallelLossLayout,gather_output:bool,mut positions:P)
-> Result<Tensor<Autodiff<B,S>,2>,C::Error>
where C:BroadcastTensorCollective<B>,K:BroadcastTensorCollective<B,Error=C::Error>,
P:FnMut(usize,Tensor<Autodiff<B,S>,3>,Tensor<Autodiff<B,S>,3>)->(Tensor<Autodiff<B,S>,3>,Tensor<Autodiff<B,S>,3>) {
self.forward_packed_with(input,packed,groups.heads.clone(),input_layout,|index,block,hidden|
block.forward_packed(hidden,packed,masks,options,groups,|query,key|positions(index,query,key)),
groups.heads.clone(),output_layout,gather_output)
}
pub fn forward_packed_causal_with_positions<C,K,P>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>,1>,labels:Tensor<Autodiff<B,S>,1,Int>,
packed:&PackedSequenceLayout,masks:Option<&[PackedDocumentAttentionMask<Autodiff<B,S>>]>,options:PackedAttentionOptions,
groups:&AttentionParallelGroups<C,K>,input_layout:&VocabParallelLossLayout,output_layout:&VocabParallelLossLayout,
criterion:&CausalCrossEntropyConfig,label_smoothing:f64,mut positions:P) -> Result<CausalLoss<Autodiff<B,S>>,C::Error>
where C:BroadcastTensorCollective<B>,K:BroadcastTensorCollective<B,Error=C::Error>,
P:FnMut(usize,Tensor<Autodiff<B,S>,3>,Tensor<Autodiff<B,S>,3>)->(Tensor<Autodiff<B,S>,3>,Tensor<Autodiff<B,S>,3>) {
self.forward_packed_causal_with(input,labels,packed,groups.heads.clone(),input_layout,|index,block,hidden|
block.forward_packed(hidden,packed,masks,options,groups,|query,key|positions(index,query,key)),
groups.heads.clone(),output_layout,criterion,label_smoothing)
}
}