use oxmera_core::{DType, Device, Error, Result, Shape};
use crate::autograd::GradFn;
use crate::autograd::is_recording;
use crate::backend::{Backend, backend_for};
use crate::tensor::Tensor;
fn square_matrix_dims(t: &Tensor, op: &'static str) -> Result<(usize, usize)> {
let d = t.dims();
if d.len() < 2 || d[d.len() - 1] != d[d.len() - 2] {
return Err(Error::InvalidArgument {
op,
detail: format!("needs a [.., n, n] tensor, got shape {:?}", t.shape()),
});
}
let n = d[d.len() - 1];
let batch: usize = d[..d.len() - 2].iter().product();
Ok((batch, n))
}
fn dispatch_cpu_fallback<T>(
t: &Tensor,
f: impl Fn(&dyn Backend, &Tensor) -> Result<T>,
upload: impl Fn(&dyn Backend, T) -> Result<T>,
) -> Result<T> {
let backend = backend_for(t.device())?;
match f(backend.as_ref(), t) {
Err(Error::NotImplemented { .. }) if t.device() != Device::Cpu => {
let cpu = backend.download(t)?;
let out = f(backend_for(Device::Cpu)?.as_ref(), &cpu)?;
upload(backend.as_ref(), out)
}
other => other,
}
}
impl Tensor {
pub fn eye(n: usize) -> Tensor {
let cells = n
.checked_mul(n)
.expect("eye: n * n overflows usize — the identity is too large to build");
let mut v = vec![0.0f32; cells];
for i in 0..n {
v[i * n + i] = 1.0;
}
Tensor::from_vec_f32(v, Shape::from([n, n])).expect("lengths match by construction")
}
pub fn eye_on(n: usize, device: Device) -> Result<Tensor> {
Tensor::eye(n).to_device(device)
}
fn eye_like(n: usize, like: &Tensor) -> Result<Tensor> {
Tensor::eye(n)
.to_dtype(like.dtype())?
.to_device(like.device())
}
pub fn diag(&self) -> Result<Tensor> {
let (_, n) = square_matrix_dims(self, "diag")?;
let eye = Tensor::eye_like(n, self)?;
self.mul(&eye)?.sum(&[self.ndim() - 1])
}
pub fn diag_embed(&self) -> Result<Tensor> {
let d = self.dims();
let Some(&n) = d.last() else {
return Err(Error::InvalidArgument {
op: "diag_embed",
detail: "needs rank >= 1".into(),
});
};
let eye = Tensor::eye_like(n, self)?;
self.unsqueeze(self.ndim())?.mul(&eye)
}
pub fn trace(&self) -> Result<Tensor> {
let diag = self.diag().map_err(|e| match e {
Error::InvalidArgument { detail, .. } => Error::InvalidArgument {
op: "trace",
detail,
},
other => other,
})?;
diag.sum(&[diag.ndim() - 1])
}
pub fn cholesky(&self) -> Result<Tensor> {
square_matrix_dims(self, "cholesky")?;
let out = dispatch_cpu_fallback(self, |be, t| be.cholesky(t), |be, l| be.upload(&l))?;
if !(is_recording() && self.is_tracked()) {
return Ok(out);
}
let l = out.clone();
let shape = self.shape().clone();
let device = self.device();
Ok(out.with_grad_fn(GradFn {
inputs: vec![self.clone()],
vjp: Box::new(move |g: &Tensor| {
let (batch, n) = square_matrix_dims(&l, "cholesky backward")?;
let dtype = l.dtype();
let lv = l
.to_device(Device::Cpu)?
.to_dtype(DType::F32)?
.to_vec_f32()?;
let gv = g
.to_device(Device::Cpu)?
.to_dtype(DType::F32)?
.to_vec_f32()?;
let grad = crate::cpu_linalg::cholesky_backward(&lv, &gv, batch, n);
let grad = Tensor::from_vec_f32(grad, shape.clone())?
.to_dtype(dtype)?
.to_device(device)?;
Ok(vec![Some(grad)])
}),
}))
}
pub fn logdet(&self) -> Result<Tensor> {
let l = self.cholesky()?;
let d = l.diag()?;
d.ln()?.sum(&[d.ndim() - 1])?.mul_scalar(2.0)
}
pub fn det(&self) -> Result<Tensor> {
self.logdet()?.exp()
}
pub fn eigh(&self) -> Result<(Tensor, Tensor)> {
square_matrix_dims(self, "eigh")?;
check_symmetric(self, "eigh")?;
dispatch_cpu_fallback(
self,
|be, t| be.eigh(t),
|be, (w, v)| Ok((be.upload(&w)?, be.upload(&v)?)),
)
}
}
pub const EIGH_SYMMETRY_TOL: f32 = 1e-5;
fn check_symmetric(t: &Tensor, op: &'static str) -> Result<()> {
let (batch, n) = square_matrix_dims(t, op)?;
if n < 2 {
return Ok(());
}
let a = t
.to_device(Device::Cpu)?
.to_dtype(DType::F32)?
.to_vec_f32()?;
let rank = t.dims().len();
for b in 0..batch {
let m = &a[b * n * n..(b + 1) * n * n];
let scale = m.iter().fold(0.0f32, |acc, v| acc.max(v.abs()));
let mut worst = (0usize, 0usize, 0.0f32);
for i in 0..n {
for j in (i + 1)..n {
let d = (m[i * n + j] - m[j * n + i]).abs();
if d > worst.2 {
worst = (i, j, d);
}
}
}
let bound = EIGH_SYMMETRY_TOL * scale.max(f32::MIN_POSITIVE);
if worst.2 > bound {
let (i, j, d) = worst;
return Err(Error::InvalidArgument {
op,
detail: format!(
"matrix {b} is not symmetric (|a[{i}][{j}] - a[{j}][{i}]| = {d:e}, \
tolerance {bound:e}); {op} needs a symmetric input — if that is \
what you meant, symmetrize it explicitly with \
a.add(&a.transpose({d0}, {d1})?)?.mul_scalar(0.5)?",
d0 = rank - 2,
d1 = rank - 1,
),
});
}
}
Ok(())
}