use crate::alphabet::{AMINO_FREQ, K, KP};
use crate::hmmfile::{P7Hmm, TDD, TDM, TIM, TII, TMD, TMI, TMM};
fn degen_set(code: usize) -> &'static [usize] {
match code {
21 => &[11, 2],
22 => &[7, 9],
23 => &[13, 3],
24 => &[8],
25 => &[1],
26 => &[0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19],
_ => &[],
}
}
pub struct ForwardFilter {
pub m: usize,
pub(crate) rfv: Vec<[f32; KP]>,
pub(crate) rfv_t: Vec<f32>,
pub(crate) tbm: Vec<f32>,
pub(crate) amm: Vec<f32>,
pub(crate) aim: Vec<f32>,
pub(crate) adm: Vec<f32>,
pub(crate) tmi: Vec<f32>,
pub(crate) tii: Vec<f32>,
pub(crate) tmd: Vec<f32>,
pub(crate) tdd: Vec<f32>,
pub ftau: f64,
pub flambda: f64,
}
impl ForwardFilter {
pub fn build(hmm: &P7Hmm) -> Self {
let m = hmm.m;
let mut rfv = vec![[0.0f32; KP]; m + 1];
for k in 1..=m {
let mut sc = [f32::NEG_INFINITY; KP];
for x in 0..K {
sc[x] = ((hmm.mat[k][x] as f64) / AMINO_FREQ[x]).ln() as f32;
}
for code in K..KP {
let set = degen_set(code);
if set.is_empty() {
continue;
}
let mut left = 0.0f32;
let mut right = 0.0f32;
for &y in set {
left += (AMINO_FREQ[y] as f32) * sc[y];
right += AMINO_FREQ[y] as f32;
}
sc[code] = left / right;
}
for x in 0..KP {
rfv[k][x] = sc[x].exp();
}
}
let mut rfv_t = vec![0.0f32; KP * (m + 1)];
for k in 1..=m {
for x in 0..KP {
rfv_t[x * (m + 1) + k] = rfv[k][x];
}
}
let mut occ = vec![0.0f32; m + 1];
if m >= 1 {
occ[1] = hmm.t[0][TMI] + hmm.t[0][TMM];
}
for k in 2..=m {
occ[k] = occ[k - 1] * (hmm.t[k - 1][TMM] + hmm.t[k - 1][TMI])
+ (1.0 - occ[k - 1]) * hmm.t[k - 1][TDM];
}
let mut z = 0.0f32;
for k in 1..=m {
z += occ[k] * (m - k + 1) as f32;
}
let mut tbm = vec![0.0f32; m + 1];
for k in 1..=m {
tbm[k] = ((occ[k] / z) as f64).ln().exp() as f32;
}
let mut amm = vec![0.0f32; m + 1];
let mut aim = vec![0.0f32; m + 1];
let mut adm = vec![0.0f32; m + 1];
for k in 1..=m {
amm[k] = hmm.t[k - 1][TMM];
aim[k] = hmm.t[k - 1][TIM];
adm[k] = hmm.t[k - 1][TDM];
}
let mut tmi = vec![0.0f32; m + 1];
let mut tii = vec![0.0f32; m + 1];
let mut tmd = vec![0.0f32; m + 1];
let mut tdd = vec![0.0f32; m + 1];
for k in 1..m {
tmi[k] = hmm.t[k][TMI];
tii[k] = hmm.t[k][TII];
tmd[k] = hmm.t[k][TMD];
tdd[k] = hmm.t[k][TDD];
}
ForwardFilter {
m,
rfv,
rfv_t,
tbm,
amm,
aim,
adm,
tmi,
tii,
tmd,
tdd,
ftau: hmm.evparam.fwd_tau,
flambda: hmm.evparam.fwd_lambda,
}
}
pub fn score(&self, dsq: &[u8], l: usize) -> f32 {
let m = self.m;
let nj = 1.0f32;
let pmove = (2.0 + nj) / (l as f32 + 2.0 + nj);
let ploop = 1.0 - pmove;
let (xf_e_move, xf_e_loop) = (0.5f32, 0.5f32);
let (xf_n_loop, xf_n_move) = (ploop, pmove);
let (xf_c_loop, xf_c_move) = (ploop, pmove);
let (xf_j_loop, xf_j_move) = (ploop, pmove);
let mut mp = vec![0.0f32; m + 1];
let mut ip = vec![0.0f32; m + 1];
let mut dp = vec![0.0f32; m + 1];
let mut mc = vec![0.0f32; m + 1];
let mut ic = vec![0.0f32; m + 1];
let mut dc = vec![0.0f32; m + 1];
let mut xn = 1.0f32;
let mut xj = 0.0f32;
let mut xb = xf_n_move;
let mut xc = 0.0f32;
let mut totscale = 0.0f64;
for i in 1..=l {
let x = dsq[i] as usize;
mc[0] = 0.0;
ic[0] = 0.0;
dc[0] = 0.0;
{
let base = x * (m + 1);
let rfv_row = &self.rfv_t[base + 1..base + m + 1];
let mpm = &mp[0..m];
let ipm = &ip[0..m];
let dpm = &dp[0..m];
let tbm = &self.tbm[1..m + 1];
let amm = &self.amm[1..m + 1];
let aim = &self.aim[1..m + 1];
let adm = &self.adm[1..m + 1];
let mc_o = &mut mc[1..m + 1];
for j in 0..m {
let sv = xb * tbm[j] + mpm[j] * amm[j] + ipm[j] * aim[j] + dpm[j] * adm[j];
mc_o[j] = sv * rfv_row[j];
}
}
{
let mpk = &mp[1..m + 1];
let ipk = &ip[1..m + 1];
let tmi = &self.tmi[1..m + 1];
let tii = &self.tii[1..m + 1];
let ic_o = &mut ic[1..m + 1];
for j in 0..m {
ic_o[j] = mpk[j] * tmi[j] + ipk[j] * tii[j];
}
}
dc[1] = 0.0;
for k in 2..=m {
dc[k] = mc[k - 1] * self.tmd[k - 1] + dc[k - 1] * self.tdd[k - 1];
}
let mut xe = 0.0f32;
for k in 1..=m {
xe += mc[k];
}
for k in 1..=m {
xe += dc[k];
}
xn *= xf_n_loop;
xc = (xc * xf_c_loop) + (xe * xf_e_move);
xj = (xj * xf_j_loop) + (xe * xf_e_loop);
xb = (xj * xf_j_move) + (xn * xf_n_move);
if xe > 1.0e4 {
let inv = 1.0 / xe;
xn *= inv;
xc *= inv;
xj *= inv;
xb *= inv;
for k in 0..=m {
mc[k] *= inv;
dc[k] *= inv;
ic[k] *= inv;
}
totscale += (xe as f64).ln();
}
std::mem::swap(&mut mp, &mut mc);
std::mem::swap(&mut ip, &mut ic);
std::mem::swap(&mut dp, &mut dc);
}
(totscale + ((xc as f64) * (xf_c_move as f64)).ln()) as f32
}
pub fn pvalue(&self, fwdsc: f32, filtersc: f32) -> f64 {
let seq_score = (fwdsc - filtersc) as f64 / std::f64::consts::LN_2;
crate::easel::exponential::esl_exp_surv(seq_score, self.ftau, self.flambda)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::bias::BiasFilter;
use crate::msv::{null_one, MsvProfile};
use crate::seqio::read_fasta;
#[test]
fn forward_filter_passes_all_golden_hits() {
let hmm = P7Hmm::read_all(&format!("{}/testdata/globins4.hmm", env!("CARGO_MANIFEST_DIR")))
.unwrap()
.pop()
.unwrap();
let mp = MsvProfile::build(&hmm);
let bf = BiasFilter::build(&hmm).unwrap();
let ff = ForwardFilter::build(&hmm);
let seqs = read_fasta(&format!("{}/testdata/globins45.fa", env!("CARGO_MANIFEST_DIR")))
.unwrap();
let (f1, f3) = (0.02_f64, 1e-5_f64);
let mut n_pass = 0;
for s in &seqs {
let l = s.len();
let usc = mp.msv_score(&s.dsq, l);
if mp.pvalue(usc, null_one(l)) > f1 {
continue;
}
let filtersc = bf.filter_score(&s.dsq, l);
if mp.pvalue(usc, filtersc) > f1 {
continue;
}
let fwdsc = ff.score(&s.dsq, l);
if ff.pvalue(fwdsc, filtersc) <= f3 {
n_pass += 1;
}
}
assert_eq!(n_pass, 45, "Fwd filter pass count {n_pass} != C's 45");
}
}