use super::*;
#[derive(Clone, Debug)]
pub struct TwoBlockRemlFitReport {
pub log_lambda_y: f64,
pub sweeps: usize,
pub converged: bool,
pub lambda_identifiable: bool,
pub loss: SaeManifoldLoss,
pub log_lambda_trajectory: Vec<f64>,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct TwoBlockRemlControls {
pub max_sweeps: usize,
pub inner_max_iter: usize,
pub step_size: f64,
pub ridge_ext_coord: f64,
pub ridge_beta: f64,
pub log_lambda_tol: f64,
}
impl SaeManifoldTerm {
pub fn run_two_block_reml_fit(
&mut self,
activation: ArrayView2<'_, f64>,
rho: &mut SaeManifoldRho,
analytic_penalties: Option<&AnalyticPenaltyRegistry>,
controls: TwoBlockRemlControls,
) -> Result<TwoBlockRemlFitReport, String> {
let TwoBlockRemlControls {
max_sweeps,
inner_max_iter,
step_size,
ridge_ext_coord,
ridge_beta,
log_lambda_tol,
} = controls;
if max_sweeps == 0 {
return Err(
"SaeManifoldTerm::run_two_block_reml_fit: max_sweeps must be ≥ 1".to_string(),
);
}
if !(log_lambda_tol.is_finite() && log_lambda_tol > 0.0) {
return Err(format!(
"SaeManifoldTerm::run_two_block_reml_fit: log_lambda_tol must be finite and \
positive; got {log_lambda_tol}"
));
}
let Some(block) = self.behavior_block().cloned() else {
return Err(
"SaeManifoldTerm::run_two_block_reml_fit: no behavior block installed \
(call set_behavior_block first)"
.to_string(),
);
};
if activation.ncols() != block.activation_dim {
return Err(format!(
"SaeManifoldTerm::run_two_block_reml_fit: activation has {} columns; behavior \
block declares p_x = {}",
activation.ncols(),
block.activation_dim
));
}
let mut block = block;
let mut trajectory = Vec::with_capacity(max_sweeps);
let mut loss: Option<SaeManifoldLoss> = None;
let mut converged = false;
let mut lambda_identifiable = true;
let mut sweeps = 0usize;
while sweeps < max_sweeps {
sweeps += 1;
let augmented = block.augmented_target(activation)?;
let sweep_loss = self.run_joint_fit_arrow_schur(
augmented.view(),
rho,
analytic_penalties,
inner_max_iter,
step_size,
ridge_ext_coord,
ridge_beta,
)?;
loss = Some(sweep_loss);
let residual = self.reconstruction_residual(augmented.view(), rho)?;
let new_log_lambda = match block.reml_updated_log_lambda_y(residual.view()) {
Ok(value) => value,
Err(_) => {
lambda_identifiable = false;
converged = true;
trajectory.push(block.log_lambda_y);
break;
}
};
let delta = (new_log_lambda - block.log_lambda_y).abs();
trajectory.push(new_log_lambda);
block = block.with_log_lambda_y(new_log_lambda)?;
self.set_behavior_block(block.clone())?;
if delta <= log_lambda_tol {
converged = true;
let augmented = block.augmented_target(activation)?;
let final_loss = self.run_joint_fit_arrow_schur(
augmented.view(),
rho,
analytic_penalties,
inner_max_iter,
step_size,
ridge_ext_coord,
ridge_beta,
)?;
loss = Some(final_loss);
break;
}
}
if !lambda_identifiable {
self.set_behavior_block(block.clone())?;
}
let loss = loss.expect("max_sweeps ≥ 1 guarantees at least one fit");
Ok(TwoBlockRemlFitReport {
log_lambda_y: block.log_lambda_y,
sweeps,
converged,
lambda_identifiable,
loss,
log_lambda_trajectory: trajectory,
})
}
}