use realfft::RealFftPlanner;
use std::cell::RefCell;
thread_local! {
static FFT_PLANNER: RefCell<RealFftPlanner<f64>> = RefCell::new(RealFftPlanner::<f64>::new());
}
fn xcorr_fft(fw: &[f64], rc: &[f64], max_lag: usize) -> Vec<f64> {
let l = fw.len();
debug_assert_eq!(l, rc.len());
let n = (l + max_lag + 1).next_power_of_two().max(2);
FFT_PLANNER.with(|p| {
let mut planner = p.borrow_mut();
let r2c = planner.plan_fft_forward(n);
let c2r = planner.plan_fft_inverse(n);
let mut a = r2c.make_input_vec();
let mut b = r2c.make_input_vec();
a[..l].copy_from_slice(fw);
b[..l].copy_from_slice(rc);
let mut fa = r2c.make_output_vec();
let mut fb = r2c.make_output_vec();
r2c.process(&mut a, &mut fa).expect("rfft fw");
r2c.process(&mut b, &mut fb).expect("rfft rc");
for (x, y) in fa.iter_mut().zip(&fb) {
*x = x.conj() * *y;
}
let mut out = c2r.make_output_vec();
c2r.process(&mut fa, &mut out).expect("irfft");
let scale = 1.0 / n as f64;
out[..=max_lag].iter().map(|v| v * scale).collect()
})
}
const ORDER: [u32; 9] = [0, 1, 2, 2, 2, 2, 2, 2, 2];
pub const CONTEXT_LENGTH: usize = 9;
pub const CONTEXT_LEFT: usize = 3;
pub const CONTEXT_RIGHT: usize = 5;
const ROWS: usize = 64;
const PRIOR: f64 = 1e-10;
const LOG_SMALL: f64 = -11.512_925_464_970_229;
#[inline]
fn base2bit(b: u8) -> u32 {
match b {
b'A' | b'a' => 0,
b'C' | b'c' => 1,
b'G' | b'g' => 2,
b'T' | b't' => 3,
_ => 0,
}
}
#[inline]
fn complement_bit(x: u32) -> u32 {
3 - x }
#[derive(Debug, Clone)]
pub struct SBModel {
probs: Vec<f64>,
probs_fp: Vec<u64>,
marginals: Vec<f64>,
shifts: [u32; CONTEXT_LENGTH],
masks: [u32; CONTEXT_LENGTH],
trained: bool,
}
impl Default for SBModel {
fn default() -> Self {
Self::new()
}
}
impl SBModel {
pub fn new() -> Self {
let mut shifts = [0u32; CONTEXT_LENGTH];
let mut masks = [0u32; CONTEXT_LENGTH];
for i in 0..CONTEXT_LENGTH {
shifts[i] = (2 * CONTEXT_LENGTH as u32) - 2 * (i as u32 + 1);
let width = 2 * (ORDER[i] + 1);
masks[i] = (1u32 << width) - 1;
}
Self {
probs: vec![PRIOR; ROWS * CONTEXT_LENGTH],
probs_fp: vec![0u64; ROWS * CONTEXT_LENGTH],
marginals: vec![PRIOR; 4 * CONTEXT_LENGTH],
shifts,
masks,
trained: false,
}
}
fn encode(context: &[u8], rev_comp: bool) -> u32 {
debug_assert_eq!(context.len(), CONTEXT_LENGTH);
let mut mer = 0u32;
if rev_comp {
for &b in context.iter().rev() {
mer = (mer << 2) | complement_bit(base2bit(b));
}
} else {
for &b in context {
mer = (mer << 2) | base2bit(b);
}
}
mer
}
#[inline]
fn index_at(&self, mer: u32, pos: usize) -> usize {
((mer >> self.shifts[pos]) & self.masks[pos]) as usize
}
pub fn dump(&self) -> &[f64] {
&self.probs
}
pub fn add_context(&mut self, context: &[u8], rev_comp: bool, weight: f64) {
debug_assert!(!self.trained, "cannot add to a normalized model");
let mer = Self::encode(context, rev_comp);
let w = crate::bias_mass_to_fp(weight);
for pos in 0..CONTEXT_LENGTH {
let idx = self.index_at(mer, pos);
self.probs_fp[pos * ROWS + idx] += w;
}
}
pub fn normalize(&mut self) {
if self.trained {
return;
}
for (p, &fp) in self.probs.iter_mut().zip(&self.probs_fp) {
*p = PRIOR + fp as f64 / crate::BIAS_WEIGHT_SCALE;
}
for pos in 0..CONTEXT_LENGTH {
let num_states = 4usize.pow(ORDER[pos]);
for s in 0..num_states {
let node = s * 4;
let base = pos * ROWS + node;
let tot: f64 = self.probs[base..base + 4].iter().sum();
if tot > 0.0 {
for j in 0..4 {
self.probs[base + j] /= tot;
self.marginals[pos * 4 + j] += self.probs[base + j];
}
}
}
for j in 0..4 {
self.marginals[pos * 4 + j] /= num_states as f64;
}
}
for p in &mut self.probs {
*p = if *p > 0.0 { p.ln() } else { LOG_SMALL };
}
self.trained = true;
}
pub fn evaluate_log(&self, context: &[u8], rev_comp: bool) -> f64 {
debug_assert!(self.trained, "evaluate_log requires a normalized model");
let mer = Self::encode(context, rev_comp);
let mut lp = 0.0;
for pos in 0..CONTEXT_LENGTH {
let idx = self.index_at(mer, pos);
lp += self.probs[pos * ROWS + idx];
}
lp
}
pub fn is_trained(&self) -> bool {
self.trained
}
pub fn combine_counts(&mut self, other: &SBModel) {
debug_assert!(!self.trained && !other.trained, "combine before normalize");
for (a, b) in self.probs_fp.iter_mut().zip(&other.probs_fp) {
*a += *b;
}
}
}
pub(crate) fn revcomp_bytes(seq: &[u8]) -> Vec<u8> {
seq.iter()
.rev()
.map(|&b| match b {
b'A' | b'a' => b'T',
b'C' | b'c' => b'G',
b'G' | b'g' => b'C',
b'T' | b't' => b'A',
_ => b'A',
})
.collect()
}
pub(crate) const MIN_ALPHA: f64 = 1e-8;
pub(crate) const MIN_CDF_MASS: f64 = 1e-10;
pub const FLD_SAMP_STRIDE: usize = 5;
pub fn fld_cdf_and_bounds(pmf_lin: &[f64]) -> (Vec<f64>, usize, usize) {
let mut cdf = vec![0.0f64; pmf_lin.len()];
let mut acc = 0.0;
let (mut lo, mut hi) = (0usize, 1usize);
let (mut lb, mut ub) = (false, false);
for i in 0..pmf_lin.len() {
acc += pmf_lin[i];
cdf[i] = acc;
if !lb && acc >= 0.005 {
lb = true;
lo = i;
}
if !ub && acc >= 0.995 {
ub = true;
hi = i;
}
}
(cdf, lo, hi)
}
#[inline]
pub(crate) fn conditional_cdf(cdf: &[f64], cdf_max_arg: usize, cdf_max_val: f64, x: i32) -> f64 {
if x > cdf_max_arg as i32 {
1.0
} else if x <= 0 {
cdf[0] / cdf_max_val
} else {
cdf[x as usize] / cdf_max_val
}
}
pub fn build_expected<'a, F>(
num_targets: usize,
seq_of: F,
alphas: &[f64],
eff_lens: &[f64],
cdf: &[f64],
) -> (SBModel, SBModel)
where
F: Fn(usize) -> &'a [u8] + Sync,
{
use rayon::prelude::*;
let k = CONTEXT_LENGTH;
let cu = CONTEXT_LEFT as i32;
let per_tid = |tid: usize| -> Option<(SBModel, SBModel)> {
if alphas[tid] < MIN_ALPHA || eff_lens[tid] <= 0.0 {
return None;
}
let seq = seq_of(tid);
let ref_len = seq.len();
if ref_len < k {
return None;
}
let cdf_max_arg = (cdf.len() - 1).min(ref_len);
let cdf_max_val = cdf[cdf_max_arg];
if cdf_max_val < MIN_CDF_MASS {
return None;
}
let weight = alphas[tid] / eff_lens[tid];
let rc = revcomp_bytes(seq);
let mut fw = SBModel::new();
let mut rc_m = SBModel::new();
for frag_start in 0..(ref_len - k) {
let max_frag_len = ref_len as i32 - (frag_start as i32 + cu);
if max_frag_len >= 0 && (max_frag_len as usize) < ref_len {
let cdensity = conditional_cdf(cdf, cdf_max_arg, cdf_max_val, max_frag_len);
let w = weight * cdensity;
fw.add_context(&seq[frag_start..frag_start + k], false, w);
rc_m.add_context(&rc[frag_start..frag_start + k], false, w);
}
}
Some((fw, rc_m))
};
let (mut exp_fw, mut exp_rc) = (0..num_targets)
.into_par_iter()
.fold(
|| (SBModel::new(), SBModel::new()),
|mut acc, tid| {
if let Some((fw, rc_m)) = per_tid(tid) {
acc.0.combine_counts(&fw);
acc.1.combine_counts(&rc_m);
}
acc
},
)
.reduce(
|| (SBModel::new(), SBModel::new()),
|mut a, b| {
a.0.combine_counts(&b.0);
a.1.combine_counts(&b.1);
a
},
);
exp_fw.normalize();
exp_rc.normalize();
(exp_fw, exp_rc)
}
#[allow(clippy::too_many_arguments)]
pub fn corrected_effective_length(
seq: &[u8],
cdf: &[f64],
fld_low: usize,
fld_high: usize,
obs_fw: &SBModel,
exp_fw: &SBModel,
obs_rc: &SBModel,
exp_rc: &SBModel,
elen: f64,
stride: usize,
) -> f64 {
let k = CONTEXT_LENGTH;
let cu = CONTEXT_LEFT; let ref_len = seq.len();
let unprocessed = (ref_len as i32 - elen as i32).max(0);
let cdf_max_arg = (cdf.len() - 1).min(ref_len);
let cdf_max_val = cdf[cdf_max_arg];
if ref_len < k || unprocessed <= 0 || cdf_max_val < MIN_CDF_MASS {
return elen;
}
let cond = |x: i32| conditional_cdf(cdf, cdf_max_arg, cdf_max_val, x);
let rc_seq = revcomp_bytes(seq);
let mut fw = vec![1.0f64; ref_len];
let mut rc = vec![1.0f64; ref_len];
for frag_start in 0..(ref_len - k) {
let read_start = frag_start + cu;
if read_start < ref_len {
fw[read_start] =
log_bias(obs_fw, exp_fw, &seq[frag_start..frag_start + k], false).exp();
rc[read_start] =
log_bias(obs_rc, exp_rc, &rc_seq[frag_start..frag_start + k], false).exp();
}
}
rc.reverse();
let stride = stride.max(1) as i32;
let max_len = (ref_len as i32).min(fld_high as i32 + 1);
let mut fl = fld_low as i32;
let mut done = fl >= max_len;
let sp = if fl > 0 { fl - 1 } else { 0 };
let mut prev_mass = cond(sp);
let mut eff = 0.0f64;
while !done {
if fl >= max_len {
done = true;
fl = max_len - 1;
}
let fl_weight = cond(fl) - prev_mass;
prev_mass = cond(fl);
let mut mass = 0.0f64;
let mut kstart = 0i32;
while kstart < ref_len as i32 - fl {
let frag_start = kstart as usize;
let frag_end = (kstart + fl - 1) as usize;
if frag_end < ref_len {
mass += fw[frag_start] * rc[frag_end];
} else {
break;
}
kstart += 1;
}
eff += fl_weight * mass;
fl += stride;
}
let offset = (unprocessed as f64).max(1.0);
eff.max(elen.min(offset))
}
#[allow(clippy::too_many_arguments)]
pub fn eff_len_from_xcorr(
a: &[f64],
b: &[f64],
cond: impl Fn(i32) -> f64,
fld_low: usize,
fld_high: usize,
elen: f64,
unprocessed: i32,
stride: usize,
no_length_threshold: bool,
) -> f64 {
let ref_len = a.len();
debug_assert_eq!(ref_len, b.len());
let max_len = (ref_len as i32).min(fld_high as i32 + 1);
if (fld_low as i32) >= max_len {
let offset = (unprocessed as f64).max(1.0);
return elen.max(elen.min(offset));
}
let max_lag = (max_len - 2).max(0) as usize;
let xc = xcorr_fft(a, b, max_lag);
let b_last = b[ref_len - 1];
let stride = stride.max(1) as i32;
let mut eff = 0.0f64;
let mut fl = fld_low as i32;
let mut done = fl >= max_len;
let sp = if fl > 0 { fl - 1 } else { 0 };
let mut prev_mass = cond(sp);
while !done {
if fl >= max_len {
done = true;
fl = max_len - 1;
}
let fl_weight = cond(fl) - prev_mass;
prev_mass = cond(fl);
if fl >= 1 {
let delta = (fl - 1) as usize;
let boundary = a[(ref_len as i32 - fl) as usize] * b_last;
eff += fl_weight * (xc[delta] - boundary);
}
fl += stride;
}
if no_length_threshold {
if eff > 1.0 {
eff
} else {
elen
}
} else {
let offset = (unprocessed as f64).max(1.0);
eff.max(elen.min(offset))
}
}
#[allow(clippy::too_many_arguments)]
pub fn corrected_effective_length_fft(
seq: &[u8],
cdf: &[f64],
fld_low: usize,
fld_high: usize,
obs_fw: &SBModel,
exp_fw: &SBModel,
obs_rc: &SBModel,
exp_rc: &SBModel,
elen: f64,
stride: usize,
no_length_threshold: bool,
) -> f64 {
let k = CONTEXT_LENGTH;
let cu = CONTEXT_LEFT;
let ref_len = seq.len();
let unprocessed = (ref_len as i32 - elen as i32).max(0);
let cdf_max_arg = (cdf.len() - 1).min(ref_len);
let cdf_max_val = cdf[cdf_max_arg];
if ref_len < k || unprocessed <= 0 || cdf_max_val < MIN_CDF_MASS {
return elen;
}
let cond = |x: i32| conditional_cdf(cdf, cdf_max_arg, cdf_max_val, x);
let rc_seq = revcomp_bytes(seq);
let mut fw = vec![1.0f64; ref_len];
let mut rc = vec![1.0f64; ref_len];
for frag_start in 0..(ref_len - k) {
let read_start = frag_start + cu;
if read_start < ref_len {
fw[read_start] =
log_bias(obs_fw, exp_fw, &seq[frag_start..frag_start + k], false).exp();
rc[read_start] =
log_bias(obs_rc, exp_rc, &rc_seq[frag_start..frag_start + k], false).exp();
}
}
rc.reverse();
eff_len_from_xcorr(
&fw,
&rc,
cond,
fld_low,
fld_high,
elen,
unprocessed,
stride,
no_length_threshold,
)
}
pub fn log_bias(observed: &SBModel, expected: &SBModel, context: &[u8], rev_comp: bool) -> f64 {
observed.evaluate_log(context, rev_comp) - expected.evaluate_log(context, rev_comp)
}
pub struct LogBiasTable {
diff: Vec<f64>,
shifts: [u32; CONTEXT_LENGTH],
masks: [u32; CONTEXT_LENGTH],
}
impl LogBiasTable {
pub fn new(observed: &SBModel, expected: &SBModel) -> Self {
debug_assert!(observed.trained && expected.trained);
let diff = observed
.probs
.iter()
.zip(&expected.probs)
.map(|(&o, &e)| o - e)
.collect();
Self {
diff,
shifts: observed.shifts,
masks: observed.masks,
}
}
#[inline]
pub fn eval(&self, context: &[u8], rev_comp: bool) -> f64 {
let mer = SBModel::encode(context, rev_comp);
let mut lp = 0.0;
for pos in 0..CONTEXT_LENGTH {
let idx = ((mer >> self.shifts[pos]) & self.masks[pos]) as usize;
lp += self.diff[pos * ROWS + idx];
}
lp
}
}
#[cfg(test)]
mod tests {
use super::*;
fn trained_pair(seq: &[u8]) -> (SBModel, SBModel) {
let rc = revcomp_bytes(seq);
let mut exp = SBModel::new();
for p in 0..=(seq.len() - CONTEXT_LENGTH) {
exp.add_context(&seq[p..p + CONTEXT_LENGTH], false, 1.0);
exp.add_context(&rc[p..p + CONTEXT_LENGTH], false, 1.0);
}
let mut obs = exp.clone();
for _ in 0..500 {
obs.add_context(b"AAACCCGGG", false, 1.0);
obs.add_context(b"TTTGGGCCC", true, 1.0);
}
obs.normalize();
exp.normalize();
(obs, exp)
}
#[test]
#[ignore = "profiling bench; run with --ignored --nocapture"]
fn bench_factor_build() {
use std::time::Instant;
let bases = b"ACGTACGTAGGCCTTAACCGGTTACGTACGTTTAGCGATCG";
let seq: Vec<u8> = (0..2000)
.map(|i| bases[(i * 7 + 3) % bases.len()])
.collect();
let (obs, exp) = trained_pair(&seq);
let rc_seq = revcomp_bytes(&seq);
let k = CONTEXT_LENGTH;
let n = seq.len() - k;
let iters = 8000usize;
let t = Instant::now();
let mut acc = 0.0f64;
for _ in 0..iters {
for fs in 0..n {
acc += log_bias(&obs, &exp, &seq[fs..fs + k], false).exp();
acc += log_bias(&obs, &exp, &rc_seq[fs..fs + k], false).exp();
}
}
let v0 = t.elapsed().as_secs_f64();
let t = Instant::now();
let mut acc1 = 0.0f64;
for _ in 0..iters {
for fs in 0..n {
acc1 += log_bias(&obs, &exp, &seq[fs..fs + k], false);
acc1 += log_bias(&obs, &exp, &rc_seq[fs..fs + k], false);
}
}
let v_noexp = t.elapsed().as_secs_f64();
let t = Instant::now();
let mut enc = 0u64;
for _ in 0..iters {
for fs in 0..n {
enc ^= SBModel::encode(&seq[fs..fs + k], false) as u64;
enc ^= SBModel::encode(&rc_seq[fs..fs + k], false) as u64;
}
}
let v_enc = t.elapsed().as_secs_f64();
let eval_mer = |m: &SBModel, mer: u32| -> f64 {
let mut lp = 0.0;
for pos in 0..CONTEXT_LENGTH {
lp += m.probs[pos * ROWS + m.index_at(mer, pos)];
}
lp
};
let t = Instant::now();
let mut acc_v1 = 0.0f64;
let mut max_d1 = 0.0f64;
for it in 0..iters {
for fs in 0..n {
let mf = SBModel::encode(&seq[fs..fs + k], false);
let mr = SBModel::encode(&rc_seq[fs..fs + k], false);
let bf = (eval_mer(&obs, mf) - eval_mer(&exp, mf)).exp();
let br = (eval_mer(&obs, mr) - eval_mer(&exp, mr)).exp();
acc_v1 += bf + br;
if it == 0 {
let rf = log_bias(&obs, &exp, &seq[fs..fs + k], false).exp();
let rr = log_bias(&obs, &exp, &rc_seq[fs..fs + k], false).exp();
max_d1 = max_d1.max((bf - rf).abs()).max((br - rr).abs());
}
}
}
let v1 = t.elapsed().as_secs_f64();
let mut diff = vec![0.0f64; ROWS * CONTEXT_LENGTH];
for (i, d) in diff.iter_mut().enumerate() {
*d = obs.probs[i] - exp.probs[i];
}
let eval_diff = |mer: u32| -> f64 {
let mut lp = 0.0;
for pos in 0..CONTEXT_LENGTH {
lp += diff[pos * ROWS + obs.index_at(mer, pos)];
}
lp
};
let t = Instant::now();
let mut acc_v2 = 0.0f64;
let mut max_d2 = 0.0f64;
for it in 0..iters {
for fs in 0..n {
let bf = eval_diff(SBModel::encode(&seq[fs..fs + k], false)).exp();
let br = eval_diff(SBModel::encode(&rc_seq[fs..fs + k], false)).exp();
acc_v2 += bf + br;
if it == 0 {
let rf = log_bias(&obs, &exp, &seq[fs..fs + k], false).exp();
let rr = log_bias(&obs, &exp, &rc_seq[fs..fs + k], false).exp();
max_d2 = max_d2.max((bf - rf).abs()).max((br - rr).abs());
}
}
}
let v2 = t.elapsed().as_secs_f64();
eprintln!("--- factor-build bench ({iters} iters x {n} pos x2) ---");
eprintln!("V0 current (log_bias+exp) : {v0:.3}s acc={acc:.3}");
eprintln!(
" no-exp (log_bias only) : {v_noexp:.3}s acc={acc1:.3} => exp cost ~{:.3}s",
v0 - v_noexp
);
eprintln!(" encode-only (2x/pos) : {v_enc:.3}s enc={enc}");
eprintln!("V1 encode-once (byte-ident) : {v1:.3}s acc={acc_v1:.3} max|Δ|={max_d1:.3e} speedup={:.2}x", v0 / v1);
eprintln!("V2 diff-table (reassoc) : {v2:.3}s acc={acc_v2:.3} max|Δ|={max_d2:.3e} speedup={:.2}x", v0 / v2);
}
#[test]
fn uniform_contexts_give_near_zero_bias() {
let ctxs: Vec<Vec<u8>> = (0..256)
.map(|i| {
let bases = b"ACGT";
(0..CONTEXT_LENGTH)
.map(|p| bases[((i >> (p * 2)) & 3) as usize])
.collect()
})
.collect();
let mut obs = SBModel::new();
let mut exp = SBModel::new();
for c in &ctxs {
obs.add_context(c, false, 1.0);
exp.add_context(c, false, 1.0);
}
obs.normalize();
exp.normalize();
for c in &ctxs {
assert!(log_bias(&obs, &exp, c, false).abs() < 1e-9);
}
}
#[test]
fn enriched_context_has_positive_bias() {
let target: Vec<u8> = b"ACGTACGTA".to_vec();
let bases = b"ACGT";
let uniform: Vec<Vec<u8>> = (0..4096)
.map(|i| {
(0..CONTEXT_LENGTH)
.map(|p| bases[((i >> (p * 2)) & 3) as usize])
.collect()
})
.collect();
let mut exp = SBModel::new();
for c in &uniform {
exp.add_context(c, false, 1.0);
}
exp.normalize();
let mut obs = SBModel::new();
for c in &uniform {
obs.add_context(c, false, 1.0);
}
for _ in 0..5000 {
obs.add_context(&target, false, 1.0); }
obs.normalize();
assert!(
log_bias(&obs, &exp, &target, false) > 0.5,
"enriched context should have positive log-bias"
);
}
#[test]
fn unbiased_correction_reduces_to_standard_eff_len() {
let bases = b"ACGTACGTAGGCCTTAACCGGTTACGTACGT";
let seq: Vec<u8> = (0..400).map(|i| bases[i % bases.len()]).collect();
let mut m = SBModel::new();
let rc = revcomp_bytes(&seq);
for p in 0..=(seq.len() - CONTEXT_LENGTH) {
m.add_context(&seq[p..p + CONTEXT_LENGTH], false, 1.0);
m.add_context(&rc[p..p + CONTEXT_LENGTH], false, 1.0);
}
let mut obs = m.clone();
let mut exp = m.clone();
obs.normalize();
exp.normalize();
let mut pmf = vec![0.0; 200];
pmf[100] = 1.0;
let (cdf, lo, hi) = fld_cdf_and_bounds(&pmf);
let eff = corrected_effective_length(&seq, &cdf, lo, hi, &obs, &exp, &obs, &exp, 300.0, 1);
assert!((eff - 300.0).abs() < 1e-6, "got {eff}");
}
#[test]
fn fft_matches_exact_scalar_corrected_eff_len() {
let bases = b"ACGTACGTAGGCCTTAACCGGTTACGTACGTTTAGCGATCG";
let seq: Vec<u8> = (0..1500)
.map(|i| bases[(i * 7 + 3) % bases.len()])
.collect();
let rc = revcomp_bytes(&seq);
let mut exp = SBModel::new();
for p in 0..=(seq.len() - CONTEXT_LENGTH) {
exp.add_context(&seq[p..p + CONTEXT_LENGTH], false, 1.0);
exp.add_context(&rc[p..p + CONTEXT_LENGTH], false, 1.0);
}
let mut obs = exp.clone();
let t1 = b"AAACCCGGG";
let t2 = b"TTTGGGCCC";
for _ in 0..500 {
obs.add_context(t1, false, 1.0);
obs.add_context(t2, true, 1.0);
}
obs.normalize();
exp.normalize();
let mut pmf = vec![0.0f64; 600];
for (l, v) in pmf.iter_mut().enumerate() {
let d = l as f64 - 250.0;
*v = (-d * d / (2.0 * 40.0 * 40.0)).exp();
}
let (cdf, lo, hi) = fld_cdf_and_bounds(&pmf);
for stride in [1usize, 5] {
let scalar = corrected_effective_length(
&seq, &cdf, lo, hi, &obs, &exp, &obs, &exp, 1200.0, stride,
);
let fft = corrected_effective_length_fft(
&seq, &cdf, lo, hi, &obs, &exp, &obs, &exp, 1200.0, stride, false,
);
let rel = (scalar - fft).abs() / scalar.abs();
assert!(
rel < 1e-9,
"FFT vs scalar mismatch at stride={stride}: scalar={scalar} fft={fft} rel={rel:.3e}"
);
}
}
#[test]
fn eff_len_from_xcorr_matches_scalar_combined_factors() {
let ref_len = 1300usize;
let a: Vec<f64> = (0..ref_len)
.map(|i| 0.5 + 1.5 * ((i as f64 * 0.013).sin() * 0.5 + 0.5))
.collect();
let b: Vec<f64> = (0..ref_len)
.map(|i| 0.4 + 1.8 * ((i as f64 * 0.021 + 1.0).cos() * 0.5 + 0.5))
.collect();
let mut pmf = vec![0.0f64; 600];
for (l, v) in pmf.iter_mut().enumerate() {
let d = l as f64 - 250.0;
*v = (-d * d / (2.0 * 40.0 * 40.0)).exp();
}
let (cdf, lo, hi) = fld_cdf_and_bounds(&pmf);
let elen = 1100.0f64;
let unprocessed = (ref_len as i32 - elen as i32).max(0);
let cdf_max_arg = (cdf.len() - 1).min(ref_len);
let cdf_max_val = cdf[cdf_max_arg];
let cond = |x: i32| conditional_cdf(&cdf, cdf_max_arg, cdf_max_val, x);
for stride in [1usize, 5] {
let max_len = (ref_len as i32).min(hi as i32 + 1);
let st = stride.max(1) as i32;
let mut fl = lo as i32;
let mut done = fl >= max_len;
let sp = if fl > 0 { fl - 1 } else { 0 };
let mut prev = cond(sp);
let mut eff = 0.0f64;
while !done {
if fl >= max_len {
done = true;
fl = max_len - 1;
}
let w = cond(fl) - prev;
prev = cond(fl);
let kmax = ref_len as i32 - fl;
let mut mass = 0.0f64;
let mut k = 0i32;
while k < kmax {
mass += a[k as usize] * b[(k + fl - 1) as usize];
k += 1;
}
eff += w * mass;
fl += st;
}
let offset = (unprocessed as f64).max(1.0);
let scalar = eff.max(elen.min(offset));
let fft = eff_len_from_xcorr(&a, &b, cond, lo, hi, elen, unprocessed, stride, false);
let rel = (scalar - fft).abs() / scalar.abs();
assert!(
rel < 1e-9,
"combined-factor FFT vs scalar mismatch at stride={stride}: scalar={scalar} fft={fft} rel={rel:.3e}"
);
}
}
#[test]
fn revcomp_encoding_is_consistent() {
let ctx: Vec<u8> = b"ACGTACGTA".to_vec();
let rc: Vec<u8> = ctx
.iter()
.rev()
.map(|&b| match b {
b'A' => b'T',
b'C' => b'G',
b'G' => b'C',
b'T' => b'A',
x => x,
})
.collect();
assert_eq!(SBModel::encode(&ctx, true), SBModel::encode(&rc, false));
}
#[test]
fn build_expected_respects_num_targets_bound() {
let bases = b"ACGTACGTAGGCCTTAACCGGTTACGTACGT";
let mut refs: Vec<Vec<u8>> = (0..5)
.map(|s| (0..200).map(|i| bases[(i + s) % bases.len()]).collect())
.collect();
refs.push(
(0..400)
.map(|i| if i % 2 == 0 { b'A' } else { b'C' })
.collect(),
);
let num_refs = refs.len();
let alphas = vec![1.0; num_refs];
let eff_lens = vec![150.0; num_refs];
let mut pmf = vec![0.0; 200];
pmf[100] = 1.0;
let (cdf, _lo, _hi) = fld_cdf_and_bounds(&pmf);
let (a_fw, _) = build_expected(5, |t| refs[t].as_slice(), &alphas, &eff_lens, &cdf);
let (b_fw, _) = build_expected(6, |t| refs[t].as_slice(), &alphas, &eff_lens, &cdf);
assert!(a_fw.is_trained() && b_fw.is_trained());
assert!(a_fw.dump().iter().all(|v| v.is_finite()));
let diff: f64 = a_fw
.dump()
.iter()
.zip(b_fw.dump())
.map(|(x, y)| (x - y).abs())
.sum();
assert!(
diff > 1e-6,
"a target beyond num_targets must not contribute (diff={diff})"
);
let mut alphas0 = alphas.clone();
alphas0[5] = 0.0;
let (c_fw, _) = build_expected(6, |t| refs[t].as_slice(), &alphas0, &eff_lens, &cdf);
let diff2: f64 = a_fw
.dump()
.iter()
.zip(c_fw.dump())
.map(|(x, y)| (x - y).abs())
.sum();
assert!(
diff2 < 1e-9,
"zero-alpha target must not contribute (diff={diff2})"
);
}
}