use super::*;
macro_rules! inference_heads {
($head:ident) => {
impl<B:Backend> $head<B> {
pub fn forward_sequence_inference<C:BroadcastTensorCollective<B>>(&self,hidden:Tensor<B,3>,visible:Tensor<B,2,Bool>,
pooling:SequencePooling,communicator:C,layout:&VocabParallelLossLayout,gather_output:bool)
-> Result<SequenceHeadOutput<B>,C::Error> {
self.forward_pooled_inference(pool_sequence(hidden,visible,pooling),communicator,layout,gather_output)
}
pub fn forward_packed_sequences_inference<C:BroadcastTensorCollective<B>>(&self,hidden:Tensor<B,2>,packed:&PackedSequenceLayout,
visible:Option<Tensor<B,1,Bool>>,pooling:SequencePooling,communicator:C,layout:&VocabParallelLossLayout,gather_output:bool)
-> Result<SequenceHeadOutput<B>,C::Error> {
self.forward_pooled_inference(pool_packed_sequences(hidden,packed,visible,pooling),communicator,layout,gather_output)
}
pub fn forward_pooled_inference<C:BroadcastTensorCollective<B>>(&self,pooled:SequencePoolOutput<B>,communicator:C,
layout:&VocabParallelLossLayout,gather_output:bool) -> Result<SequenceHeadOutput<B>,C::Error> {
Ok(SequenceHeadOutput {logits:self.forward_inference(pooled.values,communicator,layout,gather_output)?,
valid_rows:pooled.valid_rows,token_counts:pooled.token_counts})
}
pub fn forward_greedy_inference<C:BroadcastTensorCollective<B>>(&self,hidden:Tensor<B,2>,communicator:C,
layout:&VocabParallelLossLayout,visible:Option<Tensor<B,1,Bool>>) -> Result<super::super::VocabParallelGreedySelection<B>,C::Error> {
let logits = self.forward_inference(hidden,communicator.clone(),layout,false)?;
layout.greedy_indices_inference(logits,communicator,visible)
}
pub fn forward_greedy_last_inference<C:BroadcastTensorCollective<B>>(&self,hidden:Tensor<B,3>,communicator:C,
layout:&VocabParallelLossLayout,visible:Option<Tensor<B,1,Bool>>) -> Result<super::super::VocabParallelGreedySelection<B>,C::Error> {
let [batch,tokens,features] = hidden.dims();assert!(tokens > 0,"cached greedy head requires an actual last token");
self.forward_greedy_inference(hidden.slice([0..batch,tokens-1..tokens,0..features]).reshape([batch,features]),communicator,layout,visible)
}
pub fn forward_log_probabilities_inference<C:BroadcastTensorCollective<B>>(&self,hidden:Tensor<B,2>,communicator:C,
layout:&VocabParallelLossLayout,visible:Option<Tensor<B,1,Bool>>) -> Result<Tensor<B,2>,C::Error> {
let logits = self.forward_inference(hidden,communicator.clone(),layout,false)?;
layout.log_softmax_inference(logits,communicator,visible)
}
pub fn forward_probabilities_inference<C:BroadcastTensorCollective<B>>(&self,hidden:Tensor<B,2>,communicator:C,
layout:&VocabParallelLossLayout,visible:Option<Tensor<B,1,Bool>>) -> Result<Tensor<B,2>,C::Error> {
let logits = self.forward_inference(hidden,communicator.clone(),layout,false)?;
layout.softmax_inference(logits,communicator,visible)
}
}
};
}
inference_heads!(TensorParallelTransformerHead);
inference_heads!(TensorParallelAdaptedTransformerHead);
inference_heads!(VocabParallelTransformerHead);
inference_heads!(VocabParallelAdaptedTransformerHead);