use mdarray::{array, expr::Expression as _, Array, Dim, Layout, Slice};
use num_complex::{Complex64, ComplexFloat};
use num_traits::MulAdd;
use mdarray_linalg::{Contract, prelude::*, utils::identity};
trait Scalar: Clone + ComplexFloat + MulAdd<Output = Self> {}
impl<T> Scalar for T where T: Clone + ComplexFloat + MulAdd<Output = T> {}
trait Backend<T: Scalar>: Contract<T> {}
impl<T, B> Backend<T> for B
where
T: Scalar,
B: Contract<T>,
{
}
fn matrix_power<T, B, L, D>(
backend: &B,
a: &Slice<T, (D, D), L>,
mut exponent: u64,
) -> Array<T, (D, D)>
where
T: Scalar,
B: Backend<T>,
L: Layout,
D: Dim,
{
let (rows, cols) = *a.shape();
assert_eq!(rows.size(), cols.size(), "matrix must be square");
let mut result = identity::<T, D, D>(rows.size());
let mut base = a.to_array();
while exponent > 0 {
if exponent & 1 == 1 {
result = backend.matmul(&result, &base).eval();
}
exponent >>= 1;
if exponent > 0 {
base = backend.matmul(&base, &base).eval();
}
}
result
}
fn main() {
let q = array![[1.0, 1.0], [1.0, 0.0]];
let n = 21;
let naive = mdarray_linalg::Naive;
let qn_naive = matrix_power(&naive, &q, n);
let faer = mdarray_linalg_faer::Faer::default();
let qn_faer = matrix_power(&faer, &q, n);
let expected = array![[17711.0, 10946.0], [10946.0, 6765.0]];
println!("Q^{n} with Naive backend:\n{qn_naive:?}");
println!("Q^{n} with Faer backend:\n{qn_faer:?}");
println!("expected Fibonacci matrix:\n{expected:?}");
assert_eq!(qn_naive, expected);
assert_eq!(qn_faer, expected);
let complex_q = q.expr().copied().map(Complex64::from).eval();
let complex_qn_naive = matrix_power(&naive, &complex_q, n);
println!("complex Q^{n} with Naive backend:\n{complex_qn_naive:?}");
}