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
use burn_core as burn;
use burn::module::Module;
use burn::tensor::Tensor;
/// Applies the Gaussian Error Linear Units function element-wise.
///
/// See also [gelu](burn::tensor::activation::gelu)
///
/// When `approximate` is true, uses the tanh approximation:
/// `0.5 * x * (1 + tanh(sqrt(2/pi) * (x + 0.044715 * x^3)))`
#[derive(Module, Debug, Default)]
pub struct Gelu {
/// Whether to use tanh approximation.
pub approximate: bool,
}
impl Gelu {
/// Create the module with exact GELU.
pub fn new() -> Self {
Self::default()
}
/// Create the module with tanh approximation.
pub fn new_approximate() -> Self {
Self { approximate: true }
}
/// Applies the forward pass on the input tensor.
///
/// # Shapes
///
/// - input: `[..., any]`
/// - output: `[..., any]`
pub fn forward<const D: usize>(&self, input: Tensor<D>) -> Tensor<D> {
if self.approximate {
burn::tensor::activation::gelu_approximate(input)
} else {
burn::tensor::activation::gelu(input)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use burn::tensor::Tolerance;
type FT = f32;
#[test]
fn display() {
let layer = Gelu::new();
assert_eq!(alloc::format!("{layer}"), "Gelu {\n approximate: false\n}");
}
#[test]
fn forward_approximate() {
let device = Default::default();
let input = Tensor::<2>::from_data([[-1.0, 0.0, 1.0], [0.5, -0.5, 2.0]], &device);
let output = Gelu::new_approximate().forward(input);
// PyTorch: torch.nn.functional.gelu(x, approximate="tanh")
let expected = Tensor::<2>::from_data(
[
[-0.1588079929, 0.0000000000, 0.8411920071],
[0.3457140028, -0.1542859972, 1.9545977116],
],
&device,
);
output
.into_data()
.assert_approx_eq::<FT>(&expected.into_data(), Tolerance::rel_abs(1e-5, 1e-5));
}
}