Skip to main content

torsh_graph/
functional.rs

1//! Functional operations and activations for graph neural networks
2// Framework infrastructure - components designed for future use
3#![allow(dead_code)]
4/// Crate-local result alias: the error type defaults to [`TorshError`],
5/// so both `Result<T>` and `Result<T, OtherError>` stay valid.
6type Result<T, E = torsh_core::error::TorshError> = std::result::Result<T, E>;
7
8use torsh_tensor::Tensor;
9
10/// Leaky ReLU activation function
11/// LeakyReLU(x) = max(0, x) + negative_slope * min(0, x)
12pub 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
21/// ELU (Exponential Linear Unit) activation function
22/// ELU(x) = x if x > 0, alpha * (exp(x) - 1) if x <= 0
23pub 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
34/// Swish activation function (also known as SiLU)
35/// Swish(x) = x * sigmoid(x)
36pub 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
45/// GELU activation function (Gaussian Error Linear Unit)
46/// GELU(x) = 0.5 * x * (1 + tanh(sqrt(2/π) * (x + 0.044715 * x^3)))
47pub 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
64/// Mish activation function
65/// Mish(x) = x * tanh(softplus(x)) = x * tanh(ln(1 + exp(x)))
66pub 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
75/// Graph-specific normalization functions
76pub mod normalization {
77    use super::*;
78
79    /// Layer normalization for graph features
80    /// Normalizes features across the feature dimension for each node
81    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        // Create broadcast-compatible mean tensor
88        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    /// Graph normalization
108    /// Normalizes node features by the square root of node degree
109    pub fn graph_norm(input: &Tensor, edge_index: &Tensor, num_nodes: usize) -> Result<Tensor> {
110        // Convert f32 edge_index to i64 for processing
111        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        // Create degree tensor and normalize
128        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        // Apply normalization - expand degree tensor to match input shape
138        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(&degree_expanded)?)
142    }
143
144    /// Batch normalization for graphs
145    /// Normalizes features across the batch dimension
146    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
161/// Dropout function for regularization
162pub 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    // Create dropout mask using simple thresholding
172    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    // Create a mask by subtracting threshold and applying relu (values > 0 become positive, <= 0 become 0)
178    let diff_tensor = random_tensor.sub(&keep_prob_tensor)?;
179    let mask_raw = diff_tensor.relu()?;
180
181    // Normalize mask to 0 or 1 by checking if > 0, then convert to binary mask
182    let zero_tensor = torsh_tensor::creation::zeros_like(input)?;
183    let _mask_binary = mask_raw.gt(&zero_tensor)?;
184
185    // For now, let's use a simplified approach - just threshold the random tensor directly
186    // This creates a binary effect where values above threshold are kept
187    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    // Manual conversion from boolean to f32: if inverted_prob is true, use 0.0 (drop), else use 1.0 (keep)
192    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    // Apply dropout with proper scaling
201    let masked = input.mul(&keep_mask)?;
202    Ok(masked.div_scalar(keep_prob as f32)?)
203}
204
205/// Attention mechanism functions
206pub mod attention {
207    use super::*;
208
209    /// Scaled dot-product attention
210    pub fn scaled_dot_product_attention(
211        query: &Tensor,
212        key: &Tensor,
213        value: &Tensor,
214        mask: Option<&Tensor>,
215    ) -> Result<Tensor> {
216        // Compute attention scores: Q @ K^T / sqrt(d_k)
217        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        // Apply mask if provided
228        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        // Apply softmax and compute attention output
237        let attention_weights = masked_scores.softmax(-1)?;
238        Ok(attention_weights.matmul(value)?)
239    }
240
241    /// Multi-head attention mechanism
242    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        // Reshape for multi-head attention
255        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        // Apply scaled dot-product attention
281        let attention_output = scaled_dot_product_attention(&q, &k, &v, mask);
282
283        // Reshape back to original dimensions
284        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        // For x > 0, ELU(x) = x
334        assert_relative_eq!(output_vec[2], 1.0, epsilon = 1e-6);
335
336        // For x = 0, ELU(x) = 0
337        assert_relative_eq!(output_vec[1], 0.0, epsilon = 1e-6);
338
339        // For x < 0, ELU(x) = alpha * (exp(x) - 1)
340        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        // Swish(0) = 0 * sigmoid(0) = 0 * 0.5 = 0
360        assert_relative_eq!(output_vec[0], 0.0, epsilon = 1e-6);
361
362        // Swish(1) = 1 * sigmoid(1) ≈ 1 * 0.731 ≈ 0.731
363        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        // Check that each row is normalized (mean ≈ 0, std ≈ 1)
379        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        // Test training mode with 0.5 dropout
395        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        // Should have approximately 50% zeros and 50% values scaled by 2.0
402        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); // Roughly 50% with some variance
406        assert!(num_nonzeros > 30 && num_nonzeros < 70);
407
408        // Test eval mode (no dropout)
409        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}