use super::*;
#[derive(Clone, Copy)]
struct CertifiedBernoulliRow {
geometry: WorkingBernoulliGeometry,
jet: MixtureInverseLinkJet,
}
#[inline]
fn certify_bernoulli_row(
inverse_link: &InverseLink,
row: usize,
eta: f64,
y: f64,
prior_weight: f64,
) -> Result<CertifiedBernoulliRow, EstimationError> {
if matches!(inverse_link, InverseLink::Standard(StandardLink::Logit)) {
let jet5 = logit_inverse_link_jet5(eta);
let geometry = bernoulli_logit_geometry_from_jet(row, eta, y, prior_weight, jet5)?;
Ok(CertifiedBernoulliRow {
geometry,
jet: MixtureInverseLinkJet {
mu: jet5.mu,
d1: jet5.d1,
d2: jet5.d2,
d3: jet5.d3,
},
})
} else {
let jet = standard_inverse_link_jet(inverse_link, eta)?;
let omm = crate::mixture_link::inverse_link_complement_for_inverse_link(
inverse_link,
eta,
jet.mu,
);
let geometry = bernoulli_geometry_from_jet(row, eta, y, prior_weight, jet, omm)?;
Ok(CertifiedBernoulliRow { geometry, jet })
}
}
fn certify_bernoulli_rows(
y: ArrayView1<f64>,
eta: &Array1<f64>,
inverse_link: &InverseLink,
priorweights: ArrayView1<f64>,
) -> Result<Vec<CertifiedBernoulliRow>, EstimationError> {
let rows: Vec<Result<CertifiedBernoulliRow, EstimationError>> = (0..eta.len())
.into_par_iter()
.map(|i| certify_bernoulli_row(inverse_link, i, eta[i], y[i], priorweights[i]))
.collect();
rows.into_iter().collect()
}
pub fn update_glmvectors(
y: ArrayView1<f64>,
eta: &Array1<f64>,
inverse_link: &InverseLink,
priorweights: ArrayView1<f64>,
mu: &mut Array1<f64>,
weights: &mut Array1<f64>,
z: &mut Array1<f64>,
derivatives: Option<WorkingDerivativeBuffersMut<'_>>,
) -> Result<(), EstimationError> {
let link = inverse_link.link_function();
match link {
LinkFunction::Logit
| LinkFunction::Probit
| LinkFunction::CLogLog
| LinkFunction::LogLog
| LinkFunction::Cauchit
| LinkFunction::Sas
| LinkFunction::BetaLogistic => {
let certified = certify_bernoulli_rows(y, eta, inverse_link, priorweights)?;
if let Some(mut derivs) = derivatives {
let WorkingSlices {
mu: mu_s,
weights: weights_s,
z: z_s,
} = working_slices(mu, weights, z);
let WorkingDerivSlices {
c: c_s,
d: d_s,
dmu: dmu_s,
d2: d2_s,
d3: d3_s,
} = working_deriv_slices(&mut derivs);
mu_s.par_iter_mut()
.zip(weights_s.par_iter_mut())
.zip(z_s.par_iter_mut())
.zip(c_s.par_iter_mut())
.zip(d_s.par_iter_mut())
.zip(dmu_s.par_iter_mut())
.zip(d2_s.par_iter_mut())
.zip(d3_s.par_iter_mut())
.zip(certified.par_iter())
.for_each(
|((((((((mu_o, w_o), z_o), c_o), d_o), dmu_o), d2_o), d3_o), row)| {
*mu_o = row.geometry.mu;
*w_o = row.geometry.weight;
*z_o = row.geometry.z;
*c_o = row.geometry.c;
*d_o = row.geometry.d;
*dmu_o = row.jet.d1;
*d2_o = row.jet.d2;
*d3_o = row.jet.d3;
},
);
} else {
let WorkingSlices {
mu: mu_s,
weights: weights_s,
z: z_s,
} = working_slices(mu, weights, z);
mu_s.par_iter_mut()
.zip(weights_s.par_iter_mut())
.zip(z_s.par_iter_mut())
.zip(certified.par_iter())
.for_each(|(((mu_o, w_o), z_o), row)| {
*mu_o = row.geometry.mu;
*w_o = row.geometry.weight;
*z_o = row.geometry.z;
});
}
Ok(())
}
LinkFunction::Identity => {
write_identityworking_state(y, eta, priorweights, mu, weights, z, derivatives)
}
LinkFunction::Log => {
write_poisson_log_working_state(y, eta, priorweights, mu, weights, z, derivatives)
}
}
}
#[inline]
pub fn update_glmvectors_by_family(
y: ArrayView1<f64>,
eta: &Array1<f64>,
likelihood: &GlmLikelihoodSpec,
priorweights: ArrayView1<f64>,
mu: &mut Array1<f64>,
weights: &mut Array1<f64>,
z: &mut Array1<f64>,
) -> Result<(), EstimationError> {
likelihood.irls_update(y, eta, priorweights, mu, weights, z, None, None)
}
pub(crate) fn integrated_inverse_link_from_family(
spec: &LikelihoodSpec,
mixture_link_state: Option<&MixtureLinkState>,
sas_link_state: Option<&SasLinkState>,
) -> Result<InverseLink, EstimationError> {
match (&spec.response, &spec.link) {
(ResponseFamily::Binomial, InverseLink::Standard(StandardLink::Logit))
| (ResponseFamily::Binomial, InverseLink::Standard(StandardLink::Probit))
| (ResponseFamily::Binomial, InverseLink::Standard(StandardLink::CLogLog)) => {
Ok(spec.link.clone())
}
(ResponseFamily::Binomial, InverseLink::Sas(_)) => {
let state = sas_link_state.ok_or_else(|| {
EstimationError::InvalidInput(
"Integrated BinomialSas update requires explicit SasLinkState".to_string(),
)
})?;
Ok(InverseLink::Sas(*state))
}
(ResponseFamily::Binomial, InverseLink::BetaLogistic(_)) => {
let state = sas_link_state.ok_or_else(|| {
EstimationError::InvalidInput(
"Integrated BinomialBetaLogistic update requires explicit SasLinkState"
.to_string(),
)
})?;
Ok(InverseLink::BetaLogistic(*state))
}
(ResponseFamily::Binomial, InverseLink::Mixture(_)) => {
let state = mixture_link_state.ok_or_else(|| {
EstimationError::InvalidInput(
"Integrated BinomialMixture update requires explicit MixtureLinkState"
.to_string(),
)
})?;
Ok(InverseLink::Mixture(state.clone()))
}
_ => Err(EstimationError::InvalidInput(format!(
"Integrated link-runtime update is not supported for likelihood (response={:?}, link={:?})",
spec.response, spec.link
))),
}
}
#[inline]
pub fn update_glmvectors_integrated_for_link(
quadctx: &crate::quadrature::QuadratureContext,
y: ArrayView1<f64>,
eta: &Array1<f64>,
se: ArrayView1<f64>,
inverse_link: &InverseLink,
priorweights: ArrayView1<f64>,
mu: &mut Array1<f64>,
weights: &mut Array1<f64>,
z: &mut Array1<f64>,
derivatives: Option<WorkingDerivativeBuffersMut<'_>>,
) -> Result<(), EstimationError> {
let link = inverse_link.link_function();
if !matches!(
inverse_link,
InverseLink::Standard(StandardLink::Logit)
| InverseLink::Standard(StandardLink::Probit)
| InverseLink::Standard(StandardLink::CLogLog)
| InverseLink::LatentCLogLog(_)
| InverseLink::Sas(_)
| InverseLink::BetaLogistic(_)
| InverseLink::Mixture(_)
) {
crate::bail_invalid_estim!(
"Integrated link-runtime update is not supported for inverse link {:?}",
inverse_link
);
}
let certified: Vec<Result<CertifiedBernoulliRow, EstimationError>> = (0..eta.len())
.into_par_iter()
.map(|i| {
let jet = if let InverseLink::LatentCLogLog(state) = inverse_link {
crate::quadrature::latent_cloglog_inverse_link_jet(
quadctx,
eta[i],
se[i].hypot(state.latent_sd),
)?
} else if matches!(inverse_link, InverseLink::Standard(StandardLink::Logit)) {
crate::quadrature::integrated_logit_inverse_link_jet_pirls(quadctx, eta[i], se[i])?
} else {
crate::quadrature::integrated_inverse_link_jetwith_state(
quadctx,
link,
eta[i],
se[i],
inverse_link.mixture_state(),
inverse_link.sas_state(),
)?
};
let jet = MixtureInverseLinkJet {
mu: jet.mean,
d1: jet.d1,
d2: jet.d2,
d3: jet.d3,
};
let omm = 1.0 - jet.mu;
let geometry = bernoulli_geometry_from_jet(i, eta[i], y[i], priorweights[i], jet, omm)?;
Ok(CertifiedBernoulliRow { geometry, jet })
})
.collect();
let certified: Vec<CertifiedBernoulliRow> = certified.into_iter().collect::<Result<_, _>>()?;
if let Some(mut derivs) = derivatives {
let WorkingSlices {
mu: mu_s,
weights: weights_s,
z: z_s,
} = working_slices(mu, weights, z);
let WorkingDerivSlices {
c: c_s,
d: d_s,
dmu: dmu_s,
d2: d2_s,
d3: d3_s,
} = working_deriv_slices(&mut derivs);
mu_s.par_iter_mut()
.zip(weights_s.par_iter_mut())
.zip(z_s.par_iter_mut())
.zip(c_s.par_iter_mut())
.zip(d_s.par_iter_mut())
.zip(dmu_s.par_iter_mut())
.zip(d2_s.par_iter_mut())
.zip(d3_s.par_iter_mut())
.zip(certified.par_iter())
.for_each(
|((((((((mu_o, w_o), z_o), c_o), d_o), dmu_o), d2_o), d3_o), row)| {
*mu_o = row.geometry.mu;
*w_o = row.geometry.weight;
*z_o = row.geometry.z;
*c_o = row.geometry.c;
*d_o = row.geometry.d;
*dmu_o = row.jet.d1;
*d2_o = row.jet.d2;
*d3_o = row.jet.d3;
},
);
} else {
let WorkingSlices {
mu: mu_s,
weights: weights_s,
z: z_s,
} = working_slices(mu, weights, z);
mu_s.par_iter_mut()
.zip(weights_s.par_iter_mut())
.zip(z_s.par_iter_mut())
.zip(certified.par_iter())
.for_each(|(((mu_o, w_o), z_o), row)| {
*mu_o = row.geometry.mu;
*w_o = row.geometry.weight;
*z_o = row.geometry.z;
});
}
Ok(())
}
#[inline]
pub fn update_glmvectors_integrated_by_family(
quadctx: &crate::quadrature::QuadratureContext,
y: ArrayView1<f64>,
eta: &Array1<f64>,
se: ArrayView1<f64>,
spec: &LikelihoodSpec,
priorweights: ArrayView1<f64>,
mu: &mut Array1<f64>,
weights: &mut Array1<f64>,
z: &mut Array1<f64>,
derivatives: Option<WorkingDerivativeBuffersMut<'_>>,
mixture_link_state: Option<&MixtureLinkState>,
sas_link_state: Option<&SasLinkState>,
) -> Result<(), EstimationError> {
let inverse_link =
integrated_inverse_link_from_family(spec, mixture_link_state, sas_link_state)?;
update_glmvectors_integrated_for_link(
quadctx,
y,
eta,
se,
&inverse_link,
priorweights,
mu,
weights,
z,
derivatives,
)
}
pub(crate) fn computeworkingweight_derivatives_from_eta(
likelihood: &GlmLikelihoodSpec,
inverse_link: &InverseLink,
eta: &Array1<f64>,
priorweights: ArrayView1<f64>,
) -> Result<
(
Array1<f64>,
Array1<f64>,
Array1<f64>,
Array1<f64>,
Array1<f64>,
),
EstimationError,
> {
let n = eta.len();
let mut c = Array1::<f64>::zeros(n);
let mut d = Array1::<f64>::zeros(n);
let mut dmu_deta = Array1::<f64>::zeros(n);
let mut d2mu_deta2 = Array1::<f64>::zeros(n);
let mut d3mu_deta3 = Array1::<f64>::zeros(n);
match &likelihood.spec.response {
ResponseFamily::Gaussian => {
dmu_deta.fill(1.0);
}
ResponseFamily::Poisson => {
log_link_working_state::write_log_link_eta_curvature(
&log_link_working_state::LogLinkRule {
weight: log_link_working_state::WorkingWeight::PoissonIdentity,
curvature: log_link_working_state::WorkingCurvature::Proportional {
c_ratio: 1.0,
d_ratio: 1.0,
},
},
eta,
priorweights,
WorkingDerivativeBuffersMut {
c: &mut c,
d: &mut d,
dmu_deta: &mut dmu_deta,
d2mu_deta2: &mut d2mu_deta2,
d3mu_deta3: &mut d3mu_deta3,
},
)?;
}
ResponseFamily::Tweedie { p } => {
let p = *p;
let phi = fixed_glm_dispersion(likelihood)?;
if !is_valid_tweedie_power(p) {
crate::bail_invalid_estim!(
"Tweedie variance power must be finite and strictly between 1 and 2; got {p}",
p = p
);
}
if !(phi.is_finite() && phi > 0.0) {
crate::bail_invalid_estim!(
"Tweedie dispersion phi must be finite and > 0; got {phi}",
phi = phi
);
}
let exponent = 2.0 - p;
log_link_working_state::write_log_link_eta_curvature(
&log_link_working_state::LogLinkRule {
weight: log_link_working_state::WorkingWeight::TweediePower { p, phi },
curvature: log_link_working_state::WorkingCurvature::Proportional {
c_ratio: exponent,
d_ratio: exponent * exponent,
},
},
eta,
priorweights,
WorkingDerivativeBuffersMut {
c: &mut c,
d: &mut d,
dmu_deta: &mut dmu_deta,
d2mu_deta2: &mut d2mu_deta2,
d3mu_deta3: &mut d3mu_deta3,
},
)?;
}
ResponseFamily::NegativeBinomial { theta, .. } => {
let theta = *theta;
if !valid_negbin_theta(theta) {
crate::bail_invalid_estim!(
"negative-binomial theta must be finite and > 0; got {theta}",
theta = theta
);
}
log_link_working_state::write_log_link_eta_curvature(
&log_link_working_state::LogLinkRule {
weight: log_link_working_state::WorkingWeight::NegativeBinomial { theta },
curvature: log_link_working_state::WorkingCurvature::NegativeBinomial { theta },
},
eta,
priorweights,
WorkingDerivativeBuffersMut {
c: &mut c,
d: &mut d,
dmu_deta: &mut dmu_deta,
d2mu_deta2: &mut d2mu_deta2,
d3mu_deta3: &mut d3mu_deta3,
},
)?;
}
ResponseFamily::Beta { phi } => {
let phi = *phi;
if !valid_beta_phi(phi) {
crate::bail_invalid_estim!("beta-regression phi must be finite and > 0; got {phi}");
}
let certified: Vec<Result<ExactBetaLogitRow, EstimationError>> = (0..eta.len())
.into_par_iter()
.map(|i| exact_beta_logit_row(i, eta[i], None, priorweights[i], phi))
.collect();
let certified: Vec<ExactBetaLogitRow> =
certified.into_iter().collect::<Result<_, _>>()?;
let c_s = c.as_slice_mut().expect("c must be contiguous");
let d_s = d.as_slice_mut().expect("d must be contiguous");
let dmu_s = dmu_deta
.as_slice_mut()
.expect("dmu_deta must be contiguous");
let d2_s = d2mu_deta2
.as_slice_mut()
.expect("d2mu_deta2 must be contiguous");
let d3_s = d3mu_deta3
.as_slice_mut()
.expect("d3mu_deta3 must be contiguous");
c_s.par_iter_mut()
.zip(d_s.par_iter_mut())
.zip(dmu_s.par_iter_mut())
.zip(d2_s.par_iter_mut())
.zip(d3_s.par_iter_mut())
.zip(certified.par_iter())
.for_each(|(((((c_o, d_o), dmu_o), d2_o), d3_o), row)| {
*c_o = row.c;
*d_o = row.d;
*dmu_o = row.dmu;
*d2_o = row.d2mu;
*d3_o = row.d3mu;
});
}
ResponseFamily::Gamma => {
log_link_working_state::write_log_link_eta_curvature(
&log_link_working_state::LogLinkRule {
weight: log_link_working_state::WorkingWeight::Constant { factor: 1.0 },
curvature: log_link_working_state::WorkingCurvature::Proportional {
c_ratio: 0.0,
d_ratio: 0.0,
},
},
eta,
priorweights,
WorkingDerivativeBuffersMut {
c: &mut c,
d: &mut d,
dmu_deta: &mut dmu_deta,
d2mu_deta2: &mut d2mu_deta2,
d3mu_deta3: &mut d3mu_deta3,
},
)?;
}
ResponseFamily::Binomial => {
let certified: Vec<Result<CertifiedBernoulliRow, EstimationError>> = (0..eta.len())
.into_par_iter()
.map(|i| {
let jet = if matches!(inverse_link, InverseLink::Standard(StandardLink::Logit))
{
let jet = logit_inverse_link_jet5(eta[i]);
MixtureInverseLinkJet {
mu: jet.mu,
d1: jet.d1,
d2: jet.d2,
d3: jet.d3,
}
} else {
standard_inverse_link_jet(inverse_link, eta[i])?
};
certify_bernoulli_row(inverse_link, i, eta[i], jet.mu, priorweights[i])
})
.collect();
let certified: Vec<CertifiedBernoulliRow> =
certified.into_iter().collect::<Result<_, _>>()?;
let c_s = c.as_slice_mut().expect("c must be contiguous");
let d_s = d.as_slice_mut().expect("d must be contiguous");
let dmu_s = dmu_deta
.as_slice_mut()
.expect("dmu_deta must be contiguous");
let d2_s = d2mu_deta2
.as_slice_mut()
.expect("d2mu_deta2 must be contiguous");
let d3_s = d3mu_deta3
.as_slice_mut()
.expect("d3mu_deta3 must be contiguous");
c_s.par_iter_mut()
.zip(d_s.par_iter_mut())
.zip(dmu_s.par_iter_mut())
.zip(d2_s.par_iter_mut())
.zip(d3_s.par_iter_mut())
.zip(certified.par_iter())
.for_each(|(((((c_o, d_o), dmu_o), d2_o), d3_o), row)| {
*c_o = row.geometry.c;
*d_o = row.geometry.d;
*dmu_o = row.jet.d1;
*d2_o = row.jet.d2;
*d3_o = row.jet.d3;
});
}
ResponseFamily::RoystonParmar => {
crate::bail_invalid_estim!(
"RoystonParmar is survival-specific and not a GLM IRLS family"
);
}
}
Ok((c, d, dmu_deta, d2mu_deta2, d3mu_deta3))
}