use core::f64::consts::FRAC_PI_2;
use crate::{
autodiff::Dual1,
error::{ProjError, ProjResult},
projection_generic::ProjectGeneric,
};
#[derive(Debug, Clone)]
pub struct NewtonParams {
pub max_iter: usize,
pub tol: f64,
}
impl Default for NewtonParams {
fn default() -> Self {
Self {
max_iter: 50,
tol: 1e-10,
}
}
}
pub fn newton_inverse<P>(
proj: &P,
x_target: f64,
y_target: f64,
init_lam: f64,
init_phi: f64,
params: &NewtonParams,
) -> ProjResult<(f64, f64)>
where
P: ProjectGeneric,
{
let mut lam = init_lam;
let mut phi = init_phi;
for _ in 0..params.max_iter {
let lam_d = Dual1::<2>::variable(lam, 0);
let phi_d = Dual1::<2>::variable(phi, 1);
let (x_d, y_d) = proj.project_fwd_generic(lam_d, phi_d)?;
let dx = x_d.v - x_target;
let dy = y_d.v - y_target;
if dx * dx + dy * dy < params.tol * params.tol {
return Ok((lam, phi));
}
let j00 = x_d.d[0]; let j01 = x_d.d[1]; let j10 = y_d.d[0]; let j11 = y_d.d[1];
let det = j00 * j11 - j01 * j10;
if det.abs() < 1e-14 {
return Err(ProjError::OutsideProjectionDomain);
}
let inv_det = 1.0 / det;
let delta_lam = (j11 * dx - j01 * dy) * inv_det;
let delta_phi = (-j10 * dx + j00 * dy) * inv_det;
lam -= delta_lam;
phi -= delta_phi;
phi = phi.clamp(-FRAC_PI_2 + 1e-10, FRAC_PI_2 - 1e-10);
}
Err(ProjError::NoConvergence)
}
#[cfg(test)]
mod tests {
use super::*;
struct IdentityProj;
impl crate::projection_generic::ProjectGeneric for IdentityProj {
fn project_fwd_generic<S: crate::scalar::Scalar>(
&self,
lam: S,
phi: S,
) -> ProjResult<(S, S)> {
Ok((lam, phi))
}
fn project_inv_generic<S: crate::scalar::Scalar>(&self, x: S, y: S) -> ProjResult<(S, S)> {
Ok((x, y))
}
}
#[test]
fn newton_identity_convergence() {
let proj = IdentityProj;
let params = NewtonParams::default();
let (lam, phi) =
newton_inverse(&proj, 0.5, 0.3, 0.0, 0.0, ¶ms).expect("identity inverse converges");
assert!((lam - 0.5).abs() < 1e-9, "lam={lam}");
assert!((phi - 0.3).abs() < 1e-9, "phi={phi}");
}
#[test]
fn newton_identity_at_origin() {
let proj = IdentityProj;
let params = NewtonParams::default();
let (lam, phi) =
newton_inverse(&proj, 0.0, 0.0, 0.1, 0.1, ¶ms).expect("origin inverse converges");
assert!(lam.abs() < 1e-9, "lam={lam}");
assert!(phi.abs() < 1e-9, "phi={phi}");
}
}