1use kopitiam_core::{DType, Result};
5
6use crate::storage::Storage;
7
8use super::Tensor;
9
10const SQRT_2_OVER_PI: f32 = 0.797_884_6;
12
13impl Tensor {
14 pub fn silu(&self) -> Result<Tensor> {
17 self.unary_f32(|x| x / (1.0 + (-x).exp()))
18 }
19
20 pub fn gelu(&self) -> Result<Tensor> {
34 self.unary_f32(|x| 0.5 * x * (1.0 + (SQRT_2_OVER_PI * (x + 0.044_715 * x * x * x)).tanh()))
35 }
36
37 pub fn tanh(&self) -> Result<Tensor> {
47 self.unary_f32(f32::tanh)
48 }
49
50 pub fn sigmoid(&self) -> Result<Tensor> {
59 self.unary_f32(|x| 1.0 / (1.0 + (-x).exp()))
60 }
61
62 pub fn logistic(&self) -> Result<Tensor> {
64 self.sigmoid()
65 }
66
67 fn unary_f32(&self, f: impl Fn(f32) -> f32) -> Result<Tensor> {
68 self.require_dtype(DType::F32)?;
69 let Storage::F32(data) = self.storage.as_ref() else { unreachable!() };
70 let out: Vec<f32> = self.logical_offsets().map(|i| f(data[i])).collect();
71 Tensor::from_f32(out, self.shape.clone())
72 }
73}
74
75#[cfg(test)]
76mod tests {
77 use kopitiam_core::Error;
78
79 use super::*;
80
81 fn assert_close(a: f32, b: f32, epsilon: f32) {
82 assert!((a - b).abs() < epsilon, "expected {b}, got {a}");
83 }
84
85 #[test]
86 fn silu_matches_hand_computation() {
87 let t = Tensor::from_f32(vec![0.0, 1.0, -1.0], [3]).unwrap();
89 let out = t.silu().unwrap().to_vec_f32().unwrap();
90 assert_close(out[0], 0.0, 1e-6);
91 assert_close(out[1], 0.731_058_6, 1e-5);
93 assert_close(out[2], -0.268_941_4, 1e-5);
95 }
96
97 #[test]
98 fn silu_matches_the_reference_formula_across_a_range_of_inputs() {
99 let xs = vec![-5.0, -3.0, -1.278, -1.0, -0.5, 0.0, 0.5, 1.0, 3.0, 5.0];
103 let t = Tensor::from_f32(xs.clone(), [xs.len()]).unwrap();
104 let out = t.silu().unwrap().to_vec_f32().unwrap();
105 for (x, o) in xs.iter().zip(&out) {
106 let expected = x / (1.0 + (-x).exp());
107 assert_close(*o, expected, 1e-5);
108 }
109 }
110
111 #[test]
112 fn silu_is_monotonically_increasing_on_its_increasing_branch() {
113 let t = Tensor::from_f32(vec![-1.0, 0.0, 1.0, 2.0, 3.0], [5]).unwrap();
117 let out = t.silu().unwrap().to_vec_f32().unwrap();
118 for pair in out.windows(2) {
119 assert!(pair[1] > pair[0], "silu should be increasing here: {out:?}");
120 }
121 }
122
123 #[test]
124 fn gelu_matches_hand_computation_at_zero_and_matches_known_values() {
125 let t = Tensor::from_f32(vec![0.0, 1.0, -1.0], [3]).unwrap();
127 let out = t.gelu().unwrap().to_vec_f32().unwrap();
128 assert_close(out[0], 0.0, 1e-6);
129 assert_close(out[1], 0.841_192, 1e-4);
132 assert_close(out[2], -0.158_808, 1e-4);
133 }
134
135 #[test]
136 fn gelu_approaches_the_identity_for_large_positive_x_and_zero_for_large_negative_x() {
137 let t = Tensor::from_f32(vec![10.0, -10.0], [2]).unwrap();
138 let out = t.gelu().unwrap().to_vec_f32().unwrap();
139 assert_close(out[0], 10.0, 1e-3);
140 assert_close(out[1], 0.0, 1e-3);
141 }
142
143 #[test]
144 fn activations_reject_non_f32_input() {
145 let t = Tensor::from_i32(vec![1, 2, 3], [3]).unwrap();
146 assert!(matches!(t.silu(), Err(Error::DTypeMismatch { .. })));
147 assert!(matches!(t.gelu(), Err(Error::DTypeMismatch { .. })));
148 assert!(matches!(t.tanh(), Err(Error::DTypeMismatch { .. })));
149 assert!(matches!(t.sigmoid(), Err(Error::DTypeMismatch { .. })));
150 }
151
152 #[test]
153 fn tanh_matches_known_values_and_is_odd() {
154 let t = Tensor::from_f32(vec![0.0, 1.0, -1.0], [3]).unwrap();
156 let out = t.tanh().unwrap().to_vec_f32().unwrap();
157 assert_close(out[0], 0.0, 1e-6);
158 assert_close(out[1], 0.761_594_2, 1e-6);
159 assert_close(out[2], -0.761_594_2, 1e-6);
160 }
161
162 #[test]
163 fn tanh_is_monotonic_and_saturates_towards_plus_minus_one() {
164 let t = Tensor::from_f32(vec![-20.0, -2.0, -0.5, 0.0, 0.5, 2.0, 20.0], [7]).unwrap();
165 let out = t.tanh().unwrap().to_vec_f32().unwrap();
166 for pair in out.windows(2) {
167 assert!(pair[1] > pair[0], "tanh should be strictly increasing: {out:?}");
168 }
169 assert_close(out[0], -1.0, 1e-6); assert_close(out[6], 1.0, 1e-6); }
172
173 #[test]
174 fn sigmoid_matches_known_values_and_the_reflection_identity() {
175 let t = Tensor::from_f32(vec![0.0, 2.0, -2.0], [3]).unwrap();
177 let out = t.sigmoid().unwrap().to_vec_f32().unwrap();
178 assert_close(out[0], 0.5, 1e-6);
179 assert_close(out[1], 0.880_797_1, 1e-6);
180 assert_close(out[2], 1.0 - 0.880_797_1, 1e-6);
181 }
182
183 #[test]
184 fn sigmoid_is_monotonic_saturates_and_stays_in_the_open_unit_interval() {
185 let t = Tensor::from_f32(vec![-40.0, -3.0, 0.0, 3.0, 40.0], [5]).unwrap();
186 let out = t.sigmoid().unwrap().to_vec_f32().unwrap();
187 for pair in out.windows(2) {
188 assert!(pair[1] > pair[0], "sigmoid should be strictly increasing: {out:?}");
189 }
190 for &v in &out {
191 assert!((0.0..=1.0).contains(&v), "sigmoid output escaped [0, 1]: {v}");
192 }
193 assert_close(out[0], 0.0, 1e-6); assert_close(out[4], 1.0, 1e-6); }
196
197 #[test]
198 fn logistic_is_an_alias_for_sigmoid() {
199 let t = Tensor::from_f32(vec![-1.5, 0.0, 0.25, 3.0], [4]).unwrap();
200 assert_eq!(
201 t.logistic().unwrap().to_vec_f32().unwrap(),
202 t.sigmoid().unwrap().to_vec_f32().unwrap(),
203 );
204 }
205
206 #[test]
207 fn tanh_and_sigmoid_preserve_shape_and_respect_a_transposed_view() {
208 let t = Tensor::from_f32(vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0], [2, 3]).unwrap();
210 let tt = t.transpose(0, 1).unwrap(); let out = tt.tanh().unwrap();
212 assert_eq!(out.shape().dims(), &[3, 2]);
213 assert_eq!(
214 out.to_vec_f32().unwrap(),
215 [0.0f32, 3.0, 1.0, 4.0, 2.0, 5.0].iter().map(|x| x.tanh()).collect::<Vec<_>>(),
216 );
217 assert_eq!(tt.sigmoid().unwrap().shape().dims(), &[3, 2]);
218 }
219}