use crate::{
alignment::{
StatesSequence,
phmm::{
CorePhmm, DomainPhmm, GlobalPhmm, LocalPhmm, PhmmError, PhmmNumber, PhmmState, SemiLocalPhmm,
indexing::{
Begin, DpIndex, End, FirstMatch, GetCore, GetLayer, GetModule, LastMatch, PhmmIndex, PhmmIndexable, SeqIndex,
},
modules::{DomainModule, SemiLocalModule},
},
},
data::{cigar::Ciglet, mappings::ByteIndexMap},
};
use std::{cmp::Ordering, ops::Range};
impl<T: PhmmNumber, const S: usize> DomainModule<T, S> {
fn get_begin_score(&self, inserted: &[u8], mapping: &'static ByteIndexMap<S>) -> T {
if inserted.is_empty() {
self.start_to_end
} else {
let insert_to_insert = if inserted.len() > 1 {
T::cast_from(inserted.len() - 1) * self.insert_to_insert
} else {
T::ZERO
};
self.start_to_insert
+ self.insert_to_end
+ (inserted
.iter()
.map(|x| self.background_emission[mapping.to_index(*x)])
.fold(T::ZERO, |acc, elem| acc + elem)
+ insert_to_insert)
}
}
fn get_end_score(&self, inserted: &[u8], mapping: &'static ByteIndexMap<S>) -> T {
if inserted.is_empty() {
self.start_to_end
} else {
let insert_to_insert = if inserted.len() > 1 {
T::cast_from(inserted.len() - 1) * self.insert_to_insert
} else {
T::ZERO
};
self.start_to_insert
+ self.insert_to_end
+ (inserted
.iter()
.rev()
.map(|x| self.background_emission[mapping.to_index(*x)])
.fold(T::ZERO, |acc, elem| acc + elem)
+ insert_to_insert)
}
}
}
impl<T: PhmmNumber, const S: usize> DomainPhmm<T, S> {
#[inline]
fn get_begin_score(&self, inserted: &[u8], mapping: &'static ByteIndexMap<S>) -> T {
self.begin().get_begin_score(inserted, mapping)
}
#[inline]
fn get_end_score(&self, inserted: &[u8], mapping: &'static ByteIndexMap<S>) -> T {
self.end().get_end_score(inserted, mapping)
}
}
impl<T: PhmmNumber, const S: usize> LocalPhmm<T, S> {
fn get_begin_domain_score(&self, inserted: &[u8], mapping: &'static ByteIndexMap<S>) -> T {
self.begin().domain_params.get_begin_score(inserted, mapping)
}
fn get_begin_semilocal_score(&self, index: impl PhmmIndex) -> T {
self.begin().semilocal_params.get_score(index)
}
fn get_end_semilocal_score(&self, index: impl PhmmIndex) -> T {
self.end().semilocal_params.get_score(index)
}
fn get_end_domain_score(&self, inserted: &[u8], mapping: &'static ByteIndexMap<S>) -> T {
self.end().domain_params.get_end_score(inserted, mapping)
}
}
pub struct PhmmParam<T> {
pub param: T,
pub kind: PhmmParamKind<T>,
}
#[derive(Copy, Clone, Eq, PartialEq, Hash, Debug)]
pub enum ModuleLocation {
Begin,
End,
}
pub enum PhmmParamKind<T> {
EmissionMatch {
layer_idx: DpIndex,
residue_idx: usize,
},
EmissionInsert {
layer_idx: DpIndex,
residue_idx: usize,
},
Transition {
starting_layer: DpIndex,
starting_state: PhmmState,
ending_state: PhmmState,
},
LocalModule {
domain_param: T,
skipped: Vec<u8>,
semilocal_param: T,
to_layer: DpIndex,
loc: ModuleLocation,
},
DomainModule {
skipped: Vec<u8>,
loc: ModuleLocation,
},
SemiLocalModule {
core_layer: DpIndex,
loc: ModuleLocation,
},
}
impl<T: PhmmNumber> PhmmParam<T> {
fn new_emission_match(param: T, layer_idx: impl PhmmIndex, residue_idx: usize, phmm: &impl PhmmIndexable) -> Self {
Self {
param,
kind: PhmmParamKind::EmissionMatch {
layer_idx: phmm.to_dp_index(layer_idx),
residue_idx,
},
}
}
fn new_emission_insert(param: T, layer_idx: impl PhmmIndex, residue_idx: usize, phmm: &impl PhmmIndexable) -> Self {
Self {
param,
kind: PhmmParamKind::EmissionInsert {
layer_idx: phmm.to_dp_index(layer_idx),
residue_idx,
},
}
}
fn new_transition(
param: T, starting_layer: impl PhmmIndex, starting_state: PhmmState, ending_state: PhmmState,
phmm: &impl PhmmIndexable,
) -> Self {
Self {
param,
kind: PhmmParamKind::Transition {
starting_layer: phmm.to_dp_index(starting_layer),
starting_state,
ending_state,
},
}
}
fn new_local_module(
loc: ModuleLocation, domain_param: T, skipped: Vec<u8>, semilocal_param: T, to_layer: impl PhmmIndex,
phmm: &impl PhmmIndexable,
) -> Self {
Self {
param: domain_param + semilocal_param,
kind: PhmmParamKind::LocalModule {
domain_param,
skipped,
semilocal_param,
to_layer: phmm.to_dp_index(to_layer),
loc,
},
}
}
fn new_domain_module(loc: ModuleLocation, param: T, skipped: Vec<u8>) -> Self {
Self {
param,
kind: PhmmParamKind::DomainModule { skipped, loc },
}
}
fn new_semilocal_module(loc: ModuleLocation, param: T, to_layer: impl PhmmIndex, phmm: &impl PhmmIndexable) -> Self {
Self {
param,
kind: PhmmParamKind::SemiLocalModule {
core_layer: phmm.to_dp_index(to_layer),
loc,
},
}
}
}
#[inline]
fn call_f<T, F>(f: &mut F, score: &mut T, param: PhmmParam<T>)
where
T: PhmmNumber,
F: FnMut(PhmmParam<T>), {
*score += param.param;
f(param);
}
fn visit_params_core<T, const S: usize, F>(
core: &CorePhmm<T, S>, mapping: &ByteIndexMap<S>, seq_in_alignment: &[u8], ref_range: Range<usize>, ciglets: &[Ciglet],
score: &mut T, f: &mut F,
) -> Result<PhmmState, PhmmError>
where
T: PhmmNumber,
F: FnMut(PhmmParam<T>), {
use PhmmState::*;
let mut op_iter = ciglets.iter().flat_map(|ciglet| std::iter::repeat_n(ciglet.op, ciglet.inc));
let Some(op) = op_iter.next() else {
if ref_range.is_empty() {
return Ok(Match);
}
return Err(PhmmError::InvalidPath);
};
let mut i = 0;
let (mut state, mut j) = match PhmmState::from_op(op)? {
Delete => {
if ref_range.start != 0 {
return Err(PhmmError::InvalidPath);
}
(Delete, 1)
}
Match => {
let x_idx = mapping.to_index(seq_in_alignment[i]);
let layer_idx = SeqIndex(ref_range.start);
let param = core
.get_layer(layer_idx.prev_index(core))
.ok_or(PhmmError::InvalidPath)?
.emission_match[x_idx];
call_f(f, score, PhmmParam::new_emission_match(param, layer_idx, x_idx, core));
i += 1;
(Match, ref_range.start + 1)
}
Insert => {
if ref_range.start != 0 {
return Err(PhmmError::InvalidPath);
}
let x_idx = mapping.to_index(seq_in_alignment[i]);
call_f(
f,
score,
PhmmParam::new_emission_insert(core.begin_layer().emission_insert[x_idx], Begin, x_idx, core),
);
i += 1;
(Insert, 0)
}
};
for op in op_iter {
let op = PhmmState::from_op(op)?;
let layer = core.get_layer(DpIndex(j)).ok_or(PhmmError::InvalidPath)?;
match op {
Match => {
let x_idx = mapping.to_index(seq_in_alignment[i]);
call_f(
f,
score,
PhmmParam::new_transition(layer.transition[(state, Match)], DpIndex(j), state, Match, core),
);
j += 1;
call_f(
f,
score,
PhmmParam::new_emission_match(layer.emission_match[x_idx], DpIndex(j), x_idx, core),
);
state = Match;
i += 1;
}
Insert => {
let x_idx = mapping.to_index(seq_in_alignment[i]);
call_f(
f,
score,
PhmmParam::new_transition(layer.transition[(state, Insert)], DpIndex(j), state, Insert, core),
);
call_f(
f,
score,
PhmmParam::new_emission_insert(layer.emission_insert[x_idx], DpIndex(j), x_idx, core),
);
state = Insert;
i += 1;
}
Delete => {
call_f(
f,
score,
PhmmParam::new_transition(layer.transition[(state, Delete)], DpIndex(j), state, Delete, core),
);
state = Delete;
j += 1;
}
}
}
if i != seq_in_alignment.len() {
return Err(PhmmError::InvalidPath);
}
if j != ref_range.end {
return Err(PhmmError::InvalidPath);
}
Ok(state)
}
fn resolve_ambiguous_start<T: PhmmNumber, const S: usize>(
begin_module: &SemiLocalModule<T>, core: &CorePhmm<T, S>, score: T, mut ciglets: &[Ciglet],
) -> Result<(T, DpIndex, Option<T>, PhmmState), PhmmError> {
use PhmmState::*;
let first_op = ciglets.peek_op().ok_or(PhmmError::InvalidPath)?;
let first_state = PhmmState::from_op(first_op)?;
let semilocal_begin_param_through_begin = begin_module.get_score(Begin);
let transition_from_begin = core.begin_layer().transition[(Match, first_state)];
let score_through_begin = score + semilocal_begin_param_through_begin + transition_from_begin;
let (semilocal_begin_param, to_layer, transition_from_begin) = match first_state {
Delete | Insert => (
semilocal_begin_param_through_begin,
core.to_dp_index(Begin),
Some(transition_from_begin),
),
Match => {
let semilocal_begin_param_skip_begin = begin_module.get_score(FirstMatch);
let score_skipping_begin = score + semilocal_begin_param_skip_begin;
if score_through_begin <= score_skipping_begin {
(
semilocal_begin_param_through_begin,
core.to_dp_index(Begin),
Some(transition_from_begin),
)
} else {
(semilocal_begin_param_skip_begin, core.to_dp_index(FirstMatch), None)
}
}
};
Ok((semilocal_begin_param, to_layer, transition_from_begin, first_state))
}
fn resolve_ambiguous_end<T: PhmmNumber, const S: usize>(
end_module: &SemiLocalModule<T>, core: &CorePhmm<T, S>, score: T, final_state: PhmmState, domain_end_param: Option<T>,
) -> (T, DpIndex, Option<T>) {
use PhmmState::*;
let transition_to_end = core.last_match().transition[(final_state, Match)];
let semilocal_end_param_through_end = end_module.get_score(End);
let end_param_through_end = if let Some(domain_end_param) = domain_end_param {
domain_end_param + semilocal_end_param_through_end
} else {
semilocal_end_param_through_end
};
let score_through_end = score + transition_to_end + end_param_through_end;
let (transition_to_end, semilocal_end_param, from_layer) = match final_state {
Delete | Insert => (
Some(transition_to_end),
semilocal_end_param_through_end,
core.to_dp_index(End),
),
Match => {
let semilocal_end_param_skip_end = end_module.get_score(LastMatch);
let end_param_skip_end = if let Some(domain_end_param) = domain_end_param {
domain_end_param + semilocal_end_param_skip_end
} else {
semilocal_end_param_skip_end
};
let score_skipping_end = score + end_param_skip_end;
if score_through_end <= score_skipping_end {
(
Some(transition_to_end),
semilocal_end_param_through_end,
core.to_dp_index(End),
)
} else {
(None, semilocal_end_param_skip_end, core.to_dp_index(LastMatch))
}
}
};
(semilocal_end_param, from_layer, transition_to_end)
}
impl<T: PhmmNumber, const S: usize> GlobalPhmm<T, S> {
pub fn visit_params<Q, A, F>(&self, seq: Q, alignment: A, f: F) -> Result<T, PhmmError>
where
Q: AsRef<[u8]>,
A: AsRef<[Ciglet]>,
F: FnMut(PhmmParam<T>), {
self.visit_params_helper(seq.as_ref(), alignment.as_ref(), f)
}
fn visit_params_helper<F>(&self, seq: &[u8], mut ciglets: &[Ciglet], mut f: F) -> Result<T, PhmmError>
where
F: FnMut(PhmmParam<T>), {
use PhmmState::*;
let mut score = T::ZERO;
let first_op = ciglets.peek_op().ok_or(PhmmError::InvalidPath)?;
let first_state = PhmmState::from_op(first_op)?;
call_f(
&mut f,
&mut score,
PhmmParam::new_transition(
self.begin_layer().transition[(Match, first_state)],
Begin,
Match,
first_state,
self,
),
);
let final_state = visit_params_core(
self.core(),
self.mapping(),
seq,
0..self.seq_len(),
ciglets,
&mut score,
&mut f,
)?;
call_f(
&mut f,
&mut score,
PhmmParam::new_transition(
self.last_match().transition[(final_state, Match)],
LastMatch,
final_state,
Match,
self,
),
);
Ok(score)
}
}
impl<T: PhmmNumber, const S: usize> LocalPhmm<T, S> {
#[allow(clippy::too_many_lines)]
pub fn visit_params<Q, A, F>(&self, seq: Q, alignment: A, ref_range: Range<usize>, f: F) -> Result<T, PhmmError>
where
Q: AsRef<[u8]>,
A: AsRef<[Ciglet]>,
F: FnMut(PhmmParam<T>), {
self.visit_params_helper(seq.as_ref(), alignment.as_ref(), ref_range, f)
}
#[allow(clippy::too_many_lines)]
fn visit_params_helper<F>(
&self, seq: &[u8], mut ciglets: &[Ciglet], ref_range: Range<usize>, mut f: F,
) -> Result<T, PhmmError>
where
F: FnMut(PhmmParam<T>), {
use PhmmState::*;
let mut score = T::ZERO;
let (begin_seq, seq, end_seq) = {
let begin_inserted = ciglets.next_if_op(|op| op == b'S').map_or(0, |ciglet| ciglet.inc);
let end_inserted = ciglets.next_back_if_op(|op| op == b'S').map_or(0, |ciglet| ciglet.inc);
if ciglets.is_empty() {
if !ref_range.is_empty() {
return Err(PhmmError::InvalidPath);
}
return match seq.len().cmp(&(begin_inserted + end_inserted)) {
Ordering::Equal => self.visit_params_empty_alignment(seq, f),
Ordering::Less | Ordering::Greater => Err(PhmmError::InvalidPath),
};
}
let (seq, end_seq) = seq.split_at(seq.len() - end_inserted);
let (begin_seq, seq) = seq.split_at(begin_inserted);
(begin_seq, seq, end_seq)
};
let domain_begin_param = self.get_begin_domain_score(begin_seq, self.mapping());
if ref_range.start == 0 {
let (semilocal_begin_param, to_layer, transition_from_begin, first_state) =
resolve_ambiguous_start(&self.begin().semilocal_params, self.core(), score, ciglets)?;
call_f(
&mut f,
&mut score,
PhmmParam::new_local_module(
ModuleLocation::Begin,
domain_begin_param,
begin_seq.to_vec(),
semilocal_begin_param,
to_layer,
self,
),
);
if let Some(transition_from_begin) = transition_from_begin {
call_f(
&mut f,
&mut score,
PhmmParam::new_transition(transition_from_begin, Begin, Match, first_state, self),
);
}
} else {
if PhmmState::from_op(ciglets.peek_op().ok_or(PhmmError::InvalidPath)?)? != Match {
return Err(PhmmError::InvalidPath);
}
call_f(
&mut f,
&mut score,
PhmmParam::new_local_module(
ModuleLocation::Begin,
domain_begin_param,
begin_seq.to_vec(),
self.get_begin_semilocal_score(SeqIndex(ref_range.start)),
SeqIndex(ref_range.start),
self,
),
);
}
let final_state = visit_params_core(
self.core(),
self.mapping(),
seq,
ref_range.clone(),
ciglets,
&mut score,
&mut f,
)?;
let domain_end_param = self.get_end_domain_score(end_seq, self.mapping());
if ref_range.end == self.seq_len() {
let (semilocal_end_param, from_layer, transition_to_end) = resolve_ambiguous_end(
&self.end().semilocal_params,
self.core(),
score,
final_state,
Some(domain_end_param),
);
if let Some(transition_to_end) = transition_to_end {
call_f(
&mut f,
&mut score,
PhmmParam::new_transition(transition_to_end, LastMatch, final_state, Match, self),
);
}
call_f(
&mut f,
&mut score,
PhmmParam::new_local_module(
ModuleLocation::End,
domain_end_param,
end_seq.to_vec(),
semilocal_end_param,
from_layer,
self,
),
);
} else {
if final_state != Match {
return Err(PhmmError::InvalidPath);
}
call_f(
&mut f,
&mut score,
PhmmParam::new_local_module(
ModuleLocation::End,
domain_end_param,
end_seq.to_vec(),
self.get_end_semilocal_score(SeqIndex(ref_range.end - 1)),
SeqIndex(ref_range.end - 1),
self,
),
);
}
Ok(score)
}
fn visit_params_empty_alignment<F>(&self, seq: &[u8], mut f: F) -> Result<T, PhmmError>
where
F: FnMut(PhmmParam<T>), {
let mut best_i = 0;
let mut best_state = self.to_dp_index(Begin);
let mut best_score = T::INFINITY;
for i in 0..=seq.len() {
let (inserted_begin, inserted_end) = seq.split_at(i);
for through_state in [self.to_dp_index(Begin), self.to_dp_index(End)] {
let begin_score = self.begin().domain_params.get_begin_score(inserted_begin, self.mapping())
+ self.get_begin_semilocal_score(through_state);
let end_score =
self.get_end_domain_score(inserted_end, self.mapping()) + self.get_end_semilocal_score(through_state);
let score = begin_score + end_score;
if score < best_score {
best_i = i;
best_state = through_state;
best_score = score;
}
}
}
if best_score == T::INFINITY {
return Err(PhmmError::NoAlignmentFound);
}
let (inserted_begin, inserted_end) = seq.split_at(best_i);
f(PhmmParam::new_local_module(
ModuleLocation::Begin,
self.begin().domain_params.get_begin_score(inserted_begin, self.mapping()),
inserted_begin.to_vec(),
self.get_begin_semilocal_score(best_state),
best_state,
self,
));
f(PhmmParam::new_local_module(
ModuleLocation::End,
self.get_end_domain_score(inserted_end, self.mapping()),
inserted_end.to_vec(),
self.get_end_semilocal_score(best_state),
best_state,
self,
));
Ok(best_score)
}
}
impl<T: PhmmNumber, const S: usize> DomainPhmm<T, S> {
pub fn visit_params<Q, A, F>(&self, seq: Q, alignment: A, f: F) -> Result<T, PhmmError>
where
Q: AsRef<[u8]>,
A: AsRef<[Ciglet]>,
F: FnMut(PhmmParam<T>), {
self.visit_params_helper(seq.as_ref(), alignment.as_ref(), f)
}
fn visit_params_helper<F>(&self, seq: &[u8], mut ciglets: &[Ciglet], mut f: F) -> Result<T, PhmmError>
where
F: FnMut(PhmmParam<T>), {
use PhmmState::*;
let mut score = T::ZERO;
let (begin_seq, seq, end_seq) = {
let begin_inserted = ciglets.next_if_op(|op| op == b'S').map_or(0, |ciglet| ciglet.inc);
let end_inserted = ciglets.next_back_if_op(|op| op == b'S').map_or(0, |ciglet| ciglet.inc);
let (seq, end_seq) = seq.split_at(seq.len() - end_inserted);
let (begin_seq, seq) = seq.split_at(begin_inserted);
(begin_seq, seq, end_seq)
};
call_f(
&mut f,
&mut score,
PhmmParam::new_domain_module(
ModuleLocation::Begin,
self.get_begin_score(begin_seq, self.mapping()),
begin_seq.to_vec(),
),
);
let first_op = ciglets.peek_op().ok_or(PhmmError::InvalidPath)?;
let first_state = PhmmState::from_op(first_op)?;
call_f(
&mut f,
&mut score,
PhmmParam::new_transition(
self.begin_layer().transition[(Match, first_state)],
Begin,
Match,
first_state,
self,
),
);
let final_state = visit_params_core(
self.core(),
self.mapping(),
seq,
0..self.seq_len(),
ciglets,
&mut score,
&mut f,
)?;
call_f(
&mut f,
&mut score,
PhmmParam::new_transition(
self.last_match().transition[(final_state, Match)],
LastMatch,
final_state,
Match,
self,
),
);
call_f(
&mut f,
&mut score,
PhmmParam::new_domain_module(
ModuleLocation::End,
self.get_end_score(end_seq, self.mapping()),
end_seq.to_vec(),
),
);
Ok(score)
}
}
impl<T: PhmmNumber, const S: usize> SemiLocalPhmm<T, S> {
pub fn visit_params<Q, A, F>(&self, seq: Q, alignment: A, ref_range: Range<usize>, f: F) -> Result<T, PhmmError>
where
Q: AsRef<[u8]>,
A: AsRef<[Ciglet]>,
F: FnMut(PhmmParam<T>), {
self.visit_params_helper(seq.as_ref(), alignment.as_ref(), ref_range, f)
}
fn visit_params_helper<F>(
&self, seq: &[u8], mut ciglets: &[Ciglet], ref_range: Range<usize>, mut f: F,
) -> Result<T, PhmmError>
where
F: FnMut(PhmmParam<T>), {
use PhmmState::*;
let mut score = T::ZERO;
if ciglets.peek_op().is_none() {
if seq.is_empty() {
return Ok(self.visit_params_empty_alignment(f));
}
return Err(PhmmError::InvalidPath);
}
if ref_range.start == 0 {
let (begin_param, to_layer, transition_from_begin, first_state) =
resolve_ambiguous_start(self.begin(), self.core(), score, ciglets)?;
call_f(
&mut f,
&mut score,
PhmmParam::new_semilocal_module(ModuleLocation::Begin, begin_param, to_layer, self),
);
if let Some(transition_from_begin) = transition_from_begin {
call_f(
&mut f,
&mut score,
PhmmParam::new_transition(transition_from_begin, Begin, Match, first_state, self),
);
}
} else {
if PhmmState::from_op(ciglets.peek_op().ok_or(PhmmError::InvalidPath)?)? != Match {
return Err(PhmmError::InvalidPath);
}
call_f(
&mut f,
&mut score,
PhmmParam::new_semilocal_module(
ModuleLocation::Begin,
self.get_begin_score(SeqIndex(ref_range.start)),
SeqIndex(ref_range.start),
self,
),
);
}
let final_state = visit_params_core(
self.core(),
self.mapping(),
seq,
ref_range.clone(),
ciglets,
&mut score,
&mut f,
)?;
if ref_range.end == self.seq_len() {
let (end_param, from_layer, transition_to_end) =
resolve_ambiguous_end(self.end(), self.core(), score, final_state, None);
if let Some(transition_to_end) = transition_to_end {
call_f(
&mut f,
&mut score,
PhmmParam::new_transition(transition_to_end, LastMatch, final_state, Match, self),
);
}
call_f(
&mut f,
&mut score,
PhmmParam::new_semilocal_module(ModuleLocation::End, end_param, from_layer, self),
);
} else {
if final_state != Match {
return Err(PhmmError::InvalidPath);
}
call_f(
&mut f,
&mut score,
PhmmParam::new_semilocal_module(
ModuleLocation::End,
self.get_end_score(SeqIndex(ref_range.end - 1)),
SeqIndex(ref_range.end - 1),
self,
),
);
}
Ok(score)
}
fn visit_params_empty_alignment<F>(&self, mut f: F) -> T
where
F: FnMut(PhmmParam<T>), {
let score_through_begin = self.get_begin_score(Begin) + self.get_end_score(Begin);
let score_through_end = self.get_begin_score(End) + self.get_end_score(End);
let (score, through_state) = if score_through_begin <= score_through_end {
(score_through_begin, self.to_dp_index(Begin))
} else {
(score_through_end, self.to_dp_index(End))
};
f(PhmmParam::new_semilocal_module(
ModuleLocation::Begin,
self.get_begin_score(through_state),
through_state,
self,
));
f(PhmmParam::new_semilocal_module(
ModuleLocation::End,
self.get_end_score(through_state),
through_state,
self,
));
score
}
}