use super::ViterbiTraceback;
use crate::alignment::{
Alignment, AlignmentStates,
phmm::{
InvalidModelError, LayerParams, LocalPhmm, PhmmError, PhmmNumber,
indexing::{
Begin, DpIndex, End, GetLayer, GetModule, LastBase, LastMatch, PhmmIndex, PhmmIndexable, QueryIndex,
QueryIndexable,
},
modules::PrecomputedLocalModule,
state::{
PhmmBacktrackFlags,
PhmmState::{self, Delete, Insert, Match},
PhmmStateOrEnter, PhmmTracebackState, best_state_or_enter,
},
viterbi::{ExitLocation, update_delete, update_insert},
},
};
use std::ops::Bound::{Excluded, Included};
struct LocalBestScore<T> {
score: T,
i: DpIndex,
loc: ExitLocation,
}
impl<T: PhmmNumber> LocalBestScore<T> {
#[inline]
fn update<const S: usize>(
&mut self, match_val: T, i: impl QueryIndex, j: impl PhmmIndex, seq: &[u8], end: &PrecomputedLocalModule<T, S>,
) {
let score = match_val + end.get_score(i, j);
if score < self.score {
self.score = score;
self.i = seq.to_dp_index(i);
self.loc = match end.to_seq_index(j) {
Some(loc) => ExitLocation::Match(loc),
None => ExitLocation::Begin,
}
}
}
#[allow(clippy::too_many_arguments)]
fn update_last_layer<const S: usize>(
&mut self, layer: &LayerParams<T, S>, mut match_val: T, mut delete_val: T, mut insert_val: T, i: impl QueryIndex,
seq: &[u8], begin: &PrecomputedLocalModule<T, S>, end: &PrecomputedLocalModule<T, S>,
) {
use crate::alignment::phmm::state::PhmmState::*;
self.update(match_val, i, LastMatch, seq, end);
match_val += layer.transition[(Match, Match)];
delete_val += layer.transition[(Delete, Match)];
insert_val += layer.transition[(Insert, Match)];
let enter_val = begin.get_score(i, End);
let (state, mut score) = best_state_or_enter(match_val, delete_val, insert_val, enter_val);
score += end.get_score(i, End);
if score < self.score {
self.score = score;
self.i = seq.to_dp_index(i);
self.loc = ExitLocation::End(state);
}
}
}
impl<T: PhmmNumber> Default for LocalBestScore<T> {
#[inline]
fn default() -> Self {
Self {
score: T::INFINITY,
i: DpIndex(0),
loc: ExitLocation::End(PhmmStateOrEnter::Match),
}
}
}
impl<T: PhmmNumber, const S: usize> LocalPhmm<T, S> {
#[allow(clippy::too_many_lines)]
pub fn viterbi<Q: AsRef<[u8]>>(&self, seq: Q) -> Result<Alignment<T>, PhmmError> {
let seq = seq.as_ref();
let begin_mod = self.begin().precompute_begin_mod(seq, self.mapping());
let end_mod = self.end().precompute_end_mod(seq, self.mapping());
if begin_mod.num_pseudomatch() != self.num_pseudomatch()
|| end_mod.num_pseudomatch() != self.num_pseudomatch()
|| begin_mod.domain_params.seq_len() != seq.len()
|| end_mod.domain_params.seq_len() != seq.len()
{
return Err(InvalidModelError::IncompatibleModule.into());
}
let (end, layers) = self.split_last_layer();
let query_dim = seq.len() + 1;
let phmm_dim = layers.len() + 1;
let mut v_m = vec![T::INFINITY; query_dim];
for (i, value) in v_m.iter_mut().enumerate() {
*value = begin_mod.get_score(DpIndex(i), Begin);
}
let mut v_i = vec![T::INFINITY; query_dim];
let mut v_d = vec![T::INFINITY; query_dim];
let mut j = 0;
let mut traceback = ViterbiTraceback::new(PhmmBacktrackFlags::new(), query_dim, phmm_dim);
let mut best_score: LocalBestScore<T> = LocalBestScore::default();
for layer in layers {
let mut cur_m = v_m[0];
let start = query_dim * j;
let (traceback_curr_row, traceback_next_row) =
traceback.data[start..start + 2 * query_dim].split_at_mut(query_dim);
let traceback_curr_row = &mut traceback_curr_row[0..query_dim];
let traceback_next_row = &mut traceback_next_row[0..query_dim];
for (i, x_idx) in seq.iter().map(|x| self.mapping().to_index(*x)).enumerate() {
let match_val = cur_m;
let delete_val = v_d[i];
let insert_val = v_i[i];
best_score.update(match_val, DpIndex(i), DpIndex(j), seq, &end_mod);
let (state_m, match_score) = {
let i = DpIndex(i);
let j = DpIndex(j);
let next_layer_idx = j.next_index(self);
let (state, best) = best_state_or_enter(
match_val + layer.transition[(Match, Match)],
delete_val + layer.transition[(Delete, Match)],
insert_val + layer.transition[(Insert, Match)],
begin_mod.get_score(i, next_layer_idx),
);
(state, best + layer.emission_match[x_idx])
};
traceback_next_row[i + 1].set_match(state_m);
cur_m = std::mem::replace(&mut v_m[i + 1], match_score);
let (state_i, insert_score) = update_insert(layer, x_idx, match_val, delete_val, insert_val);
traceback_curr_row[i + 1].set_insert(state_i);
v_i[i + 1] = insert_score;
let (state_d, delete_score) = update_delete(layer, match_val, delete_val, insert_val);
traceback_next_row[i].set_delete(state_d);
v_d[i] = delete_score;
}
let i = seq.len();
best_score.update(cur_m, LastBase, DpIndex(j), seq, &end_mod);
let (state_d, delete_score) = update_delete(layer, cur_m, v_d[i], v_i[i]);
traceback_next_row[i].set_delete(state_d);
v_d[i] = delete_score;
j += 1;
v_m[0] = T::INFINITY;
}
let start = query_dim * j;
let traceback_curr_row = &mut traceback.data[start..start + query_dim];
for (i, x_idx) in seq.iter().map(|x| self.mapping().to_index(*x)).enumerate() {
best_score.update_last_layer(end, v_m[i], v_d[i], v_i[i], DpIndex(i), seq, &begin_mod, &end_mod);
let (state_i, insert_score) = update_insert(end, x_idx, v_m[i], v_d[i], v_i[i]);
traceback_curr_row[i + 1].set_insert(state_i);
v_i[i + 1] = insert_score;
}
let i = seq.len();
best_score.update_last_layer(end, v_m[i], v_d[i], v_i[i], LastBase, seq, &begin_mod, &end_mod);
if best_score.score == T::INFINITY {
return Err(PhmmError::NoAlignmentFound);
}
let LocalBestScore {
score,
i: DpIndex(end_i),
loc,
} = best_score;
let (state, end_j) = match loc {
ExitLocation::Begin => {
let mut states = AlignmentStates::new();
states.soft_clip(seq.len());
return Ok(Alignment {
score,
ref_range: 0..0,
query_range: 0..0,
states,
ref_len: self.seq_len(),
query_len: seq.len(),
});
}
ExitLocation::Match(j) => (Match, self.get_dp_index(j)),
ExitLocation::End(state) => {
let Some(state) = PhmmState::get_from(state) else {
let mut states = AlignmentStates::new();
states.soft_clip(seq.len());
return Ok(Alignment {
score,
ref_range: self.get_seq_range(End..End),
query_range: 0..0,
states,
ref_len: self.seq_len(),
query_len: seq.len(),
});
};
(state, self.get_dp_index(LastMatch))
}
};
let mut state = PhmmTracebackState::from(state);
let mut i = end_i;
let mut j = end_j;
let mut states = AlignmentStates::new();
states.soft_clip(seq.len() - i);
while j > 0 || !state.is_match() {
let next_state = traceback.get(i, j).get_prev_state(state);
if state.is_match() {
states.add_state(b'M');
i -= 1;
j -= 1;
} else if state.is_delete() {
states.add_state(b'D');
j -= 1;
} else {
states.add_state(b'I');
i -= 1;
}
if next_state.is_enter() {
break;
}
state = next_state;
}
states.soft_clip(i);
states.make_reverse();
let start_j = Excluded(DpIndex(j));
let end_j = Included(DpIndex(end_j));
let start_i = Excluded(DpIndex(i));
let end_i = Included(DpIndex(end_i));
Ok(Alignment {
score,
ref_range: self.get_seq_range((start_j, end_j)),
query_range: seq.get_seq_range(start_i, end_i),
states,
ref_len: self.seq_len(),
query_len: seq.len(),
})
}
}