use tenferro_cpu::CpuBackend;
use tenferro_linalg::{QrGauge, QrOptions, TracedTensorLinalgExt};
use tenferro_runtime::{GraphCompiler, Runtime, Tensor, TracedTensor};
use super::SunError;
fn linalg_err(context: &str, e: impl std::fmt::Display) -> SunError {
SunError::Linalg(format!("{context}: {e}"))
}
fn traced(rows: usize, cols: usize, data: Vec<f64>) -> Result<TracedTensor, SunError> {
TracedTensor::from_vec_col_major(vec![rows, cols], data)
.map_err(|e| linalg_err("build input", e))
}
fn run(outputs: &[&TracedTensor]) -> Result<Vec<Tensor>, SunError> {
let mut compiler = GraphCompiler::new();
let program = compiler
.compile_many(outputs)
.map_err(|e| linalg_err("compile", e))?;
let backend = CpuBackend::new();
let mut builder = Runtime::builder();
builder
.register_engine(
tenferro_cpu::runtime_engine_registration(&backend)
.map_err(|e| linalg_err("register", e))?,
)
.map_err(|e| linalg_err("register", e))?;
let engine_id = tenferro_cpu::runtime_engine_id().map_err(|e| linalg_err("register", e))?;
builder
.install_extension_module(
tenferro_linalg::extension_module::<CpuBackend>(engine_id)
.map_err(|e| linalg_err("register", e))?,
)
.map_err(|e| linalg_err("register", e))?;
let runtime = builder.build().map_err(|e| linalg_err("register", e))?;
runtime
.run_compiled(&program, &[])
.map_err(|e| linalg_err("run", e))
}
fn f64_out(t: &Tensor) -> Result<Vec<f64>, SunError> {
Ok(t.as_slice::<f64>()
.map_err(|e| linalg_err("read", e))?
.to_vec())
}
pub(crate) struct Mat {
pub rows: usize,
pub cols: usize,
pub data: Vec<f64>,
}
impl Mat {
pub fn zeros(rows: usize, cols: usize) -> Self {
Mat {
rows,
cols,
data: vec![0.0; rows * cols],
}
}
#[inline]
pub fn add(&mut self, i: usize, j: usize, v: f64) {
self.data[i + j * self.rows] += v;
}
#[inline]
pub fn at(&self, i: usize, j: usize) -> f64 {
self.data[i + j * self.rows]
}
fn traced(&self) -> Result<TracedTensor, SunError> {
traced(self.rows, self.cols, self.data.clone())
}
}
pub(crate) fn nullspace(a: &Mat, atol: f64) -> Result<Mat, SunError> {
let n = a.cols;
if a.rows == 0 || n == 0 {
let mut id = Mat::zeros(n, n);
for i in 0..n {
id.data[i + i * n] = 1.0;
}
return Ok(id);
}
let ta = a.traced()?;
let (_u, s, vh) = ta.svd_full().map_err(|e| linalg_err("svd_full", e))?;
let out = run(&[&s, &vh])?;
let sv = f64_out(&out[0])?;
let vhd = f64_out(&out[1])?; let rank = sv.iter().filter(|&&x| x > atol).count();
let k = n - rank;
let mut ns = Mat::zeros(n, k);
for alpha in 0..k {
let t = rank + alpha;
for j in 0..n {
ns.data[j + alpha * n] = vhd[t + j * n];
}
}
Ok(ns)
}
pub(crate) fn qr_positive_q(a: &Mat) -> Result<Mat, SunError> {
let ta = a.traced()?;
let (q, _r) = ta
.qr_with_options(QrOptions::default().gauge(QrGauge::PositiveDiagonal))
.map_err(|e| linalg_err("qr", e))?;
let out = run(&[&q])?;
let qd = f64_out(&out[0])?;
let cols = a.rows.min(a.cols);
Ok(Mat {
rows: a.rows,
cols,
data: qd,
})
}
pub(crate) fn lstsq(a: &Mat, b: &Mat) -> Result<Mat, SunError> {
let ta = a.traced()?;
let tb = b.traced()?;
let x = ta.lstsq(&tb).map_err(|e| linalg_err("lstsq", e))?;
let out = run(&[&x])?;
let xd = f64_out(&out[0])?;
Ok(Mat {
rows: a.cols,
cols: b.cols,
data: xd,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn nullspace_of_1x2_recovers_kernel() {
let a = Mat {
rows: 1,
cols: 2,
data: vec![1.0, 1.0],
};
let ns = nullspace(&a, 1e-13).unwrap();
assert_eq!((ns.rows, ns.cols), (2, 1));
let (v0, v1) = (ns.at(0, 0), ns.at(1, 0));
assert!((v0 + v1).abs() < 1e-10, "A v = {}", v0 + v1);
assert!((v0 * v0 + v1 * v1 - 1.0).abs() < 1e-10, "not unit");
}
#[test]
fn qr_positive_q_is_orthonormal_with_positive_r_diagonal() {
let a = Mat {
rows: 2,
cols: 1,
data: vec![-3.0, 4.0],
};
let q = qr_positive_q(&a).unwrap();
assert_eq!((q.rows, q.cols), (2, 1));
let r00 = q.at(0, 0) * a.at(0, 0) + q.at(1, 0) * a.at(1, 0);
assert!(r00 >= 0.0, "R diagonal not positive: {r00}");
assert!((q.at(0, 0).powi(2) + q.at(1, 0).powi(2) - 1.0).abs() < 1e-12);
}
#[test]
fn lstsq_recovers_consistent_tall_system() {
let a = Mat {
rows: 3,
cols: 2,
data: vec![1.0, 1.0, 1.0, 0.0, 1.0, 2.0], };
let x_true = [2.0, -1.0];
let mut bdata = vec![0.0; 3];
for (i, bi) in bdata.iter_mut().enumerate() {
for (j, xj) in x_true.iter().enumerate() {
*bi += a.at(i, j) * xj;
}
}
let b = Mat {
rows: 3,
cols: 1,
data: bdata,
};
let x = lstsq(&a, &b).unwrap();
assert_eq!((x.rows, x.cols), (2, 1));
assert!((x.at(0, 0) - 2.0).abs() < 1e-9);
assert!((x.at(1, 0) + 1.0).abs() < 1e-9);
}
}