1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
use crate::activation::{Activation, ActivationConfig};
use crate::{Dropout, DropoutConfig, Linear, LinearConfig};
use ruda_model::config::Config;
use ruda_model::module::{Content, DisplaySettings, Initializer, Module, ModuleDisplay};
use ruda_model::tensor::{Tensor, backend::Backend};
/// Configuration to create a [position-wise feed-forward](PositionWiseFeedForward) layer using the [init function](PositionWiseFeedForwardConfig::init).
#[derive(Config, Debug)]
pub struct PositionWiseFeedForwardConfig {
/// The size of the input and output features.
pub d_model: usize,
/// The size of the hidden inner features.
pub d_ff: usize,
/// The dropout rate. Default: 0.1
#[config(default = 0.1)]
pub dropout: f64,
/// The type of function used to initialize neural network parameters
#[config(
default = "Initializer::KaimingUniform{gain:1.0/num_traits::Float::sqrt(3.0), fan_out_only:false}"
)]
pub initializer: Initializer,
/// The activation function used between the two linear layers. Default: Gelu
#[config(default = "ActivationConfig::Gelu")]
pub activation: ActivationConfig,
}
/// Applies the position-wise feed-forward network to the input tensor from the paper [Attention Is All You Need](https://arxiv.org/pdf/1706.03762v7).
///
/// # Params
///
/// - linear inner: Linear layer with `d_model` input features and `d_ff` output features.
/// - linear outer: Linear layer with `d_ff` input features and `d_model` output features.
///
/// `FFN(x) = max(0, xW1 + b1)W2 + b2`
///
/// Should be created using [PositionWiseFeedForwardConfig]
///
/// # Notes
///
/// The `activation` field is currently marked `#[module(skip)]` for backward
/// compatibility with records saved before this field was introduced (when
/// the activation was always `Gelu` and had no state). This means activation
/// state is **not persisted** when saving or loading records.
///
/// For stateless activations (GELU, ReLU, etc.) this has no effect.
/// **If you are using `SwiGLU`, its learnable parameters will not be saved or
/// loaded correctly.**
#[derive(Module, Debug)]
#[module(custom_display)]
pub struct PositionWiseFeedForward<B: Backend> {
/// Linear layer with `d_model` input features and `d_ff` output features.
pub linear_inner: Linear<B>,
/// Linear layer with `d_ff` input features and `d_model` output features.
pub linear_outer: Linear<B>,
/// Dropout layer.
pub dropout: Dropout,
/// Activation function.
#[module(skip)] // for backward compatibility with previous `gelu` field name
pub activation: Activation<B>,
}
impl<B: Backend> ModuleDisplay for PositionWiseFeedForward<B> {
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 {
/// Initialize a new [position-wise feed-forward](PositionWiseFeedForward) module.
pub fn init<B: Backend>(&self, device: &B::Device) -> PositionWiseFeedForward<B> {
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<B: Backend> PositionWiseFeedForward<B> {
/// Applies the forward pass on the input tensor.
///
/// # Shapes
///
/// - tensor: `[batch_size, seq_length, d_model]`
/// - output: `[batch_size, seq_length, d_model]`
pub fn forward<const D: usize>(&self, input: Tensor<B, D>) -> Tensor<B, D> {
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::*;
use crate::TestBackend;
#[test]
fn display() {
let config = PositionWiseFeedForwardConfig::new(2, 4);
let pwff = config.init::<TestBackend>(&Default::default());
assert_eq!(
alloc::format!("{pwff}"),
"PositionWiseFeedForward {d_model: 2, d_ff: 4, prob: 0.1, params: 22}"
);
}
}