use super::*;
use crate::{loss::LossTerms,transformer::DenseTransformerNorm};
use super::super::{VocabParallelProjection,VocabParallelEmbedding};
use ruda_model::module::Param;
#[derive(Module,Debug)]
pub struct VocabParallelTransformerHead<B:Backend> {
pub projection:VocabParallelProjection<B>,
pub normalization:Option<DenseTransformerNorm<B>>,
pub dropout:Dropout,
}
impl<B:Backend> VocabParallelTransformerHead<B> {
pub fn from_projection(projection:VocabParallelProjection<B>,normalization:Option<DenseTransformerNorm<B>>,dropout:Dropout,
layout:&VocabParallelLossLayout,rank:usize) -> Self {
let [width,hidden] = projection.weight.val().dims();
assert!(hidden > 0,"native vocabulary head hidden width must be positive");
assert_eq!(width,layout.interval(rank).len(),"native head vocabulary rows differ from actual rank interval");
if let Some(norm) = &normalization {assert_eq!(norm.width(),hidden,"native vocabulary head norm/hidden width differs");}
if let Some(bias) = &projection.bias {
assert_eq!(bias.val().dims(),[width],"native vocabulary head local bias width differs");
assert_eq!(bias.val().device(),projection.weight.val().device(),"native vocabulary head bias/weight devices differ");
}
assert!(dropout.prob.is_finite() && (0.0..=1.0).contains(&dropout.prob),"invalid native vocabulary head dropout");
Self {projection,normalization,dropout}
}
pub fn from_embedding(embedding:&VocabParallelEmbedding<B>,bias:Option<Param<Tensor<B,1>>>,normalization:Option<DenseTransformerNorm<B>>,
dropout:Dropout,layout:&VocabParallelLossLayout,rank:usize) -> Self {
assert_eq!(embedding.vocabulary_start,layout.interval(rank).start,"embedding/head vocabulary starts differ");
assert_eq!(embedding.vocabulary_size,layout.vocabulary_size(),"embedding/head logical vocabularies differ");
Self::from_projection(VocabParallelProjection::from_embedding(embedding,bias),normalization,dropout,layout,rank)
}
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> {
let hidden = if let Some(norm) = &self.normalization {norm.forward(hidden)} else {hidden};
self.projection.forward_inference_with_layout(self.dropout.forward(hidden),communicator,layout,gather_output)
}
}
impl<B:Backend,S:CheckpointStrategy> VocabParallelTransformerHead<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> {
self.forward_with_dropout(hidden,communicator,layout,gather_output,|dropout,input|dropout.forward(input))
}
pub fn forward_with_dropout<C,F,const D:usize>(&self,hidden:Tensor<Autodiff<B,S>,D>,communicator:C,
layout:&VocabParallelLossLayout,gather_output:bool,dropout:F) -> 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> {
let hidden = if let Some(norm) = &self.normalization {norm.forward(hidden)} else {hidden};
self.projection.forward_with_layout(transformed(&self.dropout,hidden,dropout),communicator,layout,gather_output)
}
}
super::training::head_objectives!(VocabParallelTransformerHead);