#[inline]
fn softmax_entropy_log_plus_one(probability: f64) -> f64 {
if probability > 0.0 {
probability.ln() + 1.0
} else {
0.0
}
}
#[inline]
pub(crate) fn softmax_majorizer_log_mean(a: &[f64]) -> f64 {
a.iter()
.map(|&a_i| a_i * softmax_entropy_log_plus_one(a_i))
.sum()
}
#[inline]
fn softmax_dense_entropy_hessian_entry(a: &[f64], kk: usize, jj: usize, m: f64, scale: f64) -> f64 {
let l_kk = softmax_entropy_log_plus_one(a[kk]);
let l_jj = softmax_entropy_log_plus_one(a[jj]);
let indicator = if kk == jj { 1.0 } else { 0.0 };
scale * a[kk] * (indicator * (m - l_kk - 1.0) + a[jj] * (l_kk + l_jj + 1.0 - 2.0 * m))
}
#[inline]
pub(crate) fn active_softmax_gershgorin_majorizer_entry(a: &[f64], kk: usize, m: f64, scale: f64) -> f64 {
let l_kk = softmax_entropy_log_plus_one(a[kk]);
let h_kk = scale * a[kk] * ((m - l_kk - 1.0) + a[kk] * (2.0 * l_kk + 1.0 - 2.0 * m));
let mut sum_sq = h_kk * h_kk;
for (jj, &a_jj) in a.iter().enumerate() {
if jj == kk {
continue;
}
let l_jj = softmax_entropy_log_plus_one(a_jj);
let h_kj = scale * a[kk] * (a_jj * (l_kk + l_jj + 1.0 - 2.0 * m));
sum_sq += h_kj * h_kj;
}
let eps0 =
gam_terms::analytic_penalties::SoftmaxAssignmentSparsityPenalty::soft_abs_temperature(
a.len(),
);
let eps_sq = eps0 * eps0 * sum_sq;
let mut acc = gam_terms::analytic_penalties::soft_abs_squared_scale(h_kk, eps_sq);
for (jj, &a_jj) in a.iter().enumerate() {
if jj == kk {
continue;
}
let l_jj = softmax_entropy_log_plus_one(a_jj);
let h_kj = scale * a[kk] * (a_jj * (l_kk + l_jj + 1.0 - 2.0 * m));
acc += gam_terms::analytic_penalties::soft_abs_squared_scale(h_kj, eps_sq);
}
acc
}
#[inline]
fn active_softmax_majorizer_logit_derivative_entry(
a: &[f64],
kk: usize,
w: usize,
m: f64,
scale: f64,
inv_tau: f64,
) -> f64 {
let a_w = a[w];
let da = |r: usize| a[r] * (if r == w { 1.0 } else { 0.0 } - a_w) * inv_tau;
let l = |r: usize| softmax_entropy_log_plus_one(a[r]);
let dl = |r: usize| if a[r] > 0.0 { da(r) / a[r] } else { 0.0 };
let dm: f64 = (0..a.len()).map(|r| da(r) * l(r) + a[r] * dl(r)).sum();
let l_kk = l(kk);
let da_kk = da(kk);
let dl_kk = dl(kk);
let hessian_entry = |jj: usize| -> (f64, f64) {
let indicator = if kk == jj { 1.0 } else { 0.0 };
let l_jj = l(jj);
let bracket = indicator * (m - l_kk - 1.0) + a[jj] * (l_kk + l_jj + 1.0 - 2.0 * m);
let dbracket = indicator * (dm - dl_kk)
+ da(jj) * (l_kk + l_jj + 1.0 - 2.0 * m)
+ a[jj] * (dl_kk + dl(jj) - 2.0 * dm);
(
scale * a[kk] * bracket,
scale * (da_kk * bracket + a[kk] * dbracket),
)
};
let (h_kk, dh_kk) = hessian_entry(kk);
let mut sum_sq = h_kk * h_kk;
let mut cross = h_kk * dh_kk;
for jj in 0..a.len() {
if jj == kk {
continue;
}
let (h_kj, dh_kj) = hessian_entry(jj);
sum_sq += h_kj * h_kj;
cross += h_kj * dh_kj;
}
let eps0 =
gam_terms::analytic_penalties::SoftmaxAssignmentSparsityPenalty::soft_abs_temperature(
a.len(),
);
let eps0_sq = eps0 * eps0;
let eps_sq = eps0_sq * sum_sq;
let mut acc = 0.0_f64;
let mut inv_envelope_sum = 0.0_f64;
let s_kk = gam_terms::analytic_penalties::soft_abs_squared_scale(h_kk, eps_sq);
if s_kk != 0.0 {
acc += (h_kk / s_kk) * dh_kk;
inv_envelope_sum += 1.0 / s_kk;
}
for jj in 0..a.len() {
if jj == kk {
continue;
}
let (h_kj, dh_kj) = hessian_entry(jj);
let s_kj = gam_terms::analytic_penalties::soft_abs_squared_scale(h_kj, eps_sq);
if s_kj == 0.0 {
continue;
}
acc += (h_kj / s_kj) * dh_kj;
inv_envelope_sum += 1.0 / s_kj;
}
acc + eps0_sq * cross * inv_envelope_sum
}
pub(crate) fn softmax_sparse_curvature_rho_derivative_block(
a: &[f64],
slot_atoms: &[usize],
m: f64,
scale: f64,
weight: f64,
operator: EvidenceOperator,
) -> Array2<f64> {
let slots = slot_atoms.len();
let mut out = Array2::<f64>::zeros((slots, slots));
match operator {
EvidenceOperator::Majorizer => {
for (slot, &atom) in slot_atoms.iter().enumerate() {
out[[slot, slot]] =
weight * active_softmax_gershgorin_majorizer_entry(a, atom, m, scale);
}
}
EvidenceOperator::ExactObservedInformation => {
for (row_slot, &row_atom) in slot_atoms.iter().enumerate() {
for (col_slot, &col_atom) in slot_atoms.iter().enumerate() {
out[[row_slot, col_slot]] = weight
* softmax_dense_entropy_hessian_entry(a, row_atom, col_atom, m, scale);
}
}
}
}
out
}