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, OrderedBetaBernoulliHessianDiagThirdChannels,
OrderedBetaBernoulliLogitAdjointData, OrderedBetaBernoulliPenalty,
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;
#[derive(Debug, Clone, Copy)]
pub enum AssignmentMode {
Softmax { temperature: f64, sparsity: f64 },
OrderedBetaBernoulli {
temperature: f64,
alpha: f64,
learnable_alpha: bool,
},
ThresholdGate { temperature: f64, threshold: f64 },
TopK { k: 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 ordered_beta_bernoulli(temperature: f64, alpha: f64, learnable_alpha: bool) -> Self {
Self::OrderedBetaBernoulli {
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::OrderedBetaBernoulli { 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::OrderedBetaBernoulli { 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::OrderedBetaBernoulli { alpha, .. } => {
if !(alpha.is_finite() && alpha > 0.0) {
return Err(format!(
"AssignmentMode::OrderedBetaBernoulli: 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_ordered_beta_bernoulli_alpha(
&self,
rho: &SaeManifoldRho,
per_fit_override: Option<f64>,
) -> Option<f64> {
match *self {
AssignmentMode::OrderedBetaBernoulli {
alpha,
learnable_alpha,
..
} => Some(if let Some(over) = per_fit_override {
over
} else if learnable_alpha {
resolve_learnable_weight(alpha, rho.log_lambda_sparse)
.expect("ordered Beta--Bernoulli rho must be validated before resolution")
} else {
alpha
}),
_ => None,
}
}
}
#[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 ordered_beta_bernoulli_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,
ordered_beta_bernoulli_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 ordered Beta--Bernoulli and threshold gate, 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 ordered Beta--Bernoulli or threshold gate"
.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 ordered Beta--Bernoulli or threshold gate"
.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 ordered Beta--Bernoulli or threshold gate"
.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::OrderedBetaBernoulli { .. } | 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_ordered_beta_bernoulli_alpha(
&self,
rho: &SaeManifoldRho,
) -> Option<f64> {
self.mode
.resolved_ordered_beta_bernoulli_alpha(rho, self.ordered_beta_bernoulli_alpha_override)
}
pub(crate) fn effective_alpha_is_learnable(&self) -> bool {
match self.mode {
AssignmentMode::OrderedBetaBernoulli {
learnable_alpha, ..
} => learnable_alpha && self.ordered_beta_bernoulli_alpha_override.is_none(),
_ => false,
}
}
pub(crate) fn validate_rho_domain(&self, rho: &SaeManifoldRho) -> Result<(), String> {
rho.validate_log_strength_domain()?;
if let AssignmentMode::OrderedBetaBernoulli {
alpha,
learnable_alpha: true,
..
} = self.mode
&& self.ordered_beta_bernoulli_alpha_override.is_none()
{
resolve_learnable_weight(alpha, rho.log_lambda_sparse).map_err(|error| {
format!("ordered Beta--Bernoulli learnable concentration: {error}")
})?;
}
Ok(())
}
pub(crate) fn learnable_alpha_rho_domain(&self) -> Result<Option<(f64, f64)>, String> {
let AssignmentMode::OrderedBetaBernoulli {
alpha,
learnable_alpha: true,
..
} = self.mode
else {
return Ok(None);
};
if self.ordered_beta_bernoulli_alpha_override.is_some() {
return Ok(None);
}
gam_terms::analytic_penalties::learnable_weight_coordinate_domain(alpha)
}
pub fn set_ordered_beta_bernoulli_alpha_override(&mut self, alpha: Option<f64>) {
self.ordered_beta_bernoulli_alpha_override = alpha;
}
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::OrderedBetaBernoulli { temperature, .. } => {
ordered_beta_bernoulli_row(routing, temperature)
}
AssignmentMode::ThresholdGate {
temperature,
threshold,
} => threshold_gate_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::OrderedBetaBernoulli { temperature, .. } => {
ordered_beta_bernoulli_row_into(routing, temperature, out)
}
AssignmentMode::ThresholdGate {
temperature,
threshold,
} => threshold_gate_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_ordered_beta_bernoulli_alpha(
&mut self,
rho: &SaeManifoldRho,
) -> bool {
let AssignmentMode::OrderedBetaBernoulli {
temperature,
alpha,
learnable_alpha: true,
} = self.mode
else {
return false;
};
let resolved_alpha = resolve_learnable_weight(alpha, rho.log_lambda_sparse)
.expect("ordered Beta--Bernoulli rho must be validated before persistence");
self.mode = AssignmentMode::OrderedBetaBernoulli {
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::OrderedBetaBernoulli { temperature, .. } => {
ordered_beta_bernoulli_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;
}
}
pub fn default_ordered_beta_bernoulli_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)
}
pub fn ordered_beta_bernoulli_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
}
pub fn threshold_gate_row(
logits: ArrayView1<'_, f64>,
temperature: f64,
threshold: 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] - threshold) / temperature);
}
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 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()
}
}
#[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"
);
}
}
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 ordered_beta_bernoulli_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 threshold_gate_row_into(
logits: ArrayView1<'_, f64>,
temperature: f64,
threshold: f64,
out: &mut [f64],
) {
for i in 0..logits.len() {
out[i] = gam_linalg::utils::stable_logistic((logits[i] - threshold) / temperature);
}
}
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::OrderedBetaBernoulli { 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) {
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 ordered_beta_bernoulli_prior_penalty(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
base_alpha: f64,
temperature: f64,
row_weights: Option<&[f64]>,
) -> Result<(OrderedBetaBernoulliPenalty, Array1<f64>), String> {
let learnable = assignment.effective_alpha_is_learnable();
let alpha_eff = if learnable {
base_alpha
} else {
assignment
.resolved_ordered_beta_bernoulli_alpha(rho)
.unwrap_or(base_alpha)
};
let mut penalty =
OrderedBetaBernoulliPenalty::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)
};
Ok((penalty, rho_view))
}
pub(crate) fn ordered_beta_bernoulli_logit_adjoint_data_weighted(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
row_weights: Option<&[f64]>,
) -> Result<Option<OrderedBetaBernoulliLogitAdjointData>, String> {
assignment.validate_rho_domain(rho)?;
let AssignmentMode::OrderedBetaBernoulli {
temperature, alpha, ..
} = assignment.mode
else {
return Ok(None);
};
if assignment.routing_is_frozen() {
return Ok(None);
}
for row in 0..assignment.n_obs() {
validate_finite_logits(assignment.logits.row(row), row)?;
}
let (penalty, rho_view) =
ordered_beta_bernoulli_prior_penalty(assignment, rho, alpha, temperature, row_weights)?;
let target = flat_logits(assignment.logits.view());
Ok(Some(
penalty.logit_theta_adjoint_data(target.view(), rho_view.view()),
))
}
pub(crate) fn ordered_beta_bernoulli_exact_hessian_minus_majorizer_hvp_weighted(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
row_weights: Option<&[f64]>,
direction: ArrayView1<'_, f64>,
) -> Result<Array1<f64>, String> {
assignment.validate_rho_domain(rho)?;
let AssignmentMode::OrderedBetaBernoulli {
temperature, alpha, ..
} = assignment.mode
else {
return Err(
"ordered Beta--Bernoulli exact-Hessian correction requires ordered assignment mode"
.to_string(),
);
};
let target = flat_logits(assignment.logits.view());
if direction.len() != target.len() {
return Err(format!(
"ordered Beta--Bernoulli exact-Hessian direction has length {}; expected {}",
direction.len(),
target.len()
));
}
if !direction.iter().all(|value| value.is_finite()) {
return Err("ordered Beta--Bernoulli exact-Hessian direction must be finite".to_string());
}
if assignment.routing_is_frozen() {
return Ok(Array1::<f64>::zeros(target.len()));
}
for row in 0..assignment.n_obs() {
validate_finite_logits(assignment.logits.row(row), row)?;
}
let (penalty, rho_view) =
ordered_beta_bernoulli_prior_penalty(assignment, rho, alpha, temperature, row_weights)?;
let mut delta = penalty.hvp(target.view(), rho_view.view(), direction);
let channels = penalty.psd_majorizer_logit_third_channels(target.view(), rho_view.view());
for index in 0..delta.len() {
delta[index] -= channels.diagonal_term[index].max(0.0) * direction[index];
}
Ok(delta)
}
#[cfg(test)]
mod ordered_beta_bernoulli_exact_hessian_tests {
use super::*;
#[test]
fn exact_hessian_minus_majorizer_hvp_matches_gradient_fd_and_keeps_cross_row_term() {
let n = 4usize;
let k = 2usize;
let logits =
Array2::from_shape_vec((n, k), vec![0.2, -0.3, 0.7, -0.1, 0.4, 0.5, -0.2, 0.6])
.unwrap();
let coords = vec![Array2::<f64>::zeros((n, 1)); k];
let assignment = SaeAssignment::from_blocks_with_mode(
logits,
coords,
AssignmentMode::ordered_beta_bernoulli(0.8, 1.7, false),
)
.unwrap();
let rho = SaeManifoldRho::new(1.3_f64.ln(), 0.0, vec![Array1::zeros(1); k]);
let mut direction = Array1::<f64>::zeros(n * k);
direction[0] = 0.7;
let analytic = ordered_beta_bernoulli_exact_hessian_minus_majorizer_hvp_weighted(
&assignment,
&rho,
None,
direction.view(),
)
.unwrap();
let (penalty, rho_view) =
ordered_beta_bernoulli_prior_penalty(&assignment, &rho, 1.7, 0.8, None).unwrap();
let target = flat_logits(assignment.logits.view());
let step = 1.0e-6;
let plus = &target + &(step * &direction);
let minus = &target - &(step * &direction);
let gradient_plus = penalty.grad_target(plus.view(), rho_view.view());
let gradient_minus = penalty.grad_target(minus.view(), rho_view.view());
let channels = penalty.psd_majorizer_logit_third_channels(target.view(), rho_view.view());
for index in 0..analytic.len() {
let exact_fd = (gradient_plus[index] - gradient_minus[index]) / (2.0 * step);
let expected = exact_fd - channels.diagonal_term[index].max(0.0) * direction[index];
assert!(
(analytic[index] - expected).abs() <= 2.0e-7,
"index {index}: analytic A-B={} expected={} exact_fd={exact_fd}",
analytic[index],
expected,
);
}
assert!(
analytic[2].abs() > 1.0e-6 && analytic[4].abs() > 1.0e-6,
"a one-row direction must produce the exact cross-row rank-one action: {analytic:?}"
);
}
}
pub fn assignment_prior_value(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
) -> Result<f64, String> {
assignment_prior_value_weighted(assignment, rho, None)
}
pub(crate) fn assignment_prior_value_weighted(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
row_weights: Option<&[f64]>,
) -> Result<f64, String> {
assignment.validate_rho_domain(rho)?;
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(0.0);
}
if assignment.routing_is_frozen() {
return Ok(0.0);
}
Ok(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::OrderedBetaBernoulli {
temperature, alpha, ..
} => {
let (penalty, rho_view) = ordered_beta_bernoulli_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;
}
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,
) -> Result<f64, String> {
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]>,
) -> Result<f64, String> {
assignment.validate_rho_domain(rho)?;
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(0.0);
}
if assignment.routing_is_frozen() {
return Ok(0.0);
}
Ok(match assignment.mode {
AssignmentMode::Softmax { .. } | AssignmentMode::ThresholdGate { .. } => {
return assignment_prior_value_weighted(assignment, rho, row_weights);
}
AssignmentMode::OrderedBetaBernoulli {
temperature, alpha, ..
} => {
let (penalty, rho_view) = ordered_beta_bernoulli_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> {
assignment.validate_rho_domain(rho)?;
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];
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::OrderedBetaBernoulli {
temperature, alpha, ..
} => {
let (penalty, rho_view) = ordered_beta_bernoulli_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(|| {
"ordered Beta--Bernoulli 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> {
assignment.validate_rho_domain(rho)?;
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::OrderedBetaBernoulli {
temperature, alpha, ..
} if assignment.effective_alpha_is_learnable() => {
let (penalty, rho_view) = ordered_beta_bernoulli_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> {
assignment.validate_rho_domain(rho)?;
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::OrderedBetaBernoulli {
temperature, alpha, ..
} => {
let (penalty, rho_view) = ordered_beta_bernoulli_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(|| {
"ordered Beta--Bernoulli 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];
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 ordered_beta_bernoulli_psd_majorizer_third_channels(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
) -> Result<Option<OrderedBetaBernoulliHessianDiagThirdChannels>, String> {
ordered_beta_bernoulli_psd_majorizer_third_channels_weighted(assignment, rho, None)
}
pub(crate) fn ordered_beta_bernoulli_psd_majorizer_third_channels_weighted(
assignment: &SaeAssignment,
rho: &SaeManifoldRho,
row_weights: Option<&[f64]>,
) -> Result<Option<OrderedBetaBernoulliHessianDiagThirdChannels>, String> {
assignment.validate_rho_domain(rho)?;
let AssignmentMode::OrderedBetaBernoulli {
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) =
ordered_beta_bernoulli_prior_penalty(assignment, rho, alpha, temperature, row_weights)?;
let mut channels = penalty.psd_majorizer_logit_third_channels(target.view(), rho_view.view());
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.diagonal_term[idx] = 0.0;
}
}
for atom in 0..k {
if assignment.logit_is_fixed(atom) {
channels.mass_hessian_coefficient[atom] = 0.0;
channels.mass_hessian_log_alpha_derivative[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 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 ordered_beta_bernoulli_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::ordered_beta_bernoulli(0.5, 1.0, false),
)
.unwrap()
}
#[test]
fn frozen_routing_decouples_gates_from_logit_updates_1033() {
let (n, k) = (6usize, 3usize);
let mut a = ordered_beta_bernoulli_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 = ordered_beta_bernoulli_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 = ordered_beta_bernoulli_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 ordered_alpha_domain_tests {
use super::*;
use gam_problem::{LOG_STRENGTH_MAX, LOG_STRENGTH_MIN};
fn ordered_assignment(alpha: f64) -> SaeAssignment {
SaeAssignment::from_blocks_with_mode(
Array2::<f64>::zeros((3, 2)),
vec![Array2::<f64>::zeros((3, 1)); 2],
AssignmentMode::ordered_beta_bernoulli(0.8, alpha, true),
)
.unwrap()
}
#[test]
fn learnable_ordered_alpha_tightens_sparse_rho_face_without_saturation() {
let alpha = 1.7_f64;
let assignment = ordered_assignment(alpha);
let (lower, upper) = assignment
.learnable_alpha_rho_domain()
.unwrap()
.expect("learnable ordered alpha owns the sparse rho coordinate");
assert!(upper < LOG_STRENGTH_MAX);
let legal = SaeManifoldRho::new(upper, 0.0, vec![Array1::zeros(1); 2])
.for_assignment(assignment.mode);
assignment
.validate_rho_domain(&legal)
.expect("closed effective-alpha upper face is legal");
let invalid = SaeManifoldRho::new(upper + 1.0e-6, 0.0, vec![Array1::zeros(1); 2])
.for_assignment(assignment.mode);
assert!(assignment.validate_rho_domain(&invalid).is_err());
assert_eq!(
lower,
LOG_STRENGTH_MIN - alpha.ln(),
"lower face must be shifted by the base concentration too"
);
}
}
#[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 ordered_beta_bernoulli_into_is_bit_identical() {
assert_into_matches_alloc(&build(
7,
5,
AssignmentMode::ordered_beta_bernoulli(0.6, 1.3, false),
));
assert_into_matches_alloc(&build(
7,
5,
AssignmentMode::ordered_beta_bernoulli(0.6, 1.3, true),
));
}
#[test]
fn threshold_gate_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::ordered_beta_bernoulli(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::ordered_beta_bernoulli(0.7, 1.0, false),
));
assert_into_matches_alloc(&build(5, 1, AssignmentMode::threshold_gate(0.8, 0.1)));
}
}