pub fn measure_jet_term_spec(
spec: &TermCollectionSpec,
term_idx: usize,
) -> Option<&crate::basis::MeasureJetBasisSpec> {
spec.smooth_terms
.get(term_idx)
.and_then(|term| match &term.basis {
SmoothBasisSpec::MeasureJet { spec, .. } => Some(spec),
_ => None,
})
}
pub fn measure_jet_enrolls_psi(mj: &crate::basis::MeasureJetBasisSpec) -> bool {
measure_jet_learns_length_scale(mj)
|| (mj.tau0 > 0.0 && crate::basis::measure_jet_multiscale_mode(mj))
}
pub fn measure_jet_learns_length_scale(mj: &crate::basis::MeasureJetBasisSpec) -> bool {
mj.learn_length_scale
}
pub fn freeze_measure_jet_length_scale_learning(spec: &mut TermCollectionSpec) -> usize {
let mut frozen = 0;
for term in spec.smooth_terms.iter_mut() {
if let SmoothBasisSpec::MeasureJet { spec: mj, .. } = &mut term.basis
&& mj.learn_length_scale
{
mj.learn_length_scale = false;
frozen += 1;
}
}
frozen
}
pub const MEASURE_JET_PSI_ALPHA_BOUNDS: (f64, f64) = (-1.0, 3.0);
pub const MEASURE_JET_PSI_LN_TAU_BOUNDS: (f64, f64) = (-18.420680743952367, 4.605170185988092);
pub fn measure_jet_penalty_psi_dim(mj: &crate::basis::MeasureJetBasisSpec) -> usize {
if crate::basis::measure_jet_multiscale_mode(mj) {
2
} else {
0
}
}
pub fn measure_jet_psi_dim(mj: &crate::basis::MeasureJetBasisSpec) -> usize {
usize::from(measure_jet_learns_length_scale(mj)) + measure_jet_penalty_psi_dim(mj)
}
pub fn measure_jet_psi_seed(mj: &crate::basis::MeasureJetBasisSpec) -> Vec<f64> {
let mut seed = Vec::with_capacity(measure_jet_psi_dim(mj));
if measure_jet_learns_length_scale(mj) {
let ell = if mj.length_scale > 0.0 {
mj.length_scale
} else {
1.0
};
seed.push(ell.ln());
}
if measure_jet_penalty_psi_dim(mj) > 0 {
let ln_tau = mj.tau0.max(f64::MIN_POSITIVE).ln();
seed.extend_from_slice(&[mj.alpha, ln_tau]);
}
seed
}
pub fn measure_jet_psi_bound_values(
data: ArrayView2<'_, f64>,
term: &SmoothBasisSpec,
upper: bool,
) -> Result<Vec<f64>, BasisError> {
let SmoothBasisSpec::MeasureJet {
feature_cols,
spec: mj,
input_scale,
} = term
else {
crate::bail_invalid_basis!(
"measure-jet ψ bounds requested for a {} term",
term.structural_kind()
);
};
let pick = |b: (f64, f64)| if upper { b.1 } else { b.0 };
let mut bounds = Vec::with_capacity(measure_jet_psi_dim(mj));
if measure_jet_learns_length_scale(mj) {
let mut columns = select_columns(data, feature_cols)?;
if let Some(scale) = input_scale {
scale.standardize(&mut columns);
}
let (mut lo, mut hi) = crate::basis::measure_jet_ln_range_window(columns.view(), mj)?;
if mj.length_scale > 0.0 {
let incumbent = mj.length_scale.ln();
if incumbent.is_finite() {
lo = lo.min(incumbent);
hi = hi.max(incumbent);
}
}
bounds.push(if upper { hi } else { lo });
}
if measure_jet_penalty_psi_dim(mj) > 0 {
bounds.push(pick(MEASURE_JET_PSI_ALPHA_BOUNDS));
bounds.push(pick(MEASURE_JET_PSI_LN_TAU_BOUNDS));
}
Ok(bounds)
}
pub fn apply_measure_jet_psi(
mj: &mut crate::basis::MeasureJetBasisSpec,
psi: &[f64],
) -> Result<bool, EstimationError> {
if psi.len() != measure_jet_psi_dim(mj) {
crate::bail_invalid_estim!(
"measure-jet ψ write-back dimension mismatch: got {} values for a {}-dial term",
psi.len(),
measure_jet_psi_dim(mj)
);
}
let mut changed = false;
let mut cursor = 0usize;
if measure_jet_learns_length_scale(mj) {
let next_ell = psi[cursor].exp();
cursor += 1;
if !(next_ell.is_finite() && next_ell > 0.0) {
crate::bail_invalid_estim!(
"measure-jet ψ write-back produced a non-finite/non-positive length_scale (ℓ={next_ell})"
);
}
if next_ell != mj.length_scale {
mj.length_scale = next_ell;
changed = true;
}
}
if measure_jet_penalty_psi_dim(mj) > 0 {
let next_alpha = psi[cursor];
let next_tau = psi[cursor + 1].exp();
if !(next_alpha.is_finite() && next_tau.is_finite() && next_tau > 0.0) {
crate::bail_invalid_estim!(
"measure-jet ψ write-back produced non-finite dials (alpha={next_alpha}, tau={next_tau})"
);
}
if next_alpha != mj.alpha {
mj.alpha = next_alpha;
changed = true;
}
if next_tau != mj.tau0 {
mj.tau0 = next_tau;
changed = true;
}
}
Ok(changed)
}
pub fn set_measure_jet_psi_dials(
spec: &mut TermCollectionSpec,
term_idx: usize,
psi: &[f64],
) -> Result<bool, EstimationError> {
let Some(term) = spec.smooth_terms.get_mut(term_idx) else {
crate::bail_invalid_estim!("measure-jet ψ write-back: term index {term_idx} out of range");
};
set_single_term_measure_jet_psi_dials(term, psi)
}
pub fn set_single_term_measure_jet_psi_dials(
term: &mut SmoothTermSpec,
psi: &[f64],
) -> Result<bool, EstimationError> {
let SmoothBasisSpec::MeasureJet { spec: mj, .. } = &mut term.basis else {
crate::bail_invalid_estim!("measure-jet ψ write-back targeted a non-measure-jet term");
};
apply_measure_jet_psi(mj, psi)
}