use super::*;
use crate::{attention::{DenseAttentionMask,DenseAttentionOptions},
loss::{CausalCrossEntropyConfig,CausalLoss,LossTerms,KLDivLoss},pool::SequencePooling,transformer::SequenceHeadOutput};
use ruda_model::tensor::Bool;
use super::super::{AttentionParallelGroups,VocabParallelCrossEntropy};
impl<B:Backend,S:CheckpointStrategy> TensorParallelTransformerModel<Autodiff<B,S>> {
pub fn forward_with<C,O,F>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>>,input_group:C,
input_layout:&VocabParallelLossLayout,layer:F,output_group:O,output_layout:&VocabParallelLossLayout,gather_output:bool)
-> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C:BroadcastTensorCollective<B>,O:BroadcastTensorCollective<B,Error=C::Error>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,3>)->Result<Tensor<Autodiff<B,S>,3>,C::Error> {
let hidden = self.forward_hidden_with(input,input_group,input_layout,layer)?;
self.head.forward(hidden,output_group,output_layout,gather_output)
}
pub fn forward_with_dropouts<C,O,F,I,H,A>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>>,input_group:C,
input_layout:&VocabParallelLossLayout,layer:F,output_group:O,output_layout:&VocabParallelLossLayout,gather_output:bool,
input_dropout:I,head_dropout:H,adapter_dropout:A) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C:BroadcastTensorCollective<B>,O:BroadcastTensorCollective<B,Error=C::Error>,
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>,
H:FnOnce(&Dropout,Tensor<Autodiff<B,S>,3>)->Tensor<Autodiff<B,S>,3>,
A:FnOnce(&Dropout,Tensor<Autodiff<B,S>,3>)->Tensor<Autodiff<B,S>,3> {
let hidden = self.forward_hidden_with_dropout(input,input_group,input_layout,layer,input_dropout)?;
self.head.forward_with_dropouts(hidden,output_group,output_layout,gather_output,head_dropout,adapter_dropout)
}
pub fn forward<C,K>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>>,masks:DenseAttentionMask<Autodiff<B,S>>,
options:DenseAttentionOptions,groups:&AttentionParallelGroups<C,K>,input_layout:&VocabParallelLossLayout,
output_layout:&VocabParallelLossLayout,gather_output:bool) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C:BroadcastTensorCollective<B>,K:BroadcastTensorCollective<B,Error=C::Error> {
self.forward_with_positions(input,masks,options,groups,input_layout,output_layout,gather_output,|_,query,key|(query,key))
}
pub fn forward_with_positions<C,K,P>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>>,masks:DenseAttentionMask<Autodiff<B,S>>,
options:DenseAttentionOptions,groups:&AttentionParallelGroups<C,K>,input_layout:&VocabParallelLossLayout,
output_layout:&VocabParallelLossLayout,gather_output:bool,mut positions:P) -> Result<Tensor<Autodiff<B,S>,3>,C::Error>
where C:BroadcastTensorCollective<B>,K:BroadcastTensorCollective<B,Error=C::Error>,
P:FnMut(usize,Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>)->(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>) {
self.forward_with(input,groups.heads.clone(),input_layout,|index,block,hidden|
block.forward(hidden,masks.clone(),options,groups,|query,key|positions(index,query,key)),
groups.heads.clone(),output_layout,gather_output)
}
pub fn forward_causal_with<C,O,F>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>>,labels:Tensor<Autodiff<B,S>,2,Int>,
input_group:C,input_layout:&VocabParallelLossLayout,layer:F,output_group:O,output_layout:&VocabParallelLossLayout,
criterion:&CausalCrossEntropyConfig,label_smoothing:f64) -> Result<CausalLoss<Autodiff<B,S>>,C::Error>
where C:BroadcastTensorCollective<B>,O:BroadcastTensorCollective<B,Error=C::Error>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,3>)->Result<Tensor<Autodiff<B,S>,3>,C::Error> {
assert_eq!(labels.dims(),input.tokens.dims(),"parallel causal model labels/input rows differ");
assert_eq!(labels.device(),input.tokens.device(),"parallel causal model labels/input devices differ");
let hidden = self.forward_hidden_with(input,input_group,input_layout,layer)?;
self.head.forward_causal_loss(hidden,labels,criterion,output_group,output_layout,label_smoothing)
}
pub fn forward_causal_with_dropouts<C,O,F,I,H,A>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>>,labels:Tensor<Autodiff<B,S>,2,Int>,
input_group:C,input_layout:&VocabParallelLossLayout,layer:F,output_group:O,output_layout:&VocabParallelLossLayout,
criterion:&CausalCrossEntropyConfig,label_smoothing:f64,input_dropout:I,mut head_dropout:H,mut adapter_dropout:A)
-> Result<CausalLoss<Autodiff<B,S>>,C::Error>
where C:BroadcastTensorCollective<B>,O:BroadcastTensorCollective<B,Error=C::Error>,
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>,
H:FnMut(&Dropout,Tensor<Autodiff<B,S>,2>)->Tensor<Autodiff<B,S>,2>,
A:FnMut(&Dropout,Tensor<Autodiff<B,S>,2>)->Tensor<Autodiff<B,S>,2> {
assert_eq!(labels.dims(),input.tokens.dims(),"parallel causal model labels/input rows differ");
assert_eq!(labels.device(),input.tokens.device(),"parallel causal model labels/input devices differ");
let hidden = self.forward_hidden_with_dropout(input,input_group,input_layout,layer,input_dropout)?;
criterion.forward_sharded_hidden(hidden,labels,output_layout,output_group,|rows,group|
self.head.forward_with_dropouts(rows,group.clone(),output_layout,false,&mut head_dropout,&mut adapter_dropout),label_smoothing)
}
pub fn forward_packed_with<C,O,F>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>,1>,packed:&PackedSequenceLayout,
input_group:C,input_layout:&VocabParallelLossLayout,layer:F,output_group:O,output_layout:&VocabParallelLossLayout,gather_output:bool)
-> Result<Tensor<Autodiff<B,S>,2>,C::Error>
where C:BroadcastTensorCollective<B>,O:BroadcastTensorCollective<B,Error=C::Error>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,2>)->Result<Tensor<Autodiff<B,S>,2>,C::Error> {
let hidden = self.forward_packed_hidden_with(input,packed,input_group,input_layout,layer)?;
self.head.forward(hidden,output_group,output_layout,gather_output)
}
pub fn forward_packed_causal_with<C,O,F>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>,1>,labels:Tensor<Autodiff<B,S>,1,Int>,
packed:&PackedSequenceLayout,input_group:C,input_layout:&VocabParallelLossLayout,layer:F,output_group:O,
output_layout:&VocabParallelLossLayout,criterion:&CausalCrossEntropyConfig,label_smoothing:f64)
-> Result<CausalLoss<Autodiff<B,S>>,C::Error>
where C:BroadcastTensorCollective<B>,O:BroadcastTensorCollective<B,Error=C::Error>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,2>)->Result<Tensor<Autodiff<B,S>,2>,C::Error> {
assert_eq!(labels.dims(),input.tokens.dims(),"parallel packed causal model labels/input lengths differ");
assert_eq!(labels.device(),input.tokens.device(),"parallel packed causal model labels/input devices differ");
let hidden = self.forward_packed_hidden_with(input,packed,input_group,input_layout,layer)?;
self.head.forward_packed_causal_loss(hidden,labels,packed,criterion,output_group,output_layout,label_smoothing)
}
pub fn forward_packed_causal_with_dropouts<C,O,F,I,H,A>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>,1>,labels:Tensor<Autodiff<B,S>,1,Int>,
packed:&PackedSequenceLayout,input_group:C,input_layout:&VocabParallelLossLayout,layer:F,output_group:O,
output_layout:&VocabParallelLossLayout,criterion:&CausalCrossEntropyConfig,label_smoothing:f64,
input_dropout:I,mut head_dropout:H,mut adapter_dropout:A) -> Result<CausalLoss<Autodiff<B,S>>,C::Error>
where C:BroadcastTensorCollective<B>,O:BroadcastTensorCollective<B,Error=C::Error>,
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>,
H:FnMut(&Dropout,Tensor<Autodiff<B,S>,2>)->Tensor<Autodiff<B,S>,2>,
A:FnMut(&Dropout,Tensor<Autodiff<B,S>,2>)->Tensor<Autodiff<B,S>,2> {
assert_eq!(labels.dims(),input.tokens.dims(),"parallel packed causal model labels/input lengths differ");
assert_eq!(labels.device(),input.tokens.device(),"parallel packed causal model labels/input devices differ");
let hidden = self.forward_packed_hidden_with_dropout(input,packed,input_group,input_layout,layer,input_dropout)?;
criterion.forward_sharded_packed_hidden(hidden,labels,packed,output_layout,output_group,|rows,group|
self.head.forward_with_dropouts(rows,group.clone(),output_layout,false,&mut head_dropout,&mut adapter_dropout),label_smoothing)
}
pub fn forward_token_terms_with<C,O,F>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>>,labels:Tensor<Autodiff<B,S>,2,Int>,
input_group:C,input_layout:&VocabParallelLossLayout,layer:F,output_group:O,criterion:&VocabParallelCrossEntropy,
visible:Option<Tensor<Autodiff<B,S>,2,Bool>>,weights:Option<Tensor<Autodiff<B,S>,1>>)
-> Result<LossTerms<Autodiff<B,S>,2>,C::Error>
where C:BroadcastTensorCollective<B>,O:BroadcastTensorCollective<B,Error=C::Error>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,3>)->Result<Tensor<Autodiff<B,S>,3>,C::Error> {
assert_eq!(labels.dims(),input.tokens.dims(),"parallel token model labels/input rows differ");
let hidden = self.forward_hidden_with(input,input_group,input_layout,layer)?;
self.head.forward_token_terms(hidden,labels,criterion,output_group,visible,weights)
}
pub fn forward_soft_token_terms_with<C,O,F>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>>,targets:Tensor<Autodiff<B,S>,3>,
input_group:C,input_layout:&VocabParallelLossLayout,layer:F,output_group:O,criterion:&VocabParallelCrossEntropy,
visible:Option<Tensor<Autodiff<B,S>,2,Bool>>,weights:Option<Tensor<Autodiff<B,S>,1>>)
-> Result<LossTerms<Autodiff<B,S>,2>,C::Error>
where C:BroadcastTensorCollective<B>,O:BroadcastTensorCollective<B,Error=C::Error>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,3>)->Result<Tensor<Autodiff<B,S>,3>,C::Error> {
let rows = input.tokens.dims();assert_eq!([targets.dims()[0],targets.dims()[1]],rows,"parallel soft targets/input rows differ");
let hidden = self.forward_hidden_with(input,input_group,input_layout,layer)?;
self.head.forward_soft_token_terms(hidden,targets,criterion,output_group,visible,weights)
}
pub fn forward_token_kl_terms_with<C,O,F>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>>,targets:Tensor<Autodiff<B,S>,3>,
input_group:C,input_layout:&VocabParallelLossLayout,layer:F,output_group:O,output_layout:&VocabParallelLossLayout,
criterion:&KLDivLoss,visible:Option<Tensor<Autodiff<B,S>,2,Bool>>) -> Result<LossTerms<Autodiff<B,S>,2>,C::Error>
where C:BroadcastTensorCollective<B>,O:BroadcastTensorCollective<B,Error=C::Error>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,3>)->Result<Tensor<Autodiff<B,S>,3>,C::Error> {
let rows = input.tokens.dims();assert_eq!([targets.dims()[0],targets.dims()[1]],rows,"parallel KL targets/input rows differ");
let hidden = self.forward_hidden_with(input,input_group,input_layout,layer)?;
self.head.forward_token_kl_terms(hidden,targets,criterion,output_group,output_layout,visible)
}
pub fn forward_sequence_with<C,O,F>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>>,visible:Tensor<Autodiff<B,S>,2,Bool>,
pooling:SequencePooling,input_group:C,input_layout:&VocabParallelLossLayout,layer:F,output_group:O,
output_layout:&VocabParallelLossLayout,gather_output:bool) -> Result<SequenceHeadOutput<Autodiff<B,S>>,C::Error>
where C:BroadcastTensorCollective<B>,O:BroadcastTensorCollective<B,Error=C::Error>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,3>)->Result<Tensor<Autodiff<B,S>,3>,C::Error> {
assert_eq!(visible.dims(),input.tokens.dims(),"parallel sequence pool/input rows differ");
let hidden = self.forward_hidden_with(input,input_group,input_layout,layer)?;
self.head.forward_sequence(hidden,visible,pooling,output_group,output_layout,gather_output)
}
pub fn forward_packed_sequences_with<C,O,F>(&self,input:TensorParallelTransformerInput<Autodiff<B,S>,1>,packed:&PackedSequenceLayout,
visible:Option<Tensor<Autodiff<B,S>,1,Bool>>,pooling:SequencePooling,input_group:C,input_layout:&VocabParallelLossLayout,
layer:F,output_group:O,output_layout:&VocabParallelLossLayout,gather_output:bool) -> Result<SequenceHeadOutput<Autodiff<B,S>>,C::Error>
where C:BroadcastTensorCollective<B>,O:BroadcastTensorCollective<B,Error=C::Error>,
F:FnMut(usize,&TensorParallelAdaptedStackLayer<Autodiff<B,S>>,Tensor<Autodiff<B,S>,2>)->Result<Tensor<Autodiff<B,S>,2>,C::Error> {
let hidden = self.forward_packed_hidden_with(input,packed,input_group,input_layout,layer)?;
self.head.forward_packed_sequences(hidden,packed,visible,pooling,output_group,output_layout,gather_output)
}
}