1use super::*;
2use gam_problem::{ConstraintSet, KhatriRaoConeConstraints};
3
4impl CustomFamily for TransformationNormalFamily {
9 fn inner_objective_is_self_concordant(&self) -> bool {
25 true
26 }
27
28 fn evaluate(&self, block_states: &[ParameterBlockState]) -> Result<FamilyEvaluation, String> {
29 crate::block_layout::block_count::validate_block_count::<TransformationNormalError>(
30 "TransformationNormalFamily",
31 1,
32 block_states.len(),
33 )?;
34 let evaluate_start = std::time::Instant::now();
35 let beta = &block_states[0].beta;
36 let row_q_start = std::time::Instant::now();
37 let row_quantities = self.row_quantities(beta)?;
38 log::info!(
39 "[STAGE] CTN row_quantities (h, h', 1/h', powers) n={} elapsed={:.3}s",
40 row_quantities.h.len(),
41 row_q_start.elapsed().as_secs_f64(),
42 );
43 let h = row_quantities.h.as_ref();
44 let n = h.len();
45
46 let log_likelihood = row_quantities.log_likelihood;
47 let grad_start = std::time::Instant::now();
51 let (grad, hessian) = self.scop_gradient_and_negative_hessian(beta, &row_quantities)?;
52 log::info!(
53 "[STAGE] CTN gradient terms n={} p={} elapsed={:.3}s",
54 n,
55 grad.len(),
56 grad_start.elapsed().as_secs_f64(),
57 );
58
59 let hess_start = std::time::Instant::now();
60 let p_dim = hessian.nrows() as u64;
61 let n_u64 = n as u64;
62 log::info!(
63 "[STAGE] CTN hessian terms (SCOP exact dense) n={} p={} flops~{} elapsed={:.3}s",
64 n,
65 p_dim,
66 n_u64.saturating_mul(p_dim).saturating_mul(p_dim),
67 hess_start.elapsed().as_secs_f64(),
68 );
69 log::info!(
70 "[STAGE] CTN evaluate end n={} p={} elapsed={:.3}s",
71 n,
72 p_dim,
73 evaluate_start.elapsed().as_secs_f64(),
74 );
75
76 Ok(FamilyEvaluation {
77 log_likelihood,
78 blockworking_sets: vec![BlockWorkingSet::ExactNewton {
79 gradient: grad,
80 hessian: SymmetricMatrix::Dense(hessian),
81 }],
82 })
83 }
84
85 fn log_likelihood_only(&self, block_states: &[ParameterBlockState]) -> Result<f64, String> {
86 crate::block_layout::block_count::validate_block_count::<TransformationNormalError>(
87 "TransformationNormalFamily",
88 1,
89 block_states.len(),
90 )?;
91 let row_quantities = match self.row_quantities(&block_states[0].beta) {
95 Ok(rq) => rq,
96 Err(_) => return Ok(f64::NEG_INFINITY),
97 };
98 Ok(row_quantities.log_likelihood)
99 }
100
101 fn log_likelihood_only_with_options(
102 &self,
103 block_states: &[ParameterBlockState],
104 options: &BlockwiseFitOptions,
105 ) -> Result<f64, String> {
106 match self.maybe_with_outer_subsample_from_options(options) {
113 Ok(Some(masked)) => masked.log_likelihood_only(block_states),
114 Ok(None) => self.log_likelihood_only(block_states),
115 Err(e) => Err(e.into()),
116 }
117 }
118
119 fn exact_newton_joint_gradient_evaluation(
135 &self,
136 block_states: &[ParameterBlockState],
137 _: &[ParameterBlockSpec],
138 ) -> Result<Option<ExactNewtonJointGradientEvaluation>, String> {
139 crate::block_layout::block_count::validate_block_count::<TransformationNormalError>(
140 "TransformationNormalFamily",
141 1,
142 block_states.len(),
143 )?;
144 let beta = &block_states[0].beta;
145 let row_quantities = self.row_quantities(beta)?;
146 let log_likelihood = row_quantities.log_likelihood;
147 let gradient = self.scop_gradient(beta, &row_quantities)?;
148 Ok(Some(ExactNewtonJointGradientEvaluation {
149 log_likelihood,
150 gradient,
151 }))
152 }
153
154 fn exact_newton_joint_hessian_beta_dependent(&self) -> bool {
155 true
157 }
158
159 fn joint_jeffreys_term_required(&self) -> bool {
160 false
179 }
180
181 fn coefficient_hessian_cost(&self, specs: &[ParameterBlockSpec]) -> u64 {
182 let n_usize = self.response_val_basis.nrows();
197 let p_resp = self.response_val_basis.ncols() as u64;
198 let p_cov = self.covariate_design.ncols() as u64;
199 let expected_p_total = p_resp.saturating_mul(p_cov);
200 let p_total = match specs {
212 [] => expected_p_total,
213 [spec] if spec.design.ncols() as u64 == expected_p_total => spec.design.ncols() as u64,
214 _ => return u64::MAX,
215 };
216 let n = n_usize as u64;
217 crate::coefficient_cost::operator_aware_hessian_cost(
224 p_total,
225 n,
226 n.saturating_mul(p_resp.saturating_add(p_cov)),
227 n.saturating_mul(p_total.saturating_mul(p_total)),
228 )
229 }
230
231 fn coefficient_gradient_cost(&self, specs: &[ParameterBlockSpec]) -> u64 {
232 self.coefficient_hessian_cost(specs) / 2
236 }
237
238 fn outer_derivative_policy(
239 &self,
240 specs: &[crate::custom_family::ParameterBlockSpec],
241 psi_dim: usize,
242 options: &crate::custom_family::BlockwiseFitOptions,
243 ) -> crate::custom_family::OuterDerivativePolicy {
244 let capability = self.exact_outer_derivative_order(specs, options);
257 let n = specs.first().map_or(0u128, |s| s.design.nrows() as u128);
258 let p_total: u128 = specs
259 .iter()
260 .map(|s| s.design.ncols() as u128)
261 .fold(0u128, |acc, x| acc.saturating_add(x));
262 let rho_dim: u128 = specs
263 .iter()
264 .map(|s| s.penalties.len() as u128)
265 .fold(0u128, |acc, x| acc.saturating_add(x));
266 let k = rho_dim.saturating_add(psi_dim as u128).max(1);
267 let p_eff = p_total.max(1);
268 let work_grad = n.saturating_mul(k).saturating_mul(p_eff);
270 let dense_hess = work_grad.saturating_mul(p_eff);
276 let mfree_hess = work_grad.saturating_mul(rho_dim.max(1));
277 let work_hess = dense_hess.min(mfree_hess);
278 crate::custom_family::OuterDerivativePolicy {
279 capability,
280 predicted_hessian_work: work_hess,
281 predicted_gradient_work: work_grad,
282 subsample_capable: true,
292 }
293 }
294
295 fn outer_seed_config(&self, n_params: usize) -> gam_solve::seeding::SeedConfig {
296 gam_solve::seeding::SeedConfig {
297 bounds: (-12.0, 12.0),
298 max_seeds: if n_params <= 8 { 1 } else { 2 },
299 seed_budget: 1,
300 screen_max_inner_iterations: 2,
301 risk_profile: gam_solve::seeding::SeedRiskProfile::Gaussian,
302 num_auxiliary_trailing: 0,
303 over_smoothing_probe_rho: None,
304 }
305 }
306
307 fn max_feasible_step_size(
308 &self,
309 block_states: &[ParameterBlockState],
310 block_index: usize,
311 delta: &Array1<f64>,
312 ) -> Result<Option<f64>, String> {
313 if block_index != 0 {
314 return Ok(None);
315 }
316 crate::block_layout::block_count::validate_block_count::<TransformationNormalError>(
317 "TransformationNormalFamily",
318 1,
319 block_states.len(),
320 )?;
321 if delta.len() != block_states[0].beta.len() {
322 return Err(TransformationNormalError::InvalidInput {
323 reason: format!(
324 "CTN line-search step length {} != beta length {}",
325 delta.len(),
326 block_states[0].beta.len()
327 ),
328 }
329 .into());
330 }
331 Ok(None)
336 }
337
338 fn block_linear_constraints(
339 &self,
340 _: &[ParameterBlockState],
341 block_index: usize,
342 block_spec: &ParameterBlockSpec,
343 ) -> Result<Option<ConstraintSet>, String> {
344 assert!(!block_spec.name.is_empty());
345 if block_index != 0 {
346 return Ok(None);
347 }
348 let p_resp = self.response_val_basis.ncols();
359 if p_resp <= 1 {
360 return Ok(None);
361 }
362 let factor = self.covariate_dense_arc()?;
363 let cone = KhatriRaoConeConstraints::new(factor, (1..p_resp).collect(), p_resp)?;
364 if cone.ncols() != block_spec.design.ncols() {
365 return Err(format!(
366 "CTN factored monotonicity cone width {} != coefficient block width {}",
367 cone.ncols(),
368 block_spec.design.ncols(),
369 ));
370 }
371 Ok(Some(ConstraintSet::KhatriRaoCone(cone)))
372 }
373
374 fn exact_newton_hessian_directional_derivative(
375 &self,
376 block_states: &[ParameterBlockState],
377 block_index: usize,
378 d_beta: &Array1<f64>,
379 ) -> Result<Option<Array2<f64>>, String> {
380 if block_index != 0 {
381 return Ok(None);
382 }
383 let beta = &block_states[0].beta;
384 let row_quantities = self.row_quantities(beta)?;
385 let dd = self.scop_hessian_directional_derivative(beta, d_beta, &row_quantities)?;
386 Ok(Some(dd))
387 }
388
389 fn exact_newton_joint_hessian(
390 &self,
391 block_states: &[ParameterBlockState],
392 ) -> Result<Option<Array2<f64>>, String> {
393 let beta = &block_states[0].beta;
395 let row_quantities = self.row_quantities(beta)?;
396 let (_, hessian) = self.scop_gradient_and_negative_hessian(beta, &row_quantities)?;
397 Ok(Some(hessian))
398 }
399
400 fn exact_newton_joint_hessian_directional_derivative(
401 &self,
402 block_states: &[ParameterBlockState],
403 d_beta_flat: &Array1<f64>,
404 ) -> Result<Option<Array2<f64>>, String> {
405 self.exact_newton_hessian_directional_derivative(block_states, 0, d_beta_flat)
406 }
407
408 fn exact_newton_joint_hessiansecond_directional_derivative(
409 &self,
410 block_states: &[ParameterBlockState],
411 d_beta_u_flat: &Array1<f64>,
412 d_beta_v_flat: &Array1<f64>,
413 ) -> Result<Option<Array2<f64>>, String> {
414 let beta = &block_states[0].beta;
415 let row_quantities = self.row_quantities(beta)?;
416 let d2 = self.scop_hessian_second_directional_derivative(
417 beta,
418 d_beta_u_flat,
419 d_beta_v_flat,
420 &row_quantities,
421 )?;
422 Ok(Some(d2))
423 }
424
425 fn exact_newton_joint_psi_terms(
426 &self,
427 block_states: &[ParameterBlockState],
428 _: &[ParameterBlockSpec],
429 hyper_layout: &CustomFamilyHyperLayout,
430 psi_index: usize,
431 ) -> Result<Option<ExactNewtonJointPsiTerms>, String> {
432 if hyper_layout.family_axis_count() != 0 {
433 return Err(
434 "TransformationNormalFamily does not declare family-owned hyper axes".to_string(),
435 );
436 }
437 let psi_derivs = hyper_layout.design_derivative_blocks();
438 if psi_derivs.is_empty() || psi_index >= psi_derivs[0].len() {
439 return Ok(None);
440 }
441 let psi_first_start = std::time::Instant::now();
442 let deriv = &psi_derivs[0][psi_index];
443 let beta = &block_states[0].beta;
444 let row = self.row_quantities(beta)?;
445 let op = deriv
446 .implicit_operator
447 .as_ref()
448 .and_then(|op| op.as_any().downcast_ref::<TensorKroneckerPsiOperator>())
449 .ok_or_else(|| {
450 "TransformationNormalFamily requires tensor psi derivatives to remain operator-backed"
451 .to_string()
452 })?;
453 let axis = deriv.implicit_axis;
454 let op_arc = Arc::clone(
455 deriv
456 .implicit_operator
457 .as_ref()
458 .expect("validated CTN psi derivative operator disappeared"),
459 );
460 let terms = self.scop_psi_terms(beta, &row, op, op_arc, axis)?;
461
462 log::info!(
463 "[STAGE] CTN psi first-order terms axis={} psi_index={} elapsed={:.3}s",
464 deriv.implicit_axis,
465 psi_index,
466 psi_first_start.elapsed().as_secs_f64(),
467 );
468
469 Ok(Some(terms))
470 }
471
472 fn exact_newton_joint_psisecond_order_terms(
473 &self,
474 block_states: &[ParameterBlockState],
475 _: &[ParameterBlockSpec],
476 hyper_layout: &CustomFamilyHyperLayout,
477 psi_i: usize,
478 psi_j: usize,
479 ) -> Result<Option<ExactNewtonJointPsiSecondOrderTerms>, String> {
480 if hyper_layout.family_axis_count() != 0 {
481 return Err(
482 "TransformationNormalFamily does not declare family-owned hyper axes".to_string(),
483 );
484 }
485 let psi_derivs = hyper_layout.design_derivative_blocks();
486 if psi_derivs.is_empty() || psi_i >= psi_derivs[0].len() || psi_j >= psi_derivs[0].len() {
487 return Ok(None);
488 }
489 let psi_pair_start = std::time::Instant::now();
490 let deriv_i = &psi_derivs[0][psi_i];
491 let deriv_j = &psi_derivs[0][psi_j];
492 let beta = &block_states[0].beta;
493 let row = self.row_quantities(beta)?;
494 let p_resp = self.response_val_basis.ncols();
495 let p_cov = self.covariate_design.ncols();
496 let p_total = p_resp * p_cov;
497 if beta.len() != p_total {
498 return Err(TransformationNormalError::InvalidInput {
499 reason: format!(
500 "SCOP psi-psi terms beta length {} != p_resp({p_resp}) * p_cov({p_cov})",
501 beta.len()
502 ),
503 }
504 .into());
505 }
506
507 let op = deriv_i
508 .implicit_operator
509 .as_ref()
510 .and_then(|op| op.as_any().downcast_ref::<TensorKroneckerPsiOperator>())
511 .ok_or_else(|| {
512 "TransformationNormalFamily requires tensor psi derivatives to remain operator-backed"
513 .to_string()
514 })?;
515 let axis_i = deriv_i.implicit_axis;
516 let axis_j = deriv_j.implicit_axis;
517
518 let (objective_psi_psi, score_psi_psi, _) = self
519 .scop_psi_psi_value_score_hvp_from_operator(
520 beta,
521 op,
522 axis_i,
523 axis_j,
524 row.alpha.view(),
525 row.h.view(),
526 row.h_prime.view(),
527 row.endpoint_q.as_slice(),
528 None,
529 )?;
530 let hessian_psi_psi_operator: Arc<dyn HyperOperator> =
531 Arc::new(TransformationNormalPsiPsiHessianOperator::new(
532 Arc::new(self.clone()),
533 beta.clone(),
534 Arc::clone(
535 deriv_i
536 .implicit_operator
537 .as_ref()
538 .expect("validated CTN psi derivative has an implicit operator"),
539 ),
540 axis_i,
541 axis_j,
542 Arc::clone(&row.alpha),
543 Arc::clone(&row.h),
544 Arc::clone(&row.h_prime),
545 Arc::clone(&row.endpoint_q),
546 ));
547
548 if !objective_psi_psi.is_finite() || !score_psi_psi.iter().all(|v| v.is_finite()) {
554 return Err(TransformationNormalError::NonFinite {
555 reason: format!(
556 "TransformationNormalFamily exact ψ-ψ second-order terms produced \
557 non-finite values at psi_i={psi_i}, psi_j={psi_j}: \
558 obj_finite={}, score_all_finite={}. \
559 The outer evaluator should retreat from this trial point.",
560 objective_psi_psi.is_finite(),
561 score_psi_psi.iter().all(|v| v.is_finite()),
562 ),
563 }
564 .into());
565 }
566
567 log::info!(
568 "[STAGE] CTN psi-psi pair (psi_i={}, psi_j={}, axes={},{}) elapsed={:.3}s",
569 psi_i,
570 psi_j,
571 deriv_i.implicit_axis,
572 deriv_j.implicit_axis,
573 psi_pair_start.elapsed().as_secs_f64(),
574 );
575
576 Ok(Some(ExactNewtonJointPsiSecondOrderTerms {
577 objective_psi_psi,
578 score_psi_psi,
579 hessian_psi_psi: Array2::zeros((0, 0)),
580 hessian_psi_psi_operator: Some(hessian_psi_psi_operator),
581 }))
582 }
583
584 fn exact_newton_joint_psihessian_directional_derivative(
585 &self,
586 block_states: &[ParameterBlockState],
587 _: &[ParameterBlockSpec],
588 hyper_layout: &CustomFamilyHyperLayout,
589 psi_index: usize,
590 d_beta_flat: &Array1<f64>,
591 ) -> Result<Option<Array2<f64>>, String> {
592 if hyper_layout.family_axis_count() != 0 {
593 return Err(
594 "TransformationNormalFamily does not declare family-owned hyper axes".to_string(),
595 );
596 }
597 let psi_derivs = hyper_layout.design_derivative_blocks();
598 if psi_derivs.is_empty() || psi_index >= psi_derivs[0].len() {
599 return Ok(None);
600 }
601 let deriv = &psi_derivs[0][psi_index];
602 let beta = &block_states[0].beta;
603 let op = deriv
604 .implicit_operator
605 .as_ref()
606 .and_then(|op| op.as_any().downcast_ref::<TensorKroneckerPsiOperator>())
607 .ok_or_else(|| {
608 "TransformationNormalFamily requires tensor psi derivatives to remain operator-backed"
609 .to_string()
610 })?;
611 let axis = deriv.implicit_axis;
612 let row = self.row_quantities(beta)?;
613 let hess =
614 self.scop_psi_hessian_directional_derivative(beta, d_beta_flat, &row, op, axis)?;
615 Ok(Some(hess))
616 }
617
618 fn exact_newton_joint_hessian_workspace(
619 &self,
620 block_states: &[ParameterBlockState],
621 specs: &[ParameterBlockSpec],
622 ) -> Result<Option<Arc<dyn ExactNewtonJointHessianWorkspace>>, String> {
623 crate::block_layout::block_count::validate_block_count::<TransformationNormalError>(
624 "TransformationNormalFamily",
625 1,
626 block_states.len(),
627 )?;
628 if !self.inner_coefficient_hessian_hvp_available(specs) {
629 return Err(TransformationNormalError::InvalidInput {
630 reason: "TransformationNormalFamily joint Hessian workspace received incompatible block specs"
631 .to_string(),
632 }
633 .into());
634 }
635 let beta = &block_states[0].beta;
636 let row_quantities = self.row_quantities(beta)?;
637 let workspace = TransformationNormalJointHessianWorkspace::new(
641 Arc::new(self.clone()),
642 beta.clone(),
643 row_quantities.clone(),
644 )?;
645 Ok(Some(
646 Arc::new(workspace) as Arc<dyn ExactNewtonJointHessianWorkspace>
647 ))
648 }
649
650 fn exact_newton_joint_psi_workspace(
651 &self,
652 block_states: &[ParameterBlockState],
653 specs: &[ParameterBlockSpec],
654 hyper_layout: &CustomFamilyHyperLayout,
655 ) -> Result<Option<Arc<dyn ExactNewtonJointPsiWorkspace>>, String> {
656 if hyper_layout.family_axis_count() != 0 {
657 return Err(
658 "TransformationNormalFamily does not declare family-owned hyper axes".to_string(),
659 );
660 }
661 if !self.inner_coefficient_hessian_hvp_available(specs) {
662 return Err(TransformationNormalError::InvalidInput {
663 reason: "TransformationNormalFamily joint psi workspace received incompatible block specs"
664 .to_string(),
665 }
666 .into());
667 }
668 Ok(Some(Arc::new(TransformationNormalPsiWorkspace::new(
669 self.clone(),
670 block_states.to_vec(),
671 hyper_layout.design_derivative_blocks().to_vec(),
672 ))))
673 }
674
675 fn exact_newton_joint_hessian_workspace_with_options(
676 &self,
677 block_states: &[ParameterBlockState],
678 specs: &[ParameterBlockSpec],
679 options: &BlockwiseFitOptions,
680 ) -> Result<Option<Arc<dyn ExactNewtonJointHessianWorkspace>>, String> {
681 match self.maybe_with_outer_subsample_from_options(options)? {
689 Some(masked) => masked.exact_newton_joint_hessian_workspace(block_states, specs),
690 None => self.exact_newton_joint_hessian_workspace(block_states, specs),
691 }
692 }
693
694 fn exact_newton_joint_psi_workspace_with_options(
695 &self,
696 block_states: &[ParameterBlockState],
697 specs: &[ParameterBlockSpec],
698 hyper_layout: &CustomFamilyHyperLayout,
699 options: &BlockwiseFitOptions,
700 ) -> Result<Option<Arc<dyn ExactNewtonJointPsiWorkspace>>, String> {
701 if hyper_layout.family_axis_count() != 0 {
702 return Err(
703 "TransformationNormalFamily does not declare family-owned hyper axes".to_string(),
704 );
705 }
706 if !self.inner_coefficient_hessian_hvp_available(specs) {
707 return Err(TransformationNormalError::InvalidInput {
708 reason: "TransformationNormalFamily joint psi workspace received incompatible block specs"
709 .to_string(),
710 }
711 .into());
712 }
713 let family = match self.maybe_with_outer_subsample_from_options(options)? {
726 Some(masked) => masked,
727 None => self.clone(),
728 };
729 Ok(Some(Arc::new(TransformationNormalPsiWorkspace::new(
730 family,
731 block_states.to_vec(),
732 hyper_layout.design_derivative_blocks().to_vec(),
733 ))))
734 }
735
736 fn exact_newton_joint_psi_workspace_for_first_order_terms(&self) -> bool {
737 true
743 }
744
745 fn inner_coefficient_hessian_hvp_available(&self, specs: &[ParameterBlockSpec]) -> bool {
746 matches!(specs, [spec] if spec.design.ncols()
749 == self.response_val_basis.ncols().saturating_mul(self.covariate_design.ncols()))
750 }
751
752 fn outer_hyper_hessian_hvp_available(&self, specs: &[ParameterBlockSpec]) -> bool {
753 self.inner_coefficient_hessian_hvp_available(specs)
754 }
755
756 fn outer_hyper_hessian_dense_available(&self, specs: &[ParameterBlockSpec]) -> bool {
757 self.inner_coefficient_hessian_hvp_available(specs)
761 }
762}