use crate::alphabet::{K, KP};
use crate::forward::ForwardFilter;
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(crate) struct XFactors {
pub(crate) n_loop: f32,
pub(crate) n_move: f32,
pub(crate) c_loop: f32,
pub(crate) c_move: f32,
pub(crate) j_loop: f32,
pub(crate) j_move: f32,
pub(crate) e_move: f32,
pub(crate) e_loop: f32,
}
impl XFactors {
fn unihit(save_l: usize) -> Self {
let nj = 0.0f32;
let pmove = (2.0 + nj) / (save_l as f32 + 2.0 + nj);
let ploop = 1.0 - pmove;
XFactors {
n_loop: ploop,
n_move: pmove,
c_loop: ploop,
c_move: pmove,
j_loop: ploop,
j_move: pmove,
e_move: 1.0,
e_loop: 0.0,
}
}
pub(crate) fn multihit(save_l: usize) -> Self {
let nj = 1.0f32;
let pmove = (2.0 + nj) / (save_l as f32 + 2.0 + nj);
let ploop = 1.0 - pmove;
XFactors {
n_loop: ploop,
n_move: pmove,
c_loop: ploop,
c_move: pmove,
j_loop: ploop,
j_move: pmove,
e_move: 0.5,
e_loop: 0.5,
}
}
}
pub(crate) struct Omx {
pub(crate) m: usize,
pub(crate) ld: usize,
pub(crate) mmx: Vec<Vec<f32>>,
pub(crate) imx: Vec<Vec<f32>>,
pub(crate) dmx: Vec<Vec<f32>>,
pub(crate) xe: Vec<f32>,
pub(crate) xn: Vec<f32>,
pub(crate) xj: Vec<f32>,
pub(crate) xb: Vec<f32>,
pub(crate) xc: Vec<f32>,
pub(crate) scale: Vec<f32>,
}
impl Omx {
fn new(m: usize, ld: usize) -> Self {
Omx {
m,
ld,
mmx: vec![vec![0.0f32; m + 1]; ld + 1],
imx: vec![vec![0.0f32; m + 1]; ld + 1],
dmx: vec![vec![0.0f32; m + 1]; ld + 1],
xe: vec![0.0f32; ld + 1],
xn: vec![0.0f32; ld + 1],
xj: vec![0.0f32; ld + 1],
xb: vec![0.0f32; ld + 1],
xc: vec![0.0f32; ld + 1],
scale: vec![1.0f32; ld + 1],
}
}
}
pub(crate) fn forward(ff: &ForwardFilter, xf: &XFactors, sub: &[u8], ld: usize) -> Omx {
let m = ff.m;
let mut ox = Omx::new(m, ld);
ox.xe[0] = 0.0;
ox.xn[0] = 1.0;
ox.xj[0] = 0.0;
ox.xb[0] = xf.n_move;
ox.xc[0] = 0.0;
ox.scale[0] = 1.0;
let mut xn = 1.0f32;
let mut xj = 0.0f32;
let mut xb = xf.n_move;
let mut xc = 0.0f32;
let mut mp = vec![0.0f32; m + 1];
let mut ip = vec![0.0f32; m + 1];
let mut dp = vec![0.0f32; m + 1];
for i in 1..=ld {
let x = sub[i] as usize;
let mc = &mut ox.mmx[i];
let ic = &mut ox.imx[i];
let dc = &mut ox.dmx[i];
mc[0] = 0.0;
ic[0] = 0.0;
dc[0] = 0.0;
{
let base = x * (m + 1);
let rfv_row = &ff.rfv_t[base + 1..base + m + 1];
let mpm = &mp[0..m];
let ipm = &ip[0..m];
let dpm = &dp[0..m];
let tbm = &ff.tbm[1..m + 1];
let amm = &ff.amm[1..m + 1];
let aim = &ff.aim[1..m + 1];
let adm = &ff.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 = &ff.tmi[1..m + 1];
let tii = &ff.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] * ff.tmd[k - 1] + dc[k - 1] * ff.tdd[k - 1];
}
let q_n = (m + 3) / 4;
let mut lane = [0.0f32; 4];
for (r, l) in lane.iter_mut().enumerate() {
for q in 0..q_n {
let k = r * q_n + q + 1;
if k <= m {
*l += mc[k];
}
}
}
for (r, l) in lane.iter_mut().enumerate() {
for q in 0..q_n {
let k = r * q_n + q + 1;
if k <= m {
*l += dc[k];
}
}
}
let mut xe = (lane[0] + lane[1]) + (lane[2] + lane[3]);
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;
}
ox.scale[i] = xe;
xe = 1.0;
} else {
ox.scale[i] = 1.0;
}
ox.xe[i] = xe;
ox.xn[i] = xn;
ox.xj[i] = xj;
ox.xb[i] = xb;
ox.xc[i] = xc;
mp.copy_from_slice(&ox.mmx[i]);
ip.copy_from_slice(&ox.imx[i]);
dp.copy_from_slice(&ox.dmx[i]);
}
ox
}
fn backward(ff: &ForwardFilter, xf: &XFactors, sub: &[u8], ld: usize, fwd: &Omx) -> Omx {
let m = ff.m;
let mut bck = Omx::new(m, ld);
let mut xj = 0.0f32;
let mut xb = 0.0f32;
let mut xn = 0.0f32;
let mut xc = xf.c_move;
let mut xe = xc * xf.e_move;
{
let mc = &mut bck.mmx[ld];
let dc = &mut bck.dmx[ld];
let ic = &mut bck.imx[ld];
mc[m] = xe;
dc[m] = xe;
ic[m] = 0.0;
for k in (1..m).rev() {
dc[k] = xe + dc[k + 1] * ff.tdd[k];
mc[k] = xe + dc[k + 1] * ff.tmd[k];
ic[k] = 0.0;
}
}
let scl = fwd.scale[ld];
if scl > 1.0 {
let inv = 1.0 / scl;
xe *= inv;
xn *= inv;
xj *= inv;
xb *= inv;
xc *= inv;
for k in 0..=m {
bck.mmx[ld][k] *= inv;
bck.dmx[ld][k] *= inv;
bck.imx[ld][k] *= inv;
}
}
bck.scale[ld] = scl;
bck.xe[ld] = xe;
bck.xn[ld] = xn;
bck.xj[ld] = xj;
bck.xb[ld] = xb;
bck.xc[ld] = xc;
for i in (1..ld).rev() {
let x = sub[i + 1] as usize;
let mut b = 0.0f32;
for k in 1..=m {
b += bck.mmx[i + 1][k] * ff.tbm[k] * ff.rfv[k][x];
}
xb = b;
xc = xc * xf.c_loop;
xj = (xb * xf.j_move) + (xj * xf.j_loop);
xn = (xb * xf.n_move) + (xn * xf.n_loop);
xe = (xc * xf.e_move) + (xj * xf.e_loop);
{
let mm_next = &bck.mmx[i + 1];
let im_next = &bck.imx[i + 1];
let mut mc = vec![0.0f32; m + 1];
let mut dc = vec![0.0f32; m + 1];
let mut ic = vec![0.0f32; m + 1];
mc[m] = xe;
dc[m] = xe;
ic[m] = 0.0;
for k in (1..m).rev() {
let m_emit = mm_next[k + 1] * ff.rfv[k + 1][x];
mc[k] = m_emit * ff.amm[k + 1]
+ im_next[k] * ff.tmi[k]
+ xe
+ dc[k + 1] * ff.tmd[k];
ic[k] = m_emit * ff.aim[k + 1] + im_next[k] * ff.tii[k];
dc[k] = m_emit * ff.adm[k + 1] + dc[k + 1] * ff.tdd[k] + xe;
}
bck.mmx[i] = mc;
bck.dmx[i] = dc;
bck.imx[i] = ic;
}
let scl = fwd.scale[i];
if scl > 1.0 {
let inv = 1.0 / scl;
xe *= inv;
xn *= inv;
xj *= inv;
xb *= inv;
xc *= inv;
for k in 0..=m {
bck.mmx[i][k] *= inv;
bck.dmx[i][k] *= inv;
bck.imx[i][k] *= inv;
}
}
bck.scale[i] = scl;
bck.xe[i] = xe;
bck.xn[i] = xn;
bck.xj[i] = xj;
bck.xb[i] = xb;
bck.xc[i] = xc;
}
let x1 = sub[1] as usize;
let mut b0 = 0.0f32;
for k in 1..=m {
b0 += bck.mmx[1][k] * ff.rfv[k][x1] * ff.tbm[k];
}
xn = (b0 * xf.n_move) + (xn * xf.n_loop);
bck.xb[0] = b0;
bck.xc[0] = 0.0;
bck.xj[0] = 0.0;
bck.xn[0] = xn;
bck.xe[0] = 0.0;
bck.scale[0] = 1.0;
bck
}
fn backward_full(
ff: &ForwardFilter,
xf: &XFactors,
sub: &[u8],
ld: usize,
fwd: &Omx,
) -> (Omx, bool) {
let m = ff.m;
let mut bck = Omx::new(m, ld);
let mut has_own_scales = false;
let mut xj = 0.0f32;
let mut xb = 0.0f32;
let mut xn = 0.0f32;
let mut xc = xf.c_move;
let mut xe = xc * xf.e_move;
{
let mc = &mut bck.mmx[ld];
let dc = &mut bck.dmx[ld];
let ic = &mut bck.imx[ld];
mc[m] = xe;
dc[m] = xe;
ic[m] = 0.0;
for k in (1..m).rev() {
dc[k] = xe + dc[k + 1] * ff.tdd[k];
mc[k] = xe + dc[k + 1] * ff.tmd[k];
ic[k] = 0.0;
}
}
let scl = fwd.scale[ld];
if scl > 1.0 {
let inv = 1.0 / scl;
xe *= inv;
xn *= inv;
xj *= inv;
xb *= inv;
xc *= inv;
for k in 0..=m {
bck.mmx[ld][k] *= inv;
bck.dmx[ld][k] *= inv;
bck.imx[ld][k] *= inv;
}
}
bck.scale[ld] = scl;
bck.xe[ld] = xe;
bck.xn[ld] = xn;
bck.xj[ld] = xj;
bck.xb[ld] = xb;
bck.xc[ld] = xc;
for i in (1..ld).rev() {
let x = sub[i + 1] as usize;
let mut b = 0.0f32;
for k in 1..=m {
b += bck.mmx[i + 1][k] * ff.tbm[k] * ff.rfv[k][x];
}
xb = b;
xc = xc * xf.c_loop;
xj = (xb * xf.j_move) + (xj * xf.j_loop);
xn = (xb * xf.n_move) + (xn * xf.n_loop);
xe = (xc * xf.e_move) + (xj * xf.e_loop);
{
let mm_next = &bck.mmx[i + 1];
let im_next = &bck.imx[i + 1];
let mut mc = vec![0.0f32; m + 1];
let mut dc = vec![0.0f32; m + 1];
let mut ic = vec![0.0f32; m + 1];
mc[m] = xe;
dc[m] = xe;
ic[m] = 0.0;
for k in (1..m).rev() {
let m_emit = mm_next[k + 1] * ff.rfv[k + 1][x];
mc[k] = m_emit * ff.amm[k + 1]
+ im_next[k] * ff.tmi[k]
+ xe
+ dc[k + 1] * ff.tmd[k];
ic[k] = m_emit * ff.aim[k + 1] + im_next[k] * ff.tii[k];
dc[k] = m_emit * ff.adm[k + 1] + dc[k + 1] * ff.tdd[k] + xe;
}
bck.mmx[i] = mc;
bck.dmx[i] = dc;
bck.imx[i] = ic;
}
if xb > 1.0e16 {
has_own_scales = true;
}
let scl = if has_own_scales {
if xb > 1.0e4 {
xb
} else {
1.0
}
} else {
fwd.scale[i]
};
if scl > 1.0 {
let inv = 1.0 / scl;
xe *= inv;
xn *= inv;
xj *= inv;
xb *= inv;
xc *= inv;
for k in 0..=m {
bck.mmx[i][k] *= inv;
bck.dmx[i][k] *= inv;
bck.imx[i][k] *= inv;
}
}
bck.scale[i] = scl;
bck.xe[i] = xe;
bck.xn[i] = xn;
bck.xj[i] = xj;
bck.xb[i] = xb;
bck.xc[i] = xc;
}
let x1 = sub[1] as usize;
let mut b0 = 0.0f32;
for k in 1..=m {
b0 += bck.mmx[1][k] * ff.rfv[k][x1] * ff.tbm[k];
}
xn = (b0 * xf.n_move) + (xn * xf.n_loop);
bck.xb[0] = b0;
bck.xc[0] = 0.0;
bck.xj[0] = 0.0;
bck.xn[0] = xn;
bck.xe[0] = 0.0;
bck.scale[0] = 1.0;
(bck, has_own_scales)
}
pub(crate) fn domain_decoding(
ff: &ForwardFilter,
dsq: &[u8],
l: usize,
) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
let xf = XFactors::multihit(l);
let fwd = forward(ff, &xf, dsq, l);
let (bck, has_own_scales) = backward_full(ff, &xf, dsq, l, &fwd);
let mut btot = vec![0.0f32; l + 1];
let mut etot = vec![0.0f32; l + 1];
let mut mocc = vec![0.0f32; l + 1];
let mut scaleproduct = 1.0f32 / bck.xn[0];
for i in 1..=l {
btot[i] = btot[i - 1]
+ (fwd.xb[i - 1] * bck.xb[i - 1] * fwd.scale[i - 1] * scaleproduct);
if has_own_scales {
scaleproduct *= fwd.scale[i - 1] / bck.scale[i - 1];
}
etot[i] = etot[i - 1] + (fwd.xe[i] * bck.xe[i] * fwd.scale[i] * scaleproduct);
let mut njcp = fwd.xn[i - 1] * bck.xn[i] * xf.n_loop * scaleproduct;
njcp += fwd.xj[i - 1] * bck.xj[i] * xf.j_loop * scaleproduct;
njcp += fwd.xc[i - 1] * bck.xc[i] * xf.c_loop * scaleproduct;
mocc[i] = 1.0 - njcp;
}
(btot, etot, mocc)
}
fn decoding(xf: &XFactors, fwd: &Omx, bck: &Omx) -> Omx {
let m = fwd.m;
let ld = fwd.ld;
let mut pp = Omx::new(m, ld);
let mut scaleproduct = 1.0 / bck.xn[0];
let bck_has_own = false;
for i in 1..=ld {
let totrv = scaleproduct * fwd.scale[i];
{
let fm = &fwd.mmx[i][1..m + 1];
let bm = &bck.mmx[i][1..m + 1];
let pm = &mut pp.mmx[i][1..m + 1];
for j in 0..m {
pm[j] = fm[j] * bm[j] * totrv;
}
}
{
let fi = &fwd.imx[i][1..m + 1];
let bi = &bck.imx[i][1..m + 1];
let pi = &mut pp.imx[i][1..m + 1];
for j in 0..m {
pi[j] = fi[j] * bi[j] * totrv;
}
}
pp.xe[i] = 0.0;
pp.xn[i] = fwd.xn[i - 1] * bck.xn[i] * xf.n_loop * scaleproduct;
pp.xj[i] = fwd.xj[i - 1] * bck.xj[i] * xf.j_loop * scaleproduct;
pp.xc[i] = fwd.xc[i - 1] * bck.xc[i] * xf.c_loop * scaleproduct;
pp.xb[i] = 0.0;
if bck_has_own {
scaleproduct *= fwd.scale[i] / bck.scale[i];
}
}
pp
}
pub fn envelope_rescore(
ff: &ForwardFilter,
dsq: &[u8],
i: usize,
j: usize,
save_l: usize,
) -> (f32, [f32; KP]) {
let xf = XFactors::unihit(save_l);
let m = ff.m;
let ld = j - i + 1;
let mut sub = vec![255u8];
sub.extend_from_slice(&dsq[i..=j]);
sub.push(255u8);
let fwd = forward(ff, &xf, &sub, ld);
let mut totscale = 0.0f64;
for r in 1..=ld {
totscale += (fwd.scale[r] as f64).ln();
}
let envsc = (totscale + ((fwd.xc[ld] as f64) * (xf.c_move as f64)).ln()) as f32;
let bck = backward(ff, &xf, &sub, ld, &fwd);
let mut pp = decoding(&xf, &fwd, &bck);
{
let (r0, r1) = pp.mmx.split_at_mut(1);
r0[0].copy_from_slice(&r1[0]);
}
{
let (r0, r1) = pp.imx.split_at_mut(1);
r0[0].copy_from_slice(&r1[0]);
}
pp.xn[0] = pp.xn[1];
pp.xc[0] = pp.xc[1];
pp.xj[0] = pp.xj[1];
for r in 2..=ld {
for k in 0..=m {
pp.mmx[0][k] += pp.mmx[r][k];
pp.imx[0][k] += pp.imx[r][k];
}
pp.xn[0] += pp.xn[r];
pp.xc[0] += pp.xc[r];
pp.xj[0] += pp.xj[r];
}
let norm = 1.0f32 / ld as f32;
for k in 0..=m {
pp.mmx[0][k] *= norm;
pp.imx[0][k] *= norm;
}
pp.xn[0] *= norm;
pp.xc[0] *= norm;
pp.xj[0] *= norm;
let xfactor = pp.xn[0] + pp.xc[0] + pp.xj[0];
let mut null2 = [0.0f32; KP];
for x in 0..K {
let mut sv = 0.0f32;
for k in 1..m {
sv += pp.mmx[0][k] * ff.rfv[k][x];
sv += pp.imx[0][k];
}
sv += pp.mmx[0][m] * ff.rfv[m][x];
null2[x] = sv + xfactor;
}
for x in (K + 1)..=(KP - 3) {
let set = degen_set(x);
let mut result = 0.0f32;
let mut ndegen = 0.0f32;
for &y in set {
result += null2[y];
ndegen += 1.0;
}
null2[x] = if ndegen > 0.0 { result / ndegen } else { 0.0 };
}
null2[K] = 1.0;
null2[KP - 2] = 1.0;
null2[KP - 1] = 1.0;
(envsc, null2)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::hmmfile::P7Hmm;
use crate::seqio::read_fasta;
#[test]
fn envelope_null2_smoke() {
let hmm = P7Hmm::read_all(&format!("{}/testdata/globins4.hmm", env!("CARGO_MANIFEST_DIR")))
.unwrap()
.pop()
.unwrap();
let seqs = read_fasta(&format!("{}/testdata/globins45.fa", env!("CARGO_MANIFEST_DIR")))
.unwrap();
let s = seqs
.iter()
.find(|s| s.name == "MYG_ESCGI")
.expect("MYG_ESCGI not found");
let l = s.len();
assert_eq!(l, 153, "MYG_ESCGI length");
let (i, j, save_l) = (1usize, 147usize, 153usize);
let ff = ForwardFilter::build(&hmm);
let (_envsc, null2) = envelope_rescore(&ff, &s.dsq, i, j, save_l);
for x in 0..K {
assert!(null2[x].is_finite(), "null2[{x}] not finite: {}", null2[x]);
assert!(null2[x] > 0.0, "null2[{x}] not > 0: {}", null2[x]);
}
let mut domcorrection = 0.0f32;
for pos in i..=j {
domcorrection += null2[s.dsq[pos] as usize].ln();
}
println!("domcorrection (odds-space null2) = {domcorrection}");
assert!(
domcorrection > 6.5 && domcorrection < 8.5,
"domcorrection {domcorrection} out of sane range"
);
}
}