use crate::alphabet::{AMINO_FREQ, K, KP};
use crate::hmmfile::P7Hmm;
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],
_ => &[],
}
}
#[inline]
fn unbiased_byteify(scale_b: f32, sc: f32) -> u8 {
let v = -(scale_b * sc).round();
if v >= 255.0 {
255
} else if v <= 0.0 {
0
} else {
v as u8
}
}
#[inline]
fn biased_byteify(scale_b: f32, bias_b: u8, sc: f32) -> u8 {
let v = -(scale_b * sc).round();
if v > (255 - bias_b as i32) as f32 {
255
} else {
(v as i32 + bias_b as i32).clamp(0, 255) as u8
}
}
#[derive(Debug, Clone)]
pub struct MsvProfile {
pub m: usize,
pub scale_b: f32,
pub base_b: u8,
pub bias_b: u8,
pub tbm_b: u8,
pub tec_b: u8,
pub rbv: Vec<[u8; KP]>,
pub rbv_t: Vec<u8>,
pub sbv_t: Vec<i8>,
pub mu: f64,
pub lambda: f64,
}
impl MsvProfile {
pub fn build(hmm: &P7Hmm) -> Self {
let m = hmm.m;
let scale_b = (3.0_f64 / std::f64::consts::LN_2) as f32;
let base_b: u8 = 190;
let mut msc = vec![[f32::NEG_INFINITY; KP]; m + 1];
let mut maxsc = 0.0f32;
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;
if sc[x] > maxsc {
maxsc = sc[x];
}
}
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;
}
msc[k] = sc;
}
let bias_b = unbiased_byteify(scale_b, -maxsc);
let mut rbv = vec![[255u8; KP]; m + 1];
for k in 1..=m {
for x in 0..KP {
rbv[k][x] = biased_byteify(scale_b, bias_b, msc[k][x]);
}
}
let mut rbv_t = vec![255u8; KP * (m + 1)];
for k in 1..=m {
for x in 0..KP {
rbv_t[x * (m + 1) + k] = rbv[k][x];
}
}
let mut sbv_t = vec![0i8; KP * (m + 1)];
let tmp: u8 = (127i16 + bias_b as i16) as u8;
for k in 1..=m {
for x in 0..KP {
sbv_t[x * (m + 1) + k] = (tmp.saturating_sub(rbv[k][x]) ^ 127) as i8;
}
}
let tbm_b = unbiased_byteify(scale_b, (2.0_f64 / (m as f64 * (m as f64 + 1.0))).ln() as f32);
let tec_b = unbiased_byteify(scale_b, 0.5_f32.ln());
MsvProfile {
m,
scale_b,
base_b,
bias_b,
tbm_b,
tec_b,
rbv,
rbv_t,
sbv_t,
mu: hmm.evparam.msv_mu,
lambda: hmm.evparam.msv_lambda,
}
}
}
const LOG2: f64 = std::f64::consts::LN_2;
impl MsvProfile {
pub fn msv_score(&self, dsq: &[u8], l: usize) -> f32 {
let tjb_b = unbiased_byteify(self.scale_b, (3.0_f64 / (l as f64 + 3.0)).ln() as f32);
let tjbm = tjb_b.wrapping_add(self.tbm_b);
let m = self.m;
let mut prev = vec![0u8; m + 1];
let mut cur = vec![0u8; m + 1];
let mut xj: u8 = 0;
let mut xb: u8 = self.base_b.saturating_sub(tjbm);
let bias = self.bias_b;
for i in 1..=l {
let res = dsq[i] as usize;
let base = res * (m + 1);
let row = &self.rbv_t[base + 1..base + m + 1];
let prev_km1 = &prev[0..m];
let cur_k = &mut cur[1..m + 1];
for ((c, &p), &r) in cur_k.iter_mut().zip(prev_km1.iter()).zip(row.iter()) {
*c = p.max(xb).saturating_add(bias).saturating_sub(r);
}
let mut xe: u8 = 0;
for &c in cur_k.iter() {
if c > xe {
xe = c;
}
}
if xe.saturating_add(bias) == 255 {
return f32::INFINITY;
}
xe = xe.saturating_sub(self.tec_b);
if xj < xe {
xj = xe;
}
xb = self.base_b.max(xj).saturating_sub(tjbm);
std::mem::swap(&mut prev, &mut cur);
}
((xj as f32 - tjb_b as f32) - self.base_b as f32) / self.scale_b - 3.0
}
pub fn ssv_score(&self, dsq: &[u8], l: usize) -> Option<f32> {
let tjb_b = unbiased_byteify(self.scale_b, (3.0_f64 / (l as f64 + 3.0)).ln() as f32);
if tjb_b as i32 + self.tbm_b as i32 + self.tec_b as i32 + self.bias_b as i32 >= 127 {
return None;
}
let m = self.m;
let mut prev = vec![-128i8; m + 1];
let mut cur = vec![-128i8; m + 1];
let mut raw_xe: u8 = 128;
for i in 1..=l {
let res = dsq[i] as usize;
let base = res * (m + 1);
let srow = &self.sbv_t[base + 1..base + m + 1];
let prev_km1 = &prev[0..m];
let cur_k = &mut cur[1..m + 1];
for ((c, &p), &s) in cur_k.iter_mut().zip(prev_km1.iter()).zip(srow.iter()) {
let sv = p.saturating_sub(s);
*c = sv;
let u = sv as u8;
if u > raw_xe {
raw_xe = u;
}
}
std::mem::swap(&mut prev, &mut cur);
}
if raw_xe as i32 >= 255 - self.bias_b as i32 {
return None;
}
let xe: i32 =
raw_xe as i32 + (self.base_b as i32 - tjb_b as i32 - self.tbm_b as i32) - 128;
if xe >= 255 - self.bias_b as i32 {
return None;
}
if xe < self.tec_b as i32 {
return None;
}
let xj = xe - self.tec_b as i32;
if xj > self.base_b as i32 {
return None;
}
let sc = ((xj as f32 - tjb_b as f32) - self.base_b as f32) / self.scale_b - 3.0;
Some(sc)
}
pub fn pvalue(&self, msvsc_nats: f32, nullsc: f32) -> f64 {
if msvsc_nats.is_infinite() {
return 0.0;
}
let seq_score = (msvsc_nats - nullsc) as f64 / LOG2;
crate::easel::gumbel::esl_gumbel_surv(seq_score, self.mu, self.lambda)
}
}
pub fn null_one(l: usize) -> f32 {
let p1 = l as f64 / (l as f64 + 1.0);
(l as f64 * p1.ln() + (1.0 - p1).ln()) as f32
}
#[cfg(test)]
mod tests {
use super::*;
use crate::seqio::read_fasta;
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 mp = MsvProfile::build(&globins());
assert_eq!(mp.base_b, 190);
assert!((mp.scale_b - 4.328085).abs() < 1e-4, "scale_b={}", mp.scale_b);
assert_eq!(mp.m, 149);
for k in 1..=mp.m {
for x in 0..K {
let _ = mp.rbv[k][x];
}
}
}
#[test]
fn tbm_matches_formula() {
let mp = MsvProfile::build(&globins());
let m = 149.0f64;
let expect = unbiased_byteify(mp.scale_b, (2.0 / (m * (m + 1.0))).ln() as f32);
assert_eq!(mp.tbm_b, expect);
assert_eq!(mp.tec_b, unbiased_byteify(mp.scale_b, 0.5f32.ln()));
}
#[test]
fn msv_passes_all_golden_hits() {
let mp = MsvProfile::build(&globins());
let seqs = read_fasta(&format!("{}/testdata/globins45.fa", env!("CARGO_MANIFEST_DIR")))
.unwrap();
let f1 = 0.02_f64;
let mut n_pass = 0;
let mut top_p = 1.0_f64;
for s in &seqs {
let l = s.len();
let sc = mp.msv_score(&s.dsq, l);
let p = mp.pvalue(sc, null_one(l));
if p <= f1 {
n_pass += 1;
}
top_p = top_p.min(p);
if s.name == "MYG_ESCGI" {
assert!(p <= f1, "MYG_ESCGI failed MSV F1 (P={p:.3e}, score={sc:.2})");
assert!(p < 1e-10, "MYG_ESCGI MSV P unexpectedly weak: {p:.3e}");
}
}
assert_eq!(n_pass, 45, "MSV F1 pass count {n_pass} != C's 45");
assert!(top_p < 1e-10, "best MSV P too weak: {top_p:.3e}");
}
}