use salmon_core::atomic::AtomicF64;
use salmon_core::math::{log_add, LOG_0, LOG_EPSILON};
use statrs::distribution::{Binomial, ContinuousCDF, Discrete, Normal};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, RwLock};
#[derive(Debug)]
pub struct FragmentLengthDistribution {
kernel: Vec<f64>,
hist: Vec<AtomicF64>,
tot_mass: AtomicF64,
sum: AtomicF64,
min: AtomicUsize,
bin_size: usize,
cached_pmf: Vec<f64>,
cached_cmf: Vec<f64>,
have_cache: bool,
online_pmf: RwLock<Arc<Vec<f64>>>,
online_cmf: RwLock<Arc<Vec<f64>>>,
}
impl FragmentLengthDistribution {
pub fn new(
alpha: f64,
max_val: usize,
prior_mu: f64,
prior_sigma: f64,
kernel_n: usize,
kernel_p: f64,
bin_size: usize,
) -> Self {
assert!(bin_size >= 1, "bin_size must be >= 1");
let max_val = max_val / bin_size;
let kernel_n = kernel_n / bin_size;
assert!(
kernel_n.is_multiple_of(2),
"kernel_n must be even after binning"
);
let tot = alpha.ln();
let hist: Vec<AtomicF64>;
let mut sum = LOG_0;
let mut tot_mass;
if prior_mu > 0.0 {
let norm = Normal::new(
prior_mu / bin_size as f64,
prior_sigma / (bin_size * bin_size) as f64,
)
.expect("valid normal prior");
hist = (0..=max_val).map(|_| AtomicF64::new(LOG_0)).collect();
tot_mass = LOG_0;
for (i, slot) in hist.iter().enumerate() {
let norm_mass = norm.cdf(i as f64 + 0.5) - norm.cdf(i as f64 - 0.5);
let mass = if norm_mass != 0.0 {
tot + norm_mass.ln()
} else {
LOG_EPSILON
};
slot.store(mass);
sum = log_add(sum, (i as f64).ln() + mass);
tot_mass = log_add(tot_mass, mass);
}
} else {
let per = tot - (max_val as f64).ln();
hist = (0..=max_val).map(|_| AtomicF64::new(per)).collect();
hist[0].store(LOG_0);
let h1 = hist.get(1).map(|a| a.load()).unwrap_or(per);
sum = h1 + ((max_val * (max_val + 1)) as f64).ln() - 2.0_f64.ln();
tot_mass = tot;
}
let binom = Binomial::new(kernel_p, kernel_n as u64).expect("valid binomial kernel");
let kernel: Vec<f64> = (0..=kernel_n).map(|i| binom.pmf(i as u64).ln()).collect();
Self {
kernel,
hist,
tot_mass: AtomicF64::new(tot_mass),
sum: AtomicF64::new(sum),
min: AtomicUsize::new(max_val),
bin_size,
cached_pmf: Vec::new(),
cached_cmf: Vec::new(),
have_cache: false,
online_pmf: RwLock::new(Arc::new(Vec::new())),
online_cmf: RwLock::new(Arc::new(Vec::new())),
}
}
pub fn default_for_paired() -> Self {
Self::new(1.0, 1000, 0.0, 0.0, 4, 0.5, 1)
}
pub fn max_val(&self) -> usize {
(self.hist.len() - 1) * self.bin_size
}
pub fn min_val(&self) -> usize {
let m = self.min.load(Ordering::Relaxed);
if m == self.hist.len() - 1 {
1
} else {
m
}
}
pub fn add_val(&self, len: usize, mass: f64) {
let mut len = len / self.bin_size;
let max_v = self.max_val() / self.bin_size;
if len > max_v {
len = max_v;
}
self.min.fetch_min(len, Ordering::Relaxed);
let half = self.kernel.len() / 2;
let mut offset = len as isize - half as isize;
for &k in &self.kernel {
if offset > 0 && (offset as usize) < self.hist.len() {
let o = offset as usize;
let k_mass = mass + k;
self.hist[o].log_add_assign(k_mass);
self.sum.log_add_assign((o as f64).ln() + k_mass);
self.tot_mass.log_add_assign(k_mass);
}
offset += 1;
}
}
pub fn pmf(&self, len: usize) -> f64 {
if self.have_cache {
return *self
.cached_pmf
.get(len)
.unwrap_or_else(|| self.cached_pmf.last().unwrap());
}
let mut l = len / self.bin_size;
let max_v = self.max_val() / self.bin_size;
if l > max_v {
l = max_v;
}
self.hist[l].load() - self.tot_mass.load()
}
pub fn refresh_online(&self) {
if self.have_cache {
return;
}
let max_raw = self.max_val();
let max_v = max_raw / self.bin_size;
let tot = self.tot_mass.load();
let mut bin_cum = Vec::with_capacity(max_v + 1);
let mut cum = LOG_0;
for b in 0..=max_v {
cum = log_add(cum, self.hist[b].load() - tot);
bin_cum.push(cum);
}
let mut v = Vec::with_capacity(max_raw + 1);
let mut c = Vec::with_capacity(max_raw + 1);
for raw in 0..=max_raw {
let l = (raw / self.bin_size).min(max_v);
v.push(self.hist[l].load() - tot);
c.push(bin_cum[l]);
}
*self.online_pmf.write().unwrap() = Arc::new(v);
*self.online_cmf.write().unwrap() = Arc::new(c);
}
pub fn online_snapshot(&self) -> Arc<Vec<f64>> {
self.online_pmf.read().unwrap().clone()
}
pub fn online_cmf_snapshot(&self) -> Arc<Vec<f64>> {
self.online_cmf.read().unwrap().clone()
}
pub fn cmf(&self, len: usize) -> f64 {
if self.have_cache {
return *self
.cached_cmf
.get(len)
.unwrap_or_else(|| self.cached_cmf.last().unwrap());
}
let mut l = len / self.bin_size;
let max_v = self.max_val() / self.bin_size;
if l > max_v {
l = max_v;
}
let mut cum = LOG_0;
for i in 0..=l {
cum = log_add(cum, self.hist[i].load());
}
cum - self.tot_mass.load()
}
pub fn tot_mass(&self) -> f64 {
self.tot_mass.load()
}
pub fn mean(&self) -> f64 {
(self.sum.load() - self.tot_mass.load()).exp()
}
pub fn sd(&self) -> f64 {
let lp = self.log_pmf();
if lp.is_empty() {
return 0.0;
}
let mut mean = 0.0;
for (l, &p) in lp.iter().enumerate() {
mean += (l as f64) * p.exp();
}
let mut var = 0.0;
for (l, &p) in lp.iter().enumerate() {
let d = l as f64 - mean;
var += d * d * p.exp();
}
var.max(0.0).sqrt()
}
pub fn cache(&mut self) {
if self.have_cache {
return;
}
let max_v = self.max_val();
let mut pmf = Vec::with_capacity(max_v + 1);
let mut tot = LOG_0;
for i in 0..=max_v {
let p = self.pmf(i);
pmf.push(p);
tot = log_add(tot, p);
}
for p in &mut pmf {
*p -= tot;
}
let mut cmf = Vec::with_capacity(pmf.len());
let mut cum = LOG_0;
for &p in &pmf {
cum = log_add(cum, p);
cmf.push(cum);
}
self.cached_pmf = pmf;
self.cached_cmf = cmf;
self.have_cache = true;
}
pub fn log_pmf(&self) -> &[f64] {
debug_assert!(self.have_cache, "call cache() before log_pmf()");
&self.cached_pmf
}
pub fn conditional_means(&self) -> Vec<f64> {
debug_assert!(self.have_cache, "call cache() before conditional_means()");
let n = self.cached_pmf.len();
let mut cms = vec![0.0f64; n];
let mut vals = 0.0; let mut mult = 0.0; for i in 0..n {
let p = self.cached_pmf[i].exp();
vals += (i as f64) * p;
mult += p;
cms[i] = if mult > 0.0 { vals / mult } else { 0.0 };
}
cms
}
}
#[inline]
fn cmf_at(cmf: &[f64], len: i32) -> f64 {
if cmf.is_empty() {
return LOG_0;
}
let i = (len.max(0) as usize).min(cmf.len() - 1);
cmf[i]
}
pub fn ambig_frag_log_prob(cmf: &[f64], fwd: bool, pos: i32, read_len: i32, txp_len: i32) -> f64 {
if cmf.is_empty() {
return 0.0; }
let stxp = txp_len.max(0);
let max_frag_len = if fwd {
stxp - pos.clamp(0, stxp)
} else {
(pos + read_len).clamp(0, stxp)
};
let ref_cm = cmf_at(cmf, stxp);
if ref_cm <= LOG_0 {
return LOG_EPSILON;
}
cmf_at(cmf, max_frag_len) - ref_cm
}
pub fn smoothed_effective_length(cond_means: &[f64], ref_len: usize) -> f64 {
if cond_means.is_empty() {
return ref_len as f64;
}
let max_len = cond_means.len();
let cf = if ref_len >= max_len {
cond_means[max_len - 1]
} else {
cond_means[ref_len]
};
let eff = ref_len as f64 - cf;
if eff < 1.0 {
ref_len as f64
} else {
eff
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn uniform_prior_pmf_normalizes() {
let mut fld = FragmentLengthDistribution::new(1.0, 200, 0.0, 0.0, 4, 0.5, 1);
fld.cache();
let total: f64 = fld.log_pmf().iter().map(|p| p.exp()).sum();
assert!((total - 1.0).abs() < 1e-9, "pmf sums to {total}");
}
#[test]
fn gaussian_prior_mean_is_near_mu() {
let fld = FragmentLengthDistribution::new(1000.0, 1000, 250.0, 25.0, 4, 0.5, 1);
let m = fld.mean();
assert!((m - 250.0).abs() < 5.0, "mean {m} not near 250");
}
#[test]
fn observations_shift_the_distribution() {
let mut fld = FragmentLengthDistribution::new(1.0, 1000, 250.0, 25.0, 4, 0.5, 1);
for _ in 0..100_000 {
fld.add_val(400, 0.0); }
let m = fld.mean();
assert!(m > 300.0, "mean {m} did not move toward 400");
fld.cache();
let p400 = fld.pmf(400);
let p250 = fld.pmf(250);
assert!(p400 > p250, "p(400)={p400} not > p(250)={p250}");
}
#[test]
fn smoothed_efflen_shrinks_short_transcripts() {
let mut fld = FragmentLengthDistribution::new(1000.0, 1000, 250.0, 25.0, 4, 0.5, 1);
fld.cache();
let cm = fld.conditional_means();
for w in cm.windows(2) {
assert!(
w[1] >= w[0] - 1e-9,
"cond means not monotonic: {} < {}",
w[1],
w[0]
);
}
let short = smoothed_effective_length(&cm, 201);
assert!(
short < 201.0 && short > 1.0,
"short effLen {short} not shrunk"
);
let long = smoothed_effective_length(&cm, 5000);
assert!(long > 4000.0, "long effLen {long} shrunk too much");
let tiny = smoothed_effective_length(&cm, 2);
assert_eq!(tiny, 2.0, "tiny transcript should fall back to refLen");
}
#[test]
fn ambig_frag_prob_bounds_and_orientation() {
let mut fld = FragmentLengthDistribution::new(1000.0, 1000, 250.0, 25.0, 4, 0.5, 1);
fld.cache();
let cmf = fld.cached_cmf.clone();
let txp_len = 2000i32;
let ample = ambig_frag_log_prob(&cmf, true, 100, 75, txp_len);
assert!(ample > -0.01, "ample-space orphan logProb {ample} not ~0");
let crammed = ambig_frag_log_prob(&cmf, true, txp_len - 50, 75, txp_len);
assert!(
crammed < ample - 1.0,
"crammed orphan {crammed} not << ample {ample}"
);
let rc_crammed = ambig_frag_log_prob(&cmf, false, 0, 50, txp_len);
assert!(
rc_crammed < ample - 1.0,
"rc crammed orphan {rc_crammed} not << ample {ample}"
);
assert_eq!(ambig_frag_log_prob(&[], true, 100, 75, txp_len), 0.0);
}
#[test]
fn cmf_is_monotonic() {
let mut fld = FragmentLengthDistribution::new(1000.0, 500, 200.0, 30.0, 4, 0.5, 1);
fld.cache();
let mut prev = f64::NEG_INFINITY;
for l in 0..=500 {
let c = fld.cmf(l);
assert!(c >= prev - 1e-9, "cmf decreased at {l}: {c} < {prev}");
prev = c;
}
assert!((prev - 0.0).abs() < 1e-6, "cmf endpoint {prev} != log(1)");
}
}