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>,
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],
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);
for pos in 0..CONTEXT_LENGTH {
let idx = self.index_at(mer, pos);
self.probs[pos * ROWS + idx] += weight;
}
}
pub fn normalize(&mut self) {
if self.trained {
return;
}
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.iter_mut().zip(&other.probs) {
*a += *b - PRIOR; }
}
}
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))
}
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)
}
#[cfg(test)]
mod tests {
use super::*;
#[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 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})"
);
}
}