use burn::{
Tensor,
config::Config,
module::{
Module,
Param,
},
nn::{
Embedding,
EmbeddingConfig,
LayerNorm,
LayerNormConfig,
},
prelude::{
Backend,
s,
},
tensor::{
Bool,
Distribution,
Int,
},
};
use super::WHISPER_DEFAULT_D_MODEL;
use crate::{
blocks::transformers::attention::causal_mask,
burner::module::ModuleInit,
errors::BunsenResult,
kits::speech::whisper::blocks::{
ResidualDecoderAttentionBlock,
ResidualDecoderAttentionBlockConfig,
ResidualDecoderAttentionBlockMeta,
},
ops::embedding::unembed,
};
pub trait TextDecoderMeta {
fn vocab_size(&self) -> usize;
fn d_model(&self) -> usize;
fn max_context(&self) -> usize;
fn n_heads(&self) -> usize;
fn n_layers(&self) -> usize;
}
#[derive(Config, Debug)]
pub struct TextDecoderConfig {
pub vocab_size: usize,
pub d_model: usize,
pub max_context: usize,
pub n_layers: usize,
#[config(defaul_value = "WHISPER_DEFAULT_D_MODEL")]
pub d_head: usize,
#[config(default = "0.0")]
pub block_dropout: f64,
}
impl TextDecoderMeta for TextDecoderConfig {
fn vocab_size(&self) -> usize {
self.vocab_size
}
fn d_model(&self) -> usize {
self.d_model
}
fn max_context(&self) -> usize {
self.max_context
}
fn n_heads(&self) -> usize {
self.d_model / self.d_head
}
fn n_layers(&self) -> usize {
self.n_layers
}
}
pub fn attn_decoder_mask<B: Backend>(
seq_length: usize,
device: &B::Device,
) -> Tensor<B, 2> {
let mut mask = Tensor::<B, 2>::zeros([seq_length, seq_length], device);
for i in 0..(seq_length - 1) {
let values =
Tensor::<B, 2>::zeros([1, seq_length - (i + 1)], device).add_scalar(f64::NEG_INFINITY);
mask = mask.slice_assign([i..i + 1, i + 1..seq_length], values);
}
mask
}
impl<B: Backend> ModuleInit<B, TextDecoder<B>> for TextDecoderConfig {
fn try_init(
&self,
device: &B::Device,
) -> BunsenResult<TextDecoder<B>> {
Ok(TextDecoder {
token_embedding: EmbeddingConfig::new(self.vocab_size, self.d_model).init(device),
positional_embedding: Param::from_tensor(Tensor::<B, 2>::random(
[self.max_context, self.d_model],
Distribution::Normal(0.0, 1.0),
device,
)),
blocks: (0..self.n_layers)
.map(|_| {
ResidualDecoderAttentionBlockConfig::new(self.d_model)
.with_d_head(self.d_head)
.with_dropout(self.block_dropout)
.try_init(device)
})
.collect::<BunsenResult<Vec<ResidualDecoderAttentionBlock<B>>>>()?,
ln: LayerNormConfig::new(self.d_model).init(device),
})
}
}
#[derive(Module, Debug)]
pub struct TextDecoder<B: Backend> {
pub token_embedding: Embedding<B>,
pub positional_embedding: Param<Tensor<B, 2>>,
pub blocks: Vec<ResidualDecoderAttentionBlock<B>>,
pub ln: LayerNorm<B>,
}
impl<B: Backend> TextDecoderMeta for TextDecoder<B> {
fn vocab_size(&self) -> usize {
self.token_embedding.weight.val().dims()[0]
}
fn d_model(&self) -> usize {
self.blocks[0].d_model()
}
fn max_context(&self) -> usize {
self.positional_embedding.val().dims()[0]
}
fn n_heads(&self) -> usize {
self.blocks[0].n_heads()
}
fn n_layers(&self) -> usize {
self.blocks.len()
}
}
impl<B: Backend> TextDecoder<B> {
pub fn forward(
&self,
x: Tensor<B, 2, Int>,
xa: Tensor<B, 3>,
) -> Tensor<B, 3> {
let [_batch, seq_len] = x.dims();
assert!(
seq_len <= self.max_context(),
"Token sequence length {} must not exceed {}.",
seq_len,
self.max_context()
);
let x = self.embed(x);
let mask: Option<Tensor<B, 3, Bool>> = causal_mask(seq_len, 0, &x.device()).into();
let mut x = x;
for b in self.blocks.iter() {
x = b.forward(x, xa.clone(), mask.clone()).output;
}
let x = self.ln.forward(x);
unembed(&self.token_embedding, x)
}
fn embed(
&self,
x: Tensor<B, 2, Int>,
) -> Tensor<B, 3> {
let seq_len = x.dims()[1];
self.token_embedding.forward(x)
+ self
.positional_embedding
.val()
.slice(s![0..seq_len])
.unsqueeze::<3>()
}
}
#[cfg(test)]
mod tests {
use serial_test::serial;
use super::*;
use crate::{
contracts::assert_shape_contract,
support::testing::PerformanceBackend,
};
#[test]
#[serial]
fn test_text_decoder_forward() {
type B = PerformanceBackend;
let device = Default::default();
let d_model = 128;
let vocab_size = 64;
let max_context = 128;
let n_layers = 2;
let config = TextDecoderConfig::new(vocab_size, d_model, max_context, n_layers);
assert_eq!(config.vocab_size(), vocab_size);
assert_eq!(config.d_model(), d_model);
assert_eq!(config.max_context(), max_context);
let decoder: TextDecoder<B> = config.init(&device);
assert_eq!(decoder.vocab_size(), vocab_size);
assert_eq!(decoder.d_model(), d_model);
assert_eq!(decoder.max_context(), max_context);
let batch = 2;
let seq_len = max_context / 2;
let x: Tensor<B, 2, Int> = Tensor::zeros([batch, seq_len], &device);
let xa: Tensor<B, 3> =
Tensor::random([batch, seq_len, d_model], Default::default(), &device);
let output = decoder.forward(x.clone(), xa.clone());
assert_shape_contract!(
["batch", "seq", "n_vocab"],
&output,
&[("batch", batch), ("seq", seq_len), ("n_vocab", vocab_size),],
);
}
}