use burn_core as burn;
use burn::module::Module;
use burn::tensor::Tensor;
use burn::tensor::activation::sigmoid;
const COEFFICIENT: f64 = 1.702;
#[derive(Module, Debug, Default)]
pub(crate) struct QuickGelu;
impl QuickGelu {
pub(crate) fn forward<const D: usize>(&self, x: Tensor<D>) -> Tensor<D> {
let scaled = x.clone().mul_scalar(COEFFICIENT);
x * sigmoid(scaled)
}
}
#[cfg(test)]
mod tests {
use super::*;
use burn_core::tensor::Tolerance;
type FT = f32;
#[test]
fn quick_gelu_matches_formula() {
let device = Default::default();
let activation = QuickGelu;
let input = Tensor::<2>::from_floats([[-1.0, 0.0, 1.0, 2.0]], &device);
let output = activation.forward(input);
let expected = Tensor::<2>::from_floats([[-0.154_160, 0.0, 0.845_840, 1.935_624]], &device);
output
.into_data()
.assert_approx_eq::<FT>(&expected.into_data(), Tolerance::default());
}
#[test]
fn quick_gelu_preserves_shape() {
let device = Default::default();
let activation = QuickGelu;
let input = Tensor::<3>::zeros([2, 5, 8], &device);
let output = activation.forward(input);
assert_eq!(output.dims(), [2, 5, 8]);
}
}