Skip to main content

torsh_nn/layers/
advanced.rs

1//! Advanced neural network layers using SciRS2 algorithms
2
3use crate::{Module, Parameter};
4use torsh_core::error::TorshError;
5use torsh_tensor::{
6    creation::{ones, randn, zeros},
7    Tensor,
8};
9
10/// Multi-Head Attention layer
11pub struct MultiHeadAttention {
12    pub num_heads: usize,
13    pub d_model: usize,
14    pub d_k: usize,
15    pub d_v: usize,
16
17    // Weight matrices
18    pub w_q: Parameter,
19    pub w_k: Parameter,
20    pub w_v: Parameter,
21    pub w_o: Parameter,
22
23    // Optional bias terms
24    pub bias_q: Option<Parameter>,
25    pub bias_k: Option<Parameter>,
26    pub bias_v: Option<Parameter>,
27    pub bias_o: Option<Parameter>,
28
29    pub dropout: f64,
30    pub scale: f64,
31}
32
33impl MultiHeadAttention {
34    /// Create a new multi-head attention layer
35    pub fn new(
36        d_model: usize,
37        num_heads: usize,
38        dropout: f64,
39        bias: bool,
40    ) -> Result<Self, TorshError> {
41        let d_k = d_model / num_heads;
42        let d_v = d_model / num_heads;
43
44        let scale = 1.0 / (d_k as f64).sqrt();
45
46        // Initialize weight matrices with Xavier/Glorot initialization
47        let fan_in = d_model as f64;
48        let fan_out = d_model as f64;
49        let std = (2.0 / (fan_in + fan_out)).sqrt();
50
51        let w_q = Parameter::new(randn(&[d_model, d_model])?.mul_scalar(std as f32)?);
52        let w_k = Parameter::new(randn(&[d_model, d_model])?.mul_scalar(std as f32)?);
53        let w_v = Parameter::new(randn(&[d_model, d_model])?.mul_scalar(std as f32)?);
54        let w_o = Parameter::new(randn(&[d_model, d_model])?.mul_scalar(std as f32)?);
55
56        let bias_q = if bias {
57            Some(Parameter::new(zeros(&[d_model])?))
58        } else {
59            None
60        };
61        let bias_k = if bias {
62            Some(Parameter::new(zeros(&[d_model])?))
63        } else {
64            None
65        };
66        let bias_v = if bias {
67            Some(Parameter::new(zeros(&[d_model])?))
68        } else {
69            None
70        };
71        let bias_o = if bias {
72            Some(Parameter::new(zeros(&[d_model])?))
73        } else {
74            None
75        };
76
77        Ok(Self {
78            num_heads,
79            d_model,
80            d_k,
81            d_v,
82            w_q,
83            w_k,
84            w_v,
85            w_o,
86            bias_q,
87            bias_k,
88            bias_v,
89            bias_o,
90            dropout,
91            scale,
92        })
93    }
94
95    /// Apply attention mechanism
96    ///
97    /// Implements scaled dot-product attention: Attention(Q,K,V) = softmax(QK^T / sqrt(d_k))V
98    ///
99    /// # Arguments
100    /// * `q` - Query tensor [batch_size, num_heads, seq_len, d_k]
101    /// * `k` - Key tensor [batch_size, num_heads, seq_len, d_k]
102    /// * `v` - Value tensor [batch_size, num_heads, seq_len, d_v]
103    /// * `mask` - Optional attention mask
104    ///
105    /// # Returns
106    /// Attention output [batch_size, num_heads, seq_len, d_v]
107    pub fn attention(
108        &self,
109        q: &Tensor,
110        k: &Tensor,
111        v: &Tensor,
112        mask: Option<&Tensor>,
113    ) -> Result<Tensor, TorshError> {
114        // Input shape: [batch_size, num_heads, seq_len, d_k]
115        let q_shape_binding = q.shape();
116        let q_shape = q_shape_binding.dims();
117        let batch_size = q_shape[0];
118        let num_heads = q_shape[1];
119        let seq_len_q = q_shape[2];
120        let d_k = q_shape[3];
121
122        let k_shape_binding = k.shape();
123        let k_shape = k_shape_binding.dims();
124        let seq_len_k = k_shape[2];
125
126        // Step 1: Compute attention scores: Q @ K^T
127        // We need to transpose the last two dimensions of K
128        // K: [batch, heads, seq_k, d_k] -> K^T: [batch, heads, d_k, seq_k]
129
130        // Flatten batch and heads dimensions for easier processing
131        // q: [batch*heads, seq_q, d_k]
132        // k: [batch*heads, seq_k, d_k]
133        // v: [batch*heads, seq_k, d_v]
134        let batch_heads = batch_size * num_heads;
135
136        let q_flat = q.view(&[batch_heads as i32, seq_len_q as i32, d_k as i32])?;
137        let k_flat = k.view(&[batch_heads as i32, seq_len_k as i32, d_k as i32])?;
138        let v_flat = v.view(&[batch_heads as i32, seq_len_k as i32, d_k as i32])?;
139
140        // Compute Q @ K^T for each batch*head
141        // We need to do this manually since we need per-batch matmul
142        let q_data = q_flat.to_vec()?;
143        let k_data = k_flat.to_vec()?;
144        let v_data = v_flat.to_vec()?;
145
146        let mut scores_data = vec![0.0f32; batch_heads * seq_len_q * seq_len_k];
147
148        // For each batch*head
149        for bh in 0..batch_heads {
150            let q_offset = bh * seq_len_q * d_k;
151            let k_offset = bh * seq_len_k * d_k;
152            let scores_offset = bh * seq_len_q * seq_len_k;
153
154            // Compute Q @ K^T for this batch*head
155            for i in 0..seq_len_q {
156                for j in 0..seq_len_k {
157                    let mut dot_product = 0.0f32;
158                    for d in 0..d_k {
159                        let q_val = q_data[q_offset + i * d_k + d];
160                        let k_val = k_data[k_offset + j * d_k + d];
161                        dot_product += q_val * k_val;
162                    }
163                    // Scale by sqrt(d_k) for numerical stability
164                    scores_data[scores_offset + i * seq_len_k + j] =
165                        dot_product / (d_k as f32).sqrt();
166                }
167            }
168        }
169
170        // Step 2: Apply mask if provided (add large negative value to masked positions)
171        if let Some(mask_tensor) = mask {
172            let mask_data = mask_tensor.to_vec()?;
173            for i in 0..scores_data.len() {
174                if mask_data[i] == 0.0 {
175                    scores_data[i] = -1e9; // Large negative value for masked positions
176                }
177            }
178        }
179
180        // Step 3: Apply softmax over the last dimension (seq_len_k)
181        // Softmax is applied row-wise (for each query position)
182        for bh in 0..batch_heads {
183            for i in 0..seq_len_q {
184                let row_offset = bh * seq_len_q * seq_len_k + i * seq_len_k;
185
186                // Find max for numerical stability
187                let max_val = scores_data[row_offset..row_offset + seq_len_k]
188                    .iter()
189                    .fold(f32::NEG_INFINITY, |a, &b| a.max(b));
190
191                // Compute exp(x - max) and sum
192                let mut exp_sum = 0.0f32;
193                for j in 0..seq_len_k {
194                    let idx = row_offset + j;
195                    scores_data[idx] = (scores_data[idx] - max_val).exp();
196                    exp_sum += scores_data[idx];
197                }
198
199                // Normalize
200                for j in 0..seq_len_k {
201                    let idx = row_offset + j;
202                    scores_data[idx] /= exp_sum + 1e-9; // Add epsilon for numerical stability
203                }
204            }
205        }
206
207        // Step 4: Apply attention to values: attention_weights @ V
208        // scores: [batch*heads, seq_q, seq_k]
209        // v: [batch*heads, seq_k, d_k]
210        // output: [batch*heads, seq_q, d_k]
211        let mut output_data = vec![0.0f32; batch_heads * seq_len_q * d_k];
212
213        for bh in 0..batch_heads {
214            let scores_offset = bh * seq_len_q * seq_len_k;
215            let v_offset = bh * seq_len_k * d_k;
216            let output_offset = bh * seq_len_q * d_k;
217
218            for i in 0..seq_len_q {
219                for d in 0..d_k {
220                    let mut weighted_sum = 0.0f32;
221                    for j in 0..seq_len_k {
222                        let attention_weight = scores_data[scores_offset + i * seq_len_k + j];
223                        let v_val = v_data[v_offset + j * d_k + d];
224                        weighted_sum += attention_weight * v_val;
225                    }
226                    output_data[output_offset + i * d_k + d] = weighted_sum;
227                }
228            }
229        }
230
231        // Reshape back to [batch_size, num_heads, seq_len_q, d_k]
232        let output_flat = Tensor::from_vec(output_data, &[batch_heads, seq_len_q, d_k])?;
233        output_flat.view(&[
234            batch_size as i32,
235            num_heads as i32,
236            seq_len_q as i32,
237            d_k as i32,
238        ])
239    }
240}
241
242impl Module for MultiHeadAttention {
243    fn forward(&self, input: &Tensor) -> Result<Tensor, TorshError> {
244        let batch_size = input.shape().dims()[0];
245        let seq_len = input.shape().dims()[1];
246        let d_model = input.shape().dims()[2];
247
248        // Reshape input to 2D for linear transformations: [batch_size * seq_len, d_model]
249        let input_2d = input.view(&[(batch_size * seq_len) as i32, d_model as i32])?;
250
251        // Linear transformations for Q, K, V
252        let q_2d = input_2d.matmul(&self.w_q.clone_data())?;
253        let k_2d = input_2d.matmul(&self.w_k.clone_data())?;
254        let v_2d = input_2d.matmul(&self.w_v.clone_data())?;
255
256        // Reshape back to 3D: [batch_size, seq_len, d_model]
257        let q = q_2d.view(&[batch_size as i32, seq_len as i32, d_model as i32])?;
258        let k = k_2d.view(&[batch_size as i32, seq_len as i32, d_model as i32])?;
259        let v = v_2d.view(&[batch_size as i32, seq_len as i32, d_model as i32])?;
260
261        // Add bias if present
262        let q = if let Some(ref bias) = self.bias_q {
263            q.add(&bias.clone_data())?
264        } else {
265            q
266        };
267
268        let k = if let Some(ref bias) = self.bias_k {
269            k.add(&bias.clone_data())?
270        } else {
271            k
272        };
273
274        let v = if let Some(ref bias) = self.bias_v {
275            v.add(&bias.clone_data())?
276        } else {
277            v
278        };
279
280        // Reshape for multi-head attention
281        // [batch_size, seq_len, d_model] -> [batch_size, num_heads, seq_len, d_k]
282        let q = q
283            .view(&[
284                batch_size as i32,
285                seq_len as i32,
286                self.num_heads as i32,
287                self.d_k as i32,
288            ])?
289            .transpose(1, 2)?;
290        let k = k
291            .view(&[
292                batch_size as i32,
293                seq_len as i32,
294                self.num_heads as i32,
295                self.d_k as i32,
296            ])?
297            .transpose(1, 2)?;
298        let v = v
299            .view(&[
300                batch_size as i32,
301                seq_len as i32,
302                self.num_heads as i32,
303                self.d_v as i32,
304            ])?
305            .transpose(1, 2)?;
306
307        // Apply attention
308        let attention_output = self.attention(&q, &k, &v, None)?;
309
310        // Reshape back to original dimensions
311        let attention_output = attention_output.transpose(1, 2)?.contiguous()?.view(&[
312            batch_size as i32,
313            seq_len as i32,
314            self.d_model as i32,
315        ])?;
316
317        // Final linear transformation - reshape to 2D for matmul
318        let output_2d =
319            attention_output.view(&[(batch_size * seq_len) as i32, self.d_model as i32])?;
320        let output_transformed = output_2d.matmul(&self.w_o.clone_data())?;
321        let output =
322            output_transformed.view(&[batch_size as i32, seq_len as i32, self.d_model as i32])?;
323
324        if let Some(ref bias) = self.bias_o {
325            Ok(output.add(&bias.clone_data())?)
326        } else {
327            Ok(output)
328        }
329    }
330
331    fn parameters(&self) -> std::collections::HashMap<String, Parameter> {
332        let mut params = std::collections::HashMap::new();
333        params.insert("w_q".to_string(), self.w_q.clone());
334        params.insert("w_k".to_string(), self.w_k.clone());
335        params.insert("w_v".to_string(), self.w_v.clone());
336        params.insert("w_o".to_string(), self.w_o.clone());
337
338        if let Some(ref bias) = self.bias_q {
339            params.insert("bias_q".to_string(), bias.clone());
340        }
341        if let Some(ref bias) = self.bias_k {
342            params.insert("bias_k".to_string(), bias.clone());
343        }
344        if let Some(ref bias) = self.bias_v {
345            params.insert("bias_v".to_string(), bias.clone());
346        }
347        if let Some(ref bias) = self.bias_o {
348            params.insert("bias_o".to_string(), bias.clone());
349        }
350
351        params
352    }
353
354    fn train(&mut self) {
355        // Set to training mode
356    }
357
358    fn eval(&mut self) {
359        // Set to evaluation mode
360    }
361}
362
363/// Advanced Layer Normalization with learnable parameters for transformers
364pub struct AdvancedLayerNorm {
365    pub normalized_shape: Vec<usize>,
366    pub weight: Parameter,
367    pub bias: Option<Parameter>,
368    pub eps: f64,
369}
370
371impl AdvancedLayerNorm {
372    /// Create a new layer normalization layer
373    pub fn new(normalized_shape: Vec<usize>, bias: bool, eps: f64) -> Result<Self, TorshError> {
374        let num_features = normalized_shape.iter().product();
375
376        let weight = Parameter::new(ones(&[num_features])?);
377        let bias = if bias {
378            Some(Parameter::new(zeros(&[num_features])?))
379        } else {
380            None
381        };
382
383        Ok(Self {
384            normalized_shape,
385            weight,
386            bias,
387            eps,
388        })
389    }
390}
391
392impl Module for AdvancedLayerNorm {
393    fn forward(&self, input: &Tensor) -> Result<Tensor, TorshError> {
394        // Implement proper layer normalization
395        // Layer norm: (x - mean) / sqrt(var + eps) * weight + bias
396
397        let input_shape_binding = input.shape();
398        let input_shape = input_shape_binding.dims();
399        let num_features = self.normalized_shape.iter().product::<usize>();
400
401        // Verify that the normalized shape matches the last dimensions of input
402        let input_suffix = &input_shape[input_shape.len() - self.normalized_shape.len()..];
403        if input_suffix != self.normalized_shape.as_slice() {
404            return Err(TorshError::InvalidArgument(format!(
405                "Normalized shape {:?} doesn't match input shape suffix {:?}",
406                self.normalized_shape, input_suffix
407            )));
408        }
409
410        // Calculate batch dimensions
411        let batch_size: usize = input_shape[..input_shape.len() - self.normalized_shape.len()]
412            .iter()
413            .product();
414
415        // Get input data
416        let input_data = input.to_vec()?;
417        let weight_data = self.weight.clone_data().to_vec()?;
418        let bias_data = if let Some(ref bias) = self.bias {
419            Some(bias.clone_data().to_vec()?)
420        } else {
421            None
422        };
423
424        let mut output_data = vec![0.0f32; input_data.len()];
425
426        // Process each instance (normalize over the last dimensions)
427        for b in 0..batch_size {
428            let instance_offset = b * num_features;
429            let instance = &input_data[instance_offset..instance_offset + num_features];
430
431            // Compute mean
432            let mean: f32 = instance.iter().sum::<f32>() / num_features as f32;
433
434            // Compute variance
435            let variance: f32 =
436                instance.iter().map(|&x| (x - mean).powi(2)).sum::<f32>() / num_features as f32;
437
438            // Normalize and apply affine transformation
439            let inv_std = 1.0 / (variance + self.eps as f32).sqrt();
440
441            for i in 0..num_features {
442                let normalized = (instance[i] - mean) * inv_std;
443                let scaled = normalized * weight_data[i];
444                let shifted = if let Some(ref bias) = bias_data {
445                    scaled + bias[i]
446                } else {
447                    scaled
448                };
449                output_data[instance_offset + i] = shifted;
450            }
451        }
452
453        Tensor::from_vec(output_data, input_shape)
454    }
455
456    fn parameters(&self) -> std::collections::HashMap<String, Parameter> {
457        let mut params = std::collections::HashMap::new();
458        params.insert("weight".to_string(), self.weight.clone());
459        if let Some(ref bias) = self.bias {
460            params.insert("bias".to_string(), bias.clone());
461        }
462        params
463    }
464
465    fn train(&mut self) {}
466    fn eval(&mut self) {}
467}
468
469/// Positional Encoding for Transformer models
470pub struct PositionalEncoding {
471    pub encoding: Tensor,
472    pub dropout: f64,
473}
474
475impl PositionalEncoding {
476    /// Create positional encoding
477    ///
478    /// Creates sinusoidal positional encodings as described in "Attention Is All You Need"
479    /// PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
480    /// PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
481    pub fn new(d_model: usize, max_len: usize, dropout: f64) -> Result<Self, TorshError> {
482        // Create sinusoidal positional encoding
483        let mut encoding_data = vec![0.0f32; max_len * d_model];
484
485        for pos in 0..max_len {
486            for i in (0..d_model).step_by(2) {
487                let angle = pos as f32 / 10000.0_f32.powf(i as f32 / d_model as f32);
488
489                // Apply sin to even indices
490                encoding_data[pos * d_model + i] = angle.sin();
491
492                // Apply cos to odd indices
493                if i + 1 < d_model {
494                    encoding_data[pos * d_model + i + 1] = angle.cos();
495                }
496            }
497        }
498
499        let encoding = Tensor::from_vec(encoding_data, &[max_len, d_model])?;
500
501        Ok(Self { encoding, dropout })
502    }
503}
504
505impl Module for PositionalEncoding {
506    fn forward(&self, input: &Tensor) -> Result<Tensor, TorshError> {
507        // Input shape: [batch_size, seq_len, d_model]
508        let input_shape_binding = input.shape();
509        let input_shape = input_shape_binding.dims();
510        let seq_len = input_shape[1];
511        let d_model = input_shape[2];
512
513        // Get the positional encoding for the sequence length
514        // encoding shape: [max_len, d_model]
515        // We need: [seq_len, d_model]
516        let encoding_shape_binding = self.encoding.shape();
517        let encoding_shape = encoding_shape_binding.dims();
518        let max_len = encoding_shape[0];
519
520        if seq_len > max_len {
521            return Err(TorshError::InvalidArgument(format!(
522                "Sequence length {} exceeds maximum positional encoding length {}",
523                seq_len, max_len
524            )));
525        }
526
527        // Slice the encoding to match sequence length
528        let encoding_data = self.encoding.to_vec()?;
529        let seq_encoding_data: Vec<f32> = encoding_data[..seq_len * d_model].to_vec();
530
531        // Create tensor with shape [seq_len, d_model]
532        let seq_encoding = Tensor::from_vec(seq_encoding_data, &[seq_len, d_model])?;
533
534        // Add positional encoding to input
535        // input: [batch_size, seq_len, d_model]
536        // seq_encoding: [seq_len, d_model]
537        // Broadcasting: seq_encoding will be added to each batch
538
539        let input_data = input.to_vec()?;
540        let encoding_slice = seq_encoding.to_vec()?;
541
542        let batch_size = input_shape[0];
543        let mut output_data = vec![0.0f32; input_data.len()];
544
545        for b in 0..batch_size {
546            for s in 0..seq_len {
547                for d in 0..d_model {
548                    let input_idx = b * seq_len * d_model + s * d_model + d;
549                    let encoding_idx = s * d_model + d;
550                    output_data[input_idx] = input_data[input_idx] + encoding_slice[encoding_idx];
551                }
552            }
553        }
554
555        let output = Tensor::from_vec(output_data, input_shape)?;
556
557        // Note: Dropout would be applied here in training mode
558        // For now, dropout is not implemented as it requires training mode tracking
559        Ok(output)
560    }
561
562    fn parameters(&self) -> std::collections::HashMap<String, Parameter> {
563        // Positional encoding is not learnable
564        std::collections::HashMap::new()
565    }
566
567    fn train(&mut self) {}
568    fn eval(&mut self) {}
569}
570
571#[cfg(test)]
572mod tests {
573    use super::*;
574
575    #[test]
576    fn test_multi_head_attention() {
577        let mha = MultiHeadAttention::new(512, 8, 0.1, true)
578            .expect("Multi Head Attention should succeed");
579        let input = randn(&[2, 10, 512]).expect("randn should succeed"); // batch_size=2, seq_len=10, d_model=512
580
581        let output = mha.forward(&input).expect("forward pass should succeed");
582        assert_eq!(output.shape().dims(), &[2, 10, 512]);
583    }
584
585    #[test]
586    fn test_advanced_layer_norm() {
587        let ln = AdvancedLayerNorm::new(vec![512], true, 1e-5)
588            .expect("Advanced Layer Norm should succeed");
589        let input = randn(&[2, 10, 512]).expect("randn should succeed");
590
591        let output = ln.forward(&input).expect("forward pass should succeed");
592        assert_eq!(output.shape().dims(), &[2, 10, 512]);
593    }
594}