use super::*;
use crate::loss::LossTerms;
#[derive(Module,Debug)]
pub enum TensorParallelOutputHead<B:Backend> {
Linear(TensorParallelTransformerHead<B>),
AdaptedLinear(TensorParallelAdaptedTransformerHead<B>),
Vocabulary(VocabParallelTransformerHead<B>),
AdaptedVocabulary(VocabParallelAdaptedTransformerHead<B>),
}
impl<B:Backend> TensorParallelOutputHead<B> {
pub fn hidden_width(&self) -> usize {
match self {
Self::Linear(head)=>head.local.projection.weight.val().dims()[0],
Self::AdaptedLinear(head)=>head.local.projection.base.weight.val().dims()[0],
Self::Vocabulary(head)=>head.projection.weight.val().dims()[1],
Self::AdaptedVocabulary(head)=>head.projection.base.weight.val().dims()[1],
}
}
pub fn local_classes(&self) -> usize {
match self {
Self::Linear(head)=>head.local.projection.weight.val().dims()[1],
Self::AdaptedLinear(head)=>head.local.projection.base.weight.val().dims()[1],
Self::Vocabulary(head)=>head.projection.weight.val().dims()[0],
Self::AdaptedVocabulary(head)=>head.projection.base.weight.val().dims()[0],
}
}
pub fn forward_inference<C:BroadcastTensorCollective<B>,const D:usize>(&self,hidden:Tensor<B,D>,communicator:C,
layout:&VocabParallelLossLayout,gather_output:bool) -> Result<Tensor<B,D>,C::Error> {
match self {
Self::Linear(head)=>head.forward_inference(hidden,communicator,layout,gather_output),
Self::AdaptedLinear(head)=>head.forward_inference(hidden,communicator,layout,gather_output),
Self::Vocabulary(head)=>head.forward_inference(hidden,communicator,layout,gather_output),
Self::AdaptedVocabulary(head)=>head.forward_inference(hidden,communicator,layout,gather_output),
}
}
}
impl<B:Backend,S:CheckpointStrategy> TensorParallelOutputHead<Autodiff<B,S>> {
pub fn forward<C:BroadcastTensorCollective<B>,const D:usize>(&self,hidden:Tensor<Autodiff<B,S>,D>,communicator:C,
layout:&VocabParallelLossLayout,gather_output:bool) -> Result<Tensor<Autodiff<B,S>,D>,C::Error> {
match self {
Self::Linear(head)=>head.forward(hidden,communicator,layout,gather_output),
Self::AdaptedLinear(head)=>head.forward(hidden,communicator,layout,gather_output),
Self::Vocabulary(head)=>head.forward(hidden,communicator,layout,gather_output),
Self::AdaptedVocabulary(head)=>head.forward(hidden,communicator,layout,gather_output),
}
}
pub fn forward_with_dropouts<C,F,A,const D:usize>(&self,hidden:Tensor<Autodiff<B,S>,D>,communicator:C,
layout:&VocabParallelLossLayout,gather_output:bool,head_dropout:F,adapter_dropout:A)
-> Result<Tensor<Autodiff<B,S>,D>,C::Error>
where C:BroadcastTensorCollective<B>,F:FnOnce(&Dropout,Tensor<Autodiff<B,S>,D>)->Tensor<Autodiff<B,S>,D>,
A:FnOnce(&Dropout,Tensor<Autodiff<B,S>,D>)->Tensor<Autodiff<B,S>,D> {
match self {
Self::Linear(head)=>head.forward_with_dropout(hidden,communicator,layout,gather_output,head_dropout),
Self::Vocabulary(head)=>head.forward_with_dropout(hidden,communicator,layout,gather_output,head_dropout),
Self::AdaptedLinear(head)=>head.forward_with_dropouts(hidden,communicator,layout,gather_output,head_dropout,adapter_dropout),
Self::AdaptedVocabulary(head)=>head.forward_with_dropouts(hidden,communicator,layout,gather_output,head_dropout,adapter_dropout),
}
}
}
super::training::head_objectives!(TensorParallelOutputHead);