use crate::alphabet::{AMINO_FREQ, K, KP};
use crate::hmmfile::{P7Hmm, TDD, TDM, TIM, TII, TMD, TMI, TMM};
const LOG2: f64 = std::f64::consts::LN_2;
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],
_ => &[],
}
}
const NEGINF: i16 = -32768;
#[inline]
fn wordify(scale_w: f32, sc: f32) -> i16 {
if sc == f32::NEG_INFINITY {
return NEGINF;
}
let v = (scale_w * sc).round();
if v >= 32767.0 {
32767
} else if v <= -32768.0 {
NEGINF
} else {
v as i16
}
}
#[inline]
fn adds(a: i16, b: i16) -> i16 {
a.saturating_add(b)
}
#[derive(Debug, Clone)]
pub struct VitFilter {
pub m: usize,
pub scale_w: f32,
pub base_w: i16,
#[allow(dead_code)]
rwv: Vec<[i16; KP]>,
rwv_t: Vec<i16>,
tmm: Vec<i16>,
tim: Vec<i16>,
tdm: Vec<i16>,
vbm: Vec<i16>,
tmi: Vec<i16>,
tii: Vec<i16>,
tmd: Vec<i16>,
tdd: Vec<i16>,
xw_e_move: i16,
xw_e_loop: i16,
pub mu: f64,
pub lambda: f64,
}
impl VitFilter {
pub fn build(hmm: &P7Hmm) -> Self {
let m = hmm.m;
let scale_w = (500.0_f64 / LOG2) as f32;
let base_w: i16 = 12000;
let mut rwv = vec![[NEGINF; 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 {
rwv[k][x] = wordify(scale_w, sc[x]);
}
}
let mut rwv_t = vec![NEGINF; KP * (m + 1)];
for k in 1..=m {
for x in 0..KP {
rwv_t[x * (m + 1) + k] = rwv[k][x];
}
}
let mut tmm = vec![NEGINF; m + 1];
let mut tim = vec![NEGINF; m + 1];
let mut tdm = vec![NEGINF; m + 1];
let mut tmi = vec![NEGINF; m + 1];
let mut tii = vec![NEGINF; m + 1];
let mut tmd = vec![NEGINF; m + 1];
let mut tdd = vec![NEGINF; m + 1];
for j in 1..m {
tmm[j] = wordify(scale_w, (hmm.t[j][TMM] as f64).ln() as f32).min(0);
tim[j] = wordify(scale_w, (hmm.t[j][TIM] as f64).ln() as f32).min(0);
tdm[j] = wordify(scale_w, (hmm.t[j][TDM] as f64).ln() as f32).min(0);
tmi[j] = wordify(scale_w, (hmm.t[j][TMI] as f64).ln() as f32).min(0);
tii[j] = wordify(scale_w, (hmm.t[j][TII] as f64).ln() as f32).min(-1);
tmd[j] = wordify(scale_w, (hmm.t[j][TMD] as f64).ln() as f32).min(0);
tdd[j] = wordify(scale_w, (hmm.t[j][TDD] as f64).ln() as f32).min(0);
}
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 vbm = vec![NEGINF; m];
for k in 1..=m {
let logv = ((occ[k] / z) as f64).ln() as f32;
vbm[k - 1] = wordify(scale_w, logv).min(0);
}
let xw_e_move = wordify(scale_w, -(LOG2 as f32));
let xw_e_loop = wordify(scale_w, -(LOG2 as f32));
VitFilter {
m,
scale_w,
base_w,
rwv,
rwv_t,
tmm,
tim,
tdm,
vbm,
tmi,
tii,
tmd,
tdd,
xw_e_move,
xw_e_loop,
mu: hmm.evparam.vit_mu,
lambda: hmm.evparam.vit_lambda,
}
}
pub fn vit_score(&self, dsq: &[u8], l: usize) -> f32 {
let m = self.m;
let pmove = 3.0_f32 / (l as f32 + 3.0);
let xw_move = wordify(self.scale_w, pmove.ln());
let mut mp = vec![NEGINF; m + 1];
let mut ip = vec![NEGINF; m + 1];
let mut dp = vec![NEGINF; m + 1];
let mut mc = vec![NEGINF; m + 1];
let mut ic = vec![NEGINF; m + 1];
let mut dc = vec![NEGINF; m + 1];
let mut xn: i16 = self.base_w;
let mut xb: i16 = adds(xn, xw_move);
let mut xj: i16 = NEGINF;
let mut xc: i16 = NEGINF;
for i in 1..=l {
let x = dsq[i] as usize;
mc[0] = NEGINF;
ic[0] = NEGINF;
dc[0] = NEGINF;
{
let base = x * (m + 1);
let rwv_row = &self.rwv_t[base + 1..base + m + 1];
let mpm = &mp[0..m];
let ipm = &ip[0..m];
let dpm = &dp[0..m];
let vbm = &self.vbm[0..m];
let tmm = &self.tmm[0..m];
let tim = &self.tim[0..m];
let tdm = &self.tdm[0..m];
let mc_o = &mut mc[1..m + 1];
for j in 0..m {
let sc = adds(xb, vbm[j])
.max(adds(mpm[j], tmm[j]))
.max(adds(ipm[j], tim[j]))
.max(adds(dpm[j], tdm[j]));
mc_o[j] = adds(sc, rwv_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] = adds(mpk[j], tmi[j]).max(adds(ipk[j], tii[j]));
}
}
let mut xe: i16 = NEGINF;
for &v in &mc[1..m + 1] {
if v > xe {
xe = v;
}
}
for k in 1..=m {
dc[k] = adds(mc[k - 1], self.tmd[k - 1]).max(adds(dc[k - 1], self.tdd[k - 1]));
}
if xe >= 32767 {
return f32::INFINITY;
}
xn = adds(xn, 0);
xc = xc.max(adds(xe, self.xw_e_move));
xj = xj.max(adds(xe, self.xw_e_loop));
xb = adds(xj, xw_move).max(adds(xn, xw_move));
std::mem::swap(&mut mp, &mut mc);
std::mem::swap(&mut ip, &mut ic);
std::mem::swap(&mut dp, &mut dc);
}
if xc > NEGINF {
let mut ret = xc as f32 + xw_move as f32 - self.base_w as f32;
ret /= self.scale_w;
ret - 3.0
} else {
f32::NEG_INFINITY
}
}
pub fn pvalue(&self, vfsc: f32, filtersc: f32, _l: usize) -> f64 {
if vfsc.is_infinite() {
return if vfsc > 0.0 { 0.0 } else { 1.0 };
}
let seq_score = (vfsc - filtersc) as f64 / LOG2;
crate::easel::gumbel::esl_gumbel_surv(seq_score, self.mu, self.lambda)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn globins() -> P7Hmm {
let p = format!("{}/testdata/globins4.hmm", env!("CARGO_MANIFEST_DIR"));
P7Hmm::read_all(&p).unwrap().pop().unwrap()
}
#[test]
fn quantization_constants() {
let vf = VitFilter::build(&globins());
assert_eq!(vf.base_w, 12000);
assert_eq!(vf.m, 149);
assert!((vf.scale_w - 721.348).abs() < 0.01, "scale_w={}", vf.scale_w);
assert_eq!(vf.xw_e_move, -500);
assert_eq!(vf.xw_e_loop, -500);
}
}