use crate::infer::GraphExt;
use crate::op::SpdMatFn;
use crate::ops::spd_eig::{SpectralFn, spd_jacobi_sweeps, spectral_matfn_batched};
use crate::{Graph, NodeId};
impl Graph {
pub fn spd_matrix_fn_batch_graph(&mut self, x: NodeId, kind: SpdMatFn, eps: f64) -> NodeId {
let xs = self.node(x).shape.clone();
let b = xs.dim(0).unwrap_static();
let n = xs.dim(1).unwrap_static();
let sweeps = spd_jacobi_sweeps();
let f = match kind {
SpdMatFn::Logm => SpectralFn::Log,
SpdMatFn::Sqrtm => SpectralFn::Sqrt,
SpdMatFn::Invsqrtm => SpectralFn::InvSqrt,
SpdMatFn::Expm => return self.spd_matrix_fn_batch(x, SpdMatFn::Expm),
};
spectral_matfn_batched(self, x, b, n, sweeps, eps, f)
}
pub fn spd_logm_batch_graph(&mut self, x: NodeId, eps: f64) -> NodeId {
self.spd_matrix_fn_batch_graph(x, SpdMatFn::Logm, eps)
}
pub fn spd_sqrtm_batch_graph(&mut self, x: NodeId, eps: f64) -> NodeId {
self.spd_matrix_fn_batch_graph(x, SpdMatFn::Sqrtm, eps)
}
pub fn spd_invsqrtm_batch_graph(&mut self, x: NodeId, eps: f64) -> NodeId {
self.spd_matrix_fn_batch_graph(x, SpdMatFn::Invsqrtm, eps)
}
pub fn spd_log_map_graph(&mut self, base: NodeId, x: NodeId, eps: f64) -> NodeId {
let p_half = self.spd_sqrtm_batch_graph(base, eps); let p_ihalf = self.spd_invsqrtm_batch_graph(base, eps); let px = self.mm(p_ihalf, x);
let inner = self.mm(px, p_ihalf);
let log_inner = self.spd_logm_batch_graph(inner, eps);
let ph_log = self.mm(p_half, log_inner);
self.mm(ph_log, p_half)
}
pub fn spd_exp_map_graph(&mut self, base: NodeId, v: NodeId, eps: f64) -> NodeId {
let p_half = self.spd_sqrtm_batch_graph(base, eps);
let p_ihalf = self.spd_invsqrtm_batch_graph(base, eps);
let pv = self.mm(p_ihalf, v);
let inner = self.mm(pv, p_ihalf);
let exp_inner = self.spd_matrix_fn_batch_graph(inner, SpdMatFn::Expm, eps);
let ph_exp = self.mm(p_half, exp_inner);
self.mm(ph_exp, p_half)
}
}