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}