use crate::candle::data::loader_util::{bootstrap_indices, upload_columns_as_rows};
use crate::matrix::rand_util::mix_seed;
use candle_core::{Device, Tensor};
use nalgebra::DMatrix;
use rand::seq::SliceRandom;
use rand::{rngs::SmallRng, RngExt, SeedableRng};
use rayon::prelude::*;
type Mat = DMatrix<f32>;
#[derive(Clone, Copy, Debug)]
pub enum MaskSchedule {
Fixed,
Uniform { lo: f64, hi: f64 },
}
#[derive(Clone, Copy, Debug)]
pub struct MaskedDraw {
pub schedule: MaskSchedule,
pub mask_fraction: f64,
}
type RowDraw = Vec<u32>;
fn hidden_count(n_features: usize, rate: f64) -> usize {
debug_assert!(
n_features >= 2,
"a {n_features}-gene axis cannot be split into a visible and a hidden part"
);
debug_assert!(
rate > 0.0 && rate < 1.0,
"mask rate {rate} is outside the open interval (0, 1); the CLI refuses this"
);
let m = (rate * n_features as f64).round() as usize;
m.clamp(1, n_features - 1)
}
fn floyd_sample(d: usize, m: usize, rng: &mut SmallRng) -> Vec<u32> {
let mut seen = std::collections::HashSet::with_capacity(m);
let mut out = Vec::with_capacity(m);
for j in (d - m)..d {
let t = rng.random_range(0..=j);
let pick = if seen.insert(t) {
t
} else {
seen.insert(j);
j
};
out.push(pick as u32);
}
out.sort_unstable();
out
}
fn draw_row(row: usize, n_features: usize, epoch_seed: u64, draw: &MaskedDraw) -> RowDraw {
let mut rng = SmallRng::seed_from_u64(mix_seed(epoch_seed, row as u64));
let rate = match draw.schedule {
MaskSchedule::Fixed => draw.mask_fraction,
MaskSchedule::Uniform { lo, hi } => lo + (hi - lo) * rng.random::<f64>(),
};
floyd_sample(n_features, hidden_count(n_features, rate), &mut rng)
}
fn visible_from_hidden(hidden_ids: &Tensor, n_features: usize) -> candle_core::Result<Tensor> {
let (n, dh) = hidden_ids.dims2()?;
let dev = hidden_ids.device();
let ones = Tensor::ones((n, n_features), candle_core::DType::F32, dev)?;
let zeros = Tensor::zeros((n, dh), candle_core::DType::F32, dev)?;
ones.scatter(hidden_ids, &zeros, 1)
}
pub struct DenseMaskedLevel {
n_features: usize,
p: usize,
input_pd: Tensor,
null_pd: Option<Tensor>,
target_pd: Tensor,
mean_1d: Tensor,
dev: Device,
}
pub struct DenseMaskedMinibatch {
pub row_ids: Tensor,
pub x_nd: Tensor,
pub x0_nd: Option<Tensor>,
pub visible_nd: Tensor,
pub hidden_ids: Tensor,
pub hidden_weight: Option<Tensor>,
pub target_nd: Tensor,
}
impl DenseMaskedLevel {
pub fn from_mats(
input_dp: &Mat,
null_dp: Option<&Mat>,
target_dp: &Mat,
mean: &[f32],
dev: &Device,
) -> anyhow::Result<Self> {
let (d, p) = (input_dp.nrows(), input_dp.ncols());
anyhow::ensure!(
target_dp.nrows() == d && target_dp.ncols() == p,
"target rows {}×{} do not match the input's {d}×{p} (genes × samples)",
target_dp.nrows(),
target_dp.ncols()
);
if let Some(n) = null_dp {
anyhow::ensure!(
n.nrows() == d && n.ncols() == p,
"batch null is {}×{}, expected {d}×{p} (genes × samples)",
n.nrows(),
n.ncols()
);
}
anyhow::ensure!(
mean.len() == d,
"per-gene mean has {} entries, expected {d}",
mean.len()
);
let input_pd = upload_columns_as_rows(input_dp, dev)?;
let target_pd = if std::ptr::eq(input_dp, target_dp) {
input_pd.clone()
} else {
upload_columns_as_rows(target_dp, dev)?
};
Ok(Self {
n_features: d,
p,
input_pd,
null_pd: null_dp
.map(|n| upload_columns_as_rows(n, dev))
.transpose()?,
target_pd,
mean_1d: Tensor::from_vec(mean.to_vec(), (1, d), dev)?,
dev: dev.clone(),
})
}
#[must_use]
pub fn num_data(&self) -> usize {
self.p
}
#[must_use]
pub fn n_features(&self) -> usize {
self.n_features
}
#[must_use]
pub fn feature_mean_1d(&self) -> &Tensor {
&self.mean_1d
}
pub fn begin_epoch(
&self,
epoch_seed: u64,
draw: &MaskedDraw,
batch_size: usize,
) -> anyhow::Result<DenseMaskedEpoch<'_>> {
anyhow::ensure!(self.p > 0, "begin_epoch on an empty level");
anyhow::ensure!(batch_size > 0, "batch_size must be > 0");
let nbatch = self.p.div_ceil(batch_size);
let ntot = nbatch * batch_size;
let mut order: Vec<u32> = (0..self.p as u32).collect();
order.shuffle(&mut rand::rng());
order.extend(bootstrap_indices::<u32>(self.p, ntot - self.p));
Ok(DenseMaskedEpoch {
level: self,
order,
batch_size,
epoch_seed,
draw: *draw,
})
}
pub fn probe_minibatch(
&self,
n: usize,
epoch_seed: u64,
draw: &MaskedDraw,
) -> anyhow::Result<DenseMaskedMinibatch> {
anyhow::ensure!(self.p > 0, "probe_minibatch on an empty level");
let order: Vec<u32> = (0..n).map(|i| (i % self.p) as u32).collect();
DenseMaskedEpoch {
level: self,
order,
batch_size: n,
epoch_seed,
draw: *draw,
}
.batch(0)
}
}
pub struct DenseMaskedEpoch<'a> {
level: &'a DenseMaskedLevel,
order: Vec<u32>,
batch_size: usize,
epoch_seed: u64,
draw: MaskedDraw,
}
impl DenseMaskedEpoch<'_> {
#[must_use]
pub fn n_batches(&self) -> usize {
self.order.len().div_ceil(self.batch_size)
}
pub fn batch(&self, b: usize) -> anyhow::Result<DenseMaskedMinibatch> {
let lv = self.level;
let start = b * self.batch_size;
anyhow::ensure!(start < self.order.len(), "batch {b} is past the epoch");
let len = self.batch_size.min(self.order.len() - start);
let ids = &self.order[start..start + len];
let d = lv.n_features;
let draws: Vec<RowDraw> = ids
.par_iter()
.map(|&r| draw_row(r as usize, d, self.epoch_seed, &self.draw))
.collect();
let cpu = Device::Cpu;
let dh = draws.iter().map(Vec::len).max().unwrap_or(1);
let ragged = draws.iter().any(|rd| rd.len() != dh);
let mut hid = vec![0u32; len * dh];
let mut w = vec![0f32; len * dh];
for (row, rd) in draws.iter().enumerate() {
let n_hid = rd.len();
let slot = &mut hid[row * dh..(row + 1) * dh];
slot[..n_hid].copy_from_slice(rd);
slot[n_hid..].fill(rd[n_hid - 1]);
w[row * dh..row * dh + n_hid].fill(1.0);
}
let hidden_ids = Tensor::from_vec(hid, (len, dh), &cpu)?.to_device(&lv.dev)?;
let visible_nd = visible_from_hidden(&hidden_ids, d)?;
let hidden_weight = if ragged {
Some(Tensor::from_vec(w, (len, dh), &cpu)?.to_device(&lv.dev)?)
} else {
None
};
let row_ids = Tensor::from_vec(ids.to_vec(), len, &lv.dev)?;
let sel = |t: &Tensor| t.index_select(&row_ids, 0);
let x_nd = sel(&lv.input_pd)?;
let x0_nd = lv.null_pd.as_ref().map(sel).transpose()?;
let target_nd = sel(&lv.target_pd)?;
Ok(DenseMaskedMinibatch {
row_ids,
x_nd,
x0_nd,
visible_nd,
hidden_ids,
hidden_weight,
target_nd,
})
}
}
#[cfg(test)]
#[path = "masked_dense_tests.rs"]
mod masked_dense_tests;