use legume_numeric::matrix::dmatrix_io::*;
use legume_numeric::matrix::traits::*;
use log::info;
use rand::seq::index::sample;
use rand::SeedableRng;
use rand_distr::{Distribution, Normal, Uniform};
use super::core::{
generate_hierarchical_dictionary, sample_log_batch_effects, sample_poisson_triplets,
sample_theta_kn,
};
pub struct MultimodalSimArgs {
pub rows: usize,
pub cols: usize,
pub depth_per_modality: Vec<usize>,
pub factors: usize,
pub batches: usize,
pub base_scale: f32,
pub delta_scale: f32,
pub n_delta_features: usize,
pub pve_topic: f32,
pub pve_batch: f32,
pub rseed: u64,
pub shared_batch_effects: bool,
pub hierarchical_depth: Option<usize>,
pub beta_scale: f32,
pub n_housekeeping: usize,
pub housekeeping_fold: f32,
}
pub struct MultimodalSimOut {
pub w_base_kd: DMatrix<f32>,
pub w_delta_kd: Vec<DMatrix<f32>>,
pub beta_dk: Vec<DMatrix<f32>>,
pub spike_mask_kd: Vec<DMatrix<f32>>,
pub theta_kn: DMatrix<f32>,
pub ln_delta_db: Vec<DMatrix<f32>>,
pub batch_membership: Vec<usize>,
pub triplets: Vec<Vec<(u64, u64, f32)>>,
}
fn generate_batch_effects(
dd: usize,
bb: usize,
pve_batch: f32,
rng: &mut impl rand::Rng,
) -> DMatrix<f32> {
if bb <= 1 {
return DMatrix::<f32>::zeros(dd, bb.max(1));
}
sample_log_batch_effects(dd, bb, pve_batch, rng)
}
pub fn generate_multimodal_data(args: &MultimodalSimArgs) -> anyhow::Result<MultimodalSimOut> {
let dd = args.rows;
let nn = args.cols;
let kk = if let Some(depth) = args.hierarchical_depth {
1usize << (depth - 1)
} else {
args.factors
};
let bb = args.batches;
let mm = args.depth_per_modality.len();
let eps = 1e-8_f32;
anyhow::ensure!(mm >= 1, "need at least 1 modality (depth_per_modality)");
anyhow::ensure!(bb >= 1, "batches must be >= 1");
anyhow::ensure!(
args.n_delta_features <= dd,
"n_delta_features ({}) must be <= rows ({})",
args.n_delta_features,
dd
);
let mut rng = rand::rngs::StdRng::seed_from_u64(args.rseed);
let runif_batch = Uniform::new(0, bb)?;
let batch_membership: Vec<usize> = (0..nn).map(|_| runif_batch.sample(&mut rng)).collect();
let ln_delta_db: Vec<DMatrix<f32>> = if args.shared_batch_effects {
let shared = generate_batch_effects(dd, bb, args.pve_batch, &mut rng);
vec![shared; mm]
} else {
(0..mm)
.map(|_| generate_batch_effects(dd, bb, args.pve_batch, &mut rng))
.collect()
};
let w_base_kd = if let Some(tree_depth) = args.hierarchical_depth {
let (beta_dk, _node_probs) =
generate_hierarchical_dictionary(dd, tree_depth, args.beta_scale, &mut rng);
info!(
"hierarchical base dictionary: depth={}, K={} leaves",
tree_depth,
beta_dk.ncols()
);
let mut logits_kd = DMatrix::<f32>::zeros(kk, dd);
for k in 0..kk {
for d in 0..dd {
logits_kd[(k, d)] = beta_dk[(d, k)].max(eps).ln();
}
}
logits_kd
} else {
let normal = Normal::new(0.0f32, args.base_scale)?;
DMatrix::from_fn(kk, dd, |_, _| normal.sample(&mut rng))
};
let mut w_base_kd = w_base_kd;
if args.n_housekeeping > 0 && args.n_housekeeping < dd {
let hk_val = args.housekeeping_fold * args.base_scale;
for d in 0..args.n_housekeeping {
for k in 0..kk {
w_base_kd[(k, d)] = hk_val;
}
}
info!(
"injected {} housekeeping genes with logit value {:.2}",
args.n_housekeeping, hk_val
);
}
let normal_delta = Normal::new(0.0f32, args.delta_scale)?;
let mut w_delta_kd: Vec<DMatrix<f32>> = Vec::with_capacity(mm.saturating_sub(1));
let mut spike_mask_kd: Vec<DMatrix<f32>> = Vec::with_capacity(mm.saturating_sub(1));
for _ in 1..mm {
let mut delta = DMatrix::<f32>::zeros(kk, dd);
let mut mask = DMatrix::<f32>::zeros(kk, dd);
for k in 0..kk {
let chosen = sample(&mut rng, dd, args.n_delta_features);
for d in chosen {
delta[(k, d)] = normal_delta.sample(&mut rng);
mask[(k, d)] = 1.0;
}
}
w_delta_kd.push(delta);
spike_mask_kd.push(mask);
}
info!(
"generated {} delta matrices: {} non-zero features per topic out of {}",
w_delta_kd.len(),
args.n_delta_features,
dd
);
let mut beta_dk: Vec<DMatrix<f32>> = Vec::with_capacity(mm);
beta_dk.push(w_base_kd.transpose().normalize_exp_logits_columns());
for delta in &w_delta_kd {
let logits = &w_base_kd + delta;
beta_dk.push(logits.transpose().normalize_exp_logits_columns());
}
let theta_kn = sample_theta_kn(kk, nn, args.pve_topic, &mut rng)?;
let rseed = args.rseed;
let mut triplets: Vec<Vec<(u64, u64, f32)>> = Vec::with_capacity(mm);
let delta_exp: Vec<DMatrix<f32>> = ln_delta_db.iter().map(|m| m.map(|x| x.exp())).collect();
for (m, &depth_m) in args.depth_per_modality.iter().enumerate() {
let delta_ref = if bb > 1 { Some(&delta_exp[m]) } else { None };
let seed_offset = (m as u64) * (nn as u64);
let lambda_scale = depth_m as f32;
let trips = sample_poisson_triplets(
&beta_dk[m],
&theta_kn,
delta_ref,
&batch_membership,
lambda_scale,
rseed,
seed_offset,
);
info!(
"modality {}: {} non-zero triplets (depth={}, λ_scale={:.1})",
m,
trips.len(),
depth_m,
lambda_scale,
);
triplets.push(trips);
}
Ok(MultimodalSimOut {
w_base_kd,
w_delta_kd,
beta_dk,
spike_mask_kd,
theta_kn,
ln_delta_db,
batch_membership,
triplets,
})
}