use super::*;
use gam_linalg::utils::SPECTRAL_DEFLATION_REL_FLOOR;
use gam_problem::{LOG_STRENGTH_MAX, LOG_STRENGTH_MIN, checked_exp_log_strength};
#[inline]
fn entropy_log_plus_one(p: f64) -> f64 {
if p > 0.0 { p.ln() + 1.0 } else { 0.0 }
}
#[inline]
#[must_use]
pub fn soft_abs_squared_scale(x: f64, eps_sq: f64) -> f64 {
(x * x + eps_sq).sqrt().max(x.abs())
}
#[derive(Debug, Clone, Copy)]
pub enum SparsityKind {
SmoothedL1 { eps: f64 },
Hoyer,
Log { delta: f64 },
}
#[derive(Debug, Clone)]
pub struct SparsityPenalty {
pub target_tier: PenaltyTier,
pub kind: SparsityKind,
pub weight: f64,
pub weight_schedule: Option<ScalarWeightSchedule>,
learnable_smoothing: bool,
}
#[derive(Debug, Clone)]
pub struct SoftmaxAssignmentSparsityPenalty {
pub k_atoms: usize,
pub temperature: f64,
pub weight: f64,
pub weight_schedule: Option<ScalarWeightSchedule>,
pub row_weights: Option<std::sync::Arc<[f64]>>,
}
impl SoftmaxAssignmentSparsityPenalty {
#[must_use]
pub fn new(k_atoms: usize, temperature: f64) -> Self {
assert!(k_atoms > 0);
assert!(temperature > 0.0);
Self {
k_atoms,
temperature,
weight: 1.0,
weight_schedule: None,
row_weights: None,
}
}
#[must_use]
pub fn with_row_weights(mut self, weights: Option<&[f64]>) -> Self {
self.row_weights = weights.map(|w| std::sync::Arc::from(w.to_vec()));
self
}
#[must_use]
pub fn row_weight(&self, row: usize) -> f64 {
self.row_weights.as_ref().map_or(1.0, |w| w[row])
}
impl_with_weight_schedule!(weight);
fn softmax_row(&self, row: &[f64]) -> Vec<f64> {
let inv_tau = 1.0 / self.temperature;
let mut max_logit = f64::NEG_INFINITY;
for (idx, &v) in row.iter().enumerate() {
assert!(
v.is_finite(),
"SoftmaxAssignmentSparsityPenalty: non-finite logit at atom {idx}: {v}"
);
max_logit = max_logit.max(v);
}
let mut out = vec![0.0; self.k_atoms];
let mut sum = 0.0;
for i in 0..self.k_atoms {
let v = ((row[i] - max_logit) * inv_tau).exp();
out[i] = v;
sum += v;
}
assert!(
sum.is_finite() && sum > 0.0,
"SoftmaxAssignmentSparsityPenalty: non-finite softmax normalizer"
);
for v in out.iter_mut() {
*v /= sum;
}
out
}
#[must_use]
pub fn soft_abs_temperature(k_atoms: usize) -> f64 {
SPECTRAL_DEFLATION_REL_FLOOR / (k_atoms as f64)
}
pub fn psd_majorizer_abs_row_sums(&self, row: &[f64], scale: f64) -> Vec<f64> {
let a = self.softmax_row(row);
let k = self.k_atoms;
let l: Vec<f64> = (0..k).map(|i| entropy_log_plus_one(a[i])).collect();
let m: f64 = (0..k).map(|i| a[i] * l[i]).sum();
let eps0 = Self::soft_abs_temperature(k);
let eps0_sq = eps0 * eps0;
let mut d = vec![0.0_f64; k];
for kk in 0..k {
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 in 0..k {
if jj == kk {
continue;
}
let h_kj = scale * a[kk] * (a[jj] * (l[kk] + l[jj] + 1.0 - 2.0 * m));
sum_sq += h_kj * h_kj;
}
let eps_sq = eps0_sq * sum_sq;
let mut acc = soft_abs_squared_scale(h_kk, eps_sq);
for jj in 0..k {
if jj == kk {
continue;
}
let h_kj = scale * a[kk] * (a[jj] * (l[kk] + l[jj] + 1.0 - 2.0 * m));
acc += soft_abs_squared_scale(h_kj, eps_sq);
}
d[kk] = acc;
}
d
}
#[must_use]
pub fn row_psd_majorizer(&self, row_logits: &[f64], scale: f64) -> Array2<f64> {
let k = self.k_atoms;
let d = self.psd_majorizer_abs_row_sums(row_logits, scale);
let mut out = Array2::<f64>::zeros((k, k));
for kk in 0..k {
out[[kk, kk]] = d[kk];
}
out
}
}
impl AnalyticPenalty for SoftmaxAssignmentSparsityPenalty {
fn tier(&self) -> PenaltyTier {
PenaltyTier::Psi
}
fn validate_rho(&self, rho: ArrayView1<'_, f64>) -> Result<(), String> {
if rho.len() != 1 {
return Err(format!(
"softmax assignment sparsity rho length {} != 1",
rho.len()
));
}
resolve_learnable_weight(self.weight, rho[0])?;
Ok(())
}
fn rho_coordinate_domains(&self) -> Result<Vec<(f64, f64)>, String> {
Ok(vec![
learnable_weight_coordinate_domain(self.weight)?
.ok_or_else(|| "softmax assignment sparsity has zero base weight".to_string())?,
])
}
fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64 {
let lambda = validated_learnable_weight(self.weight, rho[0]);
let n = target.len() / self.k_atoms;
let values: Vec<f64> = target.iter().copied().collect();
let mut acc = 0.0;
for row in 0..n {
let start = row * self.k_atoms;
let a = self.softmax_row(&values[start..start + self.k_atoms]);
let w_row = self.row_weight(row);
for v in a {
if v > 0.0 {
acc += -w_row * v * v.ln();
}
}
}
lambda * acc
}
fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
let lambda = validated_learnable_weight(self.weight, rho[0]);
let n = target.len() / self.k_atoms;
let values: Vec<f64> = target.iter().copied().collect();
let mut out = Array1::<f64>::zeros(target.len());
let inv_tau = 1.0 / self.temperature;
for row in 0..n {
let start = row * self.k_atoms;
let a = self.softmax_row(&values[start..start + self.k_atoms]);
let w_row = self.row_weight(row);
let mut d_h_da = vec![0.0; self.k_atoms];
let mut mean = 0.0;
for k in 0..self.k_atoms {
d_h_da[k] = -lambda * entropy_log_plus_one(a[k]);
mean += a[k] * d_h_da[k];
}
for k in 0..self.k_atoms {
out[start + k] = w_row * a[k] * (d_h_da[k] - mean) * inv_tau;
}
}
out
}
fn hessian_diag(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> Option<Array1<f64>> {
assert_eq!(rho.len(), 1, "softmax entropy expects one rho parameter");
assert!(
rho.iter().all(|value| value.is_finite()),
"softmax entropy rho must be finite"
);
assert_eq!(
target.len() % self.k_atoms,
0,
"softmax entropy target length must be divisible by k_atoms"
);
let lambda = validated_learnable_weight(self.weight, rho[0]);
let inv_tau = 1.0 / self.temperature;
let scale = lambda * inv_tau * inv_tau;
let n = target.len() / self.k_atoms;
let values: Vec<f64> = target.iter().copied().collect();
let mut out = Array1::<f64>::zeros(target.len());
for row in 0..n {
let start = row * self.k_atoms;
let a = self.softmax_row(&values[start..start + self.k_atoms]);
let w_row = self.row_weight(row);
let mut mean_log_plus_one = 0.0;
for k in 0..self.k_atoms {
mean_log_plus_one += a[k] * entropy_log_plus_one(a[k]);
}
for k in 0..self.k_atoms {
let log_plus_one = entropy_log_plus_one(a[k]);
let term = (1.0 - 2.0 * a[k]) * (mean_log_plus_one - log_plus_one) + a[k] - 1.0;
out[start + k] = w_row * scale * a[k] * term;
}
}
Some(out)
}
fn hvp(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
v: ArrayView1<'_, f64>,
) -> Array1<f64> {
let lambda = validated_learnable_weight(self.weight, rho[0]);
assert_eq!(target.len(), v.len(), "hvp dimension mismatch");
let n = target.len() / self.k_atoms;
let values: Vec<f64> = target.iter().copied().collect();
let mut out = Array1::<f64>::zeros(target.len());
let inv_tau = 1.0 / self.temperature;
let scale = lambda * inv_tau * inv_tau;
for row in 0..n {
let start = row * self.k_atoms;
let a = self.softmax_row(&values[start..start + self.k_atoms]);
let w_row = self.row_weight(row);
let mut mean_log_plus_one = 0.0;
let mut mean_v = 0.0;
for k in 0..self.k_atoms {
mean_log_plus_one += a[k] * entropy_log_plus_one(a[k]);
mean_v += a[k] * v[start + k];
}
let mut mean_centered_v_log_plus_one = 0.0;
for k in 0..self.k_atoms {
let centered_v = v[start + k] - mean_v;
mean_centered_v_log_plus_one += a[k] * centered_v * entropy_log_plus_one(a[k]);
}
for k in 0..self.k_atoms {
let log_plus_one = entropy_log_plus_one(a[k]);
let centered_v = v[start + k] - mean_v;
out[start + k] = w_row
* scale
* a[k]
* (centered_v * (mean_log_plus_one - log_plus_one - 1.0)
+ mean_centered_v_log_plus_one);
}
}
out
}
fn psd_majorizer_diag(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> Option<Array1<f64>> {
assert_eq!(rho.len(), 1, "softmax entropy expects one rho parameter");
assert_eq!(
target.len() % self.k_atoms,
0,
"softmax entropy target length must be divisible by k_atoms"
);
let lambda = validated_learnable_weight(self.weight, rho[0]);
let inv_tau = 1.0 / self.temperature;
let scale = lambda * inv_tau * inv_tau;
let n = target.len() / self.k_atoms;
let values: Vec<f64> = target.iter().copied().collect();
let mut out = Array1::<f64>::zeros(target.len());
for row in 0..n {
let start = row * self.k_atoms;
let w_row = self.row_weight(row);
let d = self.psd_majorizer_abs_row_sums(&values[start..start + self.k_atoms], scale);
for k in 0..self.k_atoms {
out[start + k] = w_row * d[k];
}
}
Some(out)
}
fn grad_rho(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
Array1::from_vec(vec![self.value(target, rho)])
}
fn rho_count(&self) -> usize {
1
}
fn name(&self) -> &str {
"softmax_assignment_sparsity"
}
impl_scalar_apply_schedule!(weight);
}
impl SparsityPenalty {
#[must_use = "build error must be handled"]
pub fn smoothed_l1(target_tier: PenaltyTier, eps: f64) -> Result<Self, String> {
if !(eps.is_finite() && eps > 0.0) {
return Err(format!(
"SparsityPenalty::smoothed_l1 requires eps > 0 \
(Hessian / gradient have a `1/sqrt(x² + eps²)` factor that needs eps > 0 \
for differentiability at x = 0); got eps = {eps}"
));
}
Ok(Self {
target_tier,
kind: SparsityKind::SmoothedL1 { eps },
weight: 1.0,
weight_schedule: None,
learnable_smoothing: false,
})
}
#[must_use = "build error must be handled"]
pub fn log(target_tier: PenaltyTier, delta: f64) -> Result<Self, String> {
if !(delta.is_finite() && delta > 0.0) {
return Err(format!(
"SparsityPenalty::log requires delta > 0 \
(the log-sparsifier is log(1 + x²/δ²), undefined at δ = 0); \
got delta = {delta}"
));
}
Ok(Self {
target_tier,
kind: SparsityKind::Log { delta },
weight: 1.0,
weight_schedule: None,
learnable_smoothing: false,
})
}
#[must_use]
pub fn hoyer(target_tier: PenaltyTier) -> Self {
Self {
target_tier,
kind: SparsityKind::Hoyer,
weight: 1.0,
weight_schedule: None,
learnable_smoothing: false,
}
}
impl_with_weight_schedule!(weight);
#[must_use = "invalid learnable-smoothing requests must be handled"]
pub fn with_learnable_smoothing(mut self) -> Result<Self, String> {
if matches!(self.kind, SparsityKind::Hoyer) {
return Err("Hoyer sparsity has no smoothing coordinate to learn".to_string());
}
self.learnable_smoothing = true;
Ok(self)
}
#[must_use]
pub fn learns_smoothing(&self) -> bool {
self.learnable_smoothing
}
fn resolved(&self, rho: ArrayView1<'_, f64>) -> (f64, f64) {
let strength = validated_learnable_weight(self.weight, rho[0]);
let smoothing = match (self.learnable_smoothing, self.kind) {
(true, _) => validated_exp_log_strength(rho[1]),
(false, SparsityKind::SmoothedL1 { eps }) => eps,
(false, SparsityKind::Log { delta }) => delta,
(false, SparsityKind::Hoyer) => 0.0,
};
(strength, smoothing)
}
}
impl AnalyticPenalty for SparsityPenalty {
fn tier(&self) -> PenaltyTier {
self.target_tier
}
fn validate_rho(&self, rho: ArrayView1<'_, f64>) -> Result<(), String> {
if rho.len() != self.rho_count() {
return Err(format!(
"sparsity rho length {} != declared {}",
rho.len(),
self.rho_count()
));
}
resolve_learnable_weight(self.weight, rho[0])?;
if self.learnable_smoothing {
checked_exp_log_strength(rho[1]).map_err(|error| error.to_string())?;
}
Ok(())
}
fn rho_coordinate_domains(&self) -> Result<Vec<(f64, f64)>, String> {
let mut domains = vec![(LOG_STRENGTH_MIN, LOG_STRENGTH_MAX); self.rho_count()];
domains[0] = learnable_weight_coordinate_domain(self.weight)?
.ok_or_else(|| "sparsity has zero base weight".to_string())?;
Ok(domains)
}
fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64 {
let (lam, smooth) = self.resolved(rho);
match self.kind {
SparsityKind::SmoothedL1 { .. } => {
let mut acc = 0.0;
for &x in target.iter() {
acc += (x * x + smooth * smooth).sqrt();
}
lam * acc
}
SparsityKind::Hoyer => {
let n = target.len() as f64;
assert!(n > 1.0, "Hoyer requires n > 1");
let l1: f64 = target.iter().map(|x| x.abs()).sum();
let l2: f64 = target.iter().map(|x| x * x).sum::<f64>().sqrt();
if l2 == 0.0 {
return 0.0;
}
let h = (l1 / l2 - 1.0) / (n.sqrt() - 1.0);
lam * h
}
SparsityKind::Log { .. } => {
let mut acc = 0.0;
let d2 = smooth * smooth;
for &x in target.iter() {
acc += (1.0 + x * x / d2).ln();
}
lam * acc
}
}
}
fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
let (lam, smooth) = self.resolved(rho);
let mut g = Array1::<f64>::zeros(target.len());
match self.kind {
SparsityKind::SmoothedL1 { .. } => {
let eps2 = smooth * smooth;
for (i, &x) in target.iter().enumerate() {
g[i] = lam * x / (x * x + eps2).sqrt();
}
}
SparsityKind::Hoyer => {
let n = target.len() as f64;
assert!(n > 1.0, "Hoyer requires n > 1");
let l1: f64 = target.iter().map(|x| x.abs()).sum();
let l2: f64 = target.iter().map(|x| x * x).sum::<f64>().sqrt();
if l2 == 0.0 {
return g;
}
let denom = n.sqrt() - 1.0;
let a = lam / denom;
let inv_l2 = 1.0 / l2;
let inv_l2_cubed = inv_l2 * inv_l2 * inv_l2;
for (i, &x) in target.iter().enumerate() {
let sgn = if x > 0.0 {
1.0
} else if x < 0.0 {
-1.0
} else {
0.0
};
g[i] = a * (sgn * inv_l2 - l1 * x * inv_l2_cubed);
}
}
SparsityKind::Log { .. } => {
let d2 = smooth * smooth;
for (i, &x) in target.iter().enumerate() {
g[i] = lam * 2.0 * x / (d2 + x * x);
}
}
}
g
}
fn hessian_diag(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> Option<Array1<f64>> {
let (lam, smooth) = self.resolved(rho);
match self.kind {
SparsityKind::SmoothedL1 { .. } => {
let mut d = Array1::<f64>::zeros(target.len());
let eps2 = smooth * smooth;
for (i, &x) in target.iter().enumerate() {
let r = (x * x + eps2).sqrt();
d[i] = lam * eps2 / (r * r * r);
}
Some(d)
}
SparsityKind::Log { .. } => {
let mut d = Array1::<f64>::zeros(target.len());
let d2 = smooth * smooth;
for (i, &x) in target.iter().enumerate() {
let denom = d2 + x * x;
d[i] = lam * 2.0 * (d2 - x * x) / (denom * denom);
}
Some(d)
}
SparsityKind::Hoyer => None,
}
}
fn hvp(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
v: ArrayView1<'_, f64>,
) -> Array1<f64> {
let (lam, smooth) = self.resolved(rho);
let n_target = target.len();
assert_eq!(v.len(), n_target, "hvp dimension mismatch");
match self.kind {
SparsityKind::SmoothedL1 { .. } => {
let mut out = Array1::<f64>::zeros(n_target);
let eps2 = smooth * smooth;
for (i, &x) in target.iter().enumerate() {
let r = (x * x + eps2).sqrt();
out[i] = lam * eps2 / (r * r * r) * v[i];
}
out
}
SparsityKind::Log { .. } => {
let mut out = Array1::<f64>::zeros(n_target);
let d2 = smooth * smooth;
for (i, &x) in target.iter().enumerate() {
let denom = d2 + x * x;
out[i] = lam * 2.0 * (d2 - x * x) / (denom * denom) * v[i];
}
out
}
SparsityKind::Hoyer => {
let n = n_target as f64;
assert!(n > 1.0, "Hoyer requires n > 1");
let l1: f64 = target.iter().map(|x| x.abs()).sum();
let l2: f64 = target.iter().map(|x| x * x).sum::<f64>().sqrt();
let mut out = Array1::<f64>::zeros(n_target);
if l2 == 0.0 {
return out;
}
let a = lam / (n.sqrt() - 1.0);
let inv_l2_cubed = 1.0 / (l2 * l2 * l2);
let inv_l2_5 = inv_l2_cubed / (l2 * l2);
let mut x_dot_v = 0.0;
let mut s_dot_v = 0.0;
for i in 0..n_target {
let xi = target[i];
let si = if xi > 0.0 {
1.0
} else if xi < 0.0 {
-1.0
} else {
0.0
};
x_dot_v += xi * v[i];
s_dot_v += si * v[i];
}
for i in 0..n_target {
let xi = target[i];
let si = if xi > 0.0 {
1.0
} else if xi < 0.0 {
-1.0
} else {
0.0
};
out[i] = a
* (-si * x_dot_v * inv_l2_cubed
- xi * s_dot_v * inv_l2_cubed
- l1 * v[i] * inv_l2_cubed
+ 3.0 * l1 * xi * x_dot_v * inv_l2_5);
}
out
}
}
}
fn psd_majorizer_diag(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> Option<Array1<f64>> {
let (lam, smooth) = self.resolved(rho);
match self.kind {
SparsityKind::SmoothedL1 { .. } => self.hessian_diag(target, rho),
SparsityKind::Log { .. } => {
let mut d = Array1::<f64>::zeros(target.len());
let d2 = smooth * smooth;
for (i, &x) in target.iter().enumerate() {
d[i] = lam * 2.0 / (d2 + x * x);
}
Some(d)
}
SparsityKind::Hoyer => None,
}
}
fn grad_rho(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
let n_rho = self.rho_count();
let mut out = Array1::<f64>::zeros(n_rho);
let p_val = self.value(target, rho);
out[0] = p_val;
if self.learnable_smoothing {
let (lam, smooth) = self.resolved(rho);
let mut dp_deps = 0.0;
match self.kind {
SparsityKind::SmoothedL1 { .. } => {
for &x in target.iter() {
dp_deps += smooth / (x * x + smooth * smooth).sqrt();
}
dp_deps *= lam;
}
SparsityKind::Log { .. } => {
let d2 = smooth * smooth;
for &x in target.iter() {
dp_deps += -2.0 * x * x / (smooth * (d2 + x * x));
}
dp_deps *= lam;
}
SparsityKind::Hoyer => {}
}
out[1] = smooth * dp_deps;
}
out
}
fn rho_count(&self) -> usize {
1 + usize::from(self.learnable_smoothing)
}
fn name(&self) -> &str {
"sparsity"
}
impl_scalar_apply_schedule!(weight);
}
#[derive(Debug, Clone)]
pub struct TopKActivationPenalty {
pub target: PsiSlice,
pub k: usize,
pub latent_dim: usize,
pub weight: f64,
pub weight_schedule: Option<ScalarWeightSchedule>,
}
impl TopKActivationPenalty {
#[must_use = "build error must be handled"]
pub fn new(target: PsiSlice, k: usize, weight: f64) -> Result<Self, String> {
let latent_dim = target
.latent_dim
.ok_or_else(|| "TopKActivationPenalty::new requires target.latent_dim".to_string())?;
if latent_dim == 0 {
return Err("TopKActivationPenalty::new requires latent_dim > 0".to_string());
}
if k == 0 || k > latent_dim {
return Err(format!(
"TopKActivationPenalty::new requires 0 < k <= latent_dim; got k={k}, latent_dim={latent_dim}"
));
}
if !(weight.is_finite() && weight > 0.0) {
return Err(format!(
"TopKActivationPenalty::new requires finite weight > 0, got {weight}"
));
}
Ok(Self {
target,
k,
latent_dim,
weight,
weight_schedule: None,
})
}
impl_with_weight_schedule!(weight);
fn topk_mask_row(&self, target: ArrayView1<'_, f64>, row: usize, mask: &mut [bool]) {
mask.fill(false);
let d = self.latent_dim;
let base = row * d;
let mut order = (0..d).collect::<Vec<_>>();
order.sort_by(|&a, &b| {
target[base + b]
.abs()
.total_cmp(&target[base + a].abs())
.then_with(|| a.cmp(&b))
});
for &axis in order.iter().take(self.k) {
mask[axis] = true;
}
}
}
impl AnalyticPenalty for TopKActivationPenalty {
fn tier(&self) -> PenaltyTier {
PenaltyTier::Psi
}
fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64 {
assert_eq!(rho.len(), 0, "TopKActivationPenalty has no rho parameters");
let d = self.latent_dim;
let n_obs = target.len() / d;
let mut mask = vec![false; d];
let mut acc = 0.0;
for row in 0..n_obs {
self.topk_mask_row(target, row, &mut mask);
let base = row * d;
for axis in 0..d {
if mask[axis] {
let v = target[base + axis];
acc += 0.5 * self.weight * v * v;
}
}
}
acc
}
fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
assert_eq!(rho.len(), 0, "TopKActivationPenalty has no rho parameters");
let d = self.latent_dim;
let n_obs = target.len() / d;
let mut mask = vec![false; d];
let mut grad = Array1::<f64>::zeros(target.len());
for row in 0..n_obs {
self.topk_mask_row(target, row, &mut mask);
let base = row * d;
for axis in 0..d {
if mask[axis] {
grad[base + axis] = self.weight * target[base + axis];
}
}
}
grad
}
fn hessian_diag(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> Option<Array1<f64>> {
assert_eq!(rho.len(), 0, "TopKActivationPenalty has no rho parameters");
let d = self.latent_dim;
let n_obs = target.len() / d;
let mut mask = vec![false; d];
let mut diag = Array1::<f64>::zeros(target.len());
for row in 0..n_obs {
self.topk_mask_row(target, row, &mut mask);
let base = row * d;
for axis in 0..d {
if mask[axis] {
diag[base + axis] = self.weight;
}
}
}
Some(diag)
}
fn grad_rho(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
assert_eq!(rho.len(), 0, "TopKActivationPenalty has no rho parameters");
assert_eq!(
target.len() % self.latent_dim,
0,
"TopKActivationPenalty target length must be a multiple of latent_dim"
);
Array1::<f64>::zeros(0)
}
fn rho_count(&self) -> usize {
0
}
fn name(&self) -> &str {
"topk_activation"
}
impl_scalar_apply_schedule!(weight);
}
#[derive(Debug, Clone)]
pub struct SmoothThresholdPenalty {
pub target: PsiSlice,
pub latent_dim: usize,
pub thresholds: Array1<f64>,
pub weight: f64,
pub smoothing_eps: f64,
pub weight_schedule: Option<ScalarWeightSchedule>,
}
impl SmoothThresholdPenalty {
#[must_use = "build error must be handled"]
pub fn new(
target: PsiSlice,
thresholds: Array1<f64>,
weight: f64,
smoothing_eps: f64,
) -> Result<Self, String> {
let latent_dim = target
.latent_dim
.ok_or_else(|| "SmoothThresholdPenalty::new requires target.latent_dim".to_string())?;
if latent_dim == 0 {
return Err("SmoothThresholdPenalty::new requires latent_dim > 0".to_string());
}
if thresholds.len() != latent_dim {
return Err(format!(
"SmoothThresholdPenalty::new thresholds length {} does not match latent_dim {latent_dim}",
thresholds.len()
));
}
for (idx, &tau) in thresholds.iter().enumerate() {
if !(tau.is_finite() && tau > 0.0) {
return Err(format!(
"SmoothThresholdPenalty::new thresholds[{idx}] must be finite and > 0, got {tau}"
));
}
}
if !(weight.is_finite() && weight > 0.0) {
return Err(format!(
"SmoothThresholdPenalty::new requires finite weight > 0, got {weight}"
));
}
if !(smoothing_eps.is_finite() && smoothing_eps > 0.0) {
return Err(format!(
"SmoothThresholdPenalty::new requires finite smoothing_eps > 0, got {smoothing_eps}"
));
}
Ok(Self {
target,
latent_dim,
thresholds,
weight,
smoothing_eps,
weight_schedule: None,
})
}
impl_with_weight_schedule!(weight);
fn threshold(&self, axis: usize, rho: ArrayView1<'_, f64>) -> f64 {
validated_learnable_weight(self.thresholds[axis], rho[axis])
}
pub(crate) fn sigmoid_gate(&self, x: f64) -> f64 {
if x >= 0.0 {
1.0 / (1.0 + (-x).exp())
} else {
let ex = x.exp();
ex / (1.0 + ex)
}
}
fn true_hessian_diag_entry(&self, tau: f64, gate: f64) -> f64 {
self.weight * tau * gate * (1.0 - gate) * (1.0 - 2.0 * gate)
/ (self.smoothing_eps * self.smoothing_eps)
}
fn psd_hessian_diag_entry(&self, tau: f64, gate: f64) -> f64 {
let slope = gate * (1.0 - gate);
let reweighted_l2 = slope * slope;
let abs_exact = slope * (1.0 - 2.0 * gate).abs();
self.weight * tau * reweighted_l2.max(abs_exact) / (self.smoothing_eps * self.smoothing_eps)
}
}
#[must_use]
pub fn smooth_threshold_gate_value_grad(z: f64, tau: f64, smoothing_eps: f64) -> (f64, f64, f64) {
let g = gam_linalg::utils::stable_logistic((z - tau) / smoothing_eps);
let value = z * g;
let slope = z * g * (1.0 - g) / smoothing_eps;
let dphi_dz = g + slope;
let dphi_dtau = -slope;
(value, dphi_dz, dphi_dtau)
}
impl AnalyticPenalty for SmoothThresholdPenalty {
fn tier(&self) -> PenaltyTier {
PenaltyTier::Psi
}
fn validate_rho(&self, rho: ArrayView1<'_, f64>) -> Result<(), String> {
if rho.len() != self.latent_dim {
return Err(format!(
"smooth-threshold rho length {} != latent dimension {}",
rho.len(),
self.latent_dim
));
}
for axis in 0..self.latent_dim {
resolve_learnable_weight(self.thresholds[axis], rho[axis])?;
}
Ok(())
}
fn rho_coordinate_domains(&self) -> Result<Vec<(f64, f64)>, String> {
self.thresholds
.iter()
.map(|&threshold| {
learnable_weight_coordinate_domain(threshold)?.ok_or_else(|| {
"smooth-threshold cannot learn a zero threshold multiplicatively".to_string()
})
})
.collect()
}
fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64 {
let d = self.latent_dim;
let n_obs = target.len() / d;
let mut acc = 0.0;
for row in 0..n_obs {
let base = row * d;
for axis in 0..d {
let tau = self.threshold(axis, rho);
let gate = self.sigmoid_gate((target[base + axis] - tau) / self.smoothing_eps);
acc += self.weight * tau * gate;
}
}
acc
}
fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
let d = self.latent_dim;
let n_obs = target.len() / d;
let mut grad = Array1::<f64>::zeros(target.len());
for row in 0..n_obs {
let base = row * d;
for axis in 0..d {
let tau = self.threshold(axis, rho);
let gate = self.sigmoid_gate((target[base + axis] - tau) / self.smoothing_eps);
grad[base + axis] = self.weight * tau * gate * (1.0 - gate) / self.smoothing_eps;
}
}
grad
}
fn hessian_diag(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> Option<Array1<f64>> {
let d = self.latent_dim;
let n_obs = target.len() / d;
let mut diag = Array1::<f64>::zeros(target.len());
for row in 0..n_obs {
let base = row * d;
for axis in 0..d {
let tau = self.threshold(axis, rho);
let gate = self.sigmoid_gate((target[base + axis] - tau) / self.smoothing_eps);
diag[base + axis] = self.true_hessian_diag_entry(tau, gate);
}
}
Some(diag)
}
fn hvp(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
v: ArrayView1<'_, f64>,
) -> Array1<f64> {
assert_eq!(target.len(), v.len(), "hvp dimension mismatch");
let d = self.latent_dim;
let n_obs = target.len() / d;
let mut out = Array1::<f64>::zeros(target.len());
for row in 0..n_obs {
let base = row * d;
for axis in 0..d {
let tau = self.threshold(axis, rho);
let gate = self.sigmoid_gate((target[base + axis] - tau) / self.smoothing_eps);
out[base + axis] = self.true_hessian_diag_entry(tau, gate) * v[base + axis];
}
}
out
}
fn psd_majorizer_diag(
&self,
target: ArrayView1<'_, f64>,
rho: ArrayView1<'_, f64>,
) -> Option<Array1<f64>> {
let d = self.latent_dim;
let n_obs = target.len() / d;
let mut diag = Array1::<f64>::zeros(target.len());
for row in 0..n_obs {
let base = row * d;
for axis in 0..d {
let tau = self.threshold(axis, rho);
let gate = self.sigmoid_gate((target[base + axis] - tau) / self.smoothing_eps);
diag[base + axis] = self.psd_hessian_diag_entry(tau, gate);
}
}
Some(diag)
}
fn grad_rho(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
let d = self.latent_dim;
let n_obs = target.len() / d;
let mut out = Array1::<f64>::zeros(d);
for axis in 0..d {
let tau = self.threshold(axis, rho);
let mut g_tau = 0.0;
for row in 0..n_obs {
let x = target[row * d + axis];
let gate = self.sigmoid_gate((x - tau) / self.smoothing_eps);
g_tau += gate - tau * gate * (1.0 - gate) / self.smoothing_eps;
}
out[axis] = self.weight * tau * g_tau;
}
out
}
fn rho_count(&self) -> usize {
self.latent_dim
}
fn name(&self) -> &str {
"smooth_threshold"
}
impl_scalar_apply_schedule!(weight);
}
#[cfg(test)]
mod soft_abs_gershgorin_2339_tests {
use super::*;
use approx::assert_abs_diff_eq;
use gam_linalg::utils::splitmix64;
#[test]
fn soft_abs_envelope_dominates_absolute_value_2339() {
let mut state = 0x2339_0001_u64;
let magnitudes = [0.0_f64, 1e-300, 1e-30, 1e-12, 1e-8, 1e-3, 1.0, 7.5, 1e6];
for &eps in &[0.0_f64, 1e-16, 1e-12, 1e-8, 1e-3, 1.0] {
let eps_sq = eps * eps;
for &mag in &magnitudes {
for sign in [1.0_f64, -1.0] {
let x = sign * mag;
let env = soft_abs_squared_scale(x, eps_sq);
assert!(
env >= x.abs(),
"soft-abs must MAJORIZE |x| (#2339): σ({x}, ε²={eps_sq}) = {env} \
< |x| = {}",
x.abs()
);
assert!(
env <= x.abs() + eps + f64::EPSILON * (1.0 + x.abs()),
"soft-abs must exceed |x| by at most ε (#2339): \
σ({x}, ε²={eps_sq}) − |x| = {} > ε = {eps}",
env - x.abs()
);
}
}
assert_abs_diff_eq!(
soft_abs_squared_scale(0.0, eps_sq),
eps,
epsilon = 1e-15 * (1.0 + eps)
);
}
for _ in 0..4096 {
let x = (splitmix64(&mut state) >> 11) as f64 / ((1_u64 << 53) as f64) * 20.0 - 10.0;
let eps = (splitmix64(&mut state) >> 11) as f64 / ((1_u64 << 53) as f64) * 2.0;
let env = soft_abs_squared_scale(x, eps * eps);
assert!(
env >= x.abs() && env <= x.abs() + eps + f64::EPSILON * (1.0 + x.abs()),
"soft-abs envelope violated at x={x}, ε={eps}: got {env}"
);
}
let eps = 1e-3_f64;
for &x in &[1e-4_f64, 1e-3, 5e-3, 1e-2] {
let dipping = x * (x / eps).tanh();
assert!(
dipping < x.abs(),
"x·tanh(x/ε) is a MINORANT of |x| and must fail the majorization \
predicate (#2339): at x={x} it gives {dipping} ≥ |x|"
);
}
}
}
#[cfg(test)]
mod row_weighted_prior_991_tests {
use super::AnalyticPenalty;
use super::*;
use approx::assert_abs_diff_eq;
use ndarray::{Array1, s};
fn logits(n: usize, k: usize) -> Array1<f64> {
let mut v = Array1::<f64>::zeros(n * k);
for r in 0..n {
for a in 0..k {
v[r * k + a] =
0.35 * (r as f64) - 0.6 * (a as f64) + 0.11 * ((r * k + a) as f64).sin();
}
}
v
}
#[test]
fn weighted_value_is_per_row_reweight_of_unweighted() {
let (n, k) = (5usize, 3usize);
let temperature = 0.7_f64;
let rho = Array1::from_vec(vec![0.2_f64]);
let target = logits(n, k);
let base = SoftmaxAssignmentSparsityPenalty::new(k, temperature);
let mut per_row = vec![0.0_f64; n];
for r in 0..n {
let row = target.slice(s![r * k..r * k + k]).to_owned();
per_row[r] = base.value(row.view(), rho.view());
}
let unweighted: f64 = per_row.iter().sum();
assert_abs_diff_eq!(
base.value(target.view(), rho.view()),
unweighted,
epsilon = 1e-12
);
let w = vec![1.7_f64, 0.3, 1.1, 0.5, 1.4]; let weighted = base.clone().with_row_weights(Some(&w));
let expect: f64 = (0..n).map(|r| w[r] * per_row[r]).sum();
assert_abs_diff_eq!(
weighted.value(target.view(), rho.view()),
expect,
epsilon = 1e-12
);
assert_abs_diff_eq!(
weighted.value(target.view(), rho.view()),
(0..n).map(|r| w[r] * per_row[r]).sum::<f64>(),
epsilon = 1e-12
);
}
#[test]
fn weighted_value_grad_are_fd_consistent() {
let (n, k) = (4usize, 3usize);
let temperature = 0.9_f64;
let rho = Array1::from_vec(vec![-0.1_f64]);
let target = logits(n, k);
let w = vec![1.9_f64, 0.4, 0.8, 0.9];
let pen = SoftmaxAssignmentSparsityPenalty::new(k, temperature).with_row_weights(Some(&w));
let grad = pen.grad_target(target.view(), rho.view());
let eps = 1e-6;
for idx in 0..n * k {
let mut plus = target.clone();
let mut minus = target.clone();
plus[idx] += eps;
minus[idx] -= eps;
let fd = (pen.value(plus.view(), rho.view()) - pen.value(minus.view(), rho.view()))
/ (2.0 * eps);
assert_abs_diff_eq!(grad[idx], fd, epsilon = 1e-7);
}
}
#[test]
fn every_channel_scales_by_w_row_identically() {
let (n, k) = (4usize, 3usize);
let temperature = 0.8_f64;
let rho = Array1::from_vec(vec![0.15_f64]);
let target = logits(n, k);
let v = logits(n, k); let w = vec![1.6_f64, 0.25, 1.05, 1.1];
let base = SoftmaxAssignmentSparsityPenalty::new(k, temperature);
let wtd = base.clone().with_row_weights(Some(&w));
let g0 = base.grad_target(target.view(), rho.view());
let g1 = wtd.grad_target(target.view(), rho.view());
let d0 = base.hessian_diag(target.view(), rho.view()).unwrap();
let d1 = wtd.hessian_diag(target.view(), rho.view()).unwrap();
let m0 = base.psd_majorizer_diag(target.view(), rho.view()).unwrap();
let m1 = wtd.psd_majorizer_diag(target.view(), rho.view()).unwrap();
let h0 = base.hvp(target.view(), rho.view(), v.view());
let h1 = wtd.hvp(target.view(), rho.view(), v.view());
for r in 0..n {
for a in 0..k {
let i = r * k + a;
assert_abs_diff_eq!(g1[i], w[r] * g0[i], epsilon = 1e-12);
assert_abs_diff_eq!(d1[i], w[r] * d0[i], epsilon = 1e-12);
assert_abs_diff_eq!(m1[i], w[r] * m0[i], epsilon = 1e-12);
assert_abs_diff_eq!(h1[i], w[r] * h0[i], epsilon = 1e-12);
}
}
let r0 = base.grad_rho(target.view(), rho.view())[0];
let r1 = wtd.grad_rho(target.view(), rho.view())[0];
let expect: f64 = (0..n)
.map(|r| {
let row = target.slice(s![r * k..r * k + k]).to_owned();
w[r] * base.value(row.view(), rho.view())
})
.sum();
assert_abs_diff_eq!(r1, expect, epsilon = 1e-12);
assert!(r0.is_finite());
}
#[test]
fn none_weights_are_bit_for_bit_unweighted() {
let (n, k) = (3usize, 4usize);
let rho = Array1::from_vec(vec![0.0_f64]);
let target = logits(n, k);
let base = SoftmaxAssignmentSparsityPenalty::new(k, 1.0);
let none = base.clone().with_row_weights(None);
assert_eq!(
base.value(target.view(), rho.view()).to_bits(),
none.value(target.view(), rho.view()).to_bits()
);
let g0 = base.grad_target(target.view(), rho.view());
let g1 = none.grad_target(target.view(), rho.view());
for i in 0..n * k {
assert_eq!(g0[i].to_bits(), g1[i].to_bits());
}
}
}