use pounce_common::types::{Index, Number};
pub trait TSymScalingMethod {
fn compute_sym_t_scaling_factors(
&mut self,
n: Index,
nnz: Index,
airn: &[Index],
ajcn: &[Index],
a: &[Number],
scaling_factors: &mut [Number],
) -> bool;
fn set_slack_scaling(&mut self, _nx: Index, _s_scale: &[Number]) {}
}
#[derive(Debug, Default, Clone, Copy)]
pub struct IdentityScalingMethod;
impl TSymScalingMethod for IdentityScalingMethod {
fn compute_sym_t_scaling_factors(
&mut self,
n: Index,
_nnz: Index,
_airn: &[Index],
_ajcn: &[Index],
_a: &[Number],
scaling_factors: &mut [Number],
) -> bool {
debug_assert_eq!(scaling_factors.len(), n as usize);
for s in scaling_factors.iter_mut() {
*s = 1.0;
}
true
}
}
#[derive(Debug, Default, Clone)]
pub struct SlackBasedTSymScalingMethod {
nx: Index,
s_scale: Vec<Number>,
}
impl SlackBasedTSymScalingMethod {
pub fn new() -> Self {
Self::default()
}
}
impl TSymScalingMethod for SlackBasedTSymScalingMethod {
fn set_slack_scaling(&mut self, nx: Index, s_scale: &[Number]) {
self.nx = nx;
self.s_scale.clear();
self.s_scale.extend_from_slice(s_scale);
}
fn compute_sym_t_scaling_factors(
&mut self,
n: Index,
_nnz: Index,
_airn: &[Index],
_ajcn: &[Index],
_a: &[Number],
scaling_factors: &mut [Number],
) -> bool {
debug_assert_eq!(scaling_factors.len(), n as usize);
for s in scaling_factors.iter_mut() {
*s = 1.0;
}
let nx = self.nx as usize;
let ns = self.s_scale.len();
if ns == 0 || nx + ns > n as usize {
return true;
}
scaling_factors[nx..nx + ns].copy_from_slice(&self.s_scale);
true
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn slack_based_is_identity_before_the_first_iterate() {
let mut m = SlackBasedTSymScalingMethod::new();
let mut f = vec![0.0; 6];
assert!(m.compute_sym_t_scaling_factors(6, 0, &[], &[], &[], &mut f));
assert_eq!(f, vec![1.0; 6]);
}
#[test]
fn slack_based_writes_only_the_s_block() {
let mut m = SlackBasedTSymScalingMethod::new();
m.set_slack_scaling(2, &[0.25, 0.5]);
let mut f = vec![0.0; 6];
assert!(m.compute_sym_t_scaling_factors(6, 0, &[], &[], &[], &mut f));
assert_eq!(f, vec![1.0, 1.0, 0.25, 0.5, 1.0, 1.0]);
}
#[test]
fn slack_based_declines_a_system_it_does_not_fit() {
let mut m = SlackBasedTSymScalingMethod::new();
m.set_slack_scaling(4, &[0.25, 0.5]);
let mut f = vec![0.0; 5];
assert!(m.compute_sym_t_scaling_factors(5, 0, &[], &[], &[], &mut f));
assert_eq!(f, vec![1.0; 5], "must fall back to identity, not misplace");
}
#[test]
fn identity_writes_unit_factors() {
let mut method = IdentityScalingMethod;
let irn = [1, 2, 2];
let jcn = [1, 1, 2];
let vals = [2.0, 1.0, 3.0];
let mut s = vec![0.0; 2];
assert!(method.compute_sym_t_scaling_factors(2, 3, &irn, &jcn, &vals, &mut s));
assert_eq!(s, &[1.0, 1.0]);
}
}