1#![allow(dead_code)]
4type Result<T, E = torsh_core::error::TorshError> = std::result::Result<T, E>;
7
8use torsh_tensor::Tensor;
9
10pub fn leaky_relu(input: &Tensor, negative_slope: f64) -> Result<Tensor> {
13 let zero_tensor = torsh_tensor::creation::zeros_like(input)?;
14 let positive_part = input.maximum(&zero_tensor)?;
15 let negative_part = input
16 .minimum(&zero_tensor)?
17 .mul_scalar(negative_slope as f32)?;
18 Ok(positive_part.add(&negative_part)?)
19}
20
21pub fn elu(input: &Tensor, alpha: f64) -> Result<Tensor> {
24 let zero_tensor = torsh_tensor::creation::zeros_like(input)?;
25 let positive_part = input.maximum(&zero_tensor)?;
26 let negative_part = input.minimum(&zero_tensor)?;
27 let exp_part = negative_part
28 .exp()?
29 .sub_scalar(1.0)?
30 .mul_scalar(alpha as f32)?;
31 Ok(positive_part.add(&exp_part)?)
32}
33
34pub fn swish(input: &Tensor) -> Result<Tensor> {
37 let one_tensor = torsh_tensor::creation::ones_like(input)?;
38 let neg_input = input.neg()?;
39 let exp_neg = neg_input.exp()?;
40 let denominator = one_tensor.add(&exp_neg)?;
41 let sigmoid = one_tensor.div(&denominator)?;
42 Ok(input.mul(&sigmoid)?)
43}
44
45pub fn gelu(input: &Tensor) -> Result<Tensor> {
48 let sqrt_2_over_pi = (2.0 / std::f64::consts::PI).sqrt() as f32;
49 let coefficient = 0.044715_f32;
50
51 let x_cubed = input.pow_scalar(3.0)?;
52 let coeff_term = x_cubed.mul_scalar(coefficient)?;
53 let sum_term = input.add(&coeff_term)?;
54 let inner = sum_term.mul_scalar(sqrt_2_over_pi)?;
55 let tanh_part = inner.tanh()?;
56
57 let one_tensor = torsh_tensor::creation::ones_like(input)?;
58 let one_plus_tanh = one_tensor.add(&tanh_part)?;
59 let half_tensor = torsh_tensor::creation::ones_like(input)?.mul_scalar(0.5)?;
60
61 Ok(half_tensor.mul(input)?.mul(&one_plus_tanh)?)
62}
63
64pub fn mish(input: &Tensor) -> Result<Tensor> {
67 let one_tensor = torsh_tensor::creation::ones_like(input)?;
68 let exp_input = input.exp()?;
69 let one_plus_exp = one_tensor.add(&exp_input)?;
70 let softplus = one_plus_exp.ln()?;
71 let tanh_softplus = softplus.tanh()?;
72 Ok(input.mul(&tanh_softplus)?)
73}
74
75pub mod normalization {
77 use super::*;
78
79 pub fn layer_norm(input: &Tensor, eps: f64) -> Result<Tensor> {
82 let shape = input.shape();
83 let dims = shape.dims();
84 let last_dim = dims.len() - 1;
85 let mean = input.mean(Some(&[last_dim]), false)?;
86
87 let input_shape = dims.to_vec();
89 let mut mean_shape = input_shape.clone();
90 mean_shape[last_dim] = 1;
91
92 let mean_expanded = mean.unsqueeze(-1)?;
93 let mean_broadcasted = mean_expanded.expand(&input_shape)?;
94 let diff = input.sub(&mean_broadcasted)?;
95
96 let variance = diff.pow_scalar(2.0)?.mean(Some(&[last_dim]), false)?;
97
98 let variance_expanded = variance.unsqueeze(-1)?;
99 let variance_broadcasted = variance_expanded.expand(&input_shape)?;
100 let eps_tensor =
101 torsh_tensor::creation::ones_like(&variance_broadcasted)?.mul_scalar(eps as f32)?;
102 let variance_plus_eps = variance_broadcasted.add(&eps_tensor)?;
103 let std = variance_plus_eps.sqrt()?;
104 Ok(diff.div(&std)?)
105 }
106
107 pub fn graph_norm(input: &Tensor, edge_index: &Tensor, num_nodes: usize) -> Result<Tensor> {
110 let edge_index_i32 = edge_index.to_i32_simd()?;
112 let edge_index_i64 = edge_index_i32.to_i64_simd()?;
113 let edge_data = crate::utils::tensor_to_vec2::<i64>(&edge_index_i64)?;
114 let mut degrees = vec![0.0_f32; num_nodes];
115
116 for j in 0..edge_data[0].len() {
117 let src = edge_data[0][j] as usize;
118 let dst = edge_data[1][j] as usize;
119 if src < num_nodes {
120 degrees[src] += 1.0;
121 }
122 if dst < num_nodes {
123 degrees[dst] += 1.0;
124 }
125 }
126
127 let degree_tensor = torsh_tensor::creation::from_vec(
129 degrees
130 .iter()
131 .map(|&d| if d > 0.0 { 1.0 / d.sqrt() } else { 0.0 })
132 .collect(),
133 &[num_nodes],
134 torsh_core::device::DeviceType::Cpu,
135 )?;
136
137 let input_shape_binding = input.shape();
139 let input_shape = input_shape_binding.dims();
140 let degree_expanded = degree_tensor.unsqueeze(-1)?.expand(input_shape)?;
141 Ok(input.mul(°ree_expanded)?)
142 }
143
144 pub fn batch_norm(input: &Tensor, eps: f64) -> Result<Tensor> {
147 let mean = input.mean(Some(&[0]), true)?;
148 let input_shape_binding = input.shape();
149 let input_shape = input_shape_binding.dims();
150 let mean_expanded = mean.expand(input_shape)?;
151 let diff = input.sub(&mean_expanded)?;
152 let variance = diff.pow_scalar(2.0)?.mean(Some(&[0]), true)?;
153 let eps_tensor = torsh_tensor::creation::ones_like(&variance)?.mul_scalar(eps as f32)?;
154 let variance_plus_eps = variance.add(&eps_tensor)?;
155 let std = variance_plus_eps.sqrt()?;
156 let std_expanded = std.expand(input_shape)?;
157 Ok(diff.div(&std_expanded)?)
158 }
159}
160
161pub fn dropout(input: &Tensor, p: f64, training: bool) -> Result<Tensor> {
163 if !training || p == 0.0 {
164 return Ok(input.clone());
165 }
166
167 if p == 1.0 {
168 return Ok(torsh_tensor::creation::zeros_like(input)?);
169 }
170
171 let keep_prob = 1.0 - p;
173 let random_tensor = torsh_tensor::creation::rand_like(input)?;
174 let keep_prob_tensor =
175 torsh_tensor::creation::ones_like(input)?.mul_scalar(keep_prob as f32)?;
176
177 let diff_tensor = random_tensor.sub(&keep_prob_tensor)?;
179 let mask_raw = diff_tensor.relu()?;
180
181 let zero_tensor = torsh_tensor::creation::zeros_like(input)?;
183 let _mask_binary = mask_raw.gt(&zero_tensor)?;
184
185 let inverted_prob = random_tensor.gt(&keep_prob_tensor)?;
188 let _ones_f32 = torsh_tensor::creation::ones_like(input)?;
189 let _zeros_f32 = torsh_tensor::creation::zeros_like(input)?;
190
191 let inverted_data = inverted_prob.to_vec()?;
193 let keep_data: Vec<f32> = inverted_data
194 .iter()
195 .map(|&drop| if drop { 0.0 } else { 1.0 })
196 .collect();
197 let keep_mask =
198 torsh_tensor::creation::from_vec(keep_data, input.shape().dims(), input.device())?;
199
200 let masked = input.mul(&keep_mask)?;
202 Ok(masked.div_scalar(keep_prob as f32)?)
203}
204
205pub mod attention {
207 use super::*;
208
209 pub fn scaled_dot_product_attention(
211 query: &Tensor,
212 key: &Tensor,
213 value: &Tensor,
214 mask: Option<&Tensor>,
215 ) -> Result<Tensor> {
216 let binding = key.shape();
218 let d_k = binding.dims().last().ok_or_else(|| {
219 torsh_core::error::TorshError::InvalidArgument(
220 "key tensor must have at least one dimension".to_string(),
221 )
222 })?;
223 let scale = 1.0 / (*d_k as f64).sqrt();
224 let key_transposed = key.transpose(-2, -1)?;
225 let scores = query.matmul(&key_transposed)?.mul_scalar(scale as f32)?;
226
227 let masked_scores = if let Some(mask) = mask {
229 let large_neg = torsh_tensor::creation::ones_like(&scores)?.mul_scalar(-1e9_f32)?;
230 let mask_effect = mask.mul(&large_neg)?;
231 scores.add(&mask_effect)?
232 } else {
233 scores
234 };
235
236 let attention_weights = masked_scores.softmax(-1)?;
238 Ok(attention_weights.matmul(value)?)
239 }
240
241 pub fn multi_head_attention(
243 query: &Tensor,
244 key: &Tensor,
245 value: &Tensor,
246 num_heads: usize,
247 mask: Option<&Tensor>,
248 ) -> Result<Tensor> {
249 let batch_size = query.shape().dims()[0];
250 let seq_len = query.shape().dims()[1];
251 let d_model = query.shape().dims()[2];
252 let d_k = d_model / num_heads;
253
254 let q = query
256 .view(&[
257 batch_size as i32,
258 seq_len as i32,
259 num_heads as i32,
260 d_k as i32,
261 ])?
262 .transpose(1, 2)?;
263 let k = key
264 .view(&[
265 batch_size as i32,
266 seq_len as i32,
267 num_heads as i32,
268 d_k as i32,
269 ])?
270 .transpose(1, 2)?;
271 let v = value
272 .view(&[
273 batch_size as i32,
274 seq_len as i32,
275 num_heads as i32,
276 d_k as i32,
277 ])?
278 .transpose(1, 2)?;
279
280 let attention_output = scaled_dot_product_attention(&q, &k, &v, mask);
282
283 Ok(attention_output?.transpose(1, 2)?.contiguous()?.view(&[
285 batch_size as i32,
286 seq_len as i32,
287 d_model as i32,
288 ])?)
289 }
290}
291
292#[cfg(test)]
293mod tests {
294 use super::*;
295 use approx::assert_relative_eq;
296
297 #[test]
298 fn test_leaky_relu() {
299 let input = torsh_tensor::creation::from_vec(
300 vec![-2.0, -1.0, 0.0, 1.0, 2.0],
301 &[5],
302 torsh_core::device::DeviceType::Cpu,
303 )
304 .unwrap();
305
306 let output = leaky_relu(&input, 0.01);
307 let expected = vec![-0.02, -0.01, 0.0, 1.0, 2.0];
308
309 let output_vec = output
310 .expect("operation should succeed")
311 .to_vec()
312 .expect("conversion should succeed");
313 for (actual, expected) in output_vec.iter().zip(expected.iter()) {
314 assert_relative_eq!(actual, expected, epsilon = 1e-6);
315 }
316 }
317
318 #[test]
319 fn test_elu() {
320 let input = torsh_tensor::creation::from_vec(
321 vec![-1.0, 0.0, 1.0],
322 &[3],
323 torsh_core::device::DeviceType::Cpu,
324 )
325 .unwrap();
326
327 let output = elu(&input, 1.0);
328 let output_vec = output
329 .expect("operation should succeed")
330 .to_vec()
331 .expect("conversion should succeed");
332
333 assert_relative_eq!(output_vec[2], 1.0, epsilon = 1e-6);
335
336 assert_relative_eq!(output_vec[1], 0.0, epsilon = 1e-6);
338
339 let expected_negative = 1.0 * ((-1.0_f32).exp() - 1.0);
341 assert_relative_eq!(output_vec[0], expected_negative, epsilon = 1e-6);
342 }
343
344 #[test]
345 fn test_swish() {
346 let input = torsh_tensor::creation::from_vec(
347 vec![0.0, 1.0, -1.0],
348 &[3],
349 torsh_core::device::DeviceType::Cpu,
350 )
351 .unwrap();
352
353 let output = swish(&input);
354 let output_vec = output
355 .expect("operation should succeed")
356 .to_vec()
357 .expect("conversion should succeed");
358
359 assert_relative_eq!(output_vec[0], 0.0, epsilon = 1e-6);
361
362 let sigmoid_1 = 1.0 / (1.0 + (-1.0_f32).exp());
364 assert_relative_eq!(output_vec[1], sigmoid_1, epsilon = 1e-6);
365 }
366
367 #[test]
368 fn test_layer_norm() {
369 let input = torsh_tensor::creation::from_vec(
370 vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0],
371 &[2, 3],
372 torsh_core::device::DeviceType::Cpu,
373 )
374 .unwrap();
375
376 let output = normalization::layer_norm(&input, 1e-8).expect("operation should succeed");
377
378 let output_2d = crate::utils::tensor_to_vec2::<f32>(&output).unwrap();
380
381 for row in output_2d {
382 let mean: f32 = row.iter().sum::<f32>() / row.len() as f32;
383 let var: f32 = row.iter().map(|&x| (x - mean).powi(2)).sum::<f32>() / row.len() as f32;
384
385 assert_relative_eq!(mean, 0.0, epsilon = 1e-6);
386 assert_relative_eq!(var.sqrt(), 1.0, epsilon = 1e-6);
387 }
388 }
389
390 #[test]
391 fn test_dropout() {
392 let input = torsh_tensor::creation::ones(&[100]).unwrap();
393
394 let output_training = dropout(&input, 0.5, true);
396 let output_vec = output_training
397 .expect("operation should succeed")
398 .to_vec()
399 .expect("conversion should succeed");
400
401 let num_zeros = output_vec.iter().filter(|&&x| x == 0.0).count();
403 let num_nonzeros = output_vec.iter().filter(|&&x| x != 0.0).count();
404
405 assert!(num_zeros > 30 && num_zeros < 70); assert!(num_nonzeros > 30 && num_nonzeros < 70);
407
408 let output_eval = dropout(&input, 0.5, false);
410 let output_eval_vec = output_eval
411 .expect("operation should succeed")
412 .to_vec()
413 .expect("conversion should succeed");
414 assert!(output_eval_vec.iter().all(|&x| x == 1.0));
415 }
416}