use super::*;
pub(crate) const MIN_CONDITIONAL_PRECISION: f64 = 1.0e-12;
pub(crate) use gam_problem::{LOG_STRENGTH_MAX, LOG_STRENGTH_MIN, checked_exp_log_strength};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PenaltyTier {
Beta,
Psi,
Rho,
}
#[derive(Debug, Clone)]
pub struct PsiSlice {
pub range: std::ops::Range<usize>,
pub latent_dim: Option<usize>,
}
impl PsiSlice {
#[must_use]
pub fn full(len: usize, latent_dim: Option<usize>) -> Self {
Self {
range: 0..len,
latent_dim,
}
}
pub fn len(&self) -> usize {
self.range.len()
}
pub fn is_empty(&self) -> bool {
self.range.is_empty()
}
}
pub fn resolve_learnable_weight(base_weight: f64, rho: f64) -> Result<f64, String> {
if base_weight == 0.0 {
return Err(
"a multiplicatively learnable weight requires a nonzero base; zero would make its rho coordinate structurally dead"
.to_string(),
);
}
if !(base_weight.is_finite() && rho.is_finite()) {
return Err(format!(
"learnable weight requires finite base and coordinate; got base_weight={base_weight}, rho={rho}"
));
}
let log_base = base_weight.abs().ln();
let (lower, upper) = (LOG_STRENGTH_MIN - log_base, LOG_STRENGTH_MAX - log_base);
if !(lower..=upper).contains(&rho) {
return Err(format!(
"learnable coordinate must be in [{lower}, {upper}] so its effective log strength is in [{LOG_STRENGTH_MIN}, {LOG_STRENGTH_MAX}]; got {rho}"
));
}
let log_strength = if rho == lower {
LOG_STRENGTH_MIN
} else if rho == upper {
LOG_STRENGTH_MAX
} else {
log_base + rho
};
Ok(checked_exp_log_strength(log_strength)
.map_err(|error| error.to_string())?
.copysign(base_weight))
}
pub fn learnable_weight_coordinate_domain(base_weight: f64) -> Result<Option<(f64, f64)>, String> {
if base_weight == 0.0 {
return Ok(None);
}
if !base_weight.is_finite() {
return Err(format!(
"learnable weight domain requires a finite base; got {base_weight}"
));
}
let log_base = base_weight.abs().ln();
Ok(Some((
LOG_STRENGTH_MIN - log_base,
LOG_STRENGTH_MAX - log_base,
)))
}
pub(crate) fn validated_learnable_weight(base_weight: f64, rho: f64) -> f64 {
resolve_learnable_weight(base_weight, rho)
.expect("analytic-penalty rho must be validated before strength evaluation")
}
pub(crate) fn validated_exp_log_strength(log_strength: f64) -> f64 {
checked_exp_log_strength(log_strength)
.expect("analytic-penalty rho must be validated before precision evaluation")
}
#[derive(Debug, Clone)]
pub struct ScalarWeightSchedule {
pub w_start: f64,
pub w_end: f64,
pub kind: ScheduleKind,
pub iter_count: usize,
}
impl ScalarWeightSchedule {
#[must_use = "build error must be handled"]
pub fn new(w_start: f64, w_end: f64, kind: ScheduleKind) -> Result<Self, String> {
let schedule = Self {
w_start,
w_end,
kind,
iter_count: 0,
};
schedule.validate()?;
Ok(schedule)
}
pub fn validate(&self) -> Result<(), String> {
if !(self.w_start.is_finite() && self.w_start >= 0.0) {
return Err(format!(
"ScalarWeightSchedule: w_start must be finite and non-negative; got {}",
self.w_start
));
}
if !(self.w_end.is_finite() && self.w_end >= 0.0) {
return Err(format!(
"ScalarWeightSchedule: w_end must be finite and non-negative; got {}",
self.w_end
));
}
match &self.kind {
ScheduleKind::Geometric { rate } => {
if !(rate.is_finite() && *rate > 0.0 && *rate < 1.0) {
return Err(format!(
"ScalarWeightSchedule::Geometric: rate must be in (0, 1); got {rate}"
));
}
}
ScheduleKind::Linear { steps } => {
if *steps == 0 {
return Err("ScalarWeightSchedule::Linear: steps must be positive".into());
}
}
ScheduleKind::ReciprocalIter => {}
}
Ok(())
}
pub fn current_weight(&self, iter: usize) -> f64 {
let delta = self.w_end - self.w_start;
let raw = match &self.kind {
ScheduleKind::Geometric { rate } => self.w_end - delta * rate.powf(iter as f64),
ScheduleKind::Linear { steps } => {
if iter >= *steps {
self.w_end
} else {
let frac = iter as f64 / *steps as f64;
self.w_start + frac * delta
}
}
ScheduleKind::ReciprocalIter => self.w_end - delta / (1.0 + iter as f64),
};
raw.clamp(self.w_start.min(self.w_end), self.w_start.max(self.w_end))
}
pub fn step(&mut self) -> f64 {
let weight = self.current_weight(self.iter_count);
self.iter_count += 1;
weight
}
}
pub trait AnalyticPenalty: Send + Sync {
fn tier(&self) -> PenaltyTier;
fn validate_rho(&self, rho: ArrayView1<'_, f64>) -> Result<(), String> {
if rho.len() != self.rho_count() {
return Err(format!(
"analytic penalty `{}` rho length {} != declared {}",
self.name(),
rho.len(),
self.rho_count()
));
}
for (axis, &value) in rho.iter().enumerate() {
checked_exp_log_strength(value).map_err(|error| {
format!(
"analytic penalty `{}` rho axis {axis}: {error}",
self.name()
)
})?;
}
Ok(())
}
fn rho_coordinate_domains(&self) -> Result<Vec<(f64, f64)>, String> {
Ok(vec![(LOG_STRENGTH_MIN, LOG_STRENGTH_MAX); self.rho_count()])
}
fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64;
fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64>;
fn hessian_diag(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> Option<Array1<f64>> {
assert!(
rho.iter().all(|value| value.is_finite()),
"analytic-penalty rho must be finite"
);
if target.is_empty() {
Some(Array1::zeros(0))
} else {
None
}
}
fn hvp(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
v: ArrayView1<'_, f64>,
) -> Array1<f64> {
let diag = self.hessian_diag(target, rho).unwrap_or_else(|| {
panic!(
"AnalyticPenalty::hvp default reached for `{}`, whose Hessian is \
not diagonal (hessian_diag returned None). Such a penalty must \
override `hvp` with its closed-form Hessian-vector product; the \
default never finite-differences.",
self.name()
)
});
assert_eq!(diag.len(), v.len(), "hvp dimension mismatch");
let mut out = Array1::<f64>::zeros(v.len());
for i in 0..v.len() {
out[i] = diag[i] * v[i];
}
out
}
fn psd_majorizer_diag(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> Option<Array1<f64>> {
self.hessian_diag(target, rho)
}
fn psd_majorizer_hvp(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
v: ArrayView1<'_, f64>,
) -> Array1<f64> {
if let Some(diag) = self.psd_majorizer_diag(target, rho) {
assert_eq!(diag.len(), v.len(), "psd_majorizer_hvp dimension mismatch");
let mut out = Array1::<f64>::zeros(v.len());
for i in 0..v.len() {
out[i] = diag[i] * v[i];
}
return out;
}
self.hvp(target, rho, v)
}
fn grad_rho(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64>;
fn rho_count(&self) -> usize;
fn name(&self) -> &str;
fn apply_schedule(&mut self, iter: usize) {
assert!(
iter < 1_000_000,
"apply_schedule received implausible outer iteration {iter}",
);
}
}
pub(crate) fn advance_scalar_weight(
weight: &mut f64,
schedule: &mut Option<ScalarWeightSchedule>,
iter: usize,
) {
if let Some(schedule) = schedule.as_mut() {
*weight = schedule.current_weight(iter);
schedule.iter_count = iter + 1;
}
}
macro_rules! impl_with_weight_schedule {
($field:ident) => {
#[must_use]
pub fn with_weight_schedule(mut self, schedule: ScalarWeightSchedule) -> Self {
self.$field = schedule.current_weight(schedule.iter_count);
self.weight_schedule = Some(schedule);
self
}
};
}
macro_rules! impl_scalar_apply_schedule {
($field:ident) => {
fn apply_schedule(&mut self, iter: usize) {
advance_scalar_weight(&mut self.$field, &mut self.weight_schedule, iter);
}
};
}
macro_rules! impl_learnable_weight_grad_rho {
() => {
fn grad_rho(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
if !self.learnable_weight {
return Array1::<f64>::zeros(0);
}
let mut out = Array1::<f64>::zeros(1);
out[self.rho_index] = self.value(target, rho);
out
}
};
}
macro_rules! impl_learnable_weight_rho_count {
() => {
fn rho_count(&self) -> usize {
usize::from(self.learnable_weight)
}
};
}
macro_rules! impl_learnable_weight_domain {
($field:ident) => {
fn validate_rho(&self, rho: ArrayView1<'_, f64>) -> Result<(), String> {
if rho.len() != self.rho_count() {
return Err(format!(
"analytic penalty `{}` rho length {} != declared {}",
self.name(),
rho.len(),
self.rho_count()
));
}
if self.learnable_weight {
resolve_learnable_weight(self.$field, rho[self.rho_index]).map_err(|error| {
format!("analytic penalty `{}`: {error}", self.name())
})?;
}
Ok(())
}
fn rho_coordinate_domains(&self) -> Result<Vec<(f64, f64)>, String> {
if !self.learnable_weight {
return Ok(Vec::new());
}
let domain = learnable_weight_coordinate_domain(self.$field)?.ok_or_else(|| {
format!(
"analytic penalty `{}` cannot expose a learnable coordinate with zero base weight",
self.name()
)
})?;
Ok(vec![domain])
}
};
}