use ruda_model::{config::Config,module::Module,tensor::{Bool,Int,Tensor,backend::Backend}};
use crate::{Linear,LinearConfig,Dropout,DropoutConfig,
pool::{pool_sequence,pool_packed_sequences,SequencePooling,SequencePoolOutput},attention::PackedSequenceLayout};
use super::{DenseTransformerNorm,DenseTransformerNormConfig};
#[derive(Config,Debug)]
pub struct TransformerHeadConfig {
pub d_model: usize,
pub classes: usize,
#[config(default = true)]
pub bias: bool,
pub normalization: Option<DenseTransformerNormConfig>,
#[config(default = 0.0)]
pub dropout: f64,
}
#[derive(Module,Debug)]
pub struct TransformerHead<B: Backend> {
pub projection: Linear<B>,
pub normalization: Option<DenseTransformerNorm<B>>,
pub dropout: Dropout,
}
#[derive(Clone,Debug)]
pub struct SequenceHeadOutput<B: Backend> {
pub logits: Tensor<B,2>,
pub valid_rows: Tensor<B,1,Bool>,
pub token_counts: Tensor<B,1,Int>,
}
impl TransformerHeadConfig {
pub fn init<B: Backend>(&self,device: &B::Device) -> TransformerHead<B> {
assert!(self.d_model > 0 && self.classes > 0,"head hidden/class dimensions must be positive");
assert!(self.dropout.is_finite() && (0.0..=1.0).contains(&self.dropout),"invalid head dropout");
let normalization = self.normalization.as_ref().map(|config|config.init(device));
if let Some(norm) = &normalization { assert_eq!(norm.width(),self.d_model,"head norm/input widths differ"); }
TransformerHead {projection:LinearConfig::new(self.d_model,self.classes).with_bias(self.bias).init(device),
normalization,dropout:DropoutConfig::new(self.dropout).init()}
}
}
impl<B: Backend> TransformerHead<B> {
pub fn from_projection(projection: Linear<B>,normalization: Option<DenseTransformerNorm<B>>,dropout: Dropout) -> Self {
if let Some(norm) = &normalization {
assert_eq!(norm.width(),projection.weight.val().dims()[0],"head norm/projection widths differ");
}
assert!(dropout.prob.is_finite() && (0.0..=1.0).contains(&dropout.prob),"invalid head dropout");
Self {projection,normalization,dropout}
}
pub fn forward<const D: usize>(&self,hidden: Tensor<B,D>) -> Tensor<B,D> {
let hidden = if let Some(norm) = &self.normalization { norm.forward(hidden) } else { hidden };
self.projection.forward(self.dropout.forward(hidden))
}
pub fn forward_sequence(&self,hidden: Tensor<B,3>,visible: Tensor<B,2,Bool>,pooling: SequencePooling)
-> SequenceHeadOutput<B> {
self.forward_pooled(pool_sequence(hidden,visible,pooling))
}
pub fn forward_pooled(&self,pooled: SequencePoolOutput<B>) -> SequenceHeadOutput<B> {
SequenceHeadOutput {logits:self.forward(pooled.values),valid_rows:pooled.valid_rows,token_counts:pooled.token_counts}
}
pub fn forward_packed_sequences(&self,hidden: Tensor<B,2>,layout: &PackedSequenceLayout,
visible: Option<Tensor<B,1,Bool>>,pooling: SequencePooling) -> SequenceHeadOutput<B> {
self.forward_pooled(pool_packed_sequences(hidden,layout,visible,pooling))
}
}