use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
use crate::manifold::SaeManifoldRho;
use gam_solve::evidence::{HybridAtomCandidate, HybridAtomChoice, select_hybrid_atom};
use gam_terms::analytic_penalties::{
AnalyticPenalty, IBPAssignmentPenalty, IbpHessianDiagThirdChannels,
SoftmaxAssignmentSparsityPenalty, resolve_learnable_weight,
};
use gam_terms::latent::{LatentCoordValues, LatentIdMode, LatentManifold};
#[derive(Clone, Debug)]
pub struct SupportMeasure {
atom_idx: usize,
weights: Array1<f64>,
mass: f64,
fisher_n: f64,
}
impl SupportMeasure {
#[must_use = "support construction error must be handled"]
pub fn from_assignment(assignment: &SaeAssignment, atom_idx: usize) -> Result<Self, String> {
let assignments = assignment.assignments();
Self::from_assignment_matrix(assignments.view(), atom_idx)
}
#[must_use = "support construction error must be handled"]
pub fn from_assignment_matrix(
assignments: ArrayView2<'_, f64>,
atom_idx: usize,
) -> Result<Self, String> {
let (_n, k) = assignments.dim();
if atom_idx >= k {
return Err(format!(
"SupportMeasure::from_assignment_matrix: atom {atom_idx} out of range K={k}"
));
}
let weights = assignments.column(atom_idx).to_owned();
Self::from_weights(atom_idx, weights)
}
#[must_use = "support construction error must be handled"]
pub fn from_argmax_owners(
owners: &[usize],
atom_idx: usize,
k_atoms: usize,
) -> Result<Self, String> {
if atom_idx >= k_atoms {
return Err(format!(
"SupportMeasure::from_argmax_owners: atom {atom_idx} out of range K={k_atoms}"
));
}
let mut weights = Array1::<f64>::zeros(owners.len());
for (row, &owner) in owners.iter().enumerate() {
if owner >= k_atoms {
return Err(format!(
"SupportMeasure::from_argmax_owners: row {row} owner {owner} out of range K={k_atoms}"
));
}
if owner == atom_idx {
weights[row] = 1.0;
}
}
Self::from_weights(atom_idx, weights)
}
#[must_use = "support construction error must be handled"]
pub fn from_weights(atom_idx: usize, weights: Array1<f64>) -> Result<Self, String> {
let mut mass = 0.0_f64;
let mut fisher_n = 0.0_f64;
for (row, &w) in weights.iter().enumerate() {
if !(w.is_finite() && w >= 0.0) {
return Err(format!(
"SupportMeasure::from_weights: row {row} has invalid support weight {w}"
));
}
mass += w;
fisher_n += w * w;
}
Ok(Self {
atom_idx,
weights,
mass,
fisher_n,
})
}
pub fn atom_idx(&self) -> usize {
self.atom_idx
}
pub fn weights(&self) -> ArrayView1<'_, f64> {
self.weights.view()
}
pub fn len(&self) -> usize {
self.weights.len()
}
pub fn is_empty(&self) -> bool {
self.weights.is_empty()
}
pub fn mass(&self) -> f64 {
self.mass
}
pub fn fisher_n(&self) -> f64 {
self.fisher_n
}
pub fn ess(&self) -> f64 {
if self.fisher_n > 0.0 {
(self.mass * self.mass) / self.fisher_n
} else {
0.0
}
}
pub fn weight(&self, row: usize) -> f64 {
self.weights[row]
}
pub fn positive_rows(&self) -> Vec<usize> {
self.weights
.iter()
.enumerate()
.filter_map(|(row, &w)| if w > 0.0 { Some(row) } else { None })
.collect()
}
}
pub(crate) const SAE_ASSIGNMENT_LOGIT_STEP_CAP_TAUS: f64 = 4.0;
pub(crate) const SAE_ATOM_COLLAPSE_RESEED_BUDGET: usize = 1;
pub(crate) const SAE_ATOM_DECODER_NORM_COLLAPSE_RATIO: f64 = 1.0e-3;
pub(crate) const SAE_DICTIONARY_COCOLLAPSE_RESEED_BUDGET: usize = 3;
pub(crate) const JUMPRELU_OPTIMIZATION_LOGIT_CUTOFF: f64 = -36.0;
#[inline]
pub(crate) fn jumprelu_in_optimization_band(logit: f64, threshold: f64, temperature: f64) -> bool {
(logit - threshold) / temperature > JUMPRELU_OPTIMIZATION_LOGIT_CUTOFF
}
#[derive(Debug, Clone, Copy)]
pub enum AssignmentMode {
Softmax { temperature: f64, sparsity: f64 },
IBPMap {
temperature: f64,
alpha: f64,
learnable_alpha: bool,
},
ThresholdGate { temperature: f64, threshold: f64 },
TopK { k: usize },
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AssignmentModeRequest {
Default,
Softmax,
ThresholdGate,
IbpMap,
}
#[derive(Debug, Clone, Copy)]
pub struct AssignmentModeAdmission {
pub mode: AssignmentMode,
pub top_k: Option<usize>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RoutingPredictor {
Snapshot,
ChartGeometry,
}
impl AssignmentMode {
#[must_use]
pub fn softmax(temperature: f64) -> Self {
Self::Softmax {
temperature,
sparsity: 1.0,
}
}
#[must_use]
pub fn ibp_map(temperature: f64, alpha: f64, learnable_alpha: bool) -> Self {
Self::IBPMap {
temperature,
alpha,
learnable_alpha,
}
}
#[must_use]
pub fn threshold_gate(temperature: f64, threshold: f64) -> Self {
Self::ThresholdGate {
temperature,
threshold,
}
}
#[must_use]
pub fn top_k_support(k: usize) -> Self {
Self::TopK { k }
}
pub fn temperature(&self) -> f64 {
match *self {
AssignmentMode::Softmax { temperature, .. }
| AssignmentMode::IBPMap { temperature, .. }
| AssignmentMode::ThresholdGate { temperature, .. } => temperature,
AssignmentMode::TopK { .. } => 1.0,
}
}
pub(crate) fn set_temperature(&mut self, new_temperature: f64) -> Result<(), String> {
if !(new_temperature.is_finite() && new_temperature > 0.0) {
return Err(format!(
"AssignmentMode: temperature must be finite and positive; got {new_temperature}"
));
}
match self {
AssignmentMode::Softmax { temperature, .. }
| AssignmentMode::IBPMap { temperature, .. }
| AssignmentMode::ThresholdGate { temperature, .. } => {
*temperature = new_temperature;
}
AssignmentMode::TopK { .. } => {}
}
Ok(())
}
pub(crate) fn validate(&self) -> Result<(), String> {
let temperature = self.temperature();
if !(temperature.is_finite() && temperature > 0.0) {
return Err(format!(
"AssignmentMode: temperature must be finite and positive; got {temperature}"
));
}
match *self {
AssignmentMode::Softmax { sparsity, .. } => {
if !(sparsity.is_finite() && sparsity > 0.0) {
return Err(format!(
"AssignmentMode::Softmax: sparsity must be finite and positive; got {sparsity}"
));
}
}
AssignmentMode::IBPMap { alpha, .. } => {
if !(alpha.is_finite() && alpha > 0.0) {
return Err(format!(
"AssignmentMode::IBPMap: alpha must be finite and positive; got {alpha}"
));
}
}
AssignmentMode::ThresholdGate { threshold, .. } => {
if !threshold.is_finite() {
return Err(format!(
"AssignmentMode::ThresholdGate: threshold must be finite; got {threshold}"
));
}
}
AssignmentMode::TopK { k } => {
if k == 0 {
return Err(
"AssignmentMode::TopK: support size k must be at least 1".to_string()
);
}
}
}
Ok(())
}
pub(crate) fn resolved_ibp_alpha(
&self,
rho: &SaeManifoldRho,
per_fit_override: Option<f64>,
) -> Option<f64> {
match *self {
AssignmentMode::IBPMap {
alpha,
learnable_alpha,
..
} => Some(if let Some(over) = per_fit_override {
over
} else if learnable_alpha {
resolve_learnable_weight(alpha, rho.log_lambda_sparse)
} else {
alpha
}),
_ => None,
}
}
}
pub fn default_top_k_for_large_dictionary(n_obs: usize, k_atoms: usize) -> Option<usize> {
if n_obs == 0 || k_atoms <= 1 {
return None;
}
if n_obs >= k_atoms.saturating_mul(k_atoms) {
return None;
}
let cap = n_obs.div_ceil(k_atoms).clamp(1, k_atoms - 1);
Some(cap)
}
pub fn admit_assignment_mode_for_size(
request: AssignmentModeRequest,
n_obs: usize,
k_atoms: usize,
temperature: f64,
alpha: f64,
learnable_alpha: bool,
threshold: f64,
) -> Result<AssignmentModeAdmission, String> {
if n_obs == 0 {
return Err("admit_assignment_mode_for_size: n_obs must be positive".to_string());
}
if k_atoms == 0 {
return Err("admit_assignment_mode_for_size: k_atoms must be positive".to_string());
}
let large_k_top = default_top_k_for_large_dictionary(n_obs, k_atoms);
let admission = match request {
AssignmentModeRequest::Default | AssignmentModeRequest::Softmax => {
AssignmentModeAdmission {
mode: AssignmentMode::softmax(temperature),
top_k: large_k_top,
}
}
AssignmentModeRequest::ThresholdGate => AssignmentModeAdmission {
mode: AssignmentMode::threshold_gate(temperature, threshold),
top_k: None,
},
AssignmentModeRequest::IbpMap => {
AssignmentModeAdmission {
mode: AssignmentMode::ibp_map(temperature, alpha, learnable_alpha),
top_k: large_k_top,
}
}
};
admission.mode.validate()?;
Ok(admission)
}
#[doc(hidden)]
#[derive(Debug, Clone)]
pub struct SaeAssignment {
pub logits: Array2<f64>,
pub coords: Vec<LatentCoordValues>,
pub mode: AssignmentMode,
pub ungated: Vec<bool>,
pub frozen_logits: Option<Array2<f64>>,
pub ibp_alpha_override: Option<f64>,
}
impl SaeAssignment {
#[must_use = "build error must be handled"]
pub fn new(
logits: Array2<f64>,
coords: Vec<LatentCoordValues>,
temperature: f64,
) -> Result<Self, String> {
Self::with_mode(logits, coords, AssignmentMode::softmax(temperature))
}
#[must_use = "build error must be handled"]
pub fn with_mode(
mut logits: Array2<f64>,
coords: Vec<LatentCoordValues>,
mode: AssignmentMode,
) -> Result<Self, String> {
mode.validate()?;
let n = logits.nrows();
let k = logits.ncols();
if coords.len() != k {
return Err(format!(
"SaeAssignment::new: coords length {} must equal K={k}",
coords.len()
));
}
for (atom, coord) in coords.iter().enumerate() {
if coord.n_obs() != n {
return Err(format!(
"SaeAssignment::new: coord atom {atom} has n_obs={} but logits has {n}",
coord.n_obs()
));
}
}
for row in 0..n {
validate_finite_logits(logits.row(row), row)?;
}
if matches!(mode, AssignmentMode::Softmax { .. }) {
canonicalize_softmax_logits(&mut logits);
}
Ok(Self {
logits,
coords,
mode,
ungated: vec![false; k],
frozen_logits: None,
ibp_alpha_override: None,
})
}
#[must_use = "build error must be handled"]
pub fn with_frozen_routing(mut self, predicted: Option<Array2<f64>>) -> Result<Self, String> {
if let Some(ref p) = predicted {
if p.dim() != (self.n_obs(), self.k_atoms()) {
return Err(format!(
"SaeAssignment::with_frozen_routing: predicted shape {:?} must be ({}, {})",
p.dim(),
self.n_obs(),
self.k_atoms()
));
}
if matches!(self.mode, AssignmentMode::Softmax { .. }) {
return Err(
"SaeAssignment::with_frozen_routing: frozen routing under Softmax is rejected \
— the coupled simplex's entropy majorizer is assembled over the logits, which \
a frozen (non-optimized) routing would leave inconsistent; this separable-mode \
contract supports IBP-MAP and JumpReLU, whose per-atom gates have no \
simplex-coupled curvature to skip"
.to_string(),
);
}
for row in 0..p.nrows() {
validate_finite_logits(p.row(row), row)?;
}
}
self.frozen_logits = predicted;
Ok(self)
}
pub fn routing_is_frozen(&self) -> bool {
self.frozen_logits.is_some()
}
pub(crate) fn routing_logits_row(&self, row: usize) -> ArrayView1<'_, f64> {
match self.frozen_logits {
Some(ref f) => f.row(row),
None => self.logits.row(row),
}
}
pub(crate) fn logit_is_fixed(&self, k: usize) -> bool {
matches!(self.mode, AssignmentMode::TopK { .. })
|| self.routing_is_frozen()
|| self.ungated.get(k).copied().unwrap_or(false)
}
pub(crate) fn fixed_logit_mask(&self) -> Vec<bool> {
if matches!(self.mode, AssignmentMode::TopK { .. }) || self.routing_is_frozen() {
vec![true; self.k_atoms()]
} else {
self.ungated.clone()
}
}
#[must_use = "build error must be handled"]
pub fn freeze_routing_from_current_logits(self) -> Result<Self, String> {
let snapshot = self.logits.clone();
self.with_frozen_routing(Some(snapshot))
}
pub fn freeze_routing_in_place(&mut self) -> Result<(), String> {
if matches!(self.mode, AssignmentMode::Softmax { .. }) {
return Err(
"SaeAssignment::freeze_routing_in_place: frozen routing under Softmax is rejected \
(coupled-simplex entropy-majorizer); use IBP-MAP or JumpReLU"
.to_string(),
);
}
let snapshot = self.logits.clone();
for row in 0..snapshot.nrows() {
validate_finite_logits(snapshot.row(row), row)?;
}
self.frozen_logits = Some(snapshot);
Ok(())
}
pub fn set_frozen_routing_in_place(&mut self, predicted: Array2<f64>) -> Result<(), String> {
if predicted.dim() != (self.n_obs(), self.k_atoms()) {
return Err(format!(
"SaeAssignment::set_frozen_routing_in_place: predicted shape {:?} must be ({}, {})",
predicted.dim(),
self.n_obs(),
self.k_atoms()
));
}
if matches!(self.mode, AssignmentMode::Softmax { .. }) {
return Err(
"SaeAssignment::set_frozen_routing_in_place: frozen routing under Softmax is \
rejected (coupled-simplex entropy-majorizer); use IBP-MAP or JumpReLU"
.to_string(),
);
}
for row in 0..predicted.nrows() {
validate_finite_logits(predicted.row(row), row)?;
}
self.frozen_logits = Some(predicted);
Ok(())
}
pub fn thaw_routing(&mut self) {
self.frozen_logits = None;
}
#[must_use = "build error must be handled"]
pub fn with_ungated(mut self, flags: Vec<bool>) -> Result<Self, String> {
if flags.len() != self.k_atoms() {
return Err(format!(
"SaeAssignment::with_ungated: flags length {} must equal K={}",
flags.len(),
self.k_atoms()
));
}
if matches!(self.mode, AssignmentMode::Softmax { .. }) && flags.iter().any(|&u| u) {
return Err(
"SaeAssignment::with_ungated: an ungated atom under Softmax routing is \
rejected — the coupled simplex requires a gated-subset renormalization \
reflected in the logit-JVP and entropy majorizer, which this separable-mode \
contract does not perform; route a dense background tier as IBP-MAP or JumpReLU"
.to_string(),
);
}
self.ungated = flags;
Ok(self)
}
pub fn has_ungated(&self) -> bool {
self.ungated.iter().any(|&u| u)
}
pub fn n_obs(&self) -> usize {
self.logits.nrows()
}
pub fn k_atoms(&self) -> usize {
self.logits.ncols()
}
pub fn total_coord_dim(&self) -> usize {
self.coords.iter().map(|c| c.latent_dim()).sum()
}
pub fn assignment_coord_dim(&self) -> usize {
match self.mode {
AssignmentMode::Softmax { .. } => self.k_atoms().saturating_sub(1),
AssignmentMode::IBPMap { .. } | AssignmentMode::ThresholdGate { .. } => self.k_atoms(),
AssignmentMode::TopK { .. } => 0,
}
}
pub fn row_block_dim(&self) -> usize {
self.assignment_coord_dim() + self.total_coord_dim()
}
pub fn coord_offsets(&self) -> Vec<usize> {
let mut out = Vec::with_capacity(self.k_atoms());
let mut cursor = self.assignment_coord_dim();
for coord in &self.coords {
out.push(cursor);
cursor += coord.latent_dim();
}
out
}
pub fn assignments(&self) -> Array2<f64> {
let n = self.n_obs();
let k = self.k_atoms();
let mut out = Array2::<f64>::zeros((n, k));
for row in 0..n {
let a = self.assignments_row(row);
for atom in 0..k {
out[[row, atom]] = a[atom];
}
}
out
}
pub fn assignments_row(&self, row: usize) -> Array1<f64> {
self.try_assignments_row(row)
.expect("assignment logits must be finite")
}
pub fn try_assignments_row(&self, row: usize) -> Result<Array1<f64>, String> {
self.try_assignments_row_inner(row)
}
pub(crate) fn resolved_ibp_alpha(&self, rho: &SaeManifoldRho) -> Option<f64> {
self.mode.resolved_ibp_alpha(rho, self.ibp_alpha_override)
}
pub(crate) fn effective_alpha_is_learnable(&self) -> bool {
match self.mode {
AssignmentMode::IBPMap {
learnable_alpha, ..
} => learnable_alpha && self.ibp_alpha_override.is_none(),
_ => false,
}
}
pub fn set_ibp_alpha_override(&mut self, alpha: Option<f64>) {
self.ibp_alpha_override = alpha;
}
pub(crate) fn ibp_eb_log_alpha_step(
&self,
rho: &SaeManifoldRho,
) -> Result<Option<f64>, String> {
if !self.effective_alpha_is_learnable() {
return Ok(None);
}
let resolved = self.resolved_ibp_alpha(rho);
let Some(alpha_current) = resolved else {
return Ok(None);
};
if !(alpha_current.is_finite() && alpha_current > 0.0) {
return Ok(None);
}
let k = self.k_atoms();
let n = self.n_obs();
if k == 0 || n == 0 {
return Ok(None);
}
let mut occupancy = vec![0.0_f64; k];
let mut buf = vec![0.0_f64; k];
for row in 0..n {
self.try_assignments_row_into(row, &mut buf)?;
for (acc, &g) in occupancy.iter_mut().zip(buf.iter()) {
*acc += g;
}
}
let alpha_star = ibp_eb_geometric_alpha_fixed_point(&occupancy, n as f64, alpha_current);
if !(alpha_star.is_finite() && alpha_star > 0.0) {
return Ok(None);
}
const LOG_ALPHA_STEP_CAP: f64 = 2.0;
let step =
(alpha_star.ln() - alpha_current.ln()).clamp(-LOG_ALPHA_STEP_CAP, LOG_ALPHA_STEP_CAP);
Ok(Some(step))
}
fn try_assignments_row_inner(&self, row: usize) -> Result<Array1<f64>, String> {
let routing = self.routing_logits_row(row);
validate_finite_logits(routing, row)?;
if self.k_atoms() == 1 && matches!(self.mode, AssignmentMode::Softmax { .. }) {
return Ok(Array1::from_vec(vec![1.0]));
}
let mut row_gates = match self.mode {
AssignmentMode::Softmax { temperature, .. } => softmax_row(routing, temperature),
AssignmentMode::IBPMap { temperature, .. } => ibp_map_row(routing, temperature),
AssignmentMode::ThresholdGate {
temperature,
threshold,
} => jumprelu_row(routing, temperature, threshold),
AssignmentMode::TopK { k } => topk_row(routing, k),
};
if self.has_ungated() {
for (k, gate) in row_gates.iter_mut().enumerate() {
if self.ungated[k] {
*gate = 1.0;
}
}
}
Ok(row_gates)
}
pub(crate) fn try_assignments_row_into(
&self,
row: usize,
out: &mut [f64],
) -> Result<(), String> {
let routing = self.routing_logits_row(row);
validate_finite_logits(routing, row)?;
if self.k_atoms() == 1 && matches!(self.mode, AssignmentMode::Softmax { .. }) {
out[0] = 1.0;
return Ok(());
}
match self.mode {
AssignmentMode::Softmax { temperature, .. } => {
softmax_row_into(routing, temperature, out)
}
AssignmentMode::IBPMap { temperature, .. } => {
ibp_map_row_into(routing, temperature, out)
}
AssignmentMode::ThresholdGate {
temperature,
threshold,
} => jumprelu_row_into(routing, temperature, threshold, out),
AssignmentMode::TopK { k } => topk_row_into(routing, k, out),
};
if self.has_ungated() {
for (k, gate) in out.iter_mut().enumerate() {
if self.ungated[k] {
*gate = 1.0;
}
}
}
Ok(())
}
pub(crate) fn persist_resolved_ibp_alpha(&mut self, rho: &SaeManifoldRho) -> bool {
let AssignmentMode::IBPMap {
temperature,
alpha,
learnable_alpha: true,
} = self.mode
else {
return false;
};
let resolved_alpha = resolve_learnable_weight(alpha, rho.log_lambda_sparse);
self.mode = AssignmentMode::IBPMap {
temperature,
alpha: resolved_alpha,
learnable_alpha: false,
};
true
}
pub(crate) fn try_assignments(&self) -> Result<Array2<f64>, String> {
let n = self.n_obs();
let k = self.k_atoms();
let mut out = Array2::<f64>::zeros((n, k));
for row in 0..n {
let a = self.try_assignments_row(row)?;
for atom in 0..k {
out[[row, atom]] = a[atom];
}
}
Ok(out)
}
pub fn flatten_ext_coords(&self) -> Array1<f64> {
let n = self.n_obs();
let q = self.row_block_dim();
let k = self.k_atoms();
let assignment_dim = self.assignment_coord_dim();
let offsets = self.coord_offsets();
let mut out = Array1::<f64>::zeros(n * q);
for row in 0..n {
let base = row * q;
for atom in 0..assignment_dim {
out[base + atom] = self.logits[[row, atom]];
}
for atom in 0..k {
let d = self.coords[atom].latent_dim();
let t_row = self.coords[atom].row(row);
for axis in 0..d {
out[base + offsets[atom] + axis] = t_row[axis];
}
}
}
out
}
#[must_use = "build error must be handled"]
pub fn from_blocks_with_mode(
logits: Array2<f64>,
coord_blocks: Vec<Array2<f64>>,
mode: AssignmentMode,
) -> Result<Self, String> {
let coords = coord_blocks
.iter()
.map(|c| LatentCoordValues::from_matrix(c.view(), LatentIdMode::None))
.collect();
Self::with_mode(logits, coords, mode)
}
#[must_use = "build error must be handled"]
pub fn from_blocks_with_mode_and_manifolds(
logits: Array2<f64>,
coord_blocks: Vec<Array2<f64>>,
manifolds: Vec<LatentManifold>,
mode: AssignmentMode,
) -> Result<Self, String> {
if coord_blocks.len() != manifolds.len() {
return Err(format!(
"SaeAssignment::from_blocks_with_mode_and_manifolds: coord block length {} != manifold length {}",
coord_blocks.len(),
manifolds.len()
));
}
let coords = coord_blocks
.iter()
.zip(manifolds)
.map(|(c, manifold)| {
LatentCoordValues::from_matrix_with_manifold(c.view(), LatentIdMode::None, manifold)
})
.collect();
Self::with_mode(logits, coords, mode)
}
}
pub(crate) fn neutral_gate_weights(mode: AssignmentMode, k_atoms: usize) -> Array1<f64> {
match mode {
AssignmentMode::Softmax { .. } => Array1::from_elem(k_atoms, 1.0 / (k_atoms.max(1) as f64)),
AssignmentMode::IBPMap { temperature, .. } => {
ibp_map_row(Array1::<f64>::zeros(k_atoms).view(), temperature)
}
AssignmentMode::ThresholdGate { .. } => Array1::from_elem(k_atoms, 0.5),
AssignmentMode::TopK { k } => topk_row(Array1::<f64>::zeros(k_atoms).view(), k),
}
}
pub(crate) fn softmax_row(logits: ArrayView1<'_, f64>, temperature: f64) -> Array1<f64> {
let k = logits.len();
let inv_tau = 1.0 / temperature;
let mut max_logit = f64::NEG_INFINITY;
for &v in logits.iter() {
max_logit = max_logit.max(v);
}
let mut out = Array1::<f64>::zeros(k);
let mut sum = 0.0;
for i in 0..k {
let v = ((logits[i] - max_logit) * inv_tau).exp();
out[i] = v;
sum += v;
}
assert!(sum.is_finite() && sum > 0.0);
for v in out.iter_mut() {
*v /= sum;
}
out
}
pub(crate) fn validate_finite_logits(
logits: ArrayView1<'_, f64>,
row: usize,
) -> Result<(), String> {
for (col, &v) in logits.iter().enumerate() {
if !v.is_finite() {
return Err(format!(
"SaeAssignment: non-finite assignment logit at row {row}, atom {col}: {v}"
));
}
}
Ok(())
}
pub(crate) fn canonicalize_softmax_logits(logits: &mut Array2<f64>) {
let k = logits.ncols();
if k == 0 {
return;
}
if k == 1 {
logits.fill(0.0);
return;
}
for row in 0..logits.nrows() {
let reference = logits[[row, k - 1]];
for col in 0..k - 1 {
logits[[row, col]] -= reference;
}
logits[[row, k - 1]] = 0.0;
}
}
#[derive(Debug, Clone, Copy)]
pub enum OrderedPriorSchedule {
Geometric { alpha: f64 },
PowerLaw { c: f64, s: f64, k0: f64 },
}
pub fn ordered_prior_means(k_atoms: usize, schedule: OrderedPriorSchedule) -> Array1<f64> {
let mut out = Array1::<f64>::zeros(k_atoms);
match schedule {
OrderedPriorSchedule::Geometric { alpha } => {
let log_ratio = (alpha / (alpha + 1.0)).ln();
for k in 0..k_atoms {
let log_pi = ((k + 1) as f64) * log_ratio;
out[k] = log_pi.exp().max(f64::MIN_POSITIVE);
}
}
OrderedPriorSchedule::PowerLaw { c, s, k0 } => {
for k in 0..k_atoms {
let log_pi = c.ln() - s * ((k as f64) + k0).ln();
out[k] = log_pi.exp().clamp(f64::MIN_POSITIVE, 1.0);
}
}
}
out
}
pub fn default_ibp_concentration_for_k_atoms(k_atoms: usize) -> f64 {
let k = k_atoms.max(1) as f64;
let alpha = 1.0 / ((1.0 / k).exp() - 1.0);
alpha.max(1.0)
}
#[inline]
fn trigamma(mut x: f64) -> f64 {
if !(x.is_finite() && x > 0.0) {
return f64::NAN;
}
let mut acc = 0.0;
while x < 8.0 {
acc += 1.0 / (x * x);
x += 1.0;
}
let inv = 1.0 / x;
let inv2 = inv * inv;
acc + inv + 0.5 * inv2 + inv2 * inv / 6.0 - inv2 * inv2 * inv / 30.0
+ inv2 * inv2 * inv2 * inv / 42.0
- inv2 * inv2 * inv2 * inv2 * inv / 30.0
}
#[inline]
fn ibp_eb_atom_score_deriv(m_k: f64, n_obs: f64, a: f64) -> (f64, f64) {
let g = statrs::function::gamma::digamma(m_k + a)
- statrs::function::gamma::digamma(n_obs + a + 1.0)
+ 1.0 / a;
let gp = trigamma(m_k + a) - trigamma(n_obs + a + 1.0) - 1.0 / (a * a);
(g, gp)
}
pub fn ibp_eb_marginal_score(occupancy: &[f64], n_obs: f64, a: &[f64], da_dtheta: &[f64]) -> f64 {
let mut s = 0.0;
for k in 0..occupancy.len() {
let (g, _) = ibp_eb_atom_score_deriv(occupancy[k].clamp(0.0, n_obs), n_obs, a[k]);
s += g * da_dtheta[k];
}
s
}
pub fn ibp_eb_alpha_score_hess(occupancy: &[f64], n_obs: f64, alpha: f64) -> (f64, f64) {
let rho = alpha / (alpha + 1.0);
let one_m_rho = 1.0 - rho;
let mut s = 0.0;
let mut h = 0.0;
for (k, &m_raw) in occupancy.iter().enumerate() {
let u = (k + 1) as f64;
let m_k = m_raw.clamp(0.0, n_obs);
let mu = rho.powf(u).clamp(f64::MIN_POSITIVE, 1.0 - 1.0e-12);
let dmu = u * mu * one_m_rho; let d2mu = u * mu * one_m_rho * (u * one_m_rho - rho); let om = 1.0 - mu;
let a = mu / om;
let da = dmu / (om * om);
let d2a = 2.0 * dmu * dmu / (om * om * om) + d2mu / (om * om);
let (g, gp) = ibp_eb_atom_score_deriv(m_k, n_obs, a);
s += g * da;
h += gp * da * da + g * d2a;
}
(s, h)
}
pub fn ibp_eb_geometric_alpha_fixed_point(occupancy: &[f64], n_obs: f64, alpha_seed: f64) -> f64 {
const THETA_LO: f64 = -12.0; const THETA_HI: f64 = 16.0; const NEWTON_MAX_ITERS: usize = 100;
const STEP_TR: f64 = 1.0; const TOL: f64 = 1.0e-10;
if !(n_obs > 0.0) || occupancy.is_empty() {
return alpha_seed;
}
let seed = if alpha_seed.is_finite() && alpha_seed > 0.0 {
alpha_seed
} else {
1.0
};
let mut theta = seed.ln().clamp(THETA_LO, THETA_HI);
for _ in 0..NEWTON_MAX_ITERS {
let (s, h) = ibp_eb_alpha_score_hess(occupancy, n_obs, theta.exp());
if !s.is_finite() || s.abs() < TOL {
break;
}
let mut step = if h < -1.0e-12 {
-s / h
} else {
s.signum() * STEP_TR
};
if !step.is_finite() {
break;
}
step = step.clamp(-STEP_TR, STEP_TR);
let new_theta = (theta + step).clamp(THETA_LO, THETA_HI);
let converged = (new_theta - theta).abs() < TOL;
theta = new_theta;
if converged {
break;
}
}
theta.exp()
}
pub fn ibp_map_row(logits: ArrayView1<'_, f64>, temperature: f64) -> Array1<f64> {
let mut out = Array1::<f64>::zeros(logits.len());
for i in 0..logits.len() {
out[i] = gam_linalg::utils::stable_logistic(logits[i] / temperature);
}
out
}
#[must_use]
pub fn ibp_map_row_value_grad(
logits: ArrayView1<'_, f64>,
temperature: f64,
) -> (Array1<f64>, Array1<f64>) {
let inv_tau = 1.0 / temperature;
let mut value = Array1::<f64>::zeros(logits.len());
let mut grad = Array1::<f64>::zeros(logits.len());
for i in 0..logits.len() {
let sig = gam_linalg::utils::stable_logistic(logits[i] * inv_tau);
value[i] = sig;
grad[i] = sig * (1.0 - sig) * inv_tau;
}
(value, grad)
}
#[must_use]
pub fn ibp_map_batch_value_grad(
logits: ArrayView2<'_, f64>,
temperature: f64,
) -> (Array2<f64>, Array2<f64>) {
let (n, k) = logits.dim();
let inv_tau = 1.0 / temperature;
let mut value = Array2::<f64>::zeros((n, k));
let mut grad = Array2::<f64>::zeros((n, k));
for i in 0..n {
for j in 0..k {
let sig = gam_linalg::utils::stable_logistic(logits[[i, j]] * inv_tau);
value[[i, j]] = sig;
grad[[i, j]] = sig * (1.0 - sig) * inv_tau;
}
}
(value, grad)
}
pub fn jumprelu_row(logits: ArrayView1<'_, f64>, temperature: f64, threshold: f64) -> Array1<f64> {
let mut out = Array1::<f64>::zeros(logits.len());
for i in 0..logits.len() {
if logits[i] > threshold {
out[i] = gam_linalg::utils::stable_logistic((logits[i] - threshold) / temperature);
}
}
out
}
pub fn activation_matrix_from_logits(
logits: ArrayView2<'_, f64>,
kind: &str,
temperature: f64,
threshold: f64,
) -> Result<Array2<f64>, String> {
if !(temperature.is_finite() && temperature > 0.0) {
return Err(format!(
"activation_matrix_from_logits: temperature must be finite and positive; got {temperature}"
));
}
let (n_rows, k_atoms) = logits.dim();
let mut out = Array2::<f64>::zeros((n_rows, k_atoms));
for row in 0..n_rows {
let row_logits = logits.row(row);
validate_finite_logits(row_logits, row)?;
let activation = match kind {
"softmax" => softmax_row(row_logits, temperature),
"ibp_map" => ibp_map_row(row_logits, temperature),
"threshold_gate" => jumprelu_row(row_logits, temperature, threshold),
other => {
return Err(format!(
"activation_matrix_from_logits: unsupported assignment kind {other:?} \
(expected 'softmax', 'ibp_map', or 'threshold_gate')"
));
}
};
out.row_mut(row).assign(&activation);
}
Ok(out)
}
pub fn topk_row(logits: ArrayView1<'_, f64>, k: usize) -> Array1<f64> {
let mut out = Array1::<f64>::zeros(logits.len());
topk_row_into(
logits,
k,
out.as_slice_mut()
.expect("freshly allocated 1-D array is contiguous"),
);
out
}
pub(crate) fn topk_row_into(logits: ArrayView1<'_, f64>, k: usize, out: &mut [f64]) {
let n = logits.len();
if k >= n {
out[..n].fill(1.0);
return;
}
out[..n].fill(0.0);
let mut idx: Vec<usize> = (0..n).collect();
idx.select_nth_unstable_by(k, |&a, &b| {
logits[b]
.partial_cmp(&logits[a])
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.cmp(&b))
});
for &i in &idx[..k] {
out[i] = 1.0;
}
}
#[must_use]
pub fn jumprelu_row_value_grad(
logits: ArrayView1<'_, f64>,
temperature: f64,
thresholds: ArrayView1<'_, f64>,
) -> (Array1<f64>, Array1<f64>) {
assert_eq!(
logits.len(),
thresholds.len(),
"jumprelu_row_value_grad: logits/thresholds length mismatch"
);
let inv_tau = 1.0 / temperature;
let mut value = Array1::<f64>::zeros(logits.len());
let mut grad = Array1::<f64>::zeros(logits.len());
for i in 0..logits.len() {
let sig = gam_linalg::utils::stable_logistic((logits[i] - thresholds[i]) * inv_tau);
if logits[i] > thresholds[i] {
value[i] = sig;
}
grad[i] = sig * (1.0 - sig) * inv_tau;
}
(value, grad)
}
#[must_use]
pub fn jumprelu_batch_value_grad(
logits: ArrayView2<'_, f64>,
temperature: f64,
thresholds: ArrayView1<'_, f64>,
) -> (Array2<f64>, Array2<f64>) {
let (n, k) = logits.dim();
assert_eq!(
k,
thresholds.len(),
"jumprelu_batch_value_grad: logits columns {k} != thresholds length {}",
thresholds.len()
);
let inv_tau = 1.0 / temperature;
let mut value = Array2::<f64>::zeros((n, k));
let mut grad = Array2::<f64>::zeros((n, k));
for i in 0..n {
for j in 0..k {
let sig =
gam_linalg::utils::stable_logistic((logits[[i, j]] - thresholds[j]) * inv_tau);
if logits[[i, j]] > thresholds[j] {
value[[i, j]] = sig;
}
grad[[i, j]] = sig * (1.0 - sig) * inv_tau;
}
}
(value, grad)
}
#[must_use]
pub fn inverse_softplus(value: f64) -> f64 {
if value <= 0.0 || value.is_nan() {
f64::NAN
} else if value > 30.0 {
value + (-(-value).exp()).ln_1p()
} else {
value.exp_m1().ln()
}
}
#[must_use]
pub fn topk_activation_row_value_grad(
logits: ArrayView1<'_, f64>,
temperature: f64,
) -> (Array1<f64>, Array1<f64>) {
let inv_tau = 1.0 / temperature;
let mut value = Array1::<f64>::zeros(logits.len());
let mut grad = Array1::<f64>::zeros(logits.len());
for i in 0..logits.len() {
let scaled = logits[i] * inv_tau;
value[i] = temperature * gam_linalg::utils::stable_softplus(scaled);
grad[i] = gam_linalg::utils::stable_logistic(scaled);
}
(value, grad)
}
#[must_use]
pub fn topk_activation_batch_value_grad(
logits: ArrayView2<'_, f64>,
temperature: f64,
) -> (Array2<f64>, Array2<f64>) {
let (n, k) = logits.dim();
let inv_tau = 1.0 / temperature;
let mut value = Array2::<f64>::zeros((n, k));
let mut grad = Array2::<f64>::zeros((n, k));
for i in 0..n {
for j in 0..k {
let scaled = logits[[i, j]] * inv_tau;
value[[i, j]] = temperature * gam_linalg::utils::stable_softplus(scaled);
grad[[i, j]] = gam_linalg::utils::stable_logistic(scaled);
}
}
(value, grad)
}
#[cfg(test)]
mod ibp_map_batch_tests {
use super::*;
#[test]
fn ibp_map_batch_matches_row_kernel_bit_for_bit() {
let n = 5usize;
let k = 7usize;
let temperature = 0.41_f64;
let logits = Array2::from_shape_fn((n, k), |(i, j)| {
((i as f64) * 0.37 - (j as f64) * 0.19 + 0.11).sin()
});
let (value, grad) = ibp_map_batch_value_grad(logits.view(), temperature);
assert_eq!(value.dim(), (n, k));
assert_eq!(grad.dim(), (n, k));
for i in 0..n {
let (rv, rg) = ibp_map_row_value_grad(logits.row(i), temperature);
for j in 0..k {
assert_eq!(value[[i, j]], rv[j], "value mismatch at row {i} atom {j}");
assert_eq!(grad[[i, j]], rg[j], "grad mismatch at row {i} atom {j}");
}
}
}
}
#[cfg(test)]
mod topk_activation_tests {
use super::*;
#[test]
fn topk_activation_batch_matches_row_kernel_bit_for_bit() {
let n = 5usize;
let k = 7usize;
let temperature = 0.41_f64;
let logits = Array2::from_shape_fn((n, k), |(i, j)| {
((i as f64) * 0.37 - (j as f64) * 0.19 + 0.11).sin()
});
let (value, grad) = topk_activation_batch_value_grad(logits.view(), temperature);
assert_eq!(value.dim(), (n, k));
assert_eq!(grad.dim(), (n, k));
for i in 0..n {
let (rv, rg) = topk_activation_row_value_grad(logits.row(i), temperature);
for j in 0..k {
assert_eq!(value[[i, j]], rv[j], "value mismatch at row {i} atom {j}");
assert_eq!(grad[[i, j]], rg[j], "grad mismatch at row {i} atom {j}");
}
}
}
#[test]
fn topk_activation_is_nonnegative_and_grad_is_logistic() {
let temperature = 0.7_f64;
let logits = Array1::from(vec![-4.0_f64, -0.5, 0.0, 0.5, 4.0]);
let (value, grad) = topk_activation_row_value_grad(logits.view(), temperature);
for (&v, &g) in value.iter().zip(grad.iter()) {
assert!(v >= 0.0, "activation must be non-negative, got {v}");
assert!(
(0.0..=1.0).contains(&g),
"grad must be a logistic in [0,1], got {g}"
);
}
assert!((value[2] - temperature * 2.0_f64.ln()).abs() < 1e-15);
assert!((grad[2] - 0.5).abs() < 1e-15);
}
}
#[cfg(test)]
mod topk_support_gate_tests {
use super::*;
#[test]
fn topk_row_selects_exact_support_and_l0_is_k() {
let logits = Array1::from(vec![0.3_f64, 0.9, 0.9, -1.0, 0.5]);
let g = topk_row(logits.view(), 3);
assert_eq!(g.to_vec(), vec![0.0, 1.0, 1.0, 0.0, 1.0]);
assert_eq!(
g.iter().filter(|&&v| v == 1.0).count(),
3,
"L0 must equal k exactly"
);
assert!(
g.iter().all(|&v| v == 0.0 || v == 1.0),
"gates are hard {{0,1}}"
);
}
#[test]
fn topk_boundary_tie_breaks_toward_lower_index() {
let logits = Array1::from(vec![1.0_f64, 0.5, 0.5, 0.1]);
let g = topk_row(logits.view(), 2);
assert_eq!(
g.to_vec(),
vec![1.0, 1.0, 0.0, 0.0],
"the tied boundary atom with the LOWER index wins deterministically"
);
}
#[test]
fn topk_row_into_is_bit_identical_and_k_ge_n_is_all_active() {
let logits = Array1::from(vec![-0.2_f64, 3.0, 0.7, 0.7, -5.0, 2.2]);
for k in [1usize, 2, 4, 6, 9] {
let alloc = topk_row(logits.view(), k);
let mut buf = vec![f64::NAN; logits.len()];
topk_row_into(logits.view(), k, &mut buf);
assert_eq!(
alloc.to_vec(),
buf,
"into-twin must be bit-identical at k={k}"
);
}
let all = topk_row(logits.view(), 99);
assert!(
all.iter().all(|&v| v == 1.0),
"k >= n degenerates to all-active"
);
}
#[test]
fn topk_neutral_support_is_first_k_atoms() {
let w = neutral_gate_weights(AssignmentMode::top_k_support(3), 6);
assert_eq!(w.to_vec(), vec![1.0, 1.0, 1.0, 0.0, 0.0, 0.0]);
}
#[test]
fn topk_mode_carries_no_temperature_or_prior_knobs() {
let mode = AssignmentMode::top_k_support(4);
mode.validate().expect("k >= 1 validates");
assert!(
AssignmentMode::top_k_support(0).validate().is_err(),
"k = 0 must be rejected"
);
}
}
#[cfg(test)]
mod jumprelu_batch_tests {
use super::*;
#[test]
fn jumprelu_batch_matches_row_kernel_bit_for_bit() {
let n = 5usize;
let k = 7usize;
let temperature = 0.41_f64;
let logits = Array2::from_shape_fn((n, k), |(i, j)| {
((i as f64) * 0.37 - (j as f64) * 0.19 + 0.11).sin()
});
let thresholds = Array1::from_shape_fn(k, |j| 0.2 - 0.05 * j as f64);
let (value, grad) =
jumprelu_batch_value_grad(logits.view(), temperature, thresholds.view());
assert_eq!(value.dim(), (n, k));
assert_eq!(grad.dim(), (n, k));
for i in 0..n {
let (rv, rg) = jumprelu_row_value_grad(logits.row(i), temperature, thresholds.view());
for j in 0..k {
assert_eq!(value[[i, j]], rv[j], "value mismatch at row {i} atom {j}");
assert_eq!(grad[[i, j]], rg[j], "grad mismatch at row {i} atom {j}");
}
}
}
}
pub(crate) fn softmax_row_into(logits: ArrayView1<'_, f64>, temperature: f64, out: &mut [f64]) {
let k = logits.len();
let inv_tau = 1.0 / temperature;
let mut max_logit = f64::NEG_INFINITY;
for &v in logits.iter() {
max_logit = max_logit.max(v);
}
let mut sum = 0.0;
for i in 0..k {
let v = ((logits[i] - max_logit) * inv_tau).exp();
out[i] = v;
sum += v;
}
assert!(sum.is_finite() && sum > 0.0);
for v in out.iter_mut() {
*v /= sum;
}
}
pub(crate) fn ibp_map_row_into(logits: ArrayView1<'_, f64>, temperature: f64, out: &mut [f64]) {
for i in 0..logits.len() {
out[i] = gam_linalg::utils::stable_logistic(logits[i] / temperature);
}
}
pub(crate) fn jumprelu_row_into(
logits: ArrayView1<'_, f64>,
temperature: f64,
threshold: f64,
out: &mut [f64],
) {
for i in 0..logits.len() {
if logits[i] > threshold {
out[i] = gam_linalg::utils::stable_logistic((logits[i] - threshold) / temperature);
} else {
out[i] = 0.0;
}
}
}
pub(crate) struct ActiveAtomLogitJvp<'a> {
pub(crate) mode: AssignmentMode,
pub(crate) logit_k: f64,
pub(crate) a_k: f64,
pub(crate) decoded_k: ArrayView1<'a, f64>,
pub(crate) fitted: ArrayView1<'a, f64>,
pub(crate) compact_index: usize,
pub(crate) ungated: bool,
}
pub(crate) fn fill_active_atom_logit_jvp(
input: ActiveAtomLogitJvp<'_>,
jac_compact: &mut Array2<f64>,
) {
let ActiveAtomLogitJvp {
mode,
logit_k,
a_k,
decoded_k,
fitted,
compact_index,
ungated,
} = input;
let p = fitted.len();
if ungated {
return;
}
match mode {
AssignmentMode::Softmax { temperature, .. } => {
let inv_tau = 1.0 / temperature;
for out_col in 0..p {
jac_compact[[compact_index, out_col]] =
a_k * (decoded_k[out_col] - fitted[out_col]) * inv_tau;
}
}
AssignmentMode::IBPMap { temperature, .. } => {
let inv_tau = 1.0 / temperature;
let dz = a_k * (1.0 - a_k) * inv_tau;
for out_col in 0..p {
jac_compact[[compact_index, out_col]] = dz * decoded_k[out_col];
}
}
AssignmentMode::ThresholdGate {
temperature,
threshold,
} => {
if logit_k <= threshold {
return;
}
let inv_tau = 1.0 / temperature;
let activation = gam_linalg::utils::stable_logistic((logit_k - threshold) * inv_tau);
let da = activation * (1.0 - activation) * inv_tau;
for out_col in 0..p {
jac_compact[[compact_index, out_col]] = da * decoded_k[out_col];
}
}
AssignmentMode::TopK { .. } => {}
}
}
pub(crate) fn fill_assignment_logit_jvp_rows(
mode: AssignmentMode,
logits: ArrayView1<'_, f64>,
assignments: ArrayView1<'_, f64>,
decoded: ArrayView2<'_, f64>,
fitted: ArrayView1<'_, f64>,
ungated: &[bool],
local_jac: &mut Array2<f64>,
) {
let is_ungated = |k: usize| ungated.get(k).copied().unwrap_or(false);
match mode {
AssignmentMode::Softmax { temperature, .. } => {
if assignments.len() == 1 {
return;
}
let inv_tau = 1.0 / temperature;
for logit_col in 0..assignments.len() - 1 {
if is_ungated(logit_col) {
continue;
}
for out_col in 0..fitted.len() {
local_jac[[logit_col, out_col]] = assignments[logit_col]
* (decoded[[logit_col, out_col]] - fitted[out_col])
* inv_tau;
}
}
}
AssignmentMode::IBPMap { temperature, .. } => {
let inv_tau = 1.0 / temperature;
for logit_col in 0..assignments.len() {
if is_ungated(logit_col) {
continue;
}
let a_k = assignments[logit_col];
let dz = a_k * (1.0 - a_k) * inv_tau;
for out_col in 0..fitted.len() {
local_jac[[logit_col, out_col]] = dz * decoded[[logit_col, out_col]];
}
}
}
AssignmentMode::ThresholdGate {
temperature,
threshold,
} => {
let inv_tau = 1.0 / temperature;
for logit_col in 0..assignments.len() {
if is_ungated(logit_col) || logits[logit_col] <= threshold {
continue;
}
let activation =
gam_linalg::utils::stable_logistic((logits[logit_col] - threshold) * inv_tau);
let da = activation * (1.0 - activation) * inv_tau;
for out_col in 0..fitted.len() {
local_jac[[logit_col, out_col]] = da * decoded[[logit_col, out_col]];
}
}
}
AssignmentMode::TopK { .. } => {}
}
}
pub(crate) fn flat_logits(logits: ArrayView2<'_, f64>) -> Array1<f64> {
let mut out = Array1::<f64>::zeros(logits.len());
for row in 0..logits.nrows() {
let start = row * logits.ncols();
for col in 0..logits.ncols() {
out[start + col] = logits[[row, col]];
}
}
out
}
fn ibp_prior_penalty(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
base_alpha: f64,
temperature: f64,
row_weights: Option<&[f64]>,
) -> (IBPAssignmentPenalty, Array1<f64>) {
let learnable = assignment.effective_alpha_is_learnable();
let alpha_eff = if learnable {
base_alpha
} else {
assignment.resolved_ibp_alpha(rho).unwrap_or(base_alpha)
};
let mut penalty =
IBPAssignmentPenalty::new(assignment.k_atoms(), alpha_eff, temperature, learnable)
.with_row_weights(row_weights);
if assignment.has_ungated() {
penalty.fixed_columns = Some(assignment.ungated.clone());
}
let rho_view = if learnable {
Array1::from_vec(vec![rho.log_lambda_sparse])
} else {
penalty.weight = rho.lambda_sparse();
Array1::zeros(0)
};
(penalty, rho_view)
}
pub fn assignment_prior_value(assignment: &SaeAssignment, rho: &SaeManifoldRho) -> f64 {
assignment_prior_value_weighted(assignment, rho, None)
}
pub(crate) fn assignment_prior_value_weighted(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
row_weights: Option<&[f64]>,
) -> f64 {
for row in 0..assignment.n_obs() {
validate_finite_logits(assignment.logits.row(row), row)
.expect("assignment logits must be finite");
}
let target = flat_logits(assignment.logits.view());
if matches!(assignment.mode, AssignmentMode::Softmax { .. }) && assignment.k_atoms() == 1 {
return 0.0;
}
if assignment.routing_is_frozen() {
return 0.0;
}
match assignment.mode {
AssignmentMode::Softmax {
temperature,
sparsity,
} => {
let penalty = SoftmaxAssignmentSparsityPenalty::new(assignment.k_atoms(), temperature)
.with_row_weights(row_weights);
let rho_view = Array1::from_vec(vec![rho.log_lambda_sparse + sparsity.ln()]);
penalty.value(target.view(), rho_view.view())
}
AssignmentMode::IBPMap {
temperature, alpha, ..
} => {
let (penalty, rho_view) =
ibp_prior_penalty(assignment, rho, alpha, temperature, row_weights);
penalty.value(target.view(), rho_view.view())
}
AssignmentMode::ThresholdGate {
temperature,
threshold,
} => {
let sparsity_strength = rho.lambda_sparse();
let k = assignment.k_atoms();
let mut acc = 0.0;
for (idx, &logit) in target.iter().enumerate() {
if assignment.logit_is_fixed(idx % k) {
continue;
}
if jumprelu_in_optimization_band(logit, threshold, temperature) {
let w_row = row_weights.map_or(1.0, |w| w[idx / k]);
acc += w_row
* gam_linalg::utils::stable_logistic((logit - threshold) / temperature);
}
}
sparsity_strength * acc
}
AssignmentMode::TopK { .. } => 0.0,
}
}
pub fn assignment_prior_log_strength_derivative(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
) -> f64 {
assignment_prior_log_strength_derivative_weighted(assignment, rho, None)
}
pub(crate) fn assignment_prior_log_strength_derivative_weighted(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
row_weights: Option<&[f64]>,
) -> f64 {
for row in 0..assignment.n_obs() {
validate_finite_logits(assignment.logits.row(row), row)
.expect("assignment logits must be finite");
}
let target = flat_logits(assignment.logits.view());
if matches!(assignment.mode, AssignmentMode::Softmax { .. }) && assignment.k_atoms() == 1 {
return 0.0;
}
if assignment.routing_is_frozen() {
return 0.0;
}
match assignment.mode {
AssignmentMode::Softmax { .. } | AssignmentMode::ThresholdGate { .. } => {
assignment_prior_value_weighted(assignment, rho, row_weights)
}
AssignmentMode::IBPMap {
temperature, alpha, ..
} => {
let (penalty, rho_view) =
ibp_prior_penalty(assignment, rho, alpha, temperature, row_weights);
if penalty.learnable_alpha {
penalty.grad_rho(target.view(), rho_view.view())[0]
} else {
penalty.value(target.view(), rho_view.view())
}
}
AssignmentMode::TopK { .. } => 0.0,
}
}
pub fn assignment_prior_log_strength_hdiag(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
) -> Result<Array1<f64>, String> {
assignment_prior_log_strength_hdiag_weighted(assignment, rho, None)
}
pub(crate) fn assignment_prior_log_strength_hdiag_weighted(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
row_weights: Option<&[f64]>,
) -> Result<Array1<f64>, String> {
for row in 0..assignment.n_obs() {
validate_finite_logits(assignment.logits.row(row), row)?;
}
let target = flat_logits(assignment.logits.view());
if matches!(assignment.mode, AssignmentMode::Softmax { .. }) && assignment.k_atoms() == 1 {
return Ok(Array1::<f64>::zeros(target.len()));
}
if assignment.routing_is_frozen() {
return Ok(Array1::<f64>::zeros(target.len()));
}
match assignment.mode {
AssignmentMode::Softmax {
temperature,
sparsity,
} => {
let penalty = SoftmaxAssignmentSparsityPenalty::new(assignment.k_atoms(), temperature)
.with_row_weights(row_weights);
let rho_view = Array1::from_vec(vec![rho.log_lambda_sparse + sparsity.ln()]);
let mut d = penalty
.hessian_diag(target.view(), rho_view.view())
.ok_or_else(|| {
"softmax assignment log-strength hessian diag unavailable".to_string()
})?;
mask_fixed_logit_entries(assignment, &mut d);
Ok(d)
}
AssignmentMode::ThresholdGate {
temperature,
threshold,
} => {
let sparsity_strength = rho.lambda_sparse();
let inv_tau = 1.0 / temperature;
let inv_tau2 = inv_tau * inv_tau;
let k = assignment.k_atoms();
let mut d = Array1::<f64>::zeros(target.len());
for idx in 0..target.len() {
if assignment.logit_is_fixed(idx % k) {
continue;
}
let logit = target[idx];
if !jumprelu_in_optimization_band(logit, threshold, temperature) {
continue;
}
let activation = gam_linalg::utils::stable_logistic((logit - threshold) * inv_tau);
let slope = activation * (1.0 - activation);
let w_row = row_weights.map_or(1.0, |w| w[idx / k]);
d[idx] = w_row * sparsity_strength * slope * (1.0 - 2.0 * activation) * inv_tau2;
}
Ok(d)
}
AssignmentMode::IBPMap {
temperature, alpha, ..
} => {
let (penalty, rho_view) =
ibp_prior_penalty(assignment, rho, alpha, temperature, row_weights);
let mut d = if penalty.learnable_alpha {
penalty.hessian_diag_log_alpha_derivative(target.view(), rho_view.view())
} else {
penalty
.hessian_diag(target.view(), rho_view.view())
.ok_or_else(|| {
"IBP assignment log-strength hessian diag unavailable".to_string()
})?
};
mask_fixed_logit_entries(assignment, &mut d);
Ok(d)
}
AssignmentMode::TopK { .. } => Ok(Array1::<f64>::zeros(target.len())),
}
}
fn mask_fixed_logit_entries(assignment: &SaeAssignment, arr: &mut Array1<f64>) {
if !(assignment.has_ungated() || assignment.routing_is_frozen()) {
return;
}
let k = assignment.k_atoms();
for idx in 0..arr.len() {
if assignment.logit_is_fixed(idx % k) {
arr[idx] = 0.0;
}
}
}
pub fn assignment_prior_log_strength_target_mixed(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
) -> Result<Array1<f64>, String> {
assignment_prior_log_strength_target_mixed_weighted(assignment, rho, None)
}
pub(crate) fn assignment_prior_log_strength_target_mixed_weighted(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
row_weights: Option<&[f64]>,
) -> Result<Array1<f64>, String> {
for row in 0..assignment.n_obs() {
validate_finite_logits(assignment.logits.row(row), row)?;
}
let target = flat_logits(assignment.logits.view());
if matches!(assignment.mode, AssignmentMode::Softmax { .. }) && assignment.k_atoms() == 1 {
return Ok(Array1::<f64>::zeros(target.len()));
}
if assignment.routing_is_frozen() {
return Ok(Array1::<f64>::zeros(target.len()));
}
match assignment.mode {
AssignmentMode::IBPMap {
temperature, alpha, ..
} if assignment.effective_alpha_is_learnable() => {
let (penalty, rho_view) =
ibp_prior_penalty(assignment, rho, alpha, temperature, row_weights);
let mut d = penalty.log_alpha_target_mixed_derivative(target.view(), rho_view.view());
mask_fixed_logit_entries(assignment, &mut d);
Ok(d)
}
_ => Ok(assignment_prior_grad_hdiag_weighted(assignment, rho, row_weights)?.0),
}
}
pub fn assignment_prior_grad_hdiag(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
) -> Result<(Array1<f64>, Array1<f64>), String> {
assignment_prior_grad_hdiag_weighted(assignment, rho, None)
}
pub(crate) fn assignment_prior_grad_hdiag_weighted(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
row_weights: Option<&[f64]>,
) -> Result<(Array1<f64>, Array1<f64>), String> {
for row in 0..assignment.n_obs() {
validate_finite_logits(assignment.logits.row(row), row)?;
}
let target = flat_logits(assignment.logits.view());
let mut grad = Array1::<f64>::zeros(target.len());
let mut diag = Array1::<f64>::zeros(target.len());
if matches!(assignment.mode, AssignmentMode::Softmax { .. }) && assignment.k_atoms() == 1 {
return Ok((grad, diag));
}
let (sparsity_grad, sparsity_diag) = match assignment.mode {
AssignmentMode::Softmax {
temperature,
sparsity,
} => {
let penalty = SoftmaxAssignmentSparsityPenalty::new(assignment.k_atoms(), temperature)
.with_row_weights(row_weights);
let rho_view = Array1::from_vec(vec![rho.log_lambda_sparse + sparsity.ln()]);
let g = penalty.grad_target(target.view(), rho_view.view());
let d = penalty
.hessian_diag(target.view(), rho_view.view())
.ok_or_else(|| "softmax assignment hessian diag unavailable".to_string())?;
(g, d)
}
AssignmentMode::IBPMap {
temperature, alpha, ..
} => {
let (penalty, rho_view) =
ibp_prior_penalty(assignment, rho, alpha, temperature, row_weights);
let g = penalty.grad_target(target.view(), rho_view.view());
let d = penalty
.hessian_diag(target.view(), rho_view.view())
.ok_or_else(|| "IBP assignment hessian diag unavailable".to_string())?;
(g, d)
}
AssignmentMode::ThresholdGate {
temperature,
threshold,
} => {
let sparsity_strength = rho.lambda_sparse();
let inv_tau = 1.0 / temperature;
let inv_tau2 = inv_tau * inv_tau;
let k = assignment.k_atoms();
let mut g = Array1::<f64>::zeros(target.len());
let mut d = Array1::<f64>::zeros(target.len());
for idx in 0..target.len() {
let logit = target[idx];
if !jumprelu_in_optimization_band(logit, threshold, temperature) {
continue;
}
let activation = gam_linalg::utils::stable_logistic((logit - threshold) * inv_tau);
let slope = activation * (1.0 - activation);
let w_row = row_weights.map_or(1.0, |w| w[idx / k]);
g[idx] = w_row * sparsity_strength * slope * inv_tau;
d[idx] = w_row * sparsity_strength * slope * (1.0 - 2.0 * activation) * inv_tau2;
}
(g, d)
}
AssignmentMode::TopK { .. } => (
Array1::<f64>::zeros(target.len()),
Array1::<f64>::zeros(target.len()),
),
};
grad += &sparsity_grad;
diag += &sparsity_diag;
if assignment.has_ungated() || assignment.routing_is_frozen() {
let k = assignment.k_atoms();
for idx in 0..grad.len() {
if assignment.logit_is_fixed(idx % k) {
grad[idx] = 0.0;
diag[idx] = 0.0;
}
}
}
Ok((grad, diag))
}
pub fn ibp_assignment_third_channels(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
majorize: bool,
) -> Result<Option<IbpHessianDiagThirdChannels>, String> {
ibp_assignment_third_channels_weighted(assignment, rho, majorize, None)
}
pub(crate) fn ibp_assignment_third_channels_weighted(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
majorize: bool,
row_weights: Option<&[f64]>,
) -> Result<Option<IbpHessianDiagThirdChannels>, String> {
let AssignmentMode::IBPMap {
temperature, alpha, ..
} = assignment.mode
else {
return Ok(None);
};
for row in 0..assignment.n_obs() {
validate_finite_logits(assignment.logits.row(row), row)?;
}
let target = flat_logits(assignment.logits.view());
let (penalty, rho_view) = ibp_prior_penalty(assignment, rho, alpha, temperature, row_weights);
let mut channels =
penalty.hessian_diag_logit_third_channels(target.view(), rho_view.view(), majorize);
if assignment.has_ungated() || assignment.routing_is_frozen() {
let k = channels.k_max;
for idx in 0..channels.z_jac.len() {
if assignment.logit_is_fixed(idx % k) {
channels.z_jac[idx] = 0.0;
channels.local_logit_third[idx] = 0.0;
channels.m_channel[idx] = 0.0;
channels.logit_curvature[idx] = 0.0;
}
}
for atom in 0..k {
if assignment.logit_is_fixed(atom) {
channels.cross_row_d[atom] = 0.0;
channels.cross_row_dd[atom] = 0.0;
}
}
}
Ok(Some(channels))
}
pub fn select_hybrid_atom_parameterization(
manifold: &LatentManifold,
curved: Option<HybridAtomCandidate>,
linear: HybridAtomCandidate,
) -> HybridAtomChoice {
let curved = if manifold.is_euclidean() {
None
} else {
curved
};
let candidates: Vec<HybridAtomCandidate> = match curved {
Some(c) => vec![linear, c],
None => vec![linear],
};
select_hybrid_atom(&candidates).expect("hybrid atom slot always has the linear candidate")
}
#[cfg(test)]
mod ibp_prior_614_tests {
use super::*;
fn ratio(alpha: f64) -> f64 {
alpha / (alpha + 1.0)
}
#[test]
fn first_atom_is_shrunk_not_unity() {
for &alpha in &[0.1_f64, 0.5, 1.0, 2.0, 5.0] {
let prior = ordered_prior_means(8, OrderedPriorSchedule::Geometric { alpha });
let r = ratio(alpha);
assert!(
(prior[0] - r).abs() < 1e-12,
"π_0 must be the single-stick mean α/(α+1)={r} (was the unshrunk 1.0 in #614); got {}",
prior[0]
);
assert!(
prior[0] < 1.0,
"first atom must be shrunk (π_0<1) for alpha={alpha}; got {}",
prior[0]
);
}
}
#[test]
fn prior_is_consistent_geometric_product_mean() {
for &alpha in &[0.3_f64, 1.0, 4.0] {
let k = 12;
let prior = ordered_prior_means(k, OrderedPriorSchedule::Geometric { alpha });
let r = ratio(alpha);
for j in 0..k {
let expected = r.powi((j + 1) as i32);
assert!(
(prior[j] - expected).abs() < 1e-12 * expected.max(1.0),
"alpha={alpha} π_{j}: expected {expected}, got {}",
prior[j]
);
}
for j in 1..k {
assert!(
prior[j] < prior[j - 1],
"alpha={alpha}: prior must strictly decrease at index {j}"
);
}
}
}
#[test]
fn alpha_behaves_as_concentration() {
let lo = ordered_prior_means(8, OrderedPriorSchedule::Geometric { alpha: 0.5 });
let hi = ordered_prior_means(8, OrderedPriorSchedule::Geometric { alpha: 5.0 });
assert!(
hi[0] > lo[0],
"larger alpha must raise π_0 (concentration): {} vs {}",
hi[0],
lo[0]
);
assert!(
hi[4] > lo[4],
"larger alpha must put more mass in the tail: {} vs {}",
hi[4],
lo[4]
);
}
}
#[cfg(test)]
mod hybrid_split_tests {
use super::*;
use gam_solve::evidence::HybridAtomParam;
#[test]
fn flat_chart_drops_curved_candidate_and_keeps_linear() {
let linear = HybridAtomCandidate::linear(100.0, 2);
let curved = HybridAtomCandidate::curved(1, 1.0, 5, Some(2.0));
let choice =
select_hybrid_atom_parameterization(&LatentManifold::Euclidean, Some(curved), linear);
assert!(choice.param.is_linear());
}
#[test]
fn curveable_chart_selects_curved_when_turning_pays() {
let linear = HybridAtomCandidate::linear(100.0, 2);
let curved = HybridAtomCandidate::curved(1, 70.0, 5, Some(2.0 * std::f64::consts::PI));
let choice = select_hybrid_atom_parameterization(
&LatentManifold::Circle {
period: 2.0 * std::f64::consts::PI,
},
Some(curved),
linear,
);
assert_eq!(choice.param, HybridAtomParam::Curved { latent_dim: 1 });
}
#[test]
fn curveable_chart_falls_back_to_linear_when_no_curved_candidate() {
let linear = HybridAtomCandidate::linear(33.0, 2);
let choice = select_hybrid_atom_parameterization(
&LatentManifold::Circle {
period: 2.0 * std::f64::consts::PI,
},
None,
linear,
);
assert!(choice.param.is_linear());
assert_eq!(choice.num_parameters, 2);
}
}
#[cfg(test)]
mod frozen_routing_1033_tests {
use super::*;
fn ibp_assignment(n: usize, k: usize) -> SaeAssignment {
let logits = Array2::from_shape_fn((n, k), |(i, kk)| {
0.3 + 0.05 * (i as f64) - 0.1 * (kk as f64)
});
let coords: Vec<Array2<f64>> = (0..k)
.map(|_| Array2::from_shape_fn((n, 1), |(i, _)| (i as f64) * 0.1))
.collect();
SaeAssignment::from_blocks_with_mode(
logits,
coords,
AssignmentMode::ibp_map(0.5, 1.0, false),
)
.unwrap()
}
#[test]
fn frozen_routing_decouples_gates_from_logit_updates_1033() {
let (n, k) = (6usize, 3usize);
let mut a = ibp_assignment(n, k)
.freeze_routing_from_current_logits()
.unwrap();
assert!(a.routing_is_frozen());
let before: Vec<Array1<f64>> = (0..n).map(|r| a.try_assignments_row(r).unwrap()).collect();
a.logits.mapv_inplace(|v| v + 5.0);
let after: Vec<Array1<f64>> = (0..n).map(|r| a.try_assignments_row(r).unwrap()).collect();
for r in 0..n {
for kk in 0..k {
assert_eq!(
before[r][kk], after[r][kk],
"row {r} atom {kk}: frozen-routing gate must be UNCHANGED by a free-logit \
update (decoupled from inner-fit drift); {} vs {}",
before[r][kk], after[r][kk]
);
}
}
}
#[test]
fn frozen_routing_gates_are_rho_invariant_1033() {
let (n, k) = (5usize, 2usize);
let a = ibp_assignment(n, k)
.freeze_routing_from_current_logits()
.unwrap();
for r in 0..n {
let ga = a.try_assignments_row(r).unwrap();
let gb = a.try_assignments_row(r).unwrap();
for kk in 0..k {
assert_eq!(
ga[kk], gb[kk],
"row {r} atom {kk}: frozen-routing gate must be ρ-INVARIANT (the n-independence \
lever); {} at ρ_a vs {} at ρ_b",
ga[kk], gb[kk]
);
}
}
}
#[test]
fn frozen_routing_fixes_all_logits_and_thaw_restores_free_path_1033() {
let (n, k) = (4usize, 3usize);
let mut a = ibp_assignment(n, k)
.freeze_routing_from_current_logits()
.unwrap();
let mask = a.fixed_logit_mask();
assert_eq!(mask.len(), k);
assert!(
mask.iter().all(|&f| f),
"frozen routing must fix ALL logits"
);
for kk in 0..k {
assert!(
a.logit_is_fixed(kk),
"atom {kk} logit must be fixed under frozen routing"
);
}
a.thaw_routing();
assert!(!a.routing_is_frozen());
assert!(
a.fixed_logit_mask().iter().all(|&f| !f),
"thaw must restore the free-logit path"
);
}
#[test]
fn frozen_routing_rejects_softmax_1033() {
let (n, k) = (4usize, 3usize);
let logits = Array2::from_shape_fn((n, k), |(i, kk)| 0.1 * (i as f64) - 0.05 * (kk as f64));
let coords: Vec<Array2<f64>> = (0..k)
.map(|_| Array2::from_shape_fn((n, 1), |(i, _)| (i as f64) * 0.1))
.collect();
let a = SaeAssignment::from_blocks_with_mode(logits, coords, AssignmentMode::softmax(1.0))
.unwrap();
assert!(
a.freeze_routing_from_current_logits().is_err(),
"frozen routing under Softmax must be rejected (simplex entropy-majorizer coupling)"
);
}
}
#[cfg(test)]
mod support_measure_tests {
use super::*;
#[test]
fn support_measure_matches_hard_and_diffuse_semantics() {
let hard_weights = Array1::from_vec(vec![1.0, 1.0, 0.0, 1.0, 0.0]);
let hard = SupportMeasure::from_weights(0, hard_weights).unwrap();
assert_eq!(hard.mass(), 3.0);
assert_eq!(hard.fisher_n(), 3.0);
assert_eq!(hard.ess(), 3.0);
assert_eq!(hard.positive_rows(), vec![0usize, 1, 3]);
let from_owners = SupportMeasure::from_argmax_owners(&[0, 0, 1, 0, 1], 0, 2).unwrap();
assert_eq!(from_owners.mass(), hard.mass());
assert_eq!(from_owners.fisher_n(), hard.fisher_n());
assert_eq!(from_owners.ess(), hard.ess());
assert_eq!(from_owners.positive_rows(), hard.positive_rows());
let diffuse_weights = Array1::from_vec(vec![0.5, 0.5, 0.5, 0.5]);
let diffuse = SupportMeasure::from_weights(1, diffuse_weights).unwrap();
assert_eq!(diffuse.mass(), 2.0);
assert_eq!(diffuse.fisher_n(), 1.0);
assert_eq!(diffuse.ess(), 4.0);
}
#[test]
fn support_measure_reads_assignment_column() {
let assignments =
Array2::from_shape_vec((3, 2), vec![0.8, 0.2, 0.4, 0.6, 0.0, 1.0]).unwrap();
let support = SupportMeasure::from_assignment_matrix(assignments.view(), 1).unwrap();
assert!((support.mass() - 1.8).abs() < 1e-12);
assert!((support.fisher_n() - 1.4).abs() < 1e-12);
assert!((support.ess() - (1.8_f64 * 1.8 / 1.4)).abs() < 1e-12);
assert_eq!(support.positive_rows(), vec![0usize, 1, 2]);
}
}
#[cfg(test)]
mod fill_into_buffer_1557_tests {
use super::*;
fn build(n: usize, k: usize, mode: AssignmentMode) -> SaeAssignment {
let logits = Array2::from_shape_fn((n, k), |(i, kk)| {
0.37 + 0.11 * (i as f64) - 0.23 * (kk as f64)
});
let coords: Vec<Array2<f64>> = (0..k)
.map(|_| Array2::from_shape_fn((n, 1), |(i, _)| 0.1 + 0.05 * (i as f64)))
.collect();
SaeAssignment::from_blocks_with_mode(logits, coords, mode).unwrap()
}
fn assert_into_matches_alloc(a: &SaeAssignment) {
let n = a.n_obs();
let k = a.k_atoms();
let mut scratch = vec![f64::NAN; k];
for row in 0..n {
let allocated = a.try_assignments_row(row).unwrap();
for s in scratch.iter_mut() {
*s = f64::NAN;
}
a.try_assignments_row_into(row, &mut scratch).unwrap();
assert_eq!(allocated.len(), k);
for kk in 0..k {
assert_eq!(
allocated[kk], scratch[kk],
"row {row} atom {kk}: _into must be BIT-IDENTICAL to the allocating \
try_assignments_row; got {} vs {}",
allocated[kk], scratch[kk]
);
}
}
}
#[test]
fn softmax_into_is_bit_identical() {
assert_into_matches_alloc(&build(7, 4, AssignmentMode::softmax(0.8)));
}
#[test]
fn ibp_map_into_is_bit_identical() {
assert_into_matches_alloc(&build(7, 5, AssignmentMode::ibp_map(0.6, 1.3, false)));
assert_into_matches_alloc(&build(7, 5, AssignmentMode::ibp_map(0.6, 1.3, true)));
}
#[test]
fn jumprelu_into_is_bit_identical() {
assert_into_matches_alloc(&build(7, 5, AssignmentMode::threshold_gate(0.9, 0.2)));
}
#[test]
fn ungated_into_is_bit_identical() {
let a = build(6, 4, AssignmentMode::ibp_map(0.6, 1.1, false))
.with_ungated(vec![false, true, false, true])
.unwrap();
assert_into_matches_alloc(&a);
let j = build(6, 4, AssignmentMode::threshold_gate(0.9, 0.15))
.with_ungated(vec![true, false, true, false])
.unwrap();
assert_into_matches_alloc(&j);
}
#[test]
fn k_equals_one_into_is_bit_identical() {
assert_into_matches_alloc(&build(5, 1, AssignmentMode::softmax(1.0)));
assert_into_matches_alloc(&build(5, 1, AssignmentMode::ibp_map(0.7, 1.0, false)));
assert_into_matches_alloc(&build(5, 1, AssignmentMode::threshold_gate(0.8, 0.1)));
}
}
#[cfg(test)]
mod ibp_eb_alpha_f1_tests {
use super::*;
fn geometric_a(theta: f64, k: usize) -> f64 {
let alpha = theta.exp();
let rho = alpha / (alpha + 1.0);
let mu = rho.powf((k + 1) as f64);
mu / (1.0 - mu)
}
fn brute_marginal(occupancy: &[f64], n_obs: f64, theta: f64) -> f64 {
let mut l = 0.0;
for (k, &m) in occupancy.iter().enumerate() {
let a = geometric_a(theta, k);
l += statrs::function::gamma::ln_gamma(m + a)
- statrs::function::gamma::ln_gamma(n_obs + a + 1.0)
+ a.ln();
}
l
}
#[test]
fn geometric_delegate_is_bit_identical() {
for &alpha in &[0.3_f64, 1.0, 4.5, 37.0] {
for &k in &[1usize, 3, 8, 64] {
let old = ordered_prior_means(k, OrderedPriorSchedule::Geometric { alpha });
let via = ordered_prior_means(k, OrderedPriorSchedule::Geometric { alpha });
for j in 0..k {
assert_eq!(
old[j], via[j],
"geometric prior drift at alpha={alpha}, k={j}"
);
}
}
}
}
#[test]
fn power_law_schedule_round_trips() {
let (c, s, k0) = (0.9_f64, 1.2_f64, 1.0_f64);
let k = 32usize;
let mu = ordered_prior_means(k, OrderedPriorSchedule::PowerLaw { c, s, k0 });
for j in 0..k {
let expected = (c / ((j as f64) + k0).powf(s)).clamp(f64::MIN_POSITIVE, 1.0);
assert!(
(mu[j] - expected).abs() <= 1e-12 * expected.max(1.0),
"power-law mismatch at k={j}: {} vs {expected}",
mu[j]
);
assert!(mu[j] > 0.0 && mu[j] <= 1.0, "μ_{j}={} out of (0,1]", mu[j]);
if j > 0 {
assert!(mu[j] <= mu[j - 1], "power-law not decreasing at k={j}");
}
}
}
#[test]
fn eb_alpha_score_matches_brute_digamma() {
let n_obs = 500.0_f64;
let occupancy = vec![300.0_f64, 120.0, 60.0, 30.0, 12.0, 5.0, 2.0, 1.0];
let h = 1e-6_f64;
for &alpha in &[0.4_f64, 1.0, 3.0, 12.0] {
let theta = alpha.ln();
let (s, hess) = ibp_eb_alpha_score_hess(&occupancy, n_obs, alpha);
let l_plus = brute_marginal(&occupancy, n_obs, theta + h);
let l_minus = brute_marginal(&occupancy, n_obs, theta - h);
let s_fd = (l_plus - l_minus) / (2.0 * h);
assert!(
(s - s_fd).abs() <= 1e-4 * (1.0 + s_fd.abs()),
"score mismatch at alpha={alpha}: analytic {s} vs FD {s_fd}"
);
let (s_plus, _) = ibp_eb_alpha_score_hess(&occupancy, n_obs, (theta + h).exp());
let (s_minus, _) = ibp_eb_alpha_score_hess(&occupancy, n_obs, (theta - h).exp());
let h_fd = (s_plus - s_minus) / (2.0 * h);
assert!(
(hess - h_fd).abs() <= 1e-3 * (1.0 + h_fd.abs()),
"curvature mismatch at alpha={alpha}: analytic {hess} vs FD {h_fd}"
);
}
}
#[test]
fn eb_marginal_core_is_schedule_agnostic() {
let n_obs = 400.0_f64;
let occupancy = vec![200.0_f64, 90.0, 40.0, 18.0, 7.0, 3.0];
let alpha = 2.5_f64;
let theta = alpha.ln();
let dh = 1e-6_f64;
let (a, da): (Vec<f64>, Vec<f64>) = (0..occupancy.len())
.map(|k| {
let a0 = geometric_a(theta, k);
let da = (geometric_a(theta + dh, k) - geometric_a(theta - dh, k)) / (2.0 * dh);
(a0, da)
})
.unzip();
let core = ibp_eb_marginal_score(&occupancy, n_obs, &a, &da);
let (s, _) = ibp_eb_alpha_score_hess(&occupancy, n_obs, alpha);
assert!(
(core - s).abs() <= 1e-5 * (1.0 + s.abs()),
"schedule-agnostic core {core} != geometric score {s}"
);
}
#[test]
fn eb_fixed_point_is_stationary_and_moves() {
let n_obs = 1000.0_f64;
let occupancy = vec![500.0_f64; 8];
let alpha_star = ibp_eb_geometric_alpha_fixed_point(&occupancy, n_obs, 1.0);
assert!(
alpha_star > 5.0,
"flat occupancy must raise α; got {alpha_star}"
);
let (s, _) = ibp_eb_alpha_score_hess(&occupancy, n_obs, alpha_star);
assert!(
s.abs() < 1e-4,
"score not stationary at α*={alpha_star}: S={s}"
);
let steep: Vec<f64> = (0..8).map(|k| n_obs * 0.5_f64.powi(k as i32 + 1)).collect();
let alpha_low = ibp_eb_geometric_alpha_fixed_point(&steep, n_obs, 20.0);
assert!(
alpha_low < 5.0,
"steep occupancy must lower α; got {alpha_low}"
);
}
fn learnable_ibp(logits: Array2<f64>, alpha_base: f64) -> SaeAssignment {
let (n, k) = logits.dim();
let coords: Vec<Array2<f64>> = (0..k).map(|_| Array2::<f64>::zeros((n, 1))).collect();
SaeAssignment::from_blocks_with_mode(
logits,
coords,
AssignmentMode::ibp_map(1.0, alpha_base, true),
)
.unwrap()
}
#[test]
fn eb_log_alpha_step_moves_for_learnable_ibp_only() {
let n = 300usize;
let k = 8usize;
let logits = Array2::from_shape_fn((n, k), |(_, kk)| 2.0 * kk as f64);
let assign = learnable_ibp(logits.clone(), 1.0);
let rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::<f64>::zeros(1); k]);
let step = assign
.ibp_eb_log_alpha_step(&rho)
.unwrap()
.expect("learnable IBP must yield an EB step");
assert!(step.is_finite(), "step must be finite, got {step}");
assert!(
step > 1e-3,
"flattened occupancy must MOVE λ_sparse up (raise α); got Δ={step}"
);
let sm_coords: Vec<Array2<f64>> = (0..k).map(|_| Array2::<f64>::zeros((n, 1))).collect();
let softmax = SaeAssignment::from_blocks_with_mode(
logits.clone(),
sm_coords,
AssignmentMode::softmax(1.0),
)
.unwrap();
assert!(
softmax.ibp_eb_log_alpha_step(&rho).unwrap().is_none(),
"softmax sparsity prior has no EB α M-step"
);
let mut pinned = learnable_ibp(logits, 1.0);
pinned.set_ibp_alpha_override(Some(3.0));
assert!(
pinned.ibp_eb_log_alpha_step(&rho).unwrap().is_none(),
"override-pinned α must not take the EB M-step"
);
}
}