use tenferro_cpu::CpuBackend;
use tenferro_linalg::{QrGauge, QrOptions, TracedTensorLinalgExt};
use tenferro_runtime::{GraphCompiler, Runtime, Tensor, TracedTensor};
use super::sweep::SweepError;
fn linalg_err(context: &str, e: impl std::fmt::Display) -> SweepError {
SweepError::Linalg(format!("{context}: {e}"))
}
#[derive(Clone, Debug, PartialEq)]
pub(crate) struct Dense {
pub rows: usize,
pub cols: usize,
pub data: Vec<f64>,
}
impl Dense {
pub fn zeros(rows: usize, cols: usize) -> Self {
Dense {
rows,
cols,
data: vec![0.0; rows * cols],
}
}
#[inline]
pub fn at(&self, i: usize, j: usize) -> f64 {
self.data[i + j * self.rows]
}
#[inline]
pub fn set(&mut self, i: usize, j: usize, v: f64) {
self.data[i + j * self.rows] = v;
}
pub fn col(&self, j: usize) -> &[f64] {
&self.data[j * self.rows..(j + 1) * self.rows]
}
pub fn unit(rows: usize, k: usize) -> Self {
let mut m = Dense::zeros(rows, 1);
m.data[k] = 1.0;
m
}
pub fn transpose(&self) -> Dense {
let mut t = Dense::zeros(self.cols, self.rows);
for j in 0..self.cols {
for i in 0..self.rows {
t.data[j + i * self.cols] = self.data[i + j * self.rows];
}
}
t
}
pub fn cat_cols(&mut self, other: &Dense) {
debug_assert_eq!(self.rows, other.rows, "cat_cols row mismatch");
self.data.extend_from_slice(&other.data);
self.cols += other.cols;
}
pub fn select_cols(&self, keep: &[usize]) -> Dense {
let mut out = Dense::zeros(self.rows, keep.len());
for (jo, &j) in keep.iter().enumerate() {
out.data[jo * self.rows..(jo + 1) * self.rows].copy_from_slice(self.col(j));
}
out
}
pub fn norm(&self) -> f64 {
self.data.iter().map(|x| x * x).sum::<f64>().sqrt()
}
}
fn traced(m: &Dense) -> Result<TracedTensor, SweepError> {
TracedTensor::from_vec_col_major(vec![m.rows, m.cols], m.data.clone())
.map_err(|e| linalg_err("build input", e))
}
fn run(outputs: &[&TracedTensor]) -> Result<Vec<Tensor>, SweepError> {
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>, SweepError> {
Ok(t.as_slice::<f64>()
.map_err(|e| linalg_err("read", e))?
.to_vec())
}
pub(crate) fn matmul(a: &Dense, b: &Dense) -> Result<Dense, SweepError> {
debug_assert_eq!(a.cols, b.rows, "matmul inner-dim mismatch");
if a.rows == 0 || b.cols == 0 || a.cols == 0 {
return Ok(Dense::zeros(a.rows, b.cols));
}
let ta = traced(a)?;
let tb = traced(b)?;
let tc = ta.matmul(&tb).map_err(|e| linalg_err("matmul", e))?;
let out = run(&[&tc])?;
Ok(Dense {
rows: a.rows,
cols: b.cols,
data: f64_out(&out[0])?,
})
}
pub(crate) fn tmatmul(a: &Dense, b: &Dense) -> Result<Dense, SweepError> {
matmul(&a.transpose(), b)
}
pub(crate) fn svd(a: &Dense) -> Result<(Dense, Vec<f64>, Dense), SweepError> {
let k = a.rows.min(a.cols);
if k == 0 {
return Ok((Dense::zeros(a.rows, 0), Vec::new(), Dense::zeros(0, a.cols)));
}
let ta = traced(a)?;
let (u, s, vt) = ta.svd().map_err(|e| linalg_err("svd", e))?;
let out = run(&[&u, &s, &vt])?;
let u = Dense {
rows: a.rows,
cols: k,
data: f64_out(&out[0])?,
};
let s = f64_out(&out[1])?;
let vt = Dense {
rows: k,
cols: a.cols,
data: f64_out(&out[2])?,
};
Ok((u, s, vt))
}
pub(crate) fn qr_positive_q(a: &Dense, tol: f64) -> Result<Dense, SweepError> {
let ta = traced(a)?;
let (q, rr) = ta
.qr_with_options(QrOptions::default().gauge(QrGauge::PositiveDiagonal))
.map_err(|e| linalg_err("qr", e))?;
let out = run(&[&q, &rr])?;
let qdata = f64_out(&out[0])?;
let rdata = f64_out(&out[1])?;
let k = a.rows.min(a.cols); let q = Dense {
rows: a.rows,
cols: k,
data: qdata,
};
let keep: Vec<usize> = (0..k)
.filter(|&i| {
(0..a.cols)
.map(|c| rdata[i + c * k].powi(2))
.sum::<f64>()
.sqrt()
> tol
})
.collect();
Ok(q.select_cols(&keep))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn svd_reconstructs_and_procrustes_recovers_rotation() {
let m = Dense {
rows: 2,
cols: 2,
data: vec![2.0, 0.0, 0.0, 3.0], };
let (u, s, vt) = svd(&m).unwrap();
let mut sd = Dense::zeros(2, 2);
for (i, &sv) in s.iter().enumerate() {
sd.set(i, i, sv);
}
let recon = matmul(&matmul(&u, &sd).unwrap(), &vt).unwrap();
for i in 0..4 {
assert!((recon.data[i] - m.data[i]).abs() < 1e-12);
}
let c = std::f64::consts::FRAC_1_SQRT_2;
let rot = Dense {
rows: 2,
cols: 2,
data: vec![c, c, -c, c], };
let (u, _s, vt) = svd(&rot).unwrap();
let w = matmul(&u, &vt).unwrap();
for i in 0..4 {
assert!((w.data[i] - rot.data[i]).abs() < 1e-12);
}
}
}