use super::*;
use statrs::function::gamma::{digamma, ln_gamma};
#[derive(Debug, Clone)]
pub struct OrderedBetaBernoulliPenalty {
pub k_max: usize,
pub alpha: f64,
pub tau: f64,
pub temperature_schedule: Option<GumbelTemperatureSchedule>,
pub learnable_alpha: bool,
pub weight: f64,
pub weight_schedule: Option<ScalarWeightSchedule>,
pub fixed_columns: Option<Vec<bool>>,
pub row_weights: Option<std::sync::Arc<[f64]>>,
}
#[derive(Debug, Clone, Copy)]
struct MarginalColumnDerivatives {
mass: f64,
a: f64,
score: f64,
score_derivative: f64,
}
impl OrderedBetaBernoulliPenalty {
#[must_use]
pub fn new(k_max: usize, alpha: f64, tau: f64, learnable_alpha: bool) -> Self {
assert!(k_max > 0);
assert!(alpha.is_finite() && alpha > 0.0);
assert!(tau.is_finite() && tau > 0.0);
Self {
k_max,
alpha,
tau,
temperature_schedule: None,
learnable_alpha,
weight: 1.0,
weight_schedule: None,
fixed_columns: None,
row_weights: None,
}
}
#[must_use]
pub fn with_row_weights(mut self, weights: Option<&[f64]>) -> Self {
if let Some(weights) = weights {
assert!(
weights.iter().all(|w| w.is_finite() && *w >= 0.0),
"ordered Beta--Bernoulli row weights must be finite and nonnegative"
);
assert!(
weights.iter().any(|w| *w > 0.0),
"ordered Beta--Bernoulli row weights must contain positive mass"
);
}
self.row_weights = weights.map(|w| std::sync::Arc::from(w.to_vec()));
self
}
#[inline]
fn row_weight(&self, row: usize) -> f64 {
self.row_weights.as_ref().map_or(1.0, |w| w[row])
}
fn weighted_active_mass(&self, z: ArrayView1<'_, f64>) -> (Array1<f64>, f64) {
assert_eq!(
z.len() % self.k_max,
0,
"ordered Beta--Bernoulli target length must be divisible by k_max"
);
let n = z.len() / self.k_max;
if let Some(weights) = self.row_weights.as_ref() {
assert_eq!(
weights.len(),
n,
"ordered Beta--Bernoulli row-weight length must equal the row count"
);
}
let mut mass = Array1::<f64>::zeros(self.k_max);
let mut n_eff = 0.0;
for row in 0..n {
let w = self.row_weight(row);
n_eff += w;
let start = row * self.k_max;
for k in 0..self.k_max {
mass[k] += w * z[start + k];
}
}
(mass, n_eff)
}
#[inline]
fn column_is_fixed(&self, k: usize) -> bool {
self.fixed_columns
.as_ref()
.and_then(|m| m.get(k).copied())
.unwrap_or(false)
}
fn column_beta_shapes(&self, alpha: f64) -> Array1<f64> {
let log_ratio = -(1.0 / alpha).ln_1p();
let mut a_col = Array1::<f64>::zeros(self.k_max);
for k in 0..self.k_max {
let log_mu = ((k + 1) as f64) * log_ratio;
a_col[k] = (1.0 / (-log_mu).exp_m1().max(f64::MIN_POSITIVE)).max(f64::MIN_POSITIVE);
}
a_col
}
fn column_beta_shape_rho_deriv(&self, alpha: f64, a_col: ArrayView1<'_, f64>) -> Array1<f64> {
Array1::from_shape_fn(self.k_max, |k| {
let a = a_col[k];
((k + 1) as f64) * (a / (alpha + 1.0)) * (a + 1.0)
})
}
#[must_use]
pub fn with_temperature_schedule(mut self, schedule: GumbelTemperatureSchedule) -> Self {
self.tau = schedule.current_tau(schedule.iter_count);
self.temperature_schedule = Some(schedule);
self
}
impl_with_weight_schedule!(weight);
fn resolved_alpha(&self, rho: ArrayView1<'_, f64>) -> f64 {
if self.learnable_alpha {
validated_learnable_weight(self.alpha, rho[0])
} else {
self.alpha
}
}
fn concrete_logits(&self, target: ArrayView1<'_, f64>) -> Array1<f64> {
let tau = self.tau;
Array1::from_shape_fn(target.len(), |i| {
let x = target[i] / tau;
if x >= 0.0 {
1.0 / (1.0 + (-x).exp())
} else {
let ex = x.exp();
ex / (1.0 + ex)
}
})
}
fn marginal_columns(
&self,
z: ArrayView1<'_, f64>,
a_col: ArrayView1<'_, f64>,
) -> (Vec<MarginalColumnDerivatives>, f64) {
let (active_mass, n_eff) = self.weighted_active_mass(z);
let columns = (0..self.k_max)
.map(|k| {
let mass = active_mass[k].clamp(0.0, n_eff);
let a = a_col[k];
let active_arg = mass + a;
let inactive_arg = n_eff - mass + 1.0;
MarginalColumnDerivatives {
mass,
a,
score: -digamma(active_arg) + digamma(inactive_arg),
score_derivative: -trigamma(active_arg) - trigamma(inactive_arg),
}
})
.collect();
(columns, n_eff)
}
fn learnable_alpha_score_rho_derivs(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> (Array1<f64>, Array1<f64>) {
let mut d_score = Array1::<f64>::zeros(self.k_max);
let mut d_score_derivative = Array1::<f64>::zeros(self.k_max);
if !self.learnable_alpha {
return (d_score, d_score_derivative);
}
let alpha = self.resolved_alpha(rho);
let a_col = self.column_beta_shapes(alpha);
let da_col = self.column_beta_shape_rho_deriv(alpha, a_col.view());
let z = self.concrete_logits(target);
let (columns, _) = self.marginal_columns(z.view(), a_col.view());
for (k, column) in columns.iter().enumerate() {
if self.column_is_fixed(k) {
continue;
}
d_score[k] = -trigamma(column.mass + column.a) * da_col[k];
d_score_derivative[k] = -tetragamma(column.mass + column.a) * da_col[k];
}
(d_score, d_score_derivative)
}
#[must_use]
pub fn psd_majorizer_logit_third_channels(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> OrderedBetaBernoulliHessianDiagThirdChannels {
let alpha = self.resolved_alpha(rho);
let a_col = self.column_beta_shapes(alpha);
let z = self.concrete_logits(target);
let (columns, _) = self.marginal_columns(z.view(), a_col.view());
let n = z.len() / self.k_max;
let inv_tau = 1.0 / self.tau;
let inv_tau2 = inv_tau * inv_tau;
let mut z_jac = Array1::<f64>::zeros(target.len());
let mut local_logit_third = Array1::<f64>::zeros(target.len());
let mut m_channel = Array1::<f64>::zeros(target.len());
let mut diagonal_term = Array1::<f64>::zeros(target.len());
for row in 0..n {
let start = row * self.k_max;
let w_i = self.row_weight(row);
for k in 0..self.k_max {
if self.column_is_fixed(k) {
continue;
}
let column = columns[k];
let zk = z[start + k];
let jac = zk * (1.0 - zk) * inv_tau;
let u = w_i * jac;
let curvature = zk * (1.0 - zk) * (1.0 - 2.0 * zk) * inv_tau2;
let dz_curvature = (1.0 - 6.0 * zk + 6.0 * zk * zk) * inv_tau2;
let raw_diagonal_term = self.weight * column.score * w_i * curvature;
let diagonal_gate = f64::from(raw_diagonal_term > 0.0);
z_jac[start + k] = u;
diagonal_term[start + k] = raw_diagonal_term;
local_logit_third[start + k] =
self.weight * diagonal_gate * column.score * u * dz_curvature;
m_channel[start + k] =
self.weight * diagonal_gate * column.score_derivative * w_i * curvature;
}
}
let mut mass_hessian_coefficient = Array1::<f64>::zeros(self.k_max);
for k in 0..self.k_max {
if self.column_is_fixed(k) {
continue;
}
mass_hessian_coefficient[k] = self.weight * columns[k].score_derivative;
}
let mut mass_hessian_log_alpha_derivative = Array1::<f64>::zeros(self.k_max);
if self.learnable_alpha {
let (_, d_score_derivative) = self.learnable_alpha_score_rho_derivs(target, rho);
for k in 0..self.k_max {
if self.column_is_fixed(k) {
continue;
}
mass_hessian_log_alpha_derivative[k] = self.weight * d_score_derivative[k];
}
}
OrderedBetaBernoulliHessianDiagThirdChannels {
k_max: self.k_max,
z_jac,
local_logit_third,
m_channel,
mass_hessian_coefficient,
mass_hessian_log_alpha_derivative,
diagonal_term,
}
}
#[must_use]
pub fn logit_theta_adjoint_data(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> OrderedBetaBernoulliLogitAdjointData {
let alpha = self.resolved_alpha(rho);
let a_col = self.column_beta_shapes(alpha);
let z = self.concrete_logits(target);
let (columns, n_eff) = self.marginal_columns(z.view(), a_col.view());
let n = z.len() / self.k_max;
let mut score = vec![0.0_f64; self.k_max];
let mut score_derivative = vec![0.0_f64; self.k_max];
let mut score_second = vec![0.0_f64; self.k_max];
let mut column_fixed = vec![false; self.k_max];
for k in 0..self.k_max {
if self.column_is_fixed(k) {
column_fixed[k] = true;
continue;
}
let column = columns[k];
let active_arg = column.mass + column.a;
let inactive_arg = n_eff - column.mass + 1.0;
score[k] = column.score;
score_derivative[k] = column.score_derivative;
score_second[k] = -tetragamma(active_arg) + tetragamma(inactive_arg);
}
let row_weight = (0..n).map(|row| self.row_weight(row)).collect();
OrderedBetaBernoulliLogitAdjointData {
k_max: self.k_max,
n,
weight: self.weight,
tau: self.tau,
z: z.to_vec(),
row_weight,
score,
score_derivative,
score_second,
column_fixed,
}
}
#[must_use]
pub fn log_alpha_target_mixed_derivative(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> Array1<f64> {
let mut out = Array1::<f64>::zeros(target.len());
if !self.learnable_alpha {
return out;
}
let z = self.concrete_logits(target);
let n = z.len() / self.k_max;
let (d_score, _) = self.learnable_alpha_score_rho_derivs(target, rho);
for row in 0..n {
let start = row * self.k_max;
let w_i = self.row_weight(row);
for k in 0..self.k_max {
if self.column_is_fixed(k) {
continue;
}
let zk = z[start + k];
out[start + k] = self.weight * d_score[k] * w_i * zk * (1.0 - zk) / self.tau;
}
}
out
}
#[must_use]
pub fn hessian_diag_log_alpha_derivative(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> Array1<f64> {
let mut out = Array1::<f64>::zeros(target.len());
if !self.learnable_alpha {
return out;
}
let z = self.concrete_logits(target);
let n = z.len() / self.k_max;
let inv_tau = 1.0 / self.tau;
let inv_tau2 = inv_tau * inv_tau;
let (d_score, d_score_derivative) = self.learnable_alpha_score_rho_derivs(target, rho);
for row in 0..n {
let start = row * self.k_max;
let w_i = self.row_weight(row);
for k in 0..self.k_max {
if self.column_is_fixed(k) {
continue;
}
let zk = z[start + k];
let jac = zk * (1.0 - zk) * inv_tau;
let u = w_i * jac;
let curvature = zk * (1.0 - zk) * (1.0 - 2.0 * zk) * inv_tau2;
out[start + k] =
self.weight * (d_score_derivative[k] * u * u + d_score[k] * w_i * curvature);
}
}
out
}
}
#[derive(Debug, Clone)]
pub struct OrderedBetaBernoulliLogitAdjointData {
pub k_max: usize,
pub n: usize,
pub weight: f64,
pub tau: f64,
pub z: Vec<f64>,
pub row_weight: Vec<f64>,
pub score: Vec<f64>,
pub score_derivative: Vec<f64>,
pub score_second: Vec<f64>,
pub column_fixed: Vec<bool>,
}
#[derive(Debug, Clone)]
pub struct OrderedBetaBernoulliHessianDiagThirdChannels {
pub k_max: usize,
pub z_jac: Array1<f64>,
pub local_logit_third: Array1<f64>,
pub m_channel: Array1<f64>,
pub mass_hessian_coefficient: Array1<f64>,
pub mass_hessian_log_alpha_derivative: Array1<f64>,
pub diagonal_term: Array1<f64>,
}
impl AnalyticPenalty for OrderedBetaBernoulliPenalty {
fn tier(&self) -> PenaltyTier {
PenaltyTier::Psi
}
fn validate_rho(&self, rho: ArrayView1<'_, f64>) -> Result<(), String> {
if rho.len() != self.rho_count() {
return Err(format!(
"ordered Beta--Bernoulli rho length {} != declared {}",
rho.len(),
self.rho_count()
));
}
if self.learnable_alpha {
resolve_learnable_weight(self.alpha, rho[0])?;
}
Ok(())
}
fn rho_coordinate_domains(&self) -> Result<Vec<(f64, f64)>, String> {
if !self.learnable_alpha {
return Ok(Vec::new());
}
Ok(vec![
learnable_weight_coordinate_domain(self.alpha)?
.ok_or_else(|| "ordered Beta--Bernoulli alpha must be positive".to_string())?,
])
}
fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64 {
let alpha = self.resolved_alpha(rho);
let a_col = self.column_beta_shapes(alpha);
let z = self.concrete_logits(target);
let (columns, n_eff) = self.marginal_columns(z.view(), a_col.view());
let mut value = 0.0;
for (k, column) in columns.iter().enumerate() {
if self.column_is_fixed(k) {
continue;
}
value += -column.a.ln()
- ln_gamma(column.mass + column.a)
- ln_gamma(n_eff - column.mass + 1.0)
+ ln_gamma(n_eff + column.a + 1.0);
}
self.weight * value
}
fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
let alpha = self.resolved_alpha(rho);
let a_col = self.column_beta_shapes(alpha);
let z = self.concrete_logits(target);
let (columns, _) = self.marginal_columns(z.view(), a_col.view());
let n = z.len() / self.k_max;
let mut out = Array1::<f64>::zeros(target.len());
for row in 0..n {
let start = row * self.k_max;
let w_i = self.row_weight(row);
for k in 0..self.k_max {
if self.column_is_fixed(k) {
continue;
}
let zk = z[start + k];
out[start + k] = self.weight * columns[k].score * w_i * zk * (1.0 - zk) / self.tau;
}
}
out
}
fn hessian_diag(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> Option<Array1<f64>> {
let alpha = self.resolved_alpha(rho);
let a_col = self.column_beta_shapes(alpha);
let z = self.concrete_logits(target);
let (columns, _) = self.marginal_columns(z.view(), a_col.view());
let n = z.len() / self.k_max;
let inv_tau = 1.0 / self.tau;
let inv_tau2 = inv_tau * inv_tau;
let mut out = Array1::<f64>::zeros(target.len());
for row in 0..n {
let start = row * self.k_max;
let w_i = self.row_weight(row);
for k in 0..self.k_max {
if self.column_is_fixed(k) {
continue;
}
let zk = z[start + k];
let jac = zk * (1.0 - zk) * inv_tau;
let u = w_i * jac;
let curvature = zk * (1.0 - zk) * (1.0 - 2.0 * zk) * inv_tau2;
out[start + k] = self.weight
* (columns[k].score_derivative * u * u + columns[k].score * w_i * curvature);
}
}
Some(out)
}
fn hvp(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
v: ArrayView1<'_, f64>,
) -> Array1<f64> {
assert_eq!(
v.len(),
target.len(),
"OrderedBetaBernoulliPenalty::hvp dimension mismatch"
);
let alpha = self.resolved_alpha(rho);
let a_col = self.column_beta_shapes(alpha);
let z = self.concrete_logits(target);
let (columns, _) = self.marginal_columns(z.view(), a_col.view());
let n = z.len() / self.k_max;
let inv_tau = 1.0 / self.tau;
let inv_tau2 = inv_tau * inv_tau;
let mut contraction = Array1::<f64>::zeros(self.k_max);
for row in 0..n {
let start = row * self.k_max;
let w_i = self.row_weight(row);
for k in 0..self.k_max {
if self.column_is_fixed(k) {
continue;
}
let zk = z[start + k];
contraction[k] += w_i * zk * (1.0 - zk) * inv_tau * v[start + k];
}
}
let mut out = Array1::<f64>::zeros(target.len());
for row in 0..n {
let start = row * self.k_max;
let w_i = self.row_weight(row);
for k in 0..self.k_max {
if self.column_is_fixed(k) {
continue;
}
let zk = z[start + k];
let u = w_i * zk * (1.0 - zk) * inv_tau;
let curvature = zk * (1.0 - zk) * (1.0 - 2.0 * zk) * inv_tau2;
out[start + k] = self.weight
* (columns[k].score_derivative * u * contraction[k]
+ columns[k].score * w_i * curvature * v[start + k]);
}
}
out
}
fn grad_rho(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
if !self.learnable_alpha {
return Array1::zeros(0);
}
let alpha = self.resolved_alpha(rho);
let a_col = self.column_beta_shapes(alpha);
let da_col = self.column_beta_shape_rho_deriv(alpha, a_col.view());
let z = self.concrete_logits(target);
let (columns, n_eff) = self.marginal_columns(z.view(), a_col.view());
let mut gradient = 0.0;
for (k, column) in columns.iter().enumerate() {
if self.column_is_fixed(k) {
continue;
}
let d_l_da =
-1.0 / column.a - digamma(column.mass + column.a) + digamma(n_eff + column.a + 1.0);
gradient += d_l_da * da_col[k];
}
Array1::from_vec(vec![self.weight * gradient])
}
fn rho_count(&self) -> usize {
usize::from(self.learnable_alpha)
}
fn name(&self) -> &str {
"ordered_beta_bernoulli"
}
fn apply_schedule(&mut self, iter: usize) {
if let Some(schedule) = self.temperature_schedule.as_mut() {
self.tau = schedule.current_tau(iter);
schedule.iter_count = iter + 1;
}
advance_scalar_weight(&mut self.weight, &mut self.weight_schedule, iter);
}
}
use gam_math::special::{tetragamma, trigamma};