Skip to main content

oxiz_math/simd/
matrix_simd.rs

1//! SIMD-Optimized Matrix Operations.
2#![allow(clippy::needless_range_loop, clippy::type_complexity)] // Matrix algorithms use explicit indexing
3//!
4//! Provides cache-friendly matrix operations with SIMD-style chunking.
5
6#[allow(unused_imports)]
7use crate::prelude::*;
8use core::ops::{Add, Mul, Sub};
9use num_traits::{Float, Zero};
10
11/// SIMD-friendly matrix-vector multiplication.
12pub fn simd_matrix_vec_mul<T>(matrix: &[Vec<T>], vec: &[T]) -> Vec<T>
13where
14    T: Clone + Add<Output = T> + Mul<Output = T> + Zero,
15{
16    let rows = matrix.len();
17    if rows == 0 {
18        return Vec::new();
19    }
20
21    let cols = matrix[0].len();
22    if cols != vec.len() {
23        panic!("matrix-vector dimension mismatch");
24    }
25
26    let mut result = vec![T::zero(); rows];
27
28    // Process in chunks for cache locality
29    const CHUNK_SIZE: usize = 8;
30
31    for (i, row) in matrix.iter().enumerate() {
32        let mut sum = T::zero();
33
34        for chunk_idx in (0..cols).step_by(CHUNK_SIZE) {
35            let chunk_end = (chunk_idx + CHUNK_SIZE).min(cols);
36
37            for j in chunk_idx..chunk_end {
38                sum = sum.clone() + row[j].clone() * vec[j].clone();
39            }
40        }
41
42        result[i] = sum;
43    }
44
45    result
46}
47
48/// SIMD-friendly matrix-matrix multiplication.
49pub fn simd_matrix_mul<T>(a: &[Vec<T>], b: &[Vec<T>]) -> Vec<Vec<T>>
50where
51    T: Clone + Add<Output = T> + Mul<Output = T> + Zero,
52{
53    let rows_a = a.len();
54    if rows_a == 0 {
55        return Vec::new();
56    }
57
58    let cols_a = a[0].len();
59    let rows_b = b.len();
60    if rows_b == 0 || cols_a != rows_b {
61        panic!("matrix dimension mismatch");
62    }
63
64    let cols_b = b[0].len();
65
66    // Transpose B for cache-friendly access
67    let b_t = transpose(b);
68
69    let mut result = vec![vec![T::zero(); cols_b]; rows_a];
70
71    const TILE_SIZE: usize = 32;
72
73    // Tiled matrix multiplication
74    for i_tile in (0..rows_a).step_by(TILE_SIZE) {
75        let i_end = (i_tile + TILE_SIZE).min(rows_a);
76
77        for j_tile in (0..cols_b).step_by(TILE_SIZE) {
78            let j_end = (j_tile + TILE_SIZE).min(cols_b);
79
80            for k_tile in (0..cols_a).step_by(TILE_SIZE) {
81                let k_end = (k_tile + TILE_SIZE).min(cols_a);
82
83                // Compute tile
84                for i in i_tile..i_end {
85                    for j in j_tile..j_end {
86                        let mut sum = result[i][j].clone();
87
88                        for k in k_tile..k_end {
89                            sum = sum.clone() + a[i][k].clone() * b_t[j][k].clone();
90                        }
91
92                        result[i][j] = sum;
93                    }
94                }
95            }
96        }
97    }
98
99    result
100}
101
102/// Transpose a matrix.
103pub fn transpose<T: Clone>(matrix: &[Vec<T>]) -> Vec<Vec<T>> {
104    if matrix.is_empty() {
105        return Vec::new();
106    }
107
108    let rows = matrix.len();
109    let cols = matrix[0].len();
110
111    let mut result = vec![vec![matrix[0][0].clone(); rows]; cols];
112
113    for i in 0..rows {
114        for j in 0..cols {
115            result[j][i] = matrix[i][j].clone();
116        }
117    }
118
119    result
120}
121
122/// SIMD-friendly LU decomposition with partial pivoting.
123pub fn simd_lu_decomposition<T>(matrix: &[Vec<T>]) -> Option<(Vec<Vec<T>>, Vec<Vec<T>>, Vec<usize>)>
124where
125    T: Clone + Add<Output = T> + Sub<Output = T> + Mul<Output = T> + Float,
126{
127    let n = matrix.len();
128    if n == 0 || matrix[0].len() != n {
129        return None;
130    }
131
132    let mut a = matrix.to_vec();
133    let mut l = vec![vec![T::zero(); n]; n];
134    let mut u = vec![vec![T::zero(); n]; n];
135    let mut perm: Vec<usize> = (0..n).collect();
136
137    for i in 0..n {
138        l[i][i] = T::one();
139    }
140
141    for k in 0..n {
142        // Partial pivoting
143        let mut max_idx = k;
144        let mut max_val = a[k][k].abs();
145
146        for i in (k + 1)..n {
147            let val = a[i][k].abs();
148            if val > max_val {
149                max_val = val;
150                max_idx = i;
151            }
152        }
153
154        if max_val < T::epsilon() {
155            return None; // Singular matrix
156        }
157
158        if max_idx != k {
159            a.swap(k, max_idx);
160            perm.swap(k, max_idx);
161            if k > 0 {
162                for j in 0..k {
163                    let temp = l[k][j];
164                    l[k][j] = l[max_idx][j];
165                    l[max_idx][j] = temp;
166                }
167            }
168        }
169
170        // Compute L and U
171        for j in k..n {
172            u[k][j] = a[k][j];
173
174            for s in 0..k {
175                u[k][j] = u[k][j] - l[k][s] * u[s][j];
176            }
177        }
178
179        for i in (k + 1)..n {
180            l[i][k] = a[i][k];
181
182            for s in 0..k {
183                l[i][k] = l[i][k] - l[i][s] * u[s][k];
184            }
185
186            l[i][k] = l[i][k] / u[k][k];
187        }
188    }
189
190    Some((l, u, perm))
191}
192
193/// Solve linear system using LU decomposition.
194pub fn simd_lu_solve<T>(l: &[Vec<T>], u: &[Vec<T>], perm: &[usize], b: &[T]) -> Option<Vec<T>>
195where
196    T: Clone + Add<Output = T> + Sub<Output = T> + Mul<Output = T> + Float,
197{
198    let n = l.len();
199    if n == 0 || u.len() != n || perm.len() != n || b.len() != n {
200        return None;
201    }
202
203    // Apply permutation to b
204    let mut b_perm = vec![T::zero(); n];
205    for i in 0..n {
206        b_perm[i] = b[perm[i]];
207    }
208
209    // Forward substitution (L * y = b_perm)
210    let mut y = vec![T::zero(); n];
211    for i in 0..n {
212        let mut sum = b_perm[i];
213        for j in 0..i {
214            sum = sum - l[i][j] * y[j];
215        }
216        y[i] = sum;
217    }
218
219    // Backward substitution (U * x = y)
220    let mut x = vec![T::zero(); n];
221    for i in (0..n).rev() {
222        let mut sum = y[i];
223        for j in (i + 1)..n {
224            sum = sum - u[i][j] * x[j];
225        }
226        x[i] = sum / u[i][i];
227    }
228
229    Some(x)
230}
231
232/// Compute matrix determinant using LU decomposition.
233pub fn simd_determinant<T>(matrix: &[Vec<T>]) -> Option<T>
234where
235    T: Clone + Add<Output = T> + Sub<Output = T> + Mul<Output = T> + Float,
236{
237    let (_l, u, perm) = simd_lu_decomposition(matrix)?;
238
239    let n = u.len();
240
241    // Det = product of diagonal of U times sign of permutation
242    let mut det = T::one();
243    for i in 0..n {
244        det = det * u[i][i];
245    }
246
247    // Count inversions in permutation: pairs (i, j) where i < j but perm[i] > perm[j]
248    let mut inversions = 0;
249    for i in 0..n {
250        for j in (i + 1)..n {
251            if perm[i] > perm[j] {
252                inversions += 1;
253            }
254        }
255    }
256
257    if inversions % 2 == 1 {
258        det = T::zero() - det;
259    }
260
261    Some(det)
262}
263
264/// Matrix inversion using LU decomposition.
265pub fn simd_matrix_inverse<T>(matrix: &[Vec<T>]) -> Option<Vec<Vec<T>>>
266where
267    T: Clone + Add<Output = T> + Sub<Output = T> + Mul<Output = T> + Float,
268{
269    let n = matrix.len();
270    if n == 0 || matrix[0].len() != n {
271        return None;
272    }
273
274    let (l, u, perm) = simd_lu_decomposition(matrix)?;
275
276    let mut inverse = vec![vec![T::zero(); n]; n];
277
278    // Solve for each column of the identity matrix
279    for j in 0..n {
280        let mut e = vec![T::zero(); n];
281        e[j] = T::one();
282
283        let col = simd_lu_solve(&l, &u, &perm, &e)?;
284        for i in 0..n {
285            inverse[i][j] = col[i];
286        }
287    }
288
289    Some(inverse)
290}
291
292/// Compute QR decomposition using Gram-Schmidt.
293pub fn simd_qr_decomposition<T>(matrix: &[Vec<T>]) -> Option<(Vec<Vec<T>>, Vec<Vec<T>>)>
294where
295    T: Clone + Add<Output = T> + Sub<Output = T> + Mul<Output = T> + Float,
296{
297    let rows = matrix.len();
298    if rows == 0 {
299        return None;
300    }
301    let cols = matrix[0].len();
302
303    let mut q = vec![vec![T::zero(); cols]; rows];
304    let mut r = vec![vec![T::zero(); cols]; cols];
305
306    for j in 0..cols {
307        // Get column j
308        let mut v: Vec<T> = matrix.iter().map(|row| row[j]).collect();
309
310        // Orthogonalize against previous columns
311        for i in 0..j {
312            let q_col: Vec<T> = q.iter().map(|row| row[i]).collect();
313
314            let dot = dot_product(&q_col, &v);
315            r[i][j] = dot;
316
317            for k in 0..rows {
318                v[k] = v[k] - q_col[k] * dot;
319            }
320        }
321
322        // Normalize
323        let norm = vector_norm(&v);
324        if norm < T::epsilon() {
325            return None; // Linearly dependent columns
326        }
327
328        r[j][j] = norm;
329
330        for k in 0..rows {
331            q[k][j] = v[k] / norm;
332        }
333    }
334
335    Some((q, r))
336}
337
338fn dot_product<T>(a: &[T], b: &[T]) -> T
339where
340    T: Clone + Add<Output = T> + Mul<Output = T> + Zero,
341{
342    a.iter()
343        .zip(b.iter())
344        .map(|(x, y)| x.clone() * y.clone())
345        .fold(T::zero(), |acc, x| acc + x)
346}
347
348fn vector_norm<T>(v: &[T]) -> T
349where
350    T: Clone + Add<Output = T> + Mul<Output = T> + Float,
351{
352    let sum_sq = v.iter().map(|x| *x * *x).fold(T::zero(), |acc, x| acc + x);
353    sum_sq.sqrt()
354}
355
356#[cfg(test)]
357mod tests {
358    use super::*;
359
360    #[test]
361    fn test_transpose() {
362        let matrix = vec![vec![1, 2, 3], vec![4, 5, 6]];
363
364        let transposed = transpose(&matrix);
365
366        assert_eq!(transposed.len(), 3);
367        assert_eq!(transposed[0], vec![1, 4]);
368        assert_eq!(transposed[1], vec![2, 5]);
369        assert_eq!(transposed[2], vec![3, 6]);
370    }
371
372    #[test]
373    fn test_matrix_vec_mul() {
374        let matrix = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
375        let vec = vec![2.0, 3.0];
376
377        let result = simd_matrix_vec_mul(&matrix, &vec);
378
379        assert_eq!(result.len(), 2);
380        assert!((result[0] - 8.0).abs() < 1e-10);
381        assert!((result[1] - 18.0).abs() < 1e-10);
382    }
383
384    #[test]
385    fn test_matrix_mul() {
386        let a = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
387        let b = vec![vec![5.0, 6.0], vec![7.0, 8.0]];
388
389        let result = simd_matrix_mul(&a, &b);
390
391        assert_eq!(result.len(), 2);
392        assert!((result[0][0] - 19.0).abs() < 1e-10);
393        assert!((result[0][1] - 22.0).abs() < 1e-10);
394        assert!((result[1][0] - 43.0).abs() < 1e-10);
395        assert!((result[1][1] - 50.0).abs() < 1e-10);
396    }
397
398    #[test]
399    fn test_lu_decomposition() {
400        let matrix = vec![vec![4.0, 3.0], vec![6.0, 3.0]];
401
402        let (l, u, _perm) = simd_lu_decomposition(&matrix).expect("LU decomposition failed");
403
404        assert_eq!(l.len(), 2);
405        assert_eq!(u.len(), 2);
406    }
407
408    #[test]
409    fn test_determinant() {
410        let matrix = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
411
412        let det = simd_determinant(&matrix).expect("determinant computation failed");
413
414        assert!((det - (-2.0)).abs() < 1e-10);
415    }
416}