use oxiblas_core::scalar::{Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BalanceError {
EmptyMatrix,
NotSquare,
InvalidJob,
InvalidSide,
}
impl core::fmt::Display for BalanceError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::EmptyMatrix => write!(f, "Matrix is empty"),
Self::NotSquare => write!(f, "Matrix must be square"),
Self::InvalidJob => write!(f, "Invalid job specification"),
Self::InvalidSide => write!(f, "Invalid side specification"),
}
}
}
impl std::error::Error for BalanceError {}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BalanceJob {
None,
Permute,
Scale,
Both,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BalanceSide {
Right,
Left,
}
#[derive(Debug, Clone)]
pub struct Balance<T: Scalar> {
balanced: Mat<T>,
perm: Vec<usize>,
scale: Vec<T>,
ilo: usize,
ihi: usize,
n: usize,
job: BalanceJob,
}
impl<T: Field + Real + bytemuck::Zeroable> Balance<T> {
pub fn compute(a: MatRef<'_, T>, job: BalanceJob) -> Result<Self, BalanceError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Err(BalanceError::EmptyMatrix);
}
if m != n {
return Err(BalanceError::NotSquare);
}
let mut perm: Vec<usize> = (0..n).collect();
let mut scale = vec![T::one(); n];
let mut balanced = Mat::zeros(n, n);
for i in 0..n {
for j in 0..n {
balanced[(i, j)] = a[(i, j)];
}
}
let mut ilo = 0;
let mut ihi = n - 1;
if job == BalanceJob::Permute || job == BalanceJob::Both {
let mut converged = false;
while !converged && ilo < n && ihi < n && ilo <= ihi {
converged = true;
if ihi >= ilo {
let mut i = ihi;
loop {
let mut only_diag = true;
for j in ilo..=ihi {
if i != j && Scalar::abs(balanced[(i, j)]) > <T as Scalar>::epsilon() {
only_diag = false;
break;
}
}
if only_diag {
if i != ihi {
for k in 0..n {
let tmp = balanced[(i, k)];
balanced[(i, k)] = balanced[(ihi, k)];
balanced[(ihi, k)] = tmp;
}
for k in 0..n {
let tmp = balanced[(k, i)];
balanced[(k, i)] = balanced[(k, ihi)];
balanced[(k, ihi)] = tmp;
}
perm.swap(i, ihi);
}
if ihi == 0 {
break;
}
ihi -= 1;
converged = false;
break;
}
if i == ilo {
break;
}
i -= 1;
}
}
}
converged = false;
while !converged && ilo < n && ihi < n && ilo <= ihi {
converged = true;
for j in ilo..=ihi {
let mut only_diag = true;
for i in ilo..=ihi {
if i != j && Scalar::abs(balanced[(i, j)]) > <T as Scalar>::epsilon() {
only_diag = false;
break;
}
}
if only_diag {
if j != ilo {
for k in 0..n {
let tmp = balanced[(j, k)];
balanced[(j, k)] = balanced[(ilo, k)];
balanced[(ilo, k)] = tmp;
}
for k in 0..n {
let tmp = balanced[(k, j)];
balanced[(k, j)] = balanced[(k, ilo)];
balanced[(k, ilo)] = tmp;
}
perm.swap(j, ilo);
}
ilo += 1;
converged = false;
break;
}
}
}
}
if (job == BalanceJob::Scale || job == BalanceJob::Both) && ilo <= ihi {
let radix = T::from_f64(2.0).unwrap_or_else(T::zero);
let radix_sq = radix * radix;
let sfmin1 = T::from_f64(f64::MIN_POSITIVE).unwrap_or(<T as Scalar>::epsilon());
let sfmax1 = T::one() / sfmin1;
let max_iterations = 100;
for _iter in 0..max_iterations {
let mut no_conv = false;
for i in ilo..=ihi {
let mut row_norm = T::zero();
let mut col_norm = T::zero();
for j in ilo..=ihi {
if i != j {
row_norm = row_norm + Scalar::abs(balanced[(i, j)]);
col_norm = col_norm + Scalar::abs(balanced[(j, i)]);
}
}
if row_norm == T::zero() || col_norm == T::zero() {
continue;
}
let mut g = row_norm / radix;
let mut f = T::one();
let s = col_norm + row_norm;
while col_norm < g {
f = f * radix;
col_norm = col_norm * radix_sq;
}
g = row_norm * radix;
while col_norm >= g {
f = f / radix;
col_norm = col_norm / radix_sq;
}
let factor = T::from_f64(0.95).unwrap_or_else(T::zero);
if (col_norm + row_norm) / f < factor * s {
if f >= sfmin1 && f <= sfmax1 {
let g_inv = T::one() / f;
scale[i] = scale[i] * f;
no_conv = true;
for j in 0..n {
balanced[(i, j)] = balanced[(i, j)] * g_inv;
}
for j in 0..n {
balanced[(j, i)] = balanced[(j, i)] * f;
}
}
}
}
if !no_conv {
break;
}
}
}
Ok(Self {
balanced,
perm,
scale,
ilo,
ihi,
n,
job,
})
}
pub fn balanced(&self) -> MatRef<'_, T> {
self.balanced.as_ref()
}
pub fn balanced_mut(&mut self) -> &mut Mat<T> {
&mut self.balanced
}
pub fn permutation(&self) -> &[usize] {
&self.perm
}
pub fn scale(&self) -> &[T] {
&self.scale
}
pub fn ilo(&self) -> usize {
self.ilo
}
pub fn ihi(&self) -> usize {
self.ihi
}
pub fn n(&self) -> usize {
self.n
}
pub fn job(&self) -> BalanceJob {
self.job
}
pub fn back_transform(&self, v: &[Vec<T>], side: BalanceSide) -> Vec<Vec<T>> {
let num_vectors = v.len();
if num_vectors == 0 || self.n == 0 {
return v.to_vec();
}
let mut result: Vec<Vec<T>> = v.to_vec();
if self.job == BalanceJob::Scale || self.job == BalanceJob::Both {
match side {
BalanceSide::Right => {
for k in 0..num_vectors {
for i in self.ilo..=self.ihi.min(self.n - 1) {
result[k][i] = result[k][i] * self.scale[i];
}
}
}
BalanceSide::Left => {
for k in 0..num_vectors {
for i in self.ilo..=self.ihi.min(self.n - 1) {
result[k][i] = result[k][i] / self.scale[i];
}
}
}
}
}
if self.job == BalanceJob::Permute || self.job == BalanceJob::Both {
let mut inv_perm = vec![0usize; self.n];
for i in 0..self.n {
inv_perm[self.perm[i]] = i;
}
for k in 0..num_vectors {
let orig = result[k].clone();
for i in 0..self.n {
result[k][i] = orig[inv_perm[i]];
}
}
}
result
}
pub fn back_transform_matrix(&self, v: MatRef<'_, T>, side: BalanceSide) -> Mat<T> {
let n = v.nrows();
let num_vectors = v.ncols();
if n == 0 || num_vectors == 0 {
return Mat::zeros(n, num_vectors);
}
let mut result = Mat::zeros(n, num_vectors);
for i in 0..n {
for j in 0..num_vectors {
result[(i, j)] = v[(i, j)];
}
}
if self.job == BalanceJob::Scale || self.job == BalanceJob::Both {
match side {
BalanceSide::Right => {
for i in self.ilo..=self.ihi.min(n - 1) {
for j in 0..num_vectors {
result[(i, j)] = result[(i, j)] * self.scale[i];
}
}
}
BalanceSide::Left => {
for i in self.ilo..=self.ihi.min(n - 1) {
for j in 0..num_vectors {
result[(i, j)] = result[(i, j)] / self.scale[i];
}
}
}
}
}
if self.job == BalanceJob::Permute || self.job == BalanceJob::Both {
let mut inv_perm = vec![0usize; n];
for i in 0..n {
inv_perm[self.perm[i]] = i;
}
let copy = result.clone();
for i in 0..n {
for j in 0..num_vectors {
result[(i, j)] = copy[(inv_perm[i], j)];
}
}
}
result
}
pub fn reconstruct(&self) -> Mat<T> {
let mut a = Mat::zeros(self.n, self.n);
for i in 0..self.n {
for j in 0..self.n {
a[(i, j)] = self.balanced[(i, j)] * self.scale[i] / self.scale[j];
}
}
let mut inv_perm = vec![0usize; self.n];
for i in 0..self.n {
inv_perm[self.perm[i]] = i;
}
let copy = a.clone();
for i in 0..self.n {
for j in 0..self.n {
a[(inv_perm[i], inv_perm[j])] = copy[(i, j)];
}
}
a
}
}
pub fn gebal<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
) -> Result<(Mat<T>, usize, usize, Vec<T>), BalanceError> {
let bal = Balance::compute(a, BalanceJob::Both)?;
let balanced = bal.balanced.clone();
Ok((balanced, bal.ilo, bal.ihi, bal.scale))
}
pub fn gebak<T: Field + Real + bytemuck::Zeroable>(
job: BalanceJob,
side: BalanceSide,
ilo: usize,
ihi: usize,
scale: &[T],
v: MatRef<'_, T>,
) -> Mat<T> {
let n = v.nrows();
let num_vectors = v.ncols();
if n == 0 || num_vectors == 0 {
return Mat::zeros(n, num_vectors);
}
let mut result = Mat::zeros(n, num_vectors);
for i in 0..n {
for j in 0..num_vectors {
result[(i, j)] = v[(i, j)];
}
}
if job == BalanceJob::Scale || job == BalanceJob::Both {
match side {
BalanceSide::Right => {
for i in ilo..=ihi.min(n - 1) {
for j in 0..num_vectors {
result[(i, j)] = result[(i, j)] * scale[i];
}
}
}
BalanceSide::Left => {
for i in ilo..=ihi.min(n - 1) {
for j in 0..num_vectors {
result[(i, j)] = result[(i, j)] / scale[i];
}
}
}
}
}
result
}
#[cfg(test)]
mod tests {
use super::*;
fn approx_eq(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
#[test]
fn test_balance_identity() {
let eye = Mat::from_rows(&[&[1.0f64, 0.0, 0.0], &[0.0, 1.0, 0.0], &[0.0, 0.0, 1.0]]);
let bal = Balance::compute(eye.as_ref(), BalanceJob::Both).unwrap();
let b = bal.balanced();
for i in 0..3 {
for j in 0..3 {
let expected = if i == j { 1.0 } else { 0.0 };
assert!(approx_eq(b[(i, j)], expected, 1e-10));
}
}
for &s in bal.scale() {
assert!(approx_eq(s, 1.0, 1e-10));
}
}
#[test]
fn test_balance_diagonal() {
let diag = Mat::from_rows(&[&[2.0f64, 0.0, 0.0], &[0.0, 3.0, 0.0], &[0.0, 0.0, 5.0]]);
let bal = Balance::compute(diag.as_ref(), BalanceJob::Both).unwrap();
let b = bal.balanced();
for i in 0..3 {
for j in 0..3 {
if i != j {
assert!(approx_eq(b[(i, j)], 0.0, 1e-10));
}
}
}
}
#[test]
fn test_balance_unbalanced_matrix() {
let a = Mat::from_rows(&[&[1.0f64, 0.0, 0.0], &[100.0, 2.0, 200.0], &[0.0, 3.0, 4.0]]);
let bal = Balance::compute(a.as_ref(), BalanceJob::Scale).unwrap();
let b = bal.balanced();
let mut row_norms_orig = [0.0; 3];
let mut col_norms_orig = [0.0; 3];
let mut row_norms_bal = [0.0; 3];
let mut col_norms_bal = [0.0; 3];
for i in 0..3 {
for j in 0..3 {
row_norms_orig[i] += a[(i, j)].abs();
col_norms_orig[j] += a[(i, j)].abs();
row_norms_bal[i] += b[(i, j)].abs();
col_norms_bal[j] += b[(i, j)].abs();
}
}
let mean_orig: f64 =
(row_norms_orig.iter().sum::<f64>() + col_norms_orig.iter().sum::<f64>()) / 6.0;
let mean_bal: f64 =
(row_norms_bal.iter().sum::<f64>() + col_norms_bal.iter().sum::<f64>()) / 6.0;
let var_orig: f64 = row_norms_orig
.iter()
.chain(col_norms_orig.iter())
.map(|x| (x - mean_orig).powi(2))
.sum::<f64>()
/ 6.0;
let var_bal: f64 = row_norms_bal
.iter()
.chain(col_norms_bal.iter())
.map(|x| (x - mean_bal).powi(2))
.sum::<f64>()
/ 6.0;
assert!(
var_bal <= var_orig * 1.5,
"var_bal={}, var_orig={}",
var_bal,
var_orig
);
}
#[test]
fn test_balance_reconstruction() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 9.0]]);
let bal = Balance::compute(a.as_ref(), BalanceJob::Both).unwrap();
let reconstructed = bal.reconstruct();
for i in 0..3 {
for j in 0..3 {
assert!(
approx_eq(reconstructed[(i, j)], a[(i, j)], 1e-10),
"reconstructed[{},{}] = {}, a = {}",
i,
j,
reconstructed[(i, j)],
a[(i, j)]
);
}
}
}
#[test]
fn test_balance_eigenvalue_isolation() {
let a = Mat::from_rows(&[&[1.0f64, 0.0, 0.0], &[2.0, 3.0, 4.0], &[5.0, 6.0, 7.0]]);
let bal = Balance::compute(a.as_ref(), BalanceJob::Permute).unwrap();
assert!(
bal.ilo() >= 1 || bal.ihi() < 2,
"ilo={}, ihi={}",
bal.ilo(),
bal.ihi()
);
}
#[test]
fn test_balance_back_transform_matrix() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[100.0, 3.0]]);
let bal = Balance::compute(a.as_ref(), BalanceJob::Scale).unwrap();
let v = Mat::from_rows(&[&[1.0f64, 0.0], &[0.0, 1.0]]);
let v_transformed = bal.back_transform_matrix(v.as_ref(), BalanceSide::Right);
assert_eq!(v_transformed.nrows(), 2);
assert_eq!(v_transformed.ncols(), 2);
}
#[test]
fn test_gebal_gebak() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 9.0]]);
let (balanced, ilo, ihi, scale) = gebal(a.as_ref()).unwrap();
let v = Mat::from_rows(&[&[1.0f64, 0.0, 0.0], &[0.0, 1.0, 0.0], &[0.0, 0.0, 1.0]]);
let _v_back = gebak(
BalanceJob::Both,
BalanceSide::Right,
ilo,
ihi,
&scale,
v.as_ref(),
);
assert_eq!(balanced.nrows(), 3);
assert_eq!(balanced.ncols(), 3);
}
#[test]
fn test_balance_job_none() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let bal = Balance::compute(a.as_ref(), BalanceJob::None).unwrap();
let b = bal.balanced();
for i in 0..2 {
for j in 0..2 {
assert!(approx_eq(b[(i, j)], a[(i, j)], 1e-10));
}
}
}
#[test]
fn test_balance_4x4_reconstruction() {
let a = Mat::from_rows(&[
&[4.0f64, 1.0, -2.0, 2.0],
&[1.0, 2.0, 0.0, 1.0],
&[-2.0, 0.0, 3.0, -2.0],
&[2.0, 1.0, -2.0, -1.0],
]);
let bal = Balance::compute(a.as_ref(), BalanceJob::Both).unwrap();
let reconstructed = bal.reconstruct();
for i in 0..4 {
for j in 0..4 {
assert!(
approx_eq(reconstructed[(i, j)], a[(i, j)], 1e-10),
"reconstructed[{},{}] = {}, a = {}",
i,
j,
reconstructed[(i, j)],
a[(i, j)]
);
}
}
}
#[test]
fn test_balance_f32() {
let a = Mat::from_rows(&[&[1.0f32, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 9.0]]);
let bal = Balance::compute(a.as_ref(), BalanceJob::Both).unwrap();
let reconstructed = bal.reconstruct();
for i in 0..3 {
for j in 0..3 {
assert!(
(reconstructed[(i, j)] - a[(i, j)]).abs() < 1e-5,
"reconstructed[{},{}] = {}, a = {}",
i,
j,
reconstructed[(i, j)],
a[(i, j)]
);
}
}
}
}