mod block;
mod block_chart;
mod block_frame;
mod block_scoring_gpu;
mod block_stream;
mod code_evidence;
mod codes;
mod cofit;
mod coordinate;
#[cfg(target_os = "linux")]
mod decoder_gpu;
mod residual_reservoir;
#[cfg(target_os = "linux")]
mod score_router_backend;
mod scoring;
mod single_atom;
#[cfg(target_os = "linux")]
mod scoring_gpu;
mod split_lr_fdr;
mod stream;
mod update;
#[cfg(test)]
mod tests;
pub use block::{BlockSeedPolicy, BlockSparseConfig, BlockSparseConvergence, BlockSparseFit, BlockSparseFitError, block_gates, block_projections_row, block_sparse_dictionary_block_coords, block_sparse_dictionary_lift_block, block_sparse_dictionary_project_residual, block_sparse_dictionary_transform, coordinate_partition_frames, fit_block_sparse_dictionary, fit_block_sparse_dictionary_with_seed, reconstruct_block_sparse_rows, route_row_blocks};
pub use block_chart::{BlockChartComposeConfig, BlockChartComposeResult, BlockChartRecord, BlockSeedManifest, BlockSeedManifestConfig, BlockSeedRecord, CHART_FDR_ALPHA, ChartEvidence, MdlFeaturizerRow, block_sparse_dictionary_firings, block_sparse_dictionary_seed_manifest, compose_block_coordinate_charts};
pub use block_scoring_gpu::{BlockRoutePath, block_gate_row_cpu, route_blocks_cpu};
#[cfg(target_os = "linux")]
pub use block_scoring_gpu::{DEVICE_BLOCK_GATE_MIN_ELEMS, route_blocks_required};
pub use block_stream::{BlockEpochStats, BlockShardStats, BlockSparseStreamArtifact, BlockSparseStreamState};
pub use codes::SparseCode;
pub use cofit::{CofitConfig, CofitReport, CofitRound};
pub use coordinate::{BlockCoordinateReport, BlockMeasureCoordinateReport, FiringCoordinate, MeasureSpikeCoordinate, MeasureValuedCode, block_route_firing_coordinates, harmonic_measure_coordinates, harmonic_route_firing_coordinates, recover_measure_from_code};
pub use scoring::{ScoreRoutePath, ScoreRouteResult, ScoreRouteStats, TileScorer, top_s_online};
#[cfg(target_os = "linux")]
pub use scoring_gpu::{DEVICE_SCORE_BLOCK_MIN_ELEMS, ScoreBlockPath};
pub use split_lr_fdr::{
FdrCertificate, crossfit_ui_log_evalue, family_fdr_certificate, shell_vs_ring_log_evalue,
};
pub use stream::{EpochStats, ShardStats, SparseDictArtifact, SparseDictStreamState};
pub use update::{DecoderSolveStats, LinearBlockRemlStats, SparseDictionaryError, linear_shared_rho_fs_step};
pub(crate) use update::extend_linear_reml_schedule;
use ndarray::{Array2, ArrayView2};
#[derive(Clone, Copy, Debug)]
pub struct SparseDictConfig {
pub n_atoms: usize,
pub active: usize,
pub minibatch: usize,
pub max_epochs: usize,
pub score_tile: usize,
pub code_ridge: f32,
pub decoder_ridge: f32,
pub tolerance: f64,
pub score_mode: gam_gpu::GpuPolicy,
}
impl SparseDictConfig {
pub fn new(n_atoms: usize) -> Self {
Self {
n_atoms,
..Self::default()
}
}
}
impl Default for SparseDictConfig {
fn default() -> Self {
Self {
n_atoms: 1,
active: 1,
minibatch: 512,
max_epochs: 30,
score_tile: 4096,
code_ridge: 1.0e-6,
decoder_ridge: 1.0e-6,
tolerance: 1.0e-6,
score_mode: gam_gpu::GpuPolicy::Auto,
}
}
}
#[derive(Clone, Debug)]
pub struct SparseDictFit {
pub decoder: Array2<f32>,
pub indices: Array2<u32>,
pub codes: Array2<f32>,
pub explained_variance: f64,
pub epochs: usize,
pub convergence: SparseDictConvergence,
pub active: usize,
pub score_route_stats: ScoreRouteStats,
pub decoder_solve_stats: DecoderSolveStats,
}
#[derive(Clone, Copy, Debug)]
pub struct SparseDictConvergence {
pub inner_ev_residual: f64,
pub inner_tolerance: f64,
pub decoder_residual: f64,
pub decoder_tolerance: f64,
pub routing_residual: f64,
pub routing_tolerance: f64,
pub outer_rho_residual: f64,
pub outer_tolerance: f64,
pub selected_rho: f64,
pub outer_iterations: usize,
pub seeded_inner_runs: usize,
pub continued_inner_runs: usize,
pub accepted_births: usize,
pub live_atom_high_water: usize,
pub support_saturated: bool,
pub certified: bool,
}
impl SparseDictFit {
pub fn reconstruct(&self) -> Array2<f32> {
reconstruct_sparse_rows(self.decoder.view(), self.indices.view(), self.codes.view())
.expect("SparseDictFit stores internally validated routing")
}
}
pub fn reconstruct_sparse_rows(
decoder: ArrayView2<'_, f32>,
indices: ArrayView2<'_, u32>,
codes: ArrayView2<'_, f32>,
) -> Result<Array2<f32>, String> {
if indices.dim() != codes.dim() {
return Err(format!(
"reconstruct_sparse_rows: indices shape {:?} does not match codes shape {:?}",
indices.dim(),
codes.dim()
));
}
let n = indices.nrows();
let p = decoder.ncols();
let mut out = Array2::<f32>::zeros((n, p));
for i in 0..n {
for j in 0..indices.ncols() {
let atom = indices[[i, j]] as usize;
if atom >= decoder.nrows() {
return Err(format!(
"reconstruct_sparse_rows: atom index {atom} out of range 0..{}",
decoder.nrows()
));
}
let code = codes[[i, j]];
if code == 0.0 {
continue;
}
let row = decoder.row(atom);
for c in 0..p {
out[[i, c]] += code * row[c];
}
}
}
Ok(out)
}
#[derive(Clone, Debug)]
pub struct SparseDictTransform {
pub indices: Array2<u32>,
pub codes: Array2<f32>,
pub score_route_stats: ScoreRouteStats,
}
pub fn sparse_dictionary_transform_with_mode(
x: ArrayView2<'_, f32>,
decoder: ArrayView2<'_, f32>,
active: usize,
score_tile: usize,
code_ridge: f32,
score_mode: gam_gpu::GpuPolicy,
) -> Result<SparseDictTransform, String> {
let k = decoder.nrows();
if k == 0 {
return Err("sparse_dictionary_transform: dictionary has no atoms".to_string());
}
if x.ncols() != decoder.ncols() {
return Err(format!(
"sparse_dictionary_transform: X has P={} columns but the decoder has P={}",
x.ncols(),
decoder.ncols()
));
}
let s = active.min(k).max(1);
let scorer = TileScorer::new(s, score_tile.max(1));
let routed = scorer.route_minibatch_with_mode(x, decoder, score_mode)?;
let mut score_route_stats = ScoreRouteStats::default();
score_route_stats.record_result(&routed);
let m = x.nrows();
let mut indices = Array2::<u32>::zeros((m, s));
let mut codes = Array2::<f32>::zeros((m, s));
for (row_idx, active_pairs) in routed.selections.iter().enumerate() {
let code = codes::solve_row_codes(x.row(row_idx), decoder, active_pairs, s, code_ridge);
for j in 0..s {
indices[[row_idx, j]] = code.indices[j];
codes[[row_idx, j]] = code.codes[j];
}
}
Ok(SparseDictTransform {
indices,
codes,
score_route_stats,
})
}
pub fn fit_sparse_dictionary(
x: ArrayView2<'_, f32>,
config: &SparseDictConfig,
) -> Result<SparseDictFit, SparseDictionaryError> {
update::run_linear_reml_schedule(x, config)
}