Skip to main content

ruda_tensor/ops/
activation.rs

1use crate::tensor::FloatTensor;
2use crate::{Backend, Scalar, TensorMetadata, get_device_settings};
3use core::f64::consts::SQRT_2;
4
5pub fn gelu_backward_exact<B: Backend>(x: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B> {
6    let cdf = B::float_mul_scalar(
7        B::float_add_scalar(
8            B::float_erf(B::float_div_scalar(x.clone(), SQRT_2.into())),
9            1f32.into(),
10        ),
11        0.5f32.into(),
12    );
13    let pdf = B::float_exp(B::float_mul_scalar(
14        B::float_mul(x.clone(), x.clone()),
15        (-0.5f32).into(),
16    ));
17    let pdf = B::float_mul_scalar(
18        B::float_mul(x, pdf),
19        (core::f64::consts::FRAC_2_SQRT_PI / (2. * SQRT_2)).into(),
20    );
21    B::float_mul(B::float_add(cdf, pdf), grad)
22}
23
24/// Activation function operations.
25///
26/// This trait let backend implementations override activation functions for better performance.
27pub trait ActivationOps<B: Backend> {
28    /// Applies SiLU element-wise. Backends may fuse its arithmetic and rounding.
29    fn silu(tensor: FloatTensor<B>) -> FloatTensor<B> {
30        B::float_mul(tensor.clone(), B::sigmoid(tensor))
31    }
32
33    /// Applies the LeakyReLU activation function.
34    ///
35    /// # Arguments
36    ///
37    /// * `tensor` - The tensor.
38    /// * `negative_slope` - The negative_slope value that values smaller than 0 are multiplied with.
39    ///
40    /// # Returns
41    ///
42    /// The output tensor.
43    fn leaky_relu(tensor: FloatTensor<B>, negative_slope: Scalar) -> FloatTensor<B> {
44        let bool_dtype = get_device_settings::<B>(&B::float_device(&tensor)).bool_dtype;
45        let mask = B::float_lower_elem(tensor.clone(), 0f32.into(), bool_dtype);
46        let scaled_tensor = B::float_mul_scalar(tensor.clone(), negative_slope);
47
48        // Update the tensor where the values are `< 0` by `tensor * negative_slope`.
49        B::float_mask_where(tensor, mask, scaled_tensor)
50    }
51
52    /// Applies the ReLU activation function.
53    ///
54    /// # Arguments
55    ///
56    /// * `tensor` - The tensor.
57    ///
58    /// # Returns
59    ///
60    /// The output tensor.
61    fn relu(tensor: FloatTensor<B>) -> FloatTensor<B> {
62        let bool_dtype = get_device_settings::<B>(&B::float_device(&tensor)).bool_dtype;
63        let mask = B::float_lower_equal_elem(tensor.clone(), 0f32.into(), bool_dtype);
64
65        B::float_mask_fill(tensor, mask, 0f32.into())
66    }
67
68    /// Applies the ReLU activation function backward.
69    ///
70    /// # Arguments
71    ///
72    /// * `output` - The output tensor.
73    ///
74    /// # Returns
75    ///
76    /// The gradient.
77    fn relu_backward(output: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B> {
78        let bool_dtype = get_device_settings::<B>(&B::float_device(&output)).bool_dtype;
79        let mask = B::float_lower_equal_elem(output, 0f32.into(), bool_dtype);
80
81        B::float_mask_fill(grad, mask, 0.into())
82    }
83
84    /// Applies the Gelu activation function.
85    ///
86    /// # Arguments
87    ///
88    /// * `tensor` - The tensor.
89    ///
90    /// # Returns
91    ///
92    /// The output tensor.
93    fn gelu(tensor: FloatTensor<B>) -> FloatTensor<B> {
94        let x = B::float_div_scalar(tensor.clone(), SQRT_2.into());
95        let x = B::float_erf(x);
96        let x = B::float_add_scalar(x, 1f32.into());
97        let x = B::float_mul(tensor, x);
98
99        B::float_div_scalar(x, 2f32.into())
100    }
101    /// Applies the PReLu activation function.
102    /// # Arguments
103    /// * `tensor` - The input tensor
104    /// * `alpha` - The weight tensor
105    fn prelu(tensor: FloatTensor<B>, alpha: FloatTensor<B>) -> FloatTensor<B> {
106        let bool_dtype = get_device_settings::<B>(&B::float_device(&tensor)).bool_dtype;
107        let mask = B::float_lower_elem(tensor.clone(), 0f32.into(), bool_dtype);
108        let scaled_tensor = B::float_mul(tensor.clone(), alpha);
109        B::float_mask_where(tensor, mask, scaled_tensor)
110    }
111
112    /// Applies the Gelu activation function backward.
113    ///
114    /// # Arguments
115    ///
116    /// * `x` - The tensor.
117    /// * `grad` - The gradient.
118    ///
119    /// # Returns
120    ///
121    /// The output tensor.
122    fn gelu_backward(x: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B> {
123        gelu_backward_exact::<B>(x, grad)
124    }
125
126    /// Applies the Sigmoid activation function.
127    ///
128    /// # Arguments
129    ///
130    /// * `tensor` - The tensor.
131    ///
132    /// # Returns
133    ///
134    /// The output tensor.
135    fn sigmoid(tensor: FloatTensor<B>) -> FloatTensor<B> {
136        let dtype = tensor.dtype();
137        let tensor_full = B::float_cast(tensor, ruda_core::tensor::FloatDType::F32);
138        let tensor_tmp = B::float_exp(B::float_neg(B::float_log(B::float_add_scalar(
139            B::float_exp(B::float_neg(tensor_full)),
140            1.0.into(),
141        ))));
142
143        B::float_cast(tensor_tmp, dtype.into())
144    }
145
146    /// Applies the Sigmoid activation function backward.
147    ///
148    /// # Arguments
149    ///
150    /// * `output` - The output tensor of the sigmoid function.
151    /// * `grad` - The gradient.
152    ///
153    /// # Returns
154    ///
155    /// The output tensor.
156    fn sigmoid_backward(output: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B> {
157        let value = B::float_mul(
158            output.clone(),
159            B::float_add_scalar(B::float_neg(output), 1.0.into()),
160        );
161        B::float_mul(value, grad)
162    }
163
164    /// Applies the hard Sigmoid activation function.
165    ///
166    /// # Arguments
167    ///
168    /// * `tensor` - The tensor.
169    /// * `alpha` - The alpha value that the tensor is multiplied with.
170    /// * `beta` - The beta value that is added to the tensor
171    ///
172    /// # Returns
173    ///
174    /// The output tensor.
175    fn hard_sigmoid(tensor: FloatTensor<B>, alpha: Scalar, beta: Scalar) -> FloatTensor<B> {
176        let dtype = tensor.dtype();
177        let tensor_full = B::float_cast(tensor, ruda_core::tensor::FloatDType::F32);
178
179        let tensor_tmp = B::float_clamp(
180            B::float_add_scalar(B::float_mul_scalar(tensor_full, alpha), beta),
181            0.0.into(),
182            1.0.into(),
183        );
184
185        B::float_cast(tensor_tmp, dtype.into())
186    }
187
188    /// Applies the LogSigmoid activation function.
189    ///
190    /// # Arguments
191    ///
192    /// * `tensor` - The tensor.
193    ///
194    /// # Returns
195    ///
196    /// The output tensor.
197    fn log_sigmoid(tensor: FloatTensor<B>) -> FloatTensor<B> {
198        // To avoid overflow, we use the log-sum-exp trick.
199        //
200        // ```ignore
201        // log(sigmoid(x)) = log(1/(1 + exp(-x)))
202        //                 = log(1) - log(1 + exp(-x))
203        //                 = -log(1 + exp(-x))
204        //                 = -log(exp(0) + exp(-x))
205        // ```
206        // The `exp(t)` of even a moderate-magnitude positive number can be astronomically huge, so we
207        // subtract the `max(t, 0)` of each value (where `t = -x` in this case). This results in the
208        // following equivalence:
209        // ```ignore
210        // log(sigmoid(x)) = -(max(-x, 0) + log(exp(-max(-x, 0)) + exp(-x - max(-x, 0))))
211        // ```
212        //
213        // This extends the range of values for which we obtain accurate results.
214
215        // max(-x, 0)
216        let bool_dtype = get_device_settings::<B>(&B::float_device(&tensor)).bool_dtype;
217        let tensor_neg = B::float_neg(tensor);
218        let mask = B::float_lower_elem(tensor_neg.clone(), 0f32.into(), bool_dtype);
219        let max_elem = B::float_mask_fill(tensor_neg.clone(), mask, 0f32.into());
220        let max_elem_neg = B::float_neg(max_elem.clone());
221
222        // z = exp(-max(-x, 0)) + exp(-x - max(-x, 0))
223        let z = B::float_add(
224            B::float_exp(max_elem_neg.clone()),
225            B::float_exp(B::float_sub(tensor_neg, max_elem.clone())),
226        );
227
228        // -max(-x, 0) - log(-z)
229        B::float_sub(max_elem_neg, B::float_log(z))
230    }
231
232    /// Applies the softmax function along the given dimension.
233    ///
234    /// Uses the max-shift trick for numerical stability: the per-row `max` is detached
235    /// so no gradient flows back through it (the shift is a numerical-stability
236    /// transformation, not part of the function).
237    ///
238    /// # Arguments
239    ///
240    /// * `tensor` - The tensor.
241    /// * `dim` - The dimension along which softmax is computed.
242    ///
243    /// # Returns
244    ///
245    /// The output tensor.
246    fn softmax(tensor: FloatTensor<B>, dim: usize) -> FloatTensor<B> {
247        let max = B::float_max_dim(B::float_detach(tensor.clone()), dim);
248        let shifted = B::float_sub(tensor, max);
249        let exp = B::float_exp(shifted);
250        let sum = B::float_sum_dim(exp.clone(), dim);
251        B::float_div(exp, sum)
252    }
253
254    /// Applies the log-softmax function along the given dimension.
255    ///
256    /// Computed via the log-sum-exp trick with a detached max-shift for numerical
257    /// stability.
258    ///
259    /// # Arguments
260    ///
261    /// * `tensor` - The tensor.
262    /// * `dim` - The dimension along which log-softmax is computed.
263    ///
264    /// # Returns
265    ///
266    /// The output tensor.
267    fn log_softmax(tensor: FloatTensor<B>, dim: usize) -> FloatTensor<B> {
268        let max = B::float_max_dim(B::float_detach(tensor.clone()), dim);
269        let shifted = B::float_sub(tensor, max);
270        let log_sum_exp = B::float_log(B::float_sum_dim(B::float_exp(shifted.clone()), dim));
271        B::float_sub(shifted, log_sum_exp)
272    }
273
274    /// Applies the softmin function along the given dimension.
275    ///
276    /// Equivalent to `softmax(-tensor, dim)`.
277    ///
278    /// # Arguments
279    ///
280    /// * `tensor` - The tensor.
281    /// * `dim` - The dimension along which softmin is computed.
282    ///
283    /// # Returns
284    ///
285    /// The output tensor.
286    fn softmin(tensor: FloatTensor<B>, dim: usize) -> FloatTensor<B> {
287        Self::softmax(B::float_neg(tensor), dim)
288    }
289
290    /// Applies the LogSigmoid activation function backward.
291    ///
292    /// # Arguments
293    ///
294    /// * `x` - The input tensor.
295    /// * `grad` - The gradient.
296    ///
297    /// # Returns
298    ///
299    /// The output gradient.
300    fn log_sigmoid_backward(x: FloatTensor<B>, grad: FloatTensor<B>) -> FloatTensor<B> {
301        // Derivative of -max(-x, 0) - log(exp(-max(-x, 0)) - exp(-x - max(-x, 0)))) is
302        // -max_derive - (-max_derive * exp(-max(-x, 0)) + (-1 - max_derive) * exp(-x - max(-x, 0))) / z
303        // where z = exp(-max(-x, 0)) + exp(-x - max(-x, 0))
304        //
305        // This simplifies to:
306        // -max_derive - (z-1)/z if x is >= 0
307        // -max_derive + (z-1)/z if x is < 0
308
309        let shape = x.shape();
310        let dtype = x.dtype();
311        let device = B::float_device(&x);
312        let bool_dtype = get_device_settings::<B>(&device).bool_dtype;
313
314        // max(-x, 0)
315        let x_neg = B::float_neg(x);
316        let mask = B::float_lower_elem(x_neg.clone(), 0f32.into(), bool_dtype); // -x < 0 or x >= 0
317        let max_elem = B::float_mask_fill(x_neg.clone(), mask.clone(), 0f32.into());
318
319        // z = exp(-max(-x, 0)) + exp(-x - max(-x, 0))
320        let z = B::float_add(
321            B::float_exp(B::float_neg(max_elem.clone())),
322            B::float_exp(B::float_sub(x_neg, max_elem)),
323        );
324
325        // Derivative of max(-x, 0) is 1 if x < 0 or 0 if x >= 0
326        let ones = B::float_ones(shape, &device, dtype.into());
327        let max_derive = B::float_mask_fill(ones.clone(), mask.clone(), 0f32.into());
328        let sign = B::float_mask_fill(ones.clone(), mask, (-1f32).into());
329
330        // grad * (max_derive - sign * (1 - (1 / z)))
331        B::float_mul(
332            grad,
333            B::float_sub(
334                max_derive,
335                B::float_mul(sign, B::float_sub(ones, B::float_recip(z))),
336            ),
337        )
338    }
339}