rustyhmmer 0.1.1

Pure-Rust HMMER3 hmmsearch with byte-identical --tblout output
Documentation













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");
    }
}