Skip to main content

gam_terms/analytic_penalties/
isometry.rs

1use super::*;
2pub use gam_problem::WeightField;
3
4// ---------------------------------------------------------------------------
5// Isometry penalty
6// ---------------------------------------------------------------------------
7
8/// Choice of reference Riemannian metric `g^ref(t)` on the latent manifold.
9///
10/// `Euclidean` is the natural default: the reference metric is `I_d`, so the
11/// penalty pulls the decoder toward locally-isometric (length-preserving)
12/// behavior. `UserSupplied` lets the caller hand in a `(n_obs, d, d)` jet of
13/// per-row reference metrics (useful for warm-starting from a chart of a
14/// pre-fit GP-LVM).
15#[derive(Clone)]
16pub enum IsometryReference {
17    Euclidean,
18    UserSupplied(Arc<Array2<f64>>), // (n_obs, d*d) row-major flattened
19}
20
21impl std::fmt::Debug for IsometryReference {
22    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
23        match self {
24            IsometryReference::Euclidean => f.write_str("Euclidean"),
25            IsometryReference::UserSupplied(a) => f
26                .debug_tuple("UserSupplied")
27                .field(&format_args!("{}×{}", a.nrows(), a.ncols()))
28                .finish(),
29        }
30    }
31}
32
33/// Radial Duchon decoder metadata used to materialize
34/// `∂J_n[i, a] / ∂t_{n, c}` from `φ'(r)` and `φ''(r)` on demand.
35///
36/// `radial_coefficients[k, i]` is the decoder coefficient that maps radial
37/// basis column `k` into output channel `i`. Polynomial-tail columns are not
38/// represented here; callers whose decoder contains a non-linear polynomial
39/// tail should provide `jacobian_second_cache` directly.
40#[derive(Debug, Clone)]
41pub struct IsometryDuchonRadialSource {
42    pub centers: Arc<Array2<f64>>,
43    pub radial_coefficients: Arc<Array2<f64>>,
44    pub length_scale: Option<f64>,
45    pub nullspace_order: DuchonNullspaceOrder,
46    /// Forward hybrid spectral order `s = spec.power`. The Cartesian
47    /// derivative engine must resolve the same `(p, s, κ)` the forward
48    /// `build_duchon_basis` used, so it differentiates the exact resolved
49    /// hybrid Green's function `φ_{p,s,κ}` rather than a hard-coded `s = 0`
50    /// surrogate (issue #440).
51    pub power: usize,
52}
53
54/// Isometry-to-reference penalty (canonical-coordinate gauge term).
55///
56/// Lives on ext-coords: the target slice is a row of the `LatentCoordValues` flat
57/// vector (row-major `n_obs × d`). Owns one ρ-axis (`log μ_iso`).
58///
59/// Penalizes `½ μ Σ_n ‖g_n(t) − g^ref(t_n)‖²_F`, where the pullback metric
60/// at row `n` is
61///
62/// ```text
63///   g_n = J_n^T W_n J_n,    J_n ∈ ℝ^{p × d}
64/// ```
65///
66/// and `W_n` is a per-row low-rank PSD behavioral metric stored as
67/// `W_n = U_n U_n^T` with `U_n ∈ ℝ^{p × r}`. The canonical-coordinate
68/// statement is "one unit of motion in `t` ↦ one unit of behavioral change",
69/// so the `W_n` weighting is load-bearing.
70///
71/// In the SAE objective this is the extension-coordinate gauge fix: it prevents
72/// the latent chart from absorbing arbitrary smooth reparameterizations of the
73/// decoder manifold. ARD, sparsity, or rank penalties can then select axes or
74/// structure in a chart whose metric scale is pinned.
75///
76/// **Contraction order invariant.** Every place this struct touches `W_n`,
77/// the contraction is `(J^T U_n)(U_n^T J)` — never `J^T W_n J` with `W_n`
78/// materialized as `p × p`. Concretely we form `M_n = U_n^T J_n ∈ ℝ^{r × d}`
79/// once and then `g_n = M_n^T M_n` (`d × d`). Cost per row:
80/// `O(p · r · d + r · d²)`, independent of `p²`.
81///
82/// **When to use.** Whenever a `LatentCoord` block is in play without an
83/// auxiliary variable (`AuxPrior`) to break the diffeomorphism gauge. Fixes
84/// the audit finding that ARD is not a standalone gauge fix. With a Euclidean
85/// reference, the penalty pulls the decoder toward a local isometry, which is
86/// enough to make the inner Hessian on `t` full-rank and the IFT well-defined.
87///
88/// **Math.** Let `J_n ∈ ℝ^{p × d}` be the local decoder Jacobian. Then
89/// `g_n = J_n^T W_n J_n` and the penalty is
90/// `½ μ Σ_n ‖J_n^T W_n J_n − g^ref_n‖²_F`. Analytic gradient w.r.t. `t_n`:
91///
92/// ```text
93///   ∂P/∂t_{n,c}
94///     = μ Σ_{a,b} (g_n − g^ref_n)_{ab}
95///         [ H_{n,:,a,c}^T W_n J_{n,:,b}
96///           + J_{n,:,a}^T W_n H_{n,:,b,c} ],
97///   H_{n,i,a,c} = ∂J_{n,i,a}/∂t_{n,c}.
98/// ```
99///
100/// Gotchas:
101///
102/// * The value path returns the configured missing-cache default when the
103///   first-jet cache is absent; gradient/HVP paths need the first and second
104///   decoder jets and return zeros when the analytic jet source is unavailable.
105/// * The exact Hessian includes a residual-curvature term requiring the third
106///   decoder jet. REML/PIRLS curvature should prefer the Gauss-Newton PSD
107///   majorizer when a positive curvature block is required.
108/// * `W_n` is a metric weight, not a scalar confidence. Changing it changes the
109///   canonical units of latent motion.
110///
111/// The per-row Jacobian `J_n` is exactly the radial-derivative jet
112/// `design_gradient_wrt_t` already computes for `LatentCoordValues`; the
113/// second derivative `∂J/∂t` is built by the shared
114/// `crate::basis::radial_basis_cartesian_derivative` engine from the
115/// radial Hessian identity. A finite-difference oracle for the docstring is
116/// to central-difference `value(t ± h e_j)` against `grad_target(t)[j]`;
117/// the analytic value follows the oracle until finite-difference
118/// cancellation dominates. No autograd needed.
119///
120/// `μ = exp(ρ_iso)` is REML-selectable as one extra ρ axis.
121///
122/// `jacobian_cache_slot` and `jacobian_second_cache_slot` are interior-mutable
123/// (`RwLock<Option<Arc<…>>>`) so the SAE outer loop can refresh them in place
124/// each step without needing `&mut self` on the registry-held penalty (see
125/// `refresh_caches` and `crate::terms::sae::manifold::refresh_isometry_caches_from_atom`).
126/// Readers go through the [`Self::jacobian_cache`] / [`Self::jacobian_second_cache`]
127/// accessors, which take the read lock briefly and clone the inner `Arc`
128/// (refcount bump — no payload copy). Writers go through [`Self::refresh_caches`].
129#[derive(Debug)]
130pub struct IsometryPenalty {
131    pub target: PsiSlice,
132    pub reference: IsometryReference,
133    /// Index of this penalty's strength `log μ_iso` inside the *local* rho
134    /// view this penalty receives. Always `0` for now (single owned axis).
135    pub rho_index: usize,
136    /// Cached Jacobian `J ∈ ℝ^{n_obs × p × d}`, flattened row-major
137    /// `(n_obs, p*d)`. The owning driver refreshes this each IFT outer step
138    /// before invoking `value` / `grad_target`; in operator-only call sites
139    /// (Hessian-vector products) the cache must be live. Access through
140    /// [`Self::jacobian_cache`] / [`Self::set_jacobian_cache`].
141    pub jacobian_cache_slot: RwLock<Option<Arc<Array2<f64>>>>,
142    /// Optional cached per-row Jacobian *second derivative*
143    /// `H_n ∈ ℝ^{p × d × d}`, flattened row-major as `(n_obs, p*d*d)`.
144    /// `H_n[i, a, c] = ∂J_n[i, a] / ∂t_{n, c}`. Either this cache or
145    /// `duchon_radial_source` must be present for exact isometry
146    /// gradient/HVP calls. Access through [`Self::jacobian_second_cache`] /
147    /// [`Self::set_jacobian_second_cache`].
148    pub jacobian_second_cache_slot: RwLock<Option<Arc<Array2<f64>>>>,
149    /// Optional radial-Duchon source used to build `jacobian_second_cache`
150    /// analytically from `φ'(r)` and the public `φ''(r)` jet helper. This is
151    /// the exact chain-rule path for callers that do not pre-cache `∂J/∂t`.
152    pub duchon_radial_source: Option<Arc<IsometryDuchonRadialSource>>,
153    /// Optional cached per-row Jacobian *third derivative*
154    /// `K_n ∈ ℝ^{p × d × d × d}`, stored as an `Array3` with shape
155    /// `(n_obs, p, d * d * d)` where the third axis packs `(a, c, d)` in
156    /// row-major order `((a * d) + c) * d + dd`. `hvp` uses the full
157    /// residual-curvature Hessian (proposal §4(b)):
158    ///   B_{ab,cd} = K_{a,cd}^T W J_b + H_{a,c}^T W H_{b,d}
159    ///             + H_{a,d}^T W H_{b,c} + J_a^T W K_{b,cd}.
160    /// Either this cache or `duchon_radial_source` must be present for
161    /// analytic `hvp` calls. Interior-mutable (mirrors
162    /// `jacobian_second_cache_slot`) so the SAE outer loop can refresh `K` in
163    /// place each step. Access through [`Self::third_decoder_derivative`] /
164    /// [`Self::set_third_decoder_derivative`].
165    pub third_decoder_derivative_slot: RwLock<Option<Arc<ndarray::Array3<f64>>>>,
166    /// Output dimensionality `p` (column count of each per-row Jacobian).
167    pub p_out: usize,
168    /// Per-row behavioral metric in low-rank factored form. Defaults to
169    /// `Identity` (the unweighted `J^T J` pullback). When `Factored`, all
170    /// `g_n` contractions are done via `M_n = U_n^T J_n` (`r × d`), keeping
171    /// memory and FLOPs scaling at `O(p · r · d)` per row instead of
172    /// `O(p²)` per row.
173    pub weight: WeightField,
174    pub scalar_weight: f64,
175    pub weight_schedule: Option<ScalarWeightSchedule>,
176}
177
178pub(crate) struct IsometryHvpState<'a> {
179    d: usize,
180    n_obs: usize,
181    p: usize,
182    jac2: CowArray<'a, f64, Ix2>,
183    jac3: CowArray<'a, f64, Ix3>,
184    metric: IsometryMetricState,
185    wj_rows: Vec<Array2<f64>>,
186}
187
188#[derive(Debug, Clone)]
189struct IsometryMetricState {
190    g: Array2<f64>,
191    residual: Array2<f64>,
192    metric_grad: Array2<f64>,
193    normalizer: f64,
194    trace_denominator: f64,
195    residual_dot_g: f64,
196}
197
198impl IsometryMetricState {
199    fn residual_direction(&self, delta_g: ArrayView2<'_, f64>, d: usize) -> (Array2<f64>, f64) {
200        let n_obs = self.g.nrows();
201        let dd = d * d;
202        let mut delta_trace_sum = 0.0;
203        for n in 0..n_obs {
204            for a in 0..d {
205                delta_trace_sum += delta_g[[n, a * d + a]];
206            }
207        }
208        let delta_normalizer = delta_trace_sum / self.trace_denominator;
209        let inv_norm = 1.0 / self.normalizer;
210        let inv_norm_sq = inv_norm * inv_norm;
211        let mut delta_residual = Array2::<f64>::zeros((n_obs, dd));
212        for n in 0..n_obs {
213            for k in 0..dd {
214                delta_residual[[n, k]] =
215                    delta_g[[n, k]] * inv_norm - self.g[[n, k]] * delta_normalizer * inv_norm_sq;
216            }
217        }
218        (delta_residual, delta_normalizer)
219    }
220
221    fn metric_grad_direction(&self, delta_g: ArrayView2<'_, f64>, d: usize) -> Array2<f64> {
222        let n_obs = self.g.nrows();
223        let dd = d * d;
224        let (delta_residual, delta_normalizer) = self.residual_direction(delta_g, d);
225        let mut delta_residual_dot_g = 0.0;
226        for n in 0..n_obs {
227            for k in 0..dd {
228                delta_residual_dot_g += delta_residual[[n, k]] * self.g[[n, k]];
229                delta_residual_dot_g += self.residual[[n, k]] * delta_g[[n, k]];
230            }
231        }
232        let inv_norm = 1.0 / self.normalizer;
233        let inv_norm_sq = inv_norm * inv_norm;
234        let delta_trace_coeff = delta_residual_dot_g * inv_norm_sq / self.trace_denominator
235            - 2.0 * self.residual_dot_g * delta_normalizer * inv_norm_sq * inv_norm
236                / self.trace_denominator;
237        let mut out = Array2::<f64>::zeros((n_obs, dd));
238        for n in 0..n_obs {
239            for a in 0..d {
240                for b in 0..d {
241                    let k = a * d + b;
242                    let mut value = delta_residual[[n, k]] * inv_norm
243                        - self.residual[[n, k]] * delta_normalizer * inv_norm_sq;
244                    if a == b {
245                        value -= delta_trace_coeff;
246                    }
247                    out[[n, k]] = value;
248                }
249            }
250        }
251        out
252    }
253}
254
255/// Average trace per latent dimension `(1 / (N d)) Σ_n tr(m_n)` of a flattened
256/// `(n_obs, d*d)` row-major metric field. Shared by the decoder normalizer
257/// `gbar` and the reference normalizer `gref_bar`, so a decoder metric that is
258/// exactly proportional to a reference of arbitrary scale gives a zero residual.
259fn average_trace_per_dim(m: ArrayView2<'_, f64>, n_obs: usize, d: usize) -> f64 {
260    let denom = (n_obs * d) as f64;
261    let mut trace_sum = 0.0;
262    for n in 0..n_obs {
263        for a in 0..d {
264            trace_sum += m[[n, a * d + a]];
265        }
266    }
267    trace_sum / denom
268}
269
270fn isometry_dg_entry(
271    jac2: ArrayView2<'_, f64>,
272    wj: ArrayView2<'_, f64>,
273    n: usize,
274    d: usize,
275    p: usize,
276    a: usize,
277    b: usize,
278    c: usize,
279) -> f64 {
280    let mut s = 0.0;
281    for i in 0..p {
282        s += jac2[[n, (i * d + a) * d + c]] * wj[[i, b]];
283        s += wj[[i, a]] * jac2[[n, (i * d + b) * d + c]];
284    }
285    s
286}
287
288fn isometry_row_delta_g(
289    jac2: ArrayView2<'_, f64>,
290    wj: ArrayView2<'_, f64>,
291    v: ArrayView1<'_, f64>,
292    n: usize,
293    d: usize,
294    p: usize,
295) -> Array2<f64> {
296    let mut delta_g = Array2::<f64>::zeros((d, d));
297    for a in 0..d {
298        for b in 0..d {
299            let mut s = 0.0;
300            for c in 0..d {
301                s += isometry_dg_entry(jac2, wj, n, d, p, a, b, c) * v[n * d + c];
302            }
303            delta_g[[a, b]] = s;
304        }
305    }
306    delta_g
307}
308
309impl IsometryPenalty {
310    pub const DEFAULT_VALUE_ON_MISSING_CACHE: f64 = 0.0;
311
312    #[must_use]
313    pub fn new_euclidean(target: PsiSlice, p_out: usize) -> Self {
314        Self {
315            target,
316            reference: IsometryReference::Euclidean,
317            rho_index: 0,
318            jacobian_cache_slot: RwLock::new(None),
319            jacobian_second_cache_slot: RwLock::new(None),
320            duchon_radial_source: None,
321            third_decoder_derivative_slot: RwLock::new(None),
322            p_out,
323            weight: WeightField::Identity,
324            scalar_weight: 1.0,
325            weight_schedule: None,
326        }
327    }
328
329    /// Read-side accessor: takes the read lock briefly and clones the inner
330    /// `Arc` (refcount bump only; no payload copy). Returns `None` when the
331    /// cache has not been refreshed yet. Internally panics on poisoned lock
332    /// — the lock only wraps an `Option<Arc<…>>`, so the write side cannot
333    /// leave it in an invariant-violating state.
334    #[must_use]
335    pub fn jacobian_cache(&self) -> Option<Arc<Array2<f64>>> {
336        self.jacobian_cache_slot
337            .read()
338            .expect("IsometryPenalty::jacobian_cache_slot poisoned")
339            .clone()
340    }
341
342    /// Read the Jacobian cache under the *dimensional identity*
343    /// `(latent_dim, p_out)` of the atom currently being evaluated.
344    ///
345    /// The cache is interior-mutable and refreshed **per atom** (see
346    /// [`Self::refresh_caches`] and
347    /// `sae::manifold::refresh_isometry_caches_from_atom`). In a heterogeneous
348    /// SAE whose atoms have *mixed* latent dimensions (e.g. the standard zoo
349    /// `dims = [1,1,2,2,2,2,2,1]`) the slot can transiently hold a Jacobian
350    /// built for a **different** atom's `latent_dim` — its column count is
351    /// `p_out · d'` for that atom's `d'`, not this atom's `d`. Reshaping such a
352    /// cache at the wrong `d` would silently corrupt every downstream `J`
353    /// contraction (`projected_jacobian_row` / `weighted_jacobian_row`) or trip
354    /// the [`Self::pullback_metric`] shape invariant (`assert_eq!(jac.ncols(),
355    /// p·d)`) with a hard panic (issue #2294).
356    ///
357    /// The cache's built-for `latent_dim` is recoverable from its own shape:
358    /// with `p_out` fixed on the penalty, `ncols() == p_out · d_cache`. We
359    /// Per-atom SAE evaluation clones the registry descriptor, retargets it,
360    /// and refreshes that clone before any read, so a mismatched live cache is
361    /// an ownership/refresh invariant violation, not missing optional data.
362    /// Silently converting it to `None` would disable the isometry penalty and
363    /// change the fitted objective. Keep the mismatch hard: every caller sees
364    /// either no cache at all or a cache with the exact requested identity.
365    fn dimensioned_jacobian_cache(
366        &self,
367        method: &str,
368        latent_dim: usize,
369    ) -> Option<Arc<Array2<f64>>> {
370        let Some(jac) = self.jacobian_cache() else {
371            self.missing_cache_default(method, "jacobian_cache is None");
372            return None;
373        };
374        let expected = self
375            .p_out
376            .checked_mul(latent_dim)
377            .expect("IsometryPenalty Jacobian dimensional identity overflow");
378        assert_eq!(
379            jac.ncols(),
380            expected,
381            "IsometryPenalty::{method} stale cross-atom Jacobian cache: cache has {} columns, \
382             but this per-atom evaluation requires p_out {} × latent_dim {} = {}; the owner \
383             must clone, retarget, and refresh the penalty before evaluation",
384            jac.ncols(),
385            self.p_out,
386            latent_dim,
387            expected,
388        );
389        Some(jac)
390    }
391
392    /// Read-side accessor for the per-row Jacobian second derivative.
393    /// Mirrors [`Self::jacobian_cache`].
394    #[must_use]
395    pub fn jacobian_second_cache(&self) -> Option<Arc<Array2<f64>>> {
396        self.jacobian_second_cache_slot
397            .read()
398            .expect("IsometryPenalty::jacobian_second_cache_slot poisoned")
399            .clone()
400    }
401
402    /// Per-step refresh entry point. Takes `&self` (no `&mut`) so the SAE
403    /// outer loop can install fresh caches on an `Arc<IsometryPenalty>` held
404    /// in the analytic-penalty registry without disturbing the surrounding
405    /// dispatcher. Pass `None` for either argument to clear that cache (the
406    /// dispatcher will then either fall back to the Duchon radial source if
407    /// available, or return the zero safe default).
408    pub fn refresh_caches(&self, jac: Option<Arc<Array2<f64>>>, jac2: Option<Arc<Array2<f64>>>) {
409        *self
410            .jacobian_cache_slot
411            .write()
412            .expect("IsometryPenalty::jacobian_cache_slot poisoned") = jac;
413        *self
414            .jacobian_second_cache_slot
415            .write()
416            .expect("IsometryPenalty::jacobian_second_cache_slot poisoned") = jac2;
417    }
418
419    /// In-place writer for just the Jacobian cache (used by callers that
420    /// already own the radial Duchon source and only want to refresh `J`).
421    pub fn set_jacobian_cache(&self, jac: Option<Arc<Array2<f64>>>) {
422        *self
423            .jacobian_cache_slot
424            .write()
425            .expect("IsometryPenalty::jacobian_cache_slot poisoned") = jac;
426    }
427
428    /// In-place writer for just the Jacobian second-derivative cache.
429    pub fn set_jacobian_second_cache(&self, jac2: Option<Arc<Array2<f64>>>) {
430        *self
431            .jacobian_second_cache_slot
432            .write()
433            .expect("IsometryPenalty::jacobian_second_cache_slot poisoned") = jac2;
434    }
435
436    /// Read-side accessor for the per-row Jacobian third derivative `K`.
437    /// Mirrors [`Self::jacobian_second_cache`].
438    #[must_use]
439    pub fn third_decoder_derivative(&self) -> Option<Arc<ndarray::Array3<f64>>> {
440        self.third_decoder_derivative_slot
441            .read()
442            .expect("IsometryPenalty::third_decoder_derivative_slot poisoned")
443            .clone()
444    }
445
446    /// In-place writer for just the Jacobian third-derivative cache `K`.
447    pub fn set_third_decoder_derivative(&self, jac3: Option<Arc<ndarray::Array3<f64>>>) {
448        *self
449            .third_decoder_derivative_slot
450            .write()
451            .expect("IsometryPenalty::third_decoder_derivative_slot poisoned") = jac3;
452    }
453}
454
455impl Clone for IsometryPenalty {
456    fn clone(&self) -> Self {
457        Self {
458            target: self.target.clone(),
459            reference: self.reference.clone(),
460            rho_index: self.rho_index,
461            jacobian_cache_slot: RwLock::new(self.jacobian_cache()),
462            jacobian_second_cache_slot: RwLock::new(self.jacobian_second_cache()),
463            duchon_radial_source: self.duchon_radial_source.clone(),
464            third_decoder_derivative_slot: RwLock::new(self.third_decoder_derivative()),
465            p_out: self.p_out,
466            weight: self.weight.clone(),
467            scalar_weight: self.scalar_weight,
468            weight_schedule: self.weight_schedule.clone(),
469        }
470    }
471}
472
473impl IsometryPenalty {
474    /// Attach a cached third decoder derivative
475    /// `K_n[i, a, c, d] = ∂²J_n[i, a] / ∂t_{n, c} ∂t_{n, d}`, flattened
476    /// row-major as `(n_obs, p * d * d * d)`. The Hessian-vector product
477    /// uses the full residual-curvature term in addition to the metric
478    /// Gauss-Newton piece.
479    #[must_use]
480    pub fn with_third_decoder_derivative(self, k: Arc<ndarray::Array3<f64>>) -> Self {
481        self.set_third_decoder_derivative(Some(k));
482        self
483    }
484
485    #[must_use]
486    pub fn with_reference(mut self, reference: IsometryReference) -> Self {
487        self.reference = reference;
488        self
489    }
490
491    #[must_use]
492    pub fn with_jacobian_cache(self, j: Arc<Array2<f64>>) -> Self {
493        self.set_jacobian_cache(Some(j));
494        self
495    }
496
497    #[must_use]
498    pub fn with_jacobian_second_cache(self, h: Arc<Array2<f64>>) -> Self {
499        self.set_jacobian_second_cache(Some(h));
500        self
501    }
502
503
504    impl_with_weight_schedule!(scalar_weight);
505
506    fn missing_cache_default(&self, method: &str, detail: &str) {
507        log::warn!(
508            "IsometryPenalty::{method} missing required derivative state: {detail}; \
509             returning the zero safe default"
510        );
511    }
512
513    fn has_jacobian_cache(&self, method: &str) -> bool {
514        if self.jacobian_cache().is_some() {
515            true
516        } else {
517            self.missing_cache_default(method, "jacobian_cache is None");
518            false
519        }
520    }
521
522    fn has_jacobian_second_source(&self, method: &str) -> bool {
523        if self.jacobian_second_cache().is_some() || self.duchon_radial_source.is_some() {
524            true
525        } else {
526            self.missing_cache_default(
527                method,
528                "both jacobian_second_cache and duchon_radial_source are None",
529            );
530            false
531        }
532    }
533
534    fn has_jacobian_third_source(&self, method: &str) -> bool {
535        if self.third_decoder_derivative().is_some() || self.duchon_radial_source.is_some() {
536            true
537        } else {
538            self.missing_cache_default(
539                method,
540                "both third_decoder_derivative cache and duchon_radial_source are None",
541            );
542            false
543        }
544    }
545
546    /// Build `M_n = U_n^T J_n ∈ ℝ^{r_n × d}` for row `n`. For
547    /// `WeightField::Identity`, `r_n = p` and `M_n = J_n`.
548    ///
549    /// This is the single contraction site where `W_n` (or its `U_n` factor)
550    /// is consumed. Every value/grad/hvp path funnels through here, so the
551    /// `(J^T U)(U^T J)` ordering invariant cannot be violated by accident.
552    fn projected_jacobian_row(&self, n: usize, d: usize) -> Option<Array2<f64>> {
553        let jac = self.dimensioned_jacobian_cache("projected_jacobian_row", d)?;
554        let jac_row = jac.row(n);
555        let jac_slice = jac_row
556            .as_slice()
557            .expect("jacobian cache must be in standard row-major layout");
558        match &self.weight {
559            WeightField::Identity => {
560                let p = self.p_out;
561                let mut m = Array2::<f64>::zeros((p, d));
562                for i in 0..p {
563                    for a in 0..d {
564                        m[[i, a]] = jac_slice[i * d + a];
565                    }
566                }
567                Some(m)
568            }
569            WeightField::Factored { u, rank, p_out } => {
570                let u_row = u.row(n);
571                let u_slice = u_row
572                    .as_slice()
573                    .expect("weight factor U must be in standard row-major layout");
574                Some(WeightField::project_jac_row_with_u(
575                    u_slice, jac_slice, *p_out, *rank, d,
576                ))
577            }
578        }
579    }
580
581    /// Form `W_n J_n` without materializing `W_n`.
582    fn weighted_jacobian_row(&self, n: usize, d: usize) -> Option<Array2<f64>> {
583        let jac = self.dimensioned_jacobian_cache("weighted_jacobian_row", d)?;
584        let p = self.p_out;
585        match &self.weight {
586            WeightField::Identity => {
587                let mut out = Array2::<f64>::zeros((p, d));
588                for i in 0..p {
589                    for a in 0..d {
590                        out[[i, a]] = jac[[n, i * d + a]];
591                    }
592                }
593                Some(out)
594            }
595            WeightField::Factored { u, rank, p_out } => {
596                assert_eq!(p, *p_out);
597                let r = *rank;
598                let m_n = self.projected_jacobian_row(n, d)?;
599                let mut out = Array2::<f64>::zeros((p, d));
600                for i in 0..p {
601                    for a in 0..d {
602                        let mut s = 0.0;
603                        for k in 0..r {
604                            s += u[[n, i * r + k]] * m_n[[k, a]];
605                        }
606                        out[[i, a]] = s;
607                    }
608                }
609                Some(out)
610            }
611        }
612    }
613
614    fn weighted_dot_decoder_vectors<F, G>(&self, n: usize, p: usize, x: F, y: G) -> f64
615    where
616        F: Fn(usize) -> f64,
617        G: Fn(usize) -> f64,
618    {
619        match &self.weight {
620            WeightField::Identity => {
621                let mut s = 0.0;
622                for i in 0..p {
623                    s += x(i) * y(i);
624                }
625                s
626            }
627            WeightField::Factored { u, rank, p_out } => {
628                assert_eq!(p, *p_out);
629                let r = *rank;
630                let mut s = 0.0;
631                for k in 0..r {
632                    let mut ux = 0.0;
633                    let mut uy = 0.0;
634                    for i in 0..p {
635                        let uik = u[[n, i * r + k]];
636                        ux += uik * x(i);
637                        uy += uik * y(i);
638                    }
639                    s += ux * uy;
640                }
641                s
642            }
643        }
644    }
645
646    fn target_matrix(target: ArrayView1<'_, f64>, n_obs: usize, d: usize) -> Array2<f64> {
647        let mut out = Array2::<f64>::zeros((n_obs, d));
648        for n in 0..n_obs {
649            for a in 0..d {
650                out[[n, a]] = target[n * d + a];
651            }
652        }
653        out
654    }
655
656    /// Second-order input-location derivative tensor of the Duchon decoder,
657    /// flattened to `(n_obs, p_out · d²)` with column layout
658    /// `i·d² + (a·d + c)`.
659    ///
660    /// Thin adapter over the shared [`radial_basis_cartesian_derivative`]
661    /// engine: it owns the radial-jet evaluation and the radial→Cartesian map;
662    /// here we only forward the source geometry.
663    fn duchon_radial_jacobian_second(
664        &self,
665        target: ArrayView1<'_, f64>,
666        n_obs: usize,
667        d: usize,
668        source: &IsometryDuchonRadialSource,
669    ) -> Result<Array2<f64>, BasisError> {
670        assert_eq!(source.centers.ncols(), d);
671        assert_eq!(source.radial_coefficients.nrows(), source.centers.nrows());
672        assert_eq!(source.radial_coefficients.ncols(), self.p_out);
673        let t = Self::target_matrix(target, n_obs, d);
674        radial_basis_cartesian_derivative(
675            2,
676            t.view(),
677            source.centers.view(),
678            source.radial_coefficients.view(),
679            source.length_scale,
680            source.nullspace_order,
681            source.power,
682        )
683    }
684
685    /// Third-order input-location derivative tensor of the Duchon decoder,
686    /// shaped `(n_obs, p_out, d³)` with last-axis layout `(a·d + c)·d + e`.
687    ///
688    /// Thin adapter over the shared [`radial_basis_cartesian_derivative`]
689    /// engine; the flat `(n_obs, p_out · d³)` result is reshaped to the
690    /// `Array3` consumed by the HVP path (row-major flatten of `(p_out, d³)`
691    /// is exactly `i·d³ + idx`).
692    fn duchon_radial_jacobian_third(
693        &self,
694        target: ArrayView1<'_, f64>,
695        n_obs: usize,
696        d: usize,
697        source: &IsometryDuchonRadialSource,
698    ) -> Result<ndarray::Array3<f64>, BasisError> {
699        assert_eq!(source.centers.ncols(), d);
700        assert_eq!(source.radial_coefficients.nrows(), source.centers.nrows());
701        assert_eq!(source.radial_coefficients.ncols(), self.p_out);
702        let t = Self::target_matrix(target, n_obs, d);
703        let flat = radial_basis_cartesian_derivative(
704            3,
705            t.view(),
706            source.centers.view(),
707            source.radial_coefficients.view(),
708            source.length_scale,
709            source.nullspace_order,
710            source.power,
711        )?;
712        Ok(flat
713            .into_shape_with_order((n_obs, self.p_out, d * d * d))
714            .expect("radial_basis_cartesian_derivative order-3 output reshapes to (n_obs, p, d³)"))
715    }
716
717    fn jacobian_second<'a>(
718        &'a self,
719        target: ArrayView1<'_, f64>,
720        n_obs: usize,
721        d: usize,
722    ) -> Option<CowArray<'a, f64, Ix2>> {
723        if let Some(jac2) = self.jacobian_second_cache() {
724            // Clone the underlying Array2 to detach from the Arc — the
725            // CowArray needs to outlive the temporary Arc returned by the
726            // accessor. The clone is `n_obs × p·d²` floats, paid once per
727            // grad_target / hvp_state invocation; same per-step cost as the
728            // pre-refactor code path which also took ownership via
729            // `jac2.view().to_owned()` semantics implicitly.
730            return Some(CowArray::from((*jac2).clone()));
731        }
732        let source = self.duchon_radial_source.as_ref()?;
733        match self.duchon_radial_jacobian_second(target, n_obs, d, source) {
734            Ok(jac2) => Some(CowArray::from(jac2)),
735            Err(err) => {
736                self.missing_cache_default(
737                    "jacobian_second",
738                    &format!("failed to materialize Duchon radial second derivative: {err}"),
739                );
740                None
741            }
742        }
743    }
744
745    fn jacobian_third<'a>(
746        &'a self,
747        target: ArrayView1<'_, f64>,
748        n_obs: usize,
749        d: usize,
750    ) -> Option<CowArray<'a, f64, Ix3>> {
751        if let Some(jac3) = self.third_decoder_derivative() {
752            return Some(CowArray::from(jac3.as_ref().clone()));
753        }
754        let source = self.duchon_radial_source.as_ref()?;
755        match self.duchon_radial_jacobian_third(target, n_obs, d, source) {
756            Ok(jac3) => Some(CowArray::from(jac3)),
757            Err(err) => {
758                self.missing_cache_default(
759                    "jacobian_third",
760                    &format!("failed to materialize Duchon radial third derivative: {err}"),
761                );
762                None
763            }
764        }
765    }
766
767    pub(crate) fn hvp_state<'a>(
768        &'a self,
769        target: ArrayView1<'_, f64>,
770    ) -> Option<IsometryHvpState<'a>> {
771        let d = self
772            .target
773            .latent_dim
774            .expect("IsometryPenalty requires latent_dim on its PsiSlice");
775        let n_obs = target.len() / d;
776        if !self.has_jacobian_cache("hvp")
777            || !self.has_jacobian_second_source("hvp")
778            || !self.has_jacobian_third_source("hvp")
779        {
780            return None;
781        }
782        let p = self.p_out;
783        let jac2 = self.jacobian_second(target.view(), n_obs, d)?;
784        let jac3 = self.jacobian_third(target.view(), n_obs, d)?;
785        let g = self.pullback_metric(d)?;
786        let metric = self.normalized_metric_state(g, n_obs, d)?;
787        let mut wj_rows = Vec::with_capacity(n_obs);
788        for n in 0..n_obs {
789            wj_rows.push(self.weighted_jacobian_row(n, d)?);
790        }
791        Some(IsometryHvpState {
792            d,
793            n_obs,
794            p,
795            jac2,
796            jac3,
797            metric,
798            wj_rows,
799        })
800    }
801
802    pub(crate) fn hvp_with_precomputed_state(
803        &self,
804        state: &IsometryHvpState<'_>,
805        rho: ArrayView1<'_, f64>,
806        v: ArrayView1<'_, f64>,
807    ) -> Array1<f64> {
808        let mu = validated_learnable_weight(self.scalar_weight, rho[self.rho_index]);
809        let d = state.d;
810        let n_obs = state.n_obs;
811        let p = state.p;
812        let jac2 = &state.jac2;
813        let jac3 = &state.jac3;
814        let metric = &state.metric;
815        let mut out = Array1::<f64>::zeros(v.len());
816        let mut delta_g = Array2::<f64>::zeros((n_obs, d * d));
817        for n in 0..n_obs {
818            let wj = &state.wj_rows[n];
819            let row_delta = isometry_row_delta_g(jac2.view(), wj.view(), v, n, d, p);
820            for a in 0..d {
821                for b in 0..d {
822                    delta_g[[n, a * d + b]] = row_delta[[a, b]];
823                }
824            }
825        }
826        let delta_metric_grad = metric.metric_grad_direction(delta_g.view(), d);
827
828        for n in 0..n_obs {
829            let wj = &state.wj_rows[n];
830            for c in 0..d {
831                let mut acc = 0.0;
832                for a in 0..d {
833                    for b in 0..d {
834                        let dg = isometry_dg_entry(jac2.view(), wj.view(), n, d, p, a, b, c);
835                        acc += dg * delta_metric_grad[[n, a * d + b]];
836                    }
837                }
838                out[n * d + c] = mu * acc;
839            }
840
841            for c in 0..d {
842                let mut acc_res = 0.0;
843                for a in 0..d {
844                    for b in 0..d {
845                        let metric_grad = metric.metric_grad[[n, a * d + b]];
846                        if metric_grad == 0.0 {
847                            continue;
848                        }
849                        let mut bv = 0.0;
850                        for dd in 0..d {
851                            let vd = v[n * d + dd];
852                            if vd == 0.0 {
853                                continue;
854                            }
855                            let mut k_a_cd_w_j_b = 0.0;
856                            for i in 0..p {
857                                k_a_cd_w_j_b += jac3[[n, i, ((a * d) + c) * d + dd]] * wj[[i, b]];
858                            }
859                            let h_a_c_w_h_b_d = self.weighted_dot_decoder_vectors(
860                                n,
861                                p,
862                                |i| jac2[[n, (i * d + a) * d + c]],
863                                |i| jac2[[n, (i * d + b) * d + dd]],
864                            );
865                            let h_a_d_w_h_b_c = self.weighted_dot_decoder_vectors(
866                                n,
867                                p,
868                                |i| jac2[[n, (i * d + a) * d + dd]],
869                                |i| jac2[[n, (i * d + b) * d + c]],
870                            );
871                            let mut j_a_w_k_b_cd = 0.0;
872                            for i in 0..p {
873                                j_a_w_k_b_cd += wj[[i, a]] * jac3[[n, i, ((b * d) + c) * d + dd]];
874                            }
875                            bv +=
876                                (k_a_cd_w_j_b + h_a_c_w_h_b_d + h_a_d_w_h_b_c + j_a_w_k_b_cd) * vd;
877                        }
878                        acc_res += metric_grad * bv;
879                    }
880                }
881                out[n * d + c] += mu * acc_res;
882            }
883        }
884        out
885    }
886
887    /// Per-row pullback metric `g_n = J_n^T W_n J_n = M_n^T M_n` with
888    /// `M_n = U_n^T J_n ∈ ℝ^{r_n × d}`. Returns `(n_obs, d, d)` flattened
889    /// row-major as `(n_obs, d*d)`.
890    ///
891    /// Cost per row: `O(p · r · d)` for the `M_n` build (single pass over
892    /// `U_n` and `J_n`) plus `O(r · d²)` for `M_n^T M_n`. The `p × p` weight
893    /// `W_n` is never materialized.
894    pub fn pullback_metric(&self, latent_dim: usize) -> Option<Array2<f64>> {
895        let jac = self.dimensioned_jacobian_cache("pullback_metric", latent_dim)?;
896        let n_obs = jac.nrows();
897        // `dimensioned_jacobian_cache` enforces the load-bearing `(n, p·d)`
898        // shape contract before the reshape loop below. A stale cross-atom
899        // cache is a hard owner/refresh invariant failure; it is never converted
900        // into a zero isometry contribution (#2294).
901        let mut g_all = Array2::<f64>::zeros((n_obs, latent_dim * latent_dim));
902        for n in 0..n_obs {
903            // M_n = U_n^T J_n  (or J_n itself when W = I).
904            let m = self.projected_jacobian_row(n, latent_dim)?;
905            let r = m.nrows();
906            // g_n = M_n^T M_n: (d × d) result, contracting r.
907            for a in 0..latent_dim {
908                for b in 0..latent_dim {
909                    let mut s = 0.0;
910                    for k in 0..r {
911                        s += m[[k, a]] * m[[k, b]];
912                    }
913                    g_all[[n, a * latent_dim + b]] = s;
914                }
915            }
916        }
917        Some(g_all)
918    }
919
920    /// The scale normalizer `gbar = (1 / (N d)) Σ_n tr(g_n)` of the cached
921    /// pullback metric — the single shared denominator the scale-invariant
922    /// gauge divides every per-row metric by.
923    ///
924    /// `value` / `grad_*` / `hvp` consume this implicitly through
925    /// `Self::normalized_metric_state`; the SAE arrow-Schur assembly cannot
926    /// (it builds explicit per-row `htt` / `htbeta` / `hbb` curvature blocks
927    /// from the raw pullback `g_n`, not through the trait operators), so it
928    /// reads `gbar` here and folds `1/gbar²` into its Gauss-Newton curvature.
929    /// That `1/gbar²` factor is exactly the frozen-normalizer Gauss-Newton
930    /// block of the normalized residual `R_n = g_n/gbar − g^ref_n`: the raw
931    /// block (the GN block of the *un-normalized* `½μ‖g_n − g^ref‖²`) scales
932    /// ∝‖B‖⁴ in the decoder magnitude while the normalized gradient is
933    /// scale-free, so without the factor the joint Newton step collapses and
934    /// the proximal ridge saturates at 1e15 (#795). It stays PSD (a positive
935    /// scalar on an already-PSD Gram block), so the Schur complement is
936    /// unaffected. `None` when the metric is unavailable or degenerate, mirror-
937    /// ing `normalized_metric_state`'s non-positive-normalizer guard.
938    pub fn metric_normalizer(&self, latent_dim: usize) -> Option<f64> {
939        let g = self.pullback_metric(latent_dim)?;
940        let n_obs = g.nrows();
941        let normalizer = average_trace_per_dim(g.view(), n_obs, latent_dim);
942        (normalizer.is_finite() && normalizer > f64::MIN_POSITIVE).then_some(normalizer)
943    }
944
945    /// Reference metric per row for the normalized pullback metric, `(n_obs, d*d)`.
946    fn reference_metric(&self, n_obs: usize, d: usize) -> CowArray<'_, f64, Ix2> {
947        match &self.reference {
948            IsometryReference::Euclidean => {
949                let mut out = Array2::<f64>::zeros((n_obs, d * d));
950                for n in 0..n_obs {
951                    for a in 0..d {
952                        out[[n, a * d + a]] = 1.0;
953                    }
954                }
955                CowArray::from(out)
956            }
957            IsometryReference::UserSupplied(a) => {
958                assert_eq!(a.nrows(), n_obs);
959                assert_eq!(a.ncols(), d * d);
960                CowArray::from(a.view())
961            }
962        }
963    }
964
965    /// Shared normalized metric state for the scale-invariant isometry gauge.
966    ///
967    /// The residual is `R_n = g_n / gbar - g_ref,n / gref_bar`, with
968    /// `gbar = (1 / (N d)) Σ_n tr(g_n)` and `gref_bar = (1 / (N d)) Σ_n tr(g_ref,n)`
969    /// (`gref_bar == 1` for the `Euclidean` reference). `g_ref / gref_bar` is
970    /// constant w.r.t. the decoder coordinates, so the metric gradient/Hessian
971    /// form below is unchanged. The metric-gradient is the exact
972    /// derivative of `0.5 Σ ||R_n||²` with respect to the raw pullback metrics:
973    ///
974    /// `A_n = R_n / gbar - (Σ_l R_l:g_l) I / (gbar² N d)`.
975    ///
976    /// All value, gradient, and HVP paths consume this state so the global
977    /// normalizer's derivative is never detached.
978    fn normalized_metric_state(
979        &self,
980        g: Array2<f64>,
981        n_obs: usize,
982        d: usize,
983    ) -> Option<IsometryMetricState> {
984        let dd = d * d;
985        let trace_denominator = (n_obs * d) as f64;
986        let normalizer = average_trace_per_dim(g.view(), n_obs, d);
987        if !(normalizer.is_finite() && normalizer > f64::MIN_POSITIVE) {
988            self.missing_cache_default(
989                "normalized_metric_state",
990                &format!(
991                    "unit-average-speed normalizer is non-positive or non-finite: {normalizer}"
992                ),
993            );
994            return None;
995        }
996        let g_ref = self.reference_metric(n_obs, d);
997        // Normalize the reference by its own average trace per dim so the gauge
998        // is scale-invariant on both sides: a decoder metric proportional to the
999        // reference (up to an arbitrary global scale, common for external chart
1000        // metrics / GP-LVM warm starts) gives a zero residual. For `Euclidean`,
1001        // `ref_normalizer == 1.0` exactly, preserving the prior behavior bit-for-bit.
1002        let ref_normalizer = average_trace_per_dim(g_ref.view(), n_obs, d);
1003        if !(ref_normalizer.is_finite() && ref_normalizer > f64::MIN_POSITIVE) {
1004            self.missing_cache_default(
1005                "normalized_metric_state",
1006                &format!(
1007                    "reference-metric normalizer is non-positive or non-finite: {ref_normalizer}"
1008                ),
1009            );
1010            return None;
1011        }
1012        let mut residual = Array2::<f64>::zeros((n_obs, dd));
1013        let inv_norm = 1.0 / normalizer;
1014        let inv_ref_norm = 1.0 / ref_normalizer;
1015        for n in 0..n_obs {
1016            for k in 0..dd {
1017                residual[[n, k]] = g[[n, k]] * inv_norm - g_ref[[n, k]] * inv_ref_norm;
1018            }
1019        }
1020        let mut residual_dot_g = 0.0;
1021        for n in 0..n_obs {
1022            for k in 0..dd {
1023                residual_dot_g += residual[[n, k]] * g[[n, k]];
1024            }
1025        }
1026        let trace_coeff = residual_dot_g / (normalizer * normalizer * trace_denominator);
1027        let mut metric_grad = Array2::<f64>::zeros((n_obs, dd));
1028        for n in 0..n_obs {
1029            for a in 0..d {
1030                for b in 0..d {
1031                    let k = a * d + b;
1032                    let mut value = residual[[n, k]] * inv_norm;
1033                    if a == b {
1034                        value -= trace_coeff;
1035                    }
1036                    metric_grad[[n, k]] = value;
1037                }
1038            }
1039        }
1040        Some(IsometryMetricState {
1041            g,
1042            residual,
1043            metric_grad,
1044            normalizer,
1045            trace_denominator,
1046            residual_dot_g,
1047        })
1048    }
1049
1050    /// Exact closed-form gradient of the isometry penalty with respect to the
1051    /// cached decoder Jacobian `J ∈ ℝ^{n_obs × p × d}` (the autograd input that
1052    /// torch's `_IsometryPenaltyFn` differentiates). Returns the flattened
1053    /// `(n_obs, p*d)` layout that matches the Jacobian cache.
1054    ///
1055    /// Derivation (W-aware, reference-aware, weight-aware):
1056    ///
1057    ///   P        = ½ μ Σ_n ‖R_n‖²_F,
1058    ///   R_n      = g_n / gbar − g^ref_n,
1059    ///   gbar     = (1 / (N d)) Σ_n tr(g_n)
1060    ///   A_n      = ∂(P/μ)/∂g_n
1061    ///   ∂g_{ab}/∂J_{i,c}
1062    ///            = δ_{ca}(W J)_{i,b} + δ_{cb}(W J)_{i,a}   (W symmetric)
1063    ///   ∂P/∂J_{i,c}
1064    ///            = μ Σ_{a,b} A_{ab} ∂g_{ab}/∂J_{i,c}
1065    ///            = 2 μ Σ_b A_{cb} (W J)_{i,b}
1066    ///            = 2 μ ((W J) A)_{i,c}
1067    ///
1068    /// where `A` includes the exact derivative of the shared `gbar` normalizer.
1069    pub fn grad_jacobian(
1070        &self,
1071        target: ArrayView1<'_, f64>,
1072        rho: ArrayView1<'_, f64>,
1073    ) -> Array2<f64> {
1074        let d = self
1075            .target
1076            .latent_dim
1077            .expect("IsometryPenalty requires latent_dim on its PsiSlice");
1078        let n_obs = target.len() / d;
1079        let p = self.p_out;
1080        let mut grad = Array2::<f64>::zeros((n_obs, p * d));
1081        if !self.has_jacobian_cache("grad_jacobian") {
1082            return grad;
1083        }
1084        let Some(g) = self.pullback_metric(d) else {
1085            return grad;
1086        };
1087        let Some(metric) = self.normalized_metric_state(g, n_obs, d) else {
1088            return grad;
1089        };
1090        let mu = validated_learnable_weight(self.scalar_weight, rho[self.rho_index]);
1091        for n in 0..n_obs {
1092            let Some(wj) = self.weighted_jacobian_row(n, d) else {
1093                return Array2::<f64>::zeros((n_obs, p * d));
1094            };
1095            for i in 0..p {
1096                for c in 0..d {
1097                    let mut acc = 0.0;
1098                    for b in 0..d {
1099                        acc += metric.metric_grad[[n, c * d + b]] * wj[[i, b]];
1100                    }
1101                    grad[[n, i * d + c]] = 2.0 * mu * acc;
1102                }
1103            }
1104        }
1105        grad
1106    }
1107}
1108
1109impl AnalyticPenalty for IsometryPenalty {
1110    fn tier(&self) -> PenaltyTier {
1111        PenaltyTier::Psi
1112    }
1113
1114    fn validate_rho(&self, rho: ArrayView1<'_, f64>) -> Result<(), String> {
1115        if rho.len() != 1 {
1116            return Err(format!("isometry rho length {} != 1", rho.len()));
1117        }
1118        resolve_learnable_weight(self.scalar_weight, rho[self.rho_index])?;
1119        Ok(())
1120    }
1121
1122    fn rho_coordinate_domains(&self) -> Result<Vec<(f64, f64)>, String> {
1123        Ok(vec![
1124            learnable_weight_coordinate_domain(self.scalar_weight)?
1125                .ok_or_else(|| "isometry scalar weight must be positive".to_string())?,
1126        ])
1127    }
1128
1129    fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64 {
1130        let d = self
1131            .target
1132            .latent_dim
1133            .expect("IsometryPenalty requires latent_dim on its PsiSlice");
1134        let n_obs = target.len() / d;
1135        if !self.has_jacobian_cache("value") {
1136            return Self::DEFAULT_VALUE_ON_MISSING_CACHE;
1137        }
1138        let Some(g) = self.pullback_metric(d) else {
1139            return Self::DEFAULT_VALUE_ON_MISSING_CACHE;
1140        };
1141        let Some(metric) = self.normalized_metric_state(g, n_obs, d) else {
1142            return Self::DEFAULT_VALUE_ON_MISSING_CACHE;
1143        };
1144        let mu = validated_learnable_weight(self.scalar_weight, rho[self.rho_index]);
1145        let mut acc = 0.0;
1146        for n in 0..n_obs {
1147            for k in 0..(d * d) {
1148                let diff = metric.residual[[n, k]];
1149                acc += diff * diff;
1150            }
1151        }
1152        0.5 * mu * acc
1153    }
1154
1155    fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
1156        // Exact closed-form gradient, W-aware:
1157        //
1158        //   P     = ½ μ Σ_n ‖R_n‖²_F,   R_n = g_n / gbar − g^ref_n
1159        //   g_n   = J_n^T W_n J_n,      W_n = U_n U_n^T
1160        //   A_n   = ∂(P/μ)/∂g_n, including the exact derivative of
1161        //           gbar = (1 / (N d)) Σ_n tr(g_n)
1162        //   ∂g_{ab}/∂t_c
1163        //         = (H_{:,a,c})^T (W J)_{:,b}  +  (J_{:,a})^T W H_{:,b,c}
1164        //   ∂P/∂t_c
1165        //         = μ Σ_{a,b} A_{a,b} · ∂g_{ab}/∂t_c
1166        //
1167        // `H = ∂J/∂t` comes either from the live cache or from the radial
1168        // Duchon `φ''(r)` helper. The sign is positive: differentiating
1169        // `t - c` with respect to `t` contributes `+I`.
1170        let d = self
1171            .target
1172            .latent_dim
1173            .expect("IsometryPenalty requires latent_dim on its PsiSlice");
1174        let n_obs = target.len() / d;
1175        if !self.has_jacobian_cache("grad_target")
1176            || !self.has_jacobian_second_source("grad_target")
1177        {
1178            return Array1::<f64>::zeros(target.len());
1179        }
1180        let Some(g) = self.pullback_metric(d) else {
1181            return Array1::<f64>::zeros(target.len());
1182        };
1183        let Some(metric) = self.normalized_metric_state(g, n_obs, d) else {
1184            return Array1::<f64>::zeros(target.len());
1185        };
1186        let p = self.p_out;
1187        let mu = validated_learnable_weight(self.scalar_weight, rho[self.rho_index]);
1188        let mut grad = Array1::<f64>::zeros(target.len());
1189        let Some(jac2) = self.jacobian_second(target, n_obs, d) else {
1190            return grad;
1191        };
1192        assert_eq!(jac2.ncols(), p * d * d);
1193
1194        for n in 0..n_obs {
1195            let Some(wj) = self.weighted_jacobian_row(n, d) else {
1196                return grad;
1197            };
1198            for c in 0..d {
1199                let mut acc = 0.0;
1200                for a in 0..d {
1201                    for b in 0..d {
1202                        let mut dg = 0.0;
1203                        for i in 0..p {
1204                            dg += jac2[[n, (i * d + a) * d + c]] * wj[[i, b]];
1205                            dg += wj[[i, a]] * jac2[[n, (i * d + b) * d + c]];
1206                        }
1207                        acc += metric.metric_grad[[n, a * d + b]] * dg;
1208                    }
1209                }
1210                grad[n * d + c] = mu * acc;
1211            }
1212        }
1213        grad
1214    }
1215
1216    /// Fully analytic - wired through `radial_basis_cartesian_derivative`.
1217    fn hvp(
1218        &self,
1219        target: ArrayView1<'_, f64>,
1220        rho: ArrayView1<'_, f64>,
1221        v: ArrayView1<'_, f64>,
1222    ) -> Array1<f64> {
1223        // Fully analytic isometry Hessian-vector product wired through the
1224        // shared `radial_basis_cartesian_derivative` engine when no
1225        // third-derivative cache is supplied.
1226        //
1227        // The full Hessian of P_iso = (μ/2) Σ_n ||J^T W J / gbar - G_ref||²_F
1228        // (per proposal §4(b)) is
1229        //   μ [Dgᵀ · ∂²(0.5||R||²)/∂g² · Dg + A · ∂²g],
1230        // where R = g/gbar - G_ref and A = ∂(0.5||R||²)/∂g includes the global
1231        // gbar derivative.
1232        //   B_{ab,cd} = K_{a,cd}^T W J_b + H_{a,c}^T W H_{b,d}
1233        //             + H_{a,d}^T W H_{b,c} + J_a^T W K_{b,cd},
1234        // where K is the third decoder derivative and H is the second.
1235        let Some(state) = self.hvp_state(target) else {
1236            return Array1::<f64>::zeros(v.len());
1237        };
1238        self.hvp_with_precomputed_state(&state, rho, v)
1239    }
1240
1241    /// PSD majorizer-vector product `B_GN(target; ρ) v` for the **nonconvex**
1242    /// isometry penalty.
1243    ///
1244    /// The Gauss-Newton block differentiates the normalized residual
1245    /// `R = g/gbar - G_ref` itself and returns `μ DRᵀ DR v`. This is PSD by
1246    /// construction and includes the shared-normalizer derivative exactly;
1247    /// using only `∂g` would reintroduce scale coupling and would not be the
1248    /// Gauss-Newton operator of the objective being minimized.
1249    fn psd_majorizer_hvp(
1250        &self,
1251        target: ArrayView1<'_, f64>,
1252        rho: ArrayView1<'_, f64>,
1253        v: ArrayView1<'_, f64>,
1254    ) -> Array1<f64> {
1255        let d = self
1256            .target
1257            .latent_dim
1258            .expect("IsometryPenalty requires latent_dim on its PsiSlice");
1259        let n_obs = target.len() / d;
1260        if !self.has_jacobian_cache("psd_majorizer_hvp")
1261            || !self.has_jacobian_second_source("psd_majorizer_hvp")
1262        {
1263            return Array1::<f64>::zeros(v.len());
1264        }
1265        let Some(jac2) = self.jacobian_second(target, n_obs, d) else {
1266            return Array1::<f64>::zeros(v.len());
1267        };
1268        let Some(g) = self.pullback_metric(d) else {
1269            return Array1::<f64>::zeros(v.len());
1270        };
1271        let Some(metric) = self.normalized_metric_state(g, n_obs, d) else {
1272            return Array1::<f64>::zeros(v.len());
1273        };
1274        let p = self.p_out;
1275        let mu = validated_learnable_weight(self.scalar_weight, rho[self.rho_index]);
1276        let mut out = Array1::<f64>::zeros(v.len());
1277        let mut wj_rows = Vec::with_capacity(n_obs);
1278        for n in 0..n_obs {
1279            let Some(wj) = self.weighted_jacobian_row(n, d) else {
1280                return Array1::<f64>::zeros(v.len());
1281            };
1282            wj_rows.push(wj);
1283        }
1284        let mut delta_g = Array2::<f64>::zeros((n_obs, d * d));
1285        for n in 0..n_obs {
1286            let row_delta = isometry_row_delta_g(jac2.view(), wj_rows[n].view(), v, n, d, p);
1287            for a in 0..d {
1288                for b in 0..d {
1289                    delta_g[[n, a * d + b]] = row_delta[[a, b]];
1290                }
1291            }
1292        }
1293        let (delta_residual, _delta_normalizer) = metric.residual_direction(delta_g.view(), d);
1294        let mut g_dot_delta_residual = 0.0;
1295        for n in 0..n_obs {
1296            for k in 0..(d * d) {
1297                g_dot_delta_residual += metric.g[[n, k]] * delta_residual[[n, k]];
1298            }
1299        }
1300        let inv_norm = 1.0 / metric.normalizer;
1301        let inv_norm_sq = inv_norm * inv_norm;
1302        for n in 0..n_obs {
1303            let wj = &wj_rows[n];
1304            for c in 0..d {
1305                let mut trace_dg = 0.0;
1306                for a in 0..d {
1307                    trace_dg += isometry_dg_entry(jac2.view(), wj.view(), n, d, p, a, a, c);
1308                }
1309                let delta_normalizer_c = trace_dg / metric.trace_denominator;
1310                let mut acc = -delta_normalizer_c * inv_norm_sq * g_dot_delta_residual;
1311                for a in 0..d {
1312                    for b in 0..d {
1313                        let dg = isometry_dg_entry(jac2.view(), wj.view(), n, d, p, a, b, c);
1314                        acc += dg * inv_norm * delta_residual[[n, a * d + b]];
1315                    }
1316                }
1317                out[n * d + c] = mu * acc;
1318            }
1319        }
1320        out
1321    }
1322
1323    fn grad_rho(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
1324        // P(ρ) = ½ μ · S, where S is the (ρ-independent) Frobenius sum and
1325        // μ = exp(ρ_iso). So ∂P/∂ρ_iso = P.
1326        let mut out = Array1::<f64>::zeros(self.rho_count());
1327        out[self.rho_index] = self.value(target, rho);
1328        out
1329    }
1330
1331    fn rho_count(&self) -> usize {
1332        1
1333    }
1334
1335    fn name(&self) -> &str {
1336        "isometry"
1337    }
1338
1339    impl_scalar_apply_schedule!(scalar_weight);
1340}