use crate::maxcut_oracle::{grad, obj};
use crate::sdp_project::project;
use ndarray::{Array1, Array2};
use ndarray_linalg::Norm;
use sprs::CsMat;
#[derive(Clone, Copy)]
pub enum StepRule {
Grad(f64),
GradAdv(f64),
Coord(f64),
CoordNoStep,
}
pub fn generate_step_rule(step_rule: &str, alpha: f64) -> StepRule {
match step_rule {
"grad" => StepRule::Grad(alpha),
"grad_adv" => StepRule::GradAdv(alpha),
"coord" => StepRule::Coord(alpha),
"coord_no_step" => StepRule::CoordNoStep,
_ => StepRule::Grad(alpha),
}
}
pub fn apply_step(Q: &CsMat<f64>, V: Array2<f64>, step_rule: StepRule) -> Array2<f64> {
match step_rule {
StepRule::Grad(alpha) => make_step(Q, V, alpha),
StepRule::GradAdv(alpha) => make_step_adv(Q, V, alpha),
StepRule::Coord(alpha) => make_step_coord(Q, V, alpha),
StepRule::CoordNoStep => make_step_coord_no_step(Q, V),
}
}
pub fn make_step(Q: &CsMat<f64>, V: Array2<f64>, alpha_safe: f64) -> Array2<f64> {
let grad = grad(Q, &V);
project(V - alpha_safe * grad)
}
pub fn make_step_adv(Q: &CsMat<f64>, V: Array2<f64>, alpha_safe: f64) -> Array2<f64> {
let grad = grad(Q, &V);
let f_0 = obj(Q, &V);
let x = obj(Q, &(&V + alpha_safe * &grad)) - f_0;
let y = obj(Q, &(&V - alpha_safe * &grad)) - f_0;
let mut alpha = (0.5 * (y - x) * alpha_safe) / (x + y);
let proposed_step_val = obj(Q, &(&V - alpha * &grad));
if proposed_step_val > f_0 {
alpha = alpha_safe;
}
project(V - alpha * grad)
}
pub fn make_step_coord(Q: &CsMat<f64>, mut V: Array2<f64>, alpha_safe: f64) -> Array2<f64> {
for i in 0..Q.shape().0 {
let Q_i = Q.outer_view(i).unwrap();
let mut g_i = Array1::<f64>::zeros(V.shape()[1]);
for (k, &v) in Q_i.iter() {
if k != i {
g_i = g_i + v * &V.row(k);
}
}
g_i = &V.row(i) - alpha_safe * g_i;
g_i /= g_i.norm_l2();
V.row_mut(i).assign(&g_i);
}
V
}
pub fn make_step_coord_no_step(Q: &CsMat<f64>, mut V: Array2<f64>) -> Array2<f64> {
let mut g_i = Array1::<f64>::zeros(V.shape()[1]);
let mut temp = Array1::<f64>::zeros(V.shape()[1]);
for i in 0..Q.shape().0 {
let Q_i = Q.outer_view(i).unwrap();
for (k, &v) in Q_i.iter() {
if k != i {
temp.assign(&V.row(k));
temp *= v;
g_i -= &temp;
}
}
if g_i.norm_l2() >= 1E-24 {
g_i /= g_i.norm_l2();
V.row_mut(i).assign(&g_i);
}
g_i.fill(0.0f64);
}
V
}
#[cfg(test)]
mod tests {
#[test]
fn is_true() {
assert_eq!(1, 1);
}
}