use super::BlockSparseConfig;
use super::block::{
RowBlockCode, block_birth_evidence_margin, frame_fixed_point_residual, gram_schmidt_rows,
relative_scalar_change, route_and_code_all, seed_frames, stable_rank_symmetric,
};
use super::block_frame::polar_tied_frame_step;
use super::residual_reservoir::ResidualReservoir;
use super::update::{DEAD_DENOM, DecoderSolveStats};
use gam_linalg::faer_ndarray::with_faer_sequential;
use ndarray::{Array2, ArrayView2, Axis};
use rayon::prelude::*;
struct RowProjection {
sum: Vec<f64>,
rss: f64,
gamma_num: f64,
gamma_den: f64,
}
fn project_coded_rows(
rows: ArrayView2<'_, f32>,
decoder: ArrayView2<'_, f32>,
codes: &[RowBlockCode],
b: usize,
gamma: f32,
) -> Vec<RowProjection> {
codes
.par_iter()
.enumerate()
.map(|(row, code)| {
let p = rows.ncols();
let mut sum = vec![0.0; p];
let mut projection = vec![0.0; p];
for (slot, &block) in code.blocks.iter().enumerate() {
if code.gates[slot] == 0.0 {
continue;
}
projection.fill(0.0);
let w = &code.projections[slot * b..(slot + 1) * b];
for (axis, &weight) in w.iter().enumerate() {
let atom = decoder.row(block as usize * b + axis);
for (value, &direction) in projection.iter_mut().zip(atom.iter()) {
*value += weight * direction as f64;
}
}
for (value, contribution) in sum.iter_mut().zip(&projection) {
*value += contribution;
}
}
let mut rss = 0.0;
let mut gamma_num = 0.0;
let mut gamma_den = 0.0;
for (&x, &projected) in rows.row(row).iter().zip(&sum) {
let residual = x as f64 - gamma as f64 * projected;
rss += residual * residual;
gamma_num += x as f64 * projected;
gamma_den += projected * projected;
}
RowProjection {
sum,
rss,
gamma_num,
gamma_den,
}
})
.collect()
}
fn block_row_postings(codes: &[RowBlockCode], blocks: usize) -> Vec<Vec<(usize, usize)>> {
let mut postings = vec![Vec::new(); blocks];
for (row, code) in codes.iter().enumerate() {
for (slot, &block) in code.blocks.iter().enumerate() {
if code.gates[slot] != 0.0 {
postings[block as usize].push((row, slot));
}
}
}
postings
}
#[derive(Clone, Copy, Debug)]
pub struct BlockShardStats {
pub rows: usize,
pub rss: f64,
pub alive_blocks: usize,
}
#[derive(Clone, Copy, Debug)]
pub struct BlockEpochStats {
pub explained_variance: f64,
pub accepted_births: usize,
pub birth_pending: bool,
pub dead: usize,
pub gamma: f32,
pub gamma_residual: f64,
pub frame_residual: f64,
pub converged: bool,
pub epoch: usize,
pub decoder_solve_stats: DecoderSolveStats,
}
struct PendingBlockBirth {
block: usize,
baseline_decoder: Array2<f32>,
baseline_gamma: f32,
baseline_rss: f64,
baseline_rows: usize,
baseline_usage: Vec<usize>,
baseline_second: Vec<Array2<f64>>,
}
pub struct BlockSparseStreamState {
config: BlockSparseConfig,
g: usize,
b: usize,
k: usize,
p: usize,
decoder: Array2<f32>,
gamma: f32,
second: Vec<Array2<f64>>, coupling: Vec<Array2<f64>>, data_cross: Vec<Array2<f64>>, data_energy: Vec<f64>, negative_bound: Vec<f64>, usage: Vec<usize>,
alive_count: usize,
gamma_num: f64,
gamma_den: f64,
col_sum: Vec<f64>,
col_sumsq: Vec<f64>,
rss: f64,
row_count: usize,
reservoir: ResidualReservoir,
prev_ev: f64,
last_ev: f64,
last_ev_residual: f64,
last_gamma_residual: f64,
last_frame_residual: f64,
epochs_run: usize,
last_accepted_births: usize,
converged: bool,
last_util: Vec<f32>,
last_stable: Vec<f32>,
last_decoder_solve_stats: DecoderSolveStats,
last_second: Vec<Array2<f64>>,
last_usage: Vec<usize>,
last_rss: f64,
last_rows: usize,
pending_birth: Option<PendingBlockBirth>,
}
pub struct BlockRankCharges {
pub block: Vec<usize>,
pub n_eff: Vec<f64>,
pub d_eff: Vec<f64>,
pub delta_deviance: Vec<f64>,
pub charge: Vec<f64>,
pub margin: Vec<f64>,
pub kept: Vec<bool>,
}
impl BlockSparseStreamState {
pub fn new(seed: ArrayView2<'_, f32>, config: &BlockSparseConfig) -> Result<Self, String> {
validate_config(config)?;
if seed.nrows() == 0 || seed.ncols() == 0 {
return Err(
"BlockSparseStream requires a non-empty seed sample (N×P) to fix P and the initial \
block frames"
.to_string(),
);
}
if !seed.iter().all(|v| v.is_finite()) {
return Err("BlockSparseStream seed sample must be finite".to_string());
}
let p = seed.ncols();
if config.block_size > p {
return Err(format!(
"BlockSparseStream block_size b={} cannot exceed P={p} (a block's b orthonormal \
rows must fit in ℝ^P)",
config.block_size
));
}
let g = config.n_blocks;
let b = config.block_size;
let k = config.block_topk.min(g).max(1);
let decoder = seed_frames(seed, g, b);
let cap = config.aux_k.saturating_mul(b).max(1);
Ok(Self {
config: *config,
g,
b,
k,
p,
decoder,
gamma: 1.0,
second: (0..g).map(|_| Array2::<f64>::zeros((b, b))).collect(),
coupling: (0..g).map(|_| Array2::<f64>::zeros((p, b))).collect(),
data_cross: (0..g).map(|_| Array2::<f64>::zeros((p, b))).collect(),
data_energy: vec![0.0; g],
negative_bound: vec![0.0; g],
usage: vec![0; g],
alive_count: 0,
gamma_num: 0.0,
gamma_den: 0.0,
col_sum: vec![0.0; p],
col_sumsq: vec![0.0; p],
rss: 0.0,
row_count: 0,
reservoir: ResidualReservoir::new(cap),
prev_ev: f64::NEG_INFINITY,
last_ev: f64::NEG_INFINITY,
last_ev_residual: f64::INFINITY,
last_gamma_residual: f64::INFINITY,
last_frame_residual: f64::INFINITY,
epochs_run: 0,
last_accepted_births: 0,
converged: false,
last_util: vec![0.0; g],
last_stable: vec![0.0; g],
last_decoder_solve_stats: DecoderSolveStats::default(),
last_second: (0..g).map(|_| Array2::<f64>::zeros((b, b))).collect(),
last_usage: vec![0; g],
last_rss: 0.0,
last_rows: 0,
pending_birth: None,
})
}
pub fn new_with_decoder(
decoder: Array2<f32>,
config: &BlockSparseConfig,
) -> Result<Self, String> {
validate_config(config)?;
if decoder.nrows() != config.n_blocks * config.block_size {
return Err(format!(
"BlockSparseStream decoder rows must equal n_blocks*block_size = {}, got {}",
config.n_blocks * config.block_size,
decoder.nrows()
));
}
if decoder.ncols() == 0 {
return Err("BlockSparseStream decoder must have at least one column".to_string());
}
if !decoder.iter().all(|v| v.is_finite()) {
return Err("BlockSparseStream decoder must be finite".to_string());
}
if config.block_size > decoder.ncols() {
return Err(format!(
"BlockSparseStream block_size b={} cannot exceed P={}",
config.block_size,
decoder.ncols()
));
}
let p = decoder.ncols();
let g = config.n_blocks;
let b = config.block_size;
let k = config.block_topk.min(g).max(1);
let cap = config.aux_k.saturating_mul(b).max(1);
Ok(Self {
config: *config,
g,
b,
k,
p,
decoder,
gamma: 1.0,
second: (0..g).map(|_| Array2::<f64>::zeros((b, b))).collect(),
coupling: (0..g).map(|_| Array2::<f64>::zeros((p, b))).collect(),
data_cross: (0..g).map(|_| Array2::<f64>::zeros((p, b))).collect(),
data_energy: vec![0.0; g],
negative_bound: vec![0.0; g],
usage: vec![0; g],
alive_count: 0,
gamma_num: 0.0,
gamma_den: 0.0,
col_sum: vec![0.0; p],
col_sumsq: vec![0.0; p],
rss: 0.0,
row_count: 0,
reservoir: ResidualReservoir::new(cap),
prev_ev: f64::NEG_INFINITY,
last_ev: f64::NEG_INFINITY,
last_ev_residual: f64::INFINITY,
last_gamma_residual: f64::INFINITY,
last_frame_residual: f64::INFINITY,
epochs_run: 0,
last_accepted_births: 0,
converged: false,
last_util: vec![0.0; g],
last_stable: vec![0.0; g],
last_decoder_solve_stats: DecoderSolveStats::default(),
last_second: (0..g).map(|_| Array2::<f64>::zeros((b, b))).collect(),
last_usage: vec![0; g],
last_rss: 0.0,
last_rows: 0,
pending_birth: None,
})
}
pub fn partial_fit(&mut self, shard: ArrayView2<'_, f32>) -> Result<BlockShardStats, String> {
if shard.nrows() == 0 {
return Ok(BlockShardStats {
rows: 0,
rss: 0.0,
alive_blocks: self.alive_count,
});
}
if shard.ncols() != self.p {
return Err(format!(
"BlockSparseStream.partial_fit: shard has P={} columns but the fit was begun with \
P={}",
shard.ncols(),
self.p
));
}
if !shard.iter().all(|v| v.is_finite()) {
return Err("BlockSparseStream.partial_fit shard must be finite".to_string());
}
self.converged = false;
let shard_start = std::time::Instant::now();
let p = self.p;
let b = self.b;
let gamma = self.gamma;
let aux_on = self.config.aux_k > 0;
let mut shard_rss = 0.0f64;
for rows in shard.axis_chunks_iter(Axis(0), self.config.minibatch.max(1)) {
let codes = route_and_code_all(
rows,
self.decoder.view(),
gamma,
self.g,
b,
self.k,
self.config.minibatch,
self.config.block_tile,
)?;
let baseline_codes = self
.pending_birth
.as_ref()
.map(|pending| {
route_and_code_all(
rows,
pending.baseline_decoder.view(),
pending.baseline_gamma,
self.g,
b,
self.k,
self.config.minibatch,
self.config.block_tile,
)
})
.transpose()?;
let projected = project_coded_rows(rows, self.decoder.view(), &codes, b, gamma);
let postings = block_row_postings(&codes, self.g);
let columns_per_worker = p.div_ceil(rayon::current_num_threads()).max(1);
self.col_sum
.par_chunks_mut(columns_per_worker)
.zip(self.col_sumsq.par_chunks_mut(columns_per_worker))
.enumerate()
.for_each(|(chunk, (sums, squares))| {
let first = chunk * columns_per_worker;
for row in rows.outer_iter() {
for (offset, (sum, square)) in
sums.iter_mut().zip(squares.iter_mut()).enumerate()
{
let value = row[first + offset] as f64;
*sum += value;
*square += value * value;
}
}
});
for (row, projection) in projected.iter().enumerate() {
shard_rss += projection.rss;
self.rss += projection.rss;
self.gamma_num += projection.gamma_num;
self.gamma_den += projection.gamma_den;
if aux_on {
let residual = rows
.row(row)
.iter()
.zip(&projection.sum)
.map(|(&x, &sum)| (x as f64 - gamma as f64 * sum) as f32)
.collect();
self.reservoir
.offer(projection.rss, (self.row_count + row) as u64, residual);
}
}
self.coupling
.par_iter_mut()
.zip(self.data_cross.par_iter_mut())
.zip(self.second.par_iter_mut())
.zip(self.usage.par_iter_mut())
.zip(self.data_energy.par_iter_mut())
.zip(self.negative_bound.par_iter_mut())
.zip(postings.par_iter())
.enumerate()
.for_each(
|(
block,
((((((coupling, data_cross), second), usage), energy), bound), entries),
)| {
let mut v_coordinates = vec![0.0; b];
for &(row, slot) in entries {
let w = &codes[row].projections[slot * b..(slot + 1) * b];
let xi = rows.row(row);
v_coordinates.fill(0.0);
let mut x_norm_sq = 0.0;
let mut v_norm_sq = 0.0;
let mut x_dot_v = 0.0;
for c in 0..p {
let mut own = 0.0;
for (axis, &weight) in w.iter().enumerate() {
own += weight * self.decoder[[block * b + axis, c]] as f64;
}
let v = self.k as f64 * own - projected[row].sum[c];
let x = xi[c] as f64;
x_norm_sq += x * x;
v_norm_sq += v * v;
x_dot_v += x * v;
for (axis, &weight) in w.iter().enumerate() {
coupling[[c, axis]] += v * weight;
data_cross[[c, axis]] += x * weight;
v_coordinates[axis] +=
self.decoder[[block * b + axis, c]] as f64 * v;
}
}
for c in 0..p {
for axis in 0..b {
coupling[[c, axis]] += xi[c] as f64 * v_coordinates[axis];
}
}
*energy += x_norm_sq;
*bound += (x_norm_sq.sqrt() * v_norm_sq.sqrt() - x_dot_v).max(0.0);
for left in 0..b {
for right in 0..b {
second[[left, right]] += w[left] * w[right];
}
}
}
*usage += entries.len();
},
);
if let (Some(pending), Some(baseline_codes)) =
(self.pending_birth.as_mut(), baseline_codes.as_ref())
{
let baseline = project_coded_rows(
rows,
pending.baseline_decoder.view(),
baseline_codes,
b,
pending.baseline_gamma,
);
for projection in &baseline {
pending.baseline_rss += projection.rss;
}
let baseline_postings = block_row_postings(baseline_codes, self.g);
pending
.baseline_second
.par_iter_mut()
.zip(pending.baseline_usage.par_iter_mut())
.zip(baseline_postings.par_iter())
.for_each(|((second, usage), entries)| {
for &(row, slot) in entries {
let w = &baseline_codes[row].projections[slot * b..(slot + 1) * b];
for left in 0..b {
for right in 0..b {
second[[left, right]] += (pending.baseline_gamma as f64
* w[left])
* (pending.baseline_gamma as f64 * w[right]);
}
}
}
*usage += entries.len();
});
pending.baseline_rows += rows.nrows();
}
self.row_count += rows.nrows();
self.alive_count = self.usage.iter().filter(|&&count| count > 0).count();
}
log::info!(
"[SAE block shard] rows={} total_rows={} rss={:.6e} alive_blocks={}/{} \
shard_s={:.2}",
shard.nrows(),
self.row_count,
shard_rss,
self.alive_count,
self.g,
shard_start.elapsed().as_secs_f64(),
);
Ok(BlockShardStats {
rows: shard.nrows(),
rss: shard_rss,
alive_blocks: self.alive_count,
})
}
pub fn end_epoch(&mut self) -> Result<BlockEpochStats, String> {
if self.row_count == 0 {
return Err(
"BlockSparseStream.end_epoch: no rows were streamed this epoch (call partial_fit \
with at least one shard first)"
.to_string(),
);
}
let p = self.p;
let b = self.b;
let n = self.row_count as f64;
let mut tss = 0.0f64;
for c in 0..p {
tss += self.col_sumsq[c] - self.col_sum[c] * self.col_sum[c] / n;
}
let mut accepted_births = 0usize;
let mut rejected_birth = false;
if let Some(pending) = self.pending_birth.take() {
if pending.baseline_rows != self.row_count {
return Err(format!(
"BlockSparseStream birth transaction saw {} candidate rows but {} baseline rows",
self.row_count, pending.baseline_rows,
));
}
let selected = self.usage[pending.block] > 0;
let improvement_rss = pending.baseline_rss - self.rss;
let evidence_margin = block_birth_evidence_margin(
pending.block,
improvement_rss,
self.rss,
self.usage[pending.block],
&self.second[pending.block].mapv(|value| value * (self.gamma as f64).powi(2)),
self.decoder.view(),
self.row_count,
self.p,
self.b,
)?;
if selected && evidence_margin.is_some_and(|margin| margin > 0.0) {
accepted_births = 1;
} else {
self.decoder = pending.baseline_decoder;
self.gamma = pending.baseline_gamma;
self.rss = pending.baseline_rss;
self.usage = pending.baseline_usage;
self.second = pending.baseline_second;
self.alive_count = self.usage.iter().filter(|&&count| count > 0).count();
rejected_birth = true;
}
}
let previous_gamma = self.gamma;
let mut candidate_decoder = self.decoder.clone();
let mut gamma_residual = f64::INFINITY;
let mut frame_residual = f64::INFINITY;
if !rejected_birth {
self.gamma = if self.gamma_den == 0.0 {
0.0
} else {
(self.gamma_num / self.gamma_den) as f32
};
if !self.gamma.is_finite() || self.gamma < 0.0 {
return Err("BlockSparseStream gamma optimum is not finite and nonnegative".into());
}
gamma_residual = relative_scalar_change(previous_gamma, self.gamma);
let old = previous_gamma as f64;
let gamma = self.gamma as f64;
let correction =
(gamma - old) * ((gamma + old) * self.gamma_den - 2.0 * self.gamma_num);
let rss = self.rss + correction;
let resolution = f64::EPSILON
* (self.row_count * p) as f64
* (self.rss.abs() + correction.abs() + self.col_sumsq.iter().sum::<f64>());
if !rss.is_finite() || rss < -resolution {
return Err(format!("BlockSparseStream profiled RSS is invalid: {rss}"));
}
self.rss = rss.max(0.0);
let ridge = self.config.frame_ridge;
let outcomes: Vec<Result<f64, String>> = with_faer_sequential(|| {
self.coupling
.par_iter_mut()
.zip(self.second.par_iter_mut())
.zip(
candidate_decoder
.axis_chunks_iter_mut(Axis(0), b)
.into_par_iter(),
)
.enumerate()
.map(|(gg, ((moment, second), mut proposal))| {
second.mapv_inplace(|value| value * gamma * gamma);
if self.usage[gg] == 0 {
return Ok(0.0);
}
let data_scale = 2.0 * gamma - self.k as f64 * gamma * gamma;
let shift = (-data_scale).max(0.0) * self.data_energy[gg]
+ gamma * gamma * self.negative_bound[gg]
+ ridge;
for rr in 0..b {
for c in 0..p {
moment[[c, rr]] = data_scale * self.data_cross[gg][[c, rr]]
+ gamma * gamma * moment[[c, rr]];
}
}
polar_tied_frame_step(
self.decoder.slice(ndarray::s![gg * b..(gg + 1) * b, ..]),
moment.view_mut(),
second.view(),
(self.k - 1) as f64,
shift,
proposal.view_mut(),
)
.map_err(|error| format!("BlockSparseStream polar block {gg}: {error}"))
})
.collect()
});
let mut gradient_residual = 0.0_f64;
for outcome in outcomes {
gradient_residual = gradient_residual.max(outcome?);
}
frame_residual = frame_fixed_point_residual(
self.decoder.view(),
candidate_decoder.view(),
self.g,
b,
)?
.max(gradient_residual);
}
let ev = if tss <= 1.0e-24 {
if self.rss <= 1.0e-24 { 1.0 } else { 0.0 }
} else {
1.0 - self.rss / tss
};
let decoder_solve_stats = DecoderSolveStats::default();
let dead: usize = self.usage.iter().filter(|&&u| u == 0).count();
for gg in 0..self.g {
self.last_util[gg] = self.usage[gg] as f32 / self.row_count.max(1) as f32;
self.last_stable[gg] = stable_rank_symmetric(self.second[gg].view());
}
let improve = ev - self.prev_ev;
let stationary = !rejected_birth
&& accepted_births == 0
&& improve.abs() <= self.config.tolerance
&& gamma_residual <= self.config.tolerance
&& frame_residual
<= self
.config
.tolerance
.max(super::block_frame::STORED_FRAME_RESOLUTION)
&& self.epochs_run > 0;
if !stationary && !rejected_birth {
self.decoder = candidate_decoder;
}
let birth_pending = !rejected_birth && self.stage_birth_proposal();
let converged = stationary && !birth_pending;
self.prev_ev = ev;
self.last_ev = ev;
self.last_ev_residual = improve.abs();
self.last_gamma_residual = gamma_residual;
self.last_frame_residual = frame_residual;
self.last_accepted_births = accepted_births;
self.converged = converged;
self.last_decoder_solve_stats = decoder_solve_stats;
self.epochs_run += 1;
let epoch = self.epochs_run;
self.last_second.clone_from(&self.second);
self.last_usage.clone_from(&self.usage);
self.last_rss = self.rss;
self.last_rows = self.row_count;
self.reset_epoch();
Ok(BlockEpochStats {
explained_variance: ev,
accepted_births,
birth_pending,
dead,
gamma: self.gamma,
gamma_residual,
frame_residual,
converged,
epoch,
decoder_solve_stats,
})
}
fn stage_birth_proposal(&mut self) -> bool {
if self.config.aux_k == 0 || self.pending_birth.is_some() {
return false;
}
let Some(block) = (0..self.g)
.filter(|&candidate| self.usage[candidate] == 0)
.take(self.config.aux_k)
.next()
else {
return false;
};
let b = self.b;
let p = self.p;
let proposal = {
let ranked = self.reservoir.ranked();
if ranked.len() < b || ranked[0].norm2 <= DEAD_DENOM {
return false;
}
let mut seed = Array2::<f32>::zeros((b, p));
for row in 0..b {
for column in 0..p {
seed[[row, column]] = ranked[row].residual[column];
}
}
gram_schmidt_rows(&mut seed);
seed
};
let baseline_decoder = self.decoder.clone();
let baseline_gamma = self.gamma;
self.decoder
.slice_mut(ndarray::s![block * b..block * b + b, ..])
.assign(&proposal);
self.pending_birth = Some(PendingBlockBirth {
block,
baseline_decoder,
baseline_gamma,
baseline_rss: 0.0,
baseline_rows: 0,
baseline_usage: vec![0; self.g],
baseline_second: (0..self.g).map(|_| Array2::<f64>::zeros((b, b))).collect(),
});
true
}
fn reset_epoch(&mut self) {
for sg in self.second.iter_mut() {
sg.fill(0.0);
}
for mg in self.coupling.iter_mut() {
mg.fill(0.0);
}
for moment in &mut self.data_cross {
moment.fill(0.0);
}
self.data_energy.fill(0.0);
self.negative_bound.fill(0.0);
for u in self.usage.iter_mut() {
*u = 0;
}
self.alive_count = 0;
self.gamma_num = 0.0;
self.gamma_den = 0.0;
for c in 0..self.p {
self.col_sum[c] = 0.0;
self.col_sumsq[c] = 0.0;
}
self.rss = 0.0;
self.row_count = 0;
self.reservoir.clear();
}
pub fn finalize(&self) -> Result<BlockSparseStreamArtifact, String> {
if !self.converged || self.pending_birth.is_some() {
return Err(format!(
"BlockSparseStream.finalize: streaming fit has not converged after {} epoch(s) \
(last EV {:.6e}, EV residual {:.3e}, gamma residual {:.3e}, frame residual {:.3e} \
vs tolerance {:.3e}, {} accepted block \
birth(s) in the last epoch, birth pending={}); the stream state is a resumable \
checkpoint, not a model — run more epochs until end_epoch reports convergence",
self.epochs_run,
self.last_ev,
self.last_ev_residual,
self.last_gamma_residual,
self.last_frame_residual,
self.config.tolerance,
self.last_accepted_births,
self.pending_birth.is_some(),
));
}
Ok(BlockSparseStreamArtifact {
decoder: self.decoder.clone(),
gamma: self.gamma,
block_topk: self.k,
block_size: self.b,
block_utilization: self.last_util.clone(),
block_stable_rank: self.last_stable.clone(),
epochs: self.epochs_run,
explained_variance: self.last_ev,
decoder_solve_stats: self.last_decoder_solve_stats,
})
}
pub fn block_rank_charges(&self, n_obs: usize) -> Result<BlockRankCharges, String> {
if self.last_rows == 0 {
return Err(
"block_rank_charges: no closed epoch to certify; call end_epoch first".to_string(),
);
}
let phi_raw = self.last_rss / (self.last_rows as f64 * self.p as f64);
let phi = if phi_raw.is_finite() && phi_raw > 0.0 {
phi_raw
} else {
1.0
};
let ln_n = (n_obs.max(2) as f64).ln();
let mut out = BlockRankCharges {
block: Vec::with_capacity(self.g),
n_eff: Vec::with_capacity(self.g),
d_eff: Vec::with_capacity(self.g),
delta_deviance: Vec::with_capacity(self.g),
charge: Vec::with_capacity(self.g),
margin: Vec::with_capacity(self.g),
kept: Vec::with_capacity(self.g),
};
for gg in 0..self.g {
let n_eff = self.last_usage[gg] as f64;
let frame = self
.decoder
.slice(ndarray::s![gg * self.b..(gg + 1) * self.b, ..])
.mapv(f64::from);
let d_eff = crate::manifold::realised_rank_charge_dof(
&self.last_second[gg],
&frame,
n_eff,
self.p as f64,
phi,
0.0,
None,
)?;
let mut tr = 0.0_f64;
for i in 0..self.b {
tr += self.last_second[gg][[i, i]];
}
let delta_deviance = 0.5 * tr / phi;
let charge = 0.5 * d_eff * ln_n;
let margin = delta_deviance - charge;
out.block.push(gg);
out.n_eff.push(n_eff);
out.d_eff.push(d_eff);
out.delta_deviance.push(delta_deviance);
out.charge.push(charge);
out.margin.push(margin);
out.kept.push(margin > 0.0);
}
Ok(out)
}
pub fn decoder(&self) -> ArrayView2<'_, f32> {
self.decoder.view()
}
pub fn gamma(&self) -> f32 {
self.gamma
}
pub fn block_topk(&self) -> usize {
self.k
}
pub fn block_size(&self) -> usize {
self.b
}
pub fn epochs_run(&self) -> usize {
self.epochs_run
}
}
#[derive(Clone, Debug)]
pub struct BlockSparseStreamArtifact {
pub decoder: Array2<f32>,
pub gamma: f32,
pub block_topk: usize,
pub block_size: usize,
pub block_utilization: Vec<f32>,
pub block_stable_rank: Vec<f32>,
pub epochs: usize,
pub explained_variance: f64,
pub decoder_solve_stats: DecoderSolveStats,
}
fn validate_config(config: &BlockSparseConfig) -> Result<(), String> {
if config.n_blocks == 0 {
return Err("BlockSparseStream requires n_blocks >= 1".to_string());
}
if config.block_size == 0 {
return Err("BlockSparseStream requires block_size >= 1".to_string());
}
if config.block_topk == 0 {
return Err("BlockSparseStream requires block_topk >= 1".to_string());
}
if config.max_epochs == 0 {
return Err("BlockSparseStream requires max_epochs >= 1".to_string());
}
if !(config.frame_ridge.is_finite() && config.frame_ridge >= 0.0) {
return Err("BlockSparseStream frame_ridge must be finite and non-negative".to_string());
}
if !config.tolerance.is_finite() {
return Err("BlockSparseStream tolerance must be finite".to_string());
}
Ok(())
}
#[cfg(test)]
#[path = "block_stream_tests.rs"]
mod block_stream_tests;