use crate::Initializer;
use burn_core as burn;
use crate::activation::{Activation, ActivationConfig};
use crate::{Dropout, DropoutConfig, Linear, LinearConfig};
use burn::config::Config;
use burn::module::{Content, DisplaySettings, Module, ModuleDisplay};
use burn::tensor::{Device, Tensor, assert_shape};
#[derive(Config, Debug)]
pub struct PositionWiseFeedForwardConfig {
pub d_model: usize,
pub d_ff: usize,
#[config(default = 0.1)]
pub dropout: f64,
#[config(
default = "Initializer::KaimingUniform{gain:1.0/num_traits::Float::sqrt(3.0), fan_out_only:false}"
)]
pub initializer: Initializer,
#[config(default = "ActivationConfig::Gelu")]
pub activation: ActivationConfig,
}
#[derive(Module, Debug)]
#[module(custom_display)]
pub struct PositionWiseFeedForward {
pub linear_inner: Linear,
pub linear_outer: Linear,
pub dropout: Dropout,
#[module(skip)] pub activation: Activation,
}
impl ModuleDisplay for PositionWiseFeedForward {
fn custom_settings(&self) -> Option<DisplaySettings> {
DisplaySettings::new()
.with_new_line_after_attribute(false)
.optional()
}
fn custom_content(&self, content: Content) -> Option<Content> {
let [d_model, dff] = self.linear_inner.weight.shape().dims();
content
.add("d_model", &d_model)
.add("d_ff", &dff)
.add("prob", &self.dropout.prob)
.optional()
}
}
impl PositionWiseFeedForwardConfig {
pub fn init(&self, device: &Device) -> PositionWiseFeedForward {
PositionWiseFeedForward {
linear_inner: LinearConfig::new(self.d_model, self.d_ff)
.with_initializer(self.initializer.clone())
.init(device),
linear_outer: LinearConfig::new(self.d_ff, self.d_model)
.with_initializer(self.initializer.clone())
.init(device),
dropout: DropoutConfig::new(self.dropout).init(),
activation: self.activation.init(device),
}
}
}
impl PositionWiseFeedForward {
pub fn forward<const D: usize>(&self, input: Tensor<D>) -> Tensor<D> {
let [d_model, _] = self.linear_inner.weight.shape().dims();
assert_shape!(input, [.., d_model]);
let x = self.linear_inner.forward(input);
let x = self.activation.forward(x);
let x = self.dropout.forward(x);
self.linear_outer.forward(x)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[should_panic(expected = "assert_shape!(input, [.., d_model]): axis 2 expected 2, got 3")]
fn input_d_model_must_match() {
let device = Default::default();
let pwff = PositionWiseFeedForwardConfig::new(2, 4).init(&device);
let _ = pwff.forward(Tensor::<3>::zeros([1, 5, 3], &device));
}
#[test]
fn display() {
let config = PositionWiseFeedForwardConfig::new(2, 4);
let pwff = config.init(&Default::default());
assert_eq!(
alloc::format!("{pwff}"),
"PositionWiseFeedForward {d_model: 2, d_ff: 4, prob: 0.1, params: 22}"
);
}
}