oxiz_math/simd/
matrix_simd.rs1#![allow(clippy::needless_range_loop, clippy::type_complexity)] #[allow(unused_imports)]
7use crate::prelude::*;
8use core::ops::{Add, Mul, Sub};
9use num_traits::{Float, Zero};
10
11pub 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 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
48pub 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 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 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 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
102pub 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
122pub 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 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; }
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 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
193pub 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 let mut b_perm = vec![T::zero(); n];
205 for i in 0..n {
206 b_perm[i] = b[perm[i]];
207 }
208
209 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 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
232pub 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 let mut det = T::one();
243 for i in 0..n {
244 det = det * u[i][i];
245 }
246
247 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
264pub 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 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
292pub 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 let mut v: Vec<T> = matrix.iter().map(|row| row[j]).collect();
309
310 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 let norm = vector_norm(&v);
324 if norm < T::epsilon() {
325 return None; }
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}