1use crate::{Module, Parameter};
4use torsh_core::error::TorshError;
5use torsh_tensor::{
6 creation::{ones, randn, zeros},
7 Tensor,
8};
9
10pub struct MultiHeadAttention {
12 pub num_heads: usize,
13 pub d_model: usize,
14 pub d_k: usize,
15 pub d_v: usize,
16
17 pub w_q: Parameter,
19 pub w_k: Parameter,
20 pub w_v: Parameter,
21 pub w_o: Parameter,
22
23 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 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 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 pub fn attention(
108 &self,
109 q: &Tensor,
110 k: &Tensor,
111 v: &Tensor,
112 mask: Option<&Tensor>,
113 ) -> Result<Tensor, TorshError> {
114 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 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 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 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 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 scores_data[scores_offset + i * seq_len_k + j] =
165 dot_product / (d_k as f32).sqrt();
166 }
167 }
168 }
169
170 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; }
177 }
178 }
179
180 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 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 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 for j in 0..seq_len_k {
201 let idx = row_offset + j;
202 scores_data[idx] /= exp_sum + 1e-9; }
204 }
205 }
206
207 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 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 let input_2d = input.view(&[(batch_size * seq_len) as i32, d_model as i32])?;
250
251 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 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 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 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 let attention_output = self.attention(&q, &k, &v, None)?;
309
310 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 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 }
357
358 fn eval(&mut self) {
359 }
361}
362
363pub 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 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 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 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 let batch_size: usize = input_shape[..input_shape.len() - self.normalized_shape.len()]
412 .iter()
413 .product();
414
415 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 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 let mean: f32 = instance.iter().sum::<f32>() / num_features as f32;
433
434 let variance: f32 =
436 instance.iter().map(|&x| (x - mean).powi(2)).sum::<f32>() / num_features as f32;
437
438 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
469pub struct PositionalEncoding {
471 pub encoding: Tensor,
472 pub dropout: f64,
473}
474
475impl PositionalEncoding {
476 pub fn new(d_model: usize, max_len: usize, dropout: f64) -> Result<Self, TorshError> {
482 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 encoding_data[pos * d_model + i] = angle.sin();
491
492 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 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 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 let encoding_data = self.encoding.to_vec()?;
529 let seq_encoding_data: Vec<f32> = encoding_data[..seq_len * d_model].to_vec();
530
531 let seq_encoding = Tensor::from_vec(seq_encoding_data, &[seq_len, d_model])?;
533
534 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 Ok(output)
560 }
561
562 fn parameters(&self) -> std::collections::HashMap<String, Parameter> {
563 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"); 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}