use burn::{
Tensor,
config::Config,
module::Module,
nn::{
LayerNorm,
LayerNormConfig,
attention::{
MultiHeadAttention,
MultiHeadAttentionConfig,
},
},
prelude::Backend,
};
use super::WHISPER_DEFAULT_D_MODEL;
use crate::{
blocks::transformers::{
attention::layer_norm_self_attn,
mlp::{
Mlp,
MlpConfig,
layer_norm_mlp,
},
},
burner::module::ModuleInit,
};
pub trait ResidualEncoderAttentionBlockMeta {
fn d_model(&self) -> usize;
fn n_heads(&self) -> usize;
fn dropout(&self) -> f64;
}
#[derive(Config, Debug)]
pub struct ResidualEncoderAttentionBlockConfig {
pub d_model: usize,
#[config(defaul_value = "WHISPER_DEFAULT_D_MODEL")]
pub d_head: usize,
#[config(default = "0.0")]
pub dropout: f64,
}
impl ResidualEncoderAttentionBlockMeta for ResidualEncoderAttentionBlockConfig {
fn d_model(&self) -> usize {
self.d_model
}
fn n_heads(&self) -> usize {
self.d_model / self.d_head
}
fn dropout(&self) -> f64 {
self.dropout
}
}
impl<B: Backend> ModuleInit<B, ResidualEncoderAttentionBlock<B>>
for ResidualEncoderAttentionBlockConfig
{
fn try_init(
&self,
device: &B::Device,
) -> crate::errors::BunsenResult<ResidualEncoderAttentionBlock<B>> {
let mha_cfg =
MultiHeadAttentionConfig::new(self.d_model, self.n_heads()).with_dropout(self.dropout);
let ln_cfg = LayerNormConfig::new(self.d_model);
let mut attn = mha_cfg.init(device);
attn.key.bias = None;
Ok(ResidualEncoderAttentionBlock {
attn_ln: ln_cfg.init(device),
attn,
mlp_ln: ln_cfg.init(device),
mlp: MlpConfig::new(self.d_model).try_init(device)?,
})
}
}
#[derive(Module, Debug)]
pub struct ResidualEncoderAttentionBlock<B: Backend> {
pub attn_ln: LayerNorm<B>,
pub attn: MultiHeadAttention<B>,
pub mlp_ln: LayerNorm<B>,
pub mlp: Mlp<B>,
}
impl<B: Backend> ResidualEncoderAttentionBlockMeta for ResidualEncoderAttentionBlock<B> {
fn d_model(&self) -> usize {
self.attn.d_model
}
fn n_heads(&self) -> usize {
self.attn.n_heads
}
fn dropout(&self) -> f64 {
self.attn.dropout.prob
}
}
impl<B: Backend> ResidualEncoderAttentionBlock<B> {
pub fn forward(
&self,
x: Tensor<B, 3>,
) -> Tensor<B, 3> {
let self_attn = layer_norm_self_attn(&self.attn_ln, &self.attn, x.clone(), None);
let x = x + self_attn.context;
let mlp = layer_norm_mlp(&self.mlp_ln, &self.mlp, x.clone());
x + mlp
}
}
#[cfg(test)]
mod tests {
use burn::{
prelude::Shape,
tensor::Distribution,
};
use super::*;
use crate::contracts::assert_shape_contract;
#[test]
#[serial_test::serial]
fn test_residual_decoder_forward() {
type B = crate::support::testing::PerformanceBackend;
let device = Default::default();
let d_model = 128;
let cfg = ResidualEncoderAttentionBlockConfig::new(d_model);
let n_heads = cfg.n_heads();
assert_eq!(cfg.d_model(), d_model);
assert_eq!(cfg.n_heads(), n_heads);
let block: ResidualEncoderAttentionBlock<B> = cfg.init(&device);
assert_eq!(block.d_model(), d_model);
assert_eq!(block.n_heads(), n_heads);
let batch = 2;
let seq_len = 10;
let shape: Shape = [batch, seq_len, d_model].into();
let x: Tensor<B, 3> = Tensor::random(shape.clone(), Distribution::Default, &device);
let output = block.forward(x.clone());
let expected = {
let self_attn = layer_norm_self_attn(&block.attn_ln, &block.attn, x.clone(), None);
let x = x + self_attn.context;
let mlp = layer_norm_mlp(&block.mlp_ln, &block.mlp, x.clone());
x + mlp
};
output
.clone()
.into_data()
.assert_approx_eq::<f64>(&expected.into_data(), Default::default());
assert_shape_contract!(
["batch", "seq_len", "d_model"],
&output,
&[("batch", batch), ("seq_len", seq_len), ("d_model", d_model),],
);
}
}