use ebi_objects::{
Activity, AutomatonState, StochasticNondeterministicFiniteAutomaton,
anyhow::{Context, Ok, Result},
ebi_arithmetic::{EbiMatrix, Fraction, FractionMatrix, Inversion, One, Zero},
};
use std::collections::{HashMap, VecDeque};
pub trait TauRemoval {
fn remove_tau_transitions(&mut self) -> Result<()>;
}
impl TauRemoval for StochasticNondeterministicFiniteAutomaton {
fn remove_tau_transitions(&mut self) -> Result<()> {
let n = self.sources.len();
if n == 0 {
return Ok(());
}
let one = Fraction::one();
let zero = Fraction::zero();
let initial_state = if let Some(initial) = self.initial_state {
initial
} else {
return Ok(());
};
let mut tau_adj: Vec<Vec<(AutomatonState, Fraction)>> = vec![vec![]; n];
for (source, target, label, probability) in self.into_iter() {
if label.is_none() {
tau_adj[*source].push((*target, probability.clone()));
}
}
let mut d: Vec<Vec<Fraction>> = vec![vec![zero.clone(); n]; n];
let mut row_nonzero: Vec<bool> = vec![false; n];
for i in 0..n {
d[i][i] = one.clone();
}
let sccs = compute_scc(n, &tau_adj);
let mut comp_of = vec![0usize; n];
for (cid, comp) in sccs.iter().enumerate() {
for &v in comp {
comp_of[v] = cid;
}
}
let mut dag_succ: Vec<Vec<usize>> = vec![vec![]; sccs.len()];
for u in 0..n {
for &(v, _) in &tau_adj[u] {
let cu = comp_of[u];
let cv = comp_of[v];
if cu != cv && !dag_succ[cu].contains(&cv) {
dag_succ[cu].push(cv);
}
}
}
let mut visited = vec![false; sccs.len()];
let mut rev_topo: Vec<usize> = Vec::new();
fn dfs(u: usize, dag: &Vec<Vec<usize>>, vis: &mut Vec<bool>, out: &mut Vec<usize>) {
if vis[u] {
return;
}
vis[u] = true;
for &v in &dag[u] {
dfs(v, dag, vis, out);
}
out.push(u);
}
for cid in 0..sccs.len() {
dfs(cid, &dag_succ, &mut visited, &mut rev_topo);
}
for &cid in &rev_topo {
let comp = &sccs[cid];
let m = comp.len();
if m == 1 && tau_adj[comp[0]].is_empty() {
} else {
let mut mat = FractionMatrix::new(m, m);
for (i_idx, &i) in comp.iter().enumerate() {
mat.set(i_idx, i_idx, one.clone());
for &(j, ref w) in &tau_adj[i] {
if comp_of[j] == cid {
let j_idx = comp.iter().position(|&x| x == j).unwrap();
mat.set(i_idx, j_idx, &mat.get(i_idx, j_idx).unwrap() - w);
}
}
}
let d_local = mat
.invert()
.with_context(|| "tau-removal: singular (I-W) in SCC")?;
for (i_idx, &i) in comp.iter().enumerate() {
for (j_idx, &j) in comp.iter().enumerate() {
d[i][j] = d_local.get(i_idx, j_idx).unwrap().clone();
}
row_nonzero[i] = true;
}
}
for &q in comp {
for &(r, ref w_qr) in &tau_adj[q] {
if comp_of[r] == cid {
continue;
}
for p in 0..n {
if !row_nonzero[p] {
continue;
}
if d[p][q].is_zero() {
continue;
}
let prefix = &d[p][q] * &w_qr.clone();
if prefix.is_zero() {
continue;
}
for s in 0..n {
let d_rs = d[r][s].clone();
if d_rs.is_zero() {
continue;
}
let add = &prefix * &d_rs;
if add.is_zero() {
continue;
}
d[p][s] += add;
}
}
}
}
}
let mut new_transitions: Vec<Vec<Transition>> = vec![vec![]; n];
for q in 0..n {
for (_, target, label, probability) in self.outgoing_edges(AutomatonState::of(q)) {
if !label.is_none() {
for p in 0..n {
if d[p][q].is_zero() {
continue;
}
let prob = &d[p][q] * probability;
new_transitions[p].push(Transition {
target: *target,
label: *label,
probability: prob,
});
}
}
}
}
for p in 0..n {
let mut map: HashMap<(AutomatonState, Option<Activity>), Fraction> = HashMap::new();
for tr in new_transitions[p].drain(..) {
*map.entry((tr.target, tr.label.clone()))
.or_insert_with(Fraction::zero) += tr.probability.clone();
}
new_transitions[p] = map
.into_iter()
.map(|((tgt, lbl), prob)| Transition {
target: tgt,
label: lbl,
probability: prob,
})
.collect();
}
replace_states(self, new_transitions)?;
{
let mut reachable = vec![false; self.termination_probabilities.len()];
let mut queue = VecDeque::new();
queue.push_back(initial_state);
while let Some(u) = queue.pop_front() {
if reachable[u] {
continue;
}
reachable[u] = true;
for (_, target, label, _) in self.outgoing_edges(u) {
if label.is_none() {
continue;
}
queue.push_back(*target);
}
}
if reachable.iter().any(|&r| !r) {
let mut map = vec![None; self.termination_probabilities.len()];
let mut new_states: Vec<State> = Vec::new();
for old_idx in 0..self.termination_probabilities.len() {
if reachable[old_idx] {
let new_idx = new_states.len();
map[old_idx] = Some(AutomatonState::of(new_idx));
new_states.push(State {
transitions: Vec::new(),
});
}
}
for old_idx in 0..self.termination_probabilities.len() {
if let Some(new_src) = map[old_idx] {
for (_, target, label, probability) in
self.outgoing_edges(AutomatonState::of(old_idx))
{
if let Some(new_tgt) = map[*target] {
new_states[new_src].transitions.push(Transition {
target: new_tgt,
label: *label,
probability: probability.clone(),
});
}
}
}
}
replace_everything(self, new_states)?;
self.initial_state = Some(map[initial_state].unwrap());
}
}
Ok(())
}
}
fn compute_scc(n: usize, adj: &Vec<Vec<(AutomatonState, Fraction)>>) -> Vec<Vec<AutomatonState>> {
let mut visited = vec![false; n];
let mut order: Vec<AutomatonState> = Vec::with_capacity(n);
fn dfs1(
u: AutomatonState,
adj: &Vec<Vec<(AutomatonState, Fraction)>>,
vis: &mut Vec<bool>,
order: &mut Vec<AutomatonState>,
) {
if vis[u] {
return;
}
vis[u] = true;
for &(v, _) in &adj[u] {
dfs1(v, adj, vis, order);
}
order.push(u);
}
for v in 0..n {
dfs1(AutomatonState::of(v), adj, &mut visited, &mut order);
}
let mut radj: Vec<Vec<AutomatonState>> = vec![vec![]; n];
for u in 0..n {
for &(v, _) in &adj[u] {
radj[v].push(AutomatonState::of(u));
}
}
let mut comp_id = vec![None; n];
let mut comps: Vec<Vec<AutomatonState>> = Vec::new();
fn dfs2(
u: AutomatonState,
radj: &Vec<Vec<AutomatonState>>,
comp: &mut Vec<AutomatonState>,
comp_id: &mut Vec<Option<usize>>,
cid: usize,
) {
if comp_id[u].is_some() {
return;
}
comp_id[u] = Some(cid);
comp.push(u);
for &v in &radj[u] {
dfs2(v, radj, comp, comp_id, cid);
}
}
let mut cid = 0;
while let Some(u) = order.pop() {
if comp_id[u].is_none() {
comps.push(Vec::new());
dfs2(u, &radj, comps.last_mut().unwrap(), &mut comp_id, cid);
cid += 1;
}
}
comps
}
#[derive(Clone)]
struct Transition {
target: AutomatonState,
label: Option<Activity>,
probability: Fraction,
}
#[derive(Clone)]
struct State {
transitions: Vec<Transition>,
}
fn replace_everything(
snfa: &mut StochasticNondeterministicFiniteAutomaton,
states: Vec<State>,
) -> Result<()> {
snfa.sources.clear();
snfa.targets.clear();
snfa.activities.clear();
snfa.probabilities.clear();
snfa.termination_probabilities.clear();
for _ in 0..states.len() {
snfa.add_state();
}
for (source, state) in states.into_iter().enumerate() {
for transition in state.transitions.into_iter() {
snfa.add_transition(
AutomatonState::of(source),
transition.label,
transition.target,
transition.probability,
)?;
}
}
Ok(())
}
fn replace_states(
snfa: &mut StochasticNondeterministicFiniteAutomaton,
transitionss: Vec<Vec<Transition>>,
) -> Result<()> {
snfa.sources.clear();
snfa.targets.clear();
snfa.activities.clear();
snfa.probabilities.clear();
snfa.termination_probabilities.fill(Fraction::one());
for (source, transitions) in transitionss.into_iter().enumerate() {
for transition in transitions.into_iter() {
snfa.add_transition(
AutomatonState::of(source),
transition.label,
transition.target,
transition.probability,
)?;
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use ebi_objects::{
Activity, AutomatonState, StochasticNondeterministicFiniteAutomaton, a,
ebi_arithmetic::{Fraction, f0, f1},
};
use std::collections::HashMap;
macro_rules! frac {
($n:expr , $d:expr) => {
Fraction::from(($n, $d))
};
}
fn collect(
snfa: &StochasticNondeterministicFiniteAutomaton,
) -> HashMap<(AutomatonState, Option<Activity>, AutomatonState), Fraction> {
let mut m = HashMap::new();
for transition in 0..snfa.sources.len() {
let source = snfa.sources[transition];
let target = snfa.targets[transition];
let probability = &snfa.probabilities[transition];
let label = snfa.activities[transition];
*m.entry((source, label, target))
.or_insert_with(Fraction::zero) += probability;
}
m
}
#[test]
fn tau_removal_self_loop_and_cycle_rows_sum_to_one() {
let mut snfa = StochasticNondeterministicFiniteAutomaton::new();
for _ in 0..3 {
snfa.add_state();
}
snfa.add_transition(a!(0), None, a!(0), frac!(1, 5))
.unwrap();
snfa.add_transition(a!(1), None, a!(2), frac!(1, 4))
.unwrap();
snfa.add_transition(a!(2), None, a!(1), frac!(1, 2))
.unwrap();
let acta = Some(snfa.activity_key.process_activity("a"));
let actb = Some(snfa.activity_key.process_activity("b"));
let actc = Some(snfa.activity_key.process_activity("c"));
let actd = Some(snfa.activity_key.process_activity("d"));
snfa.add_transition(a!(0), acta, a!(0), frac!(3, 5))
.unwrap();
snfa.add_transition(a!(0), actb, a!(1), frac!(1, 5))
.unwrap();
snfa.add_transition(a!(1), actc, a!(1), frac!(3, 4))
.unwrap();
snfa.add_transition(a!(2), actd, a!(2), frac!(1, 4))
.unwrap();
snfa.remove_tau_transitions().unwrap();
assert!(snfa.activities.iter().all(|label| label.is_some()));
let mut expect = HashMap::new();
expect.insert((a!(0), acta, a!(0)), frac!(3, 4));
expect.insert((a!(0), actb, a!(1)), frac!(1, 4));
expect.insert((a!(1), actc, a!(1)), frac!(6, 7));
expect.insert((a!(1), actd, a!(2)), frac!(1, 14));
expect.insert((a!(2), actc, a!(1)), frac!(3, 7));
expect.insert((a!(2), actd, a!(2)), frac!(2, 7));
let expect_final = [f0!(), frac!(1, 14), frac!(2, 7)];
assert_eq!(collect(&snfa), expect, "visible multiset differs");
for (state, termination_probability) in snfa.termination_probabilities.iter().enumerate() {
assert_eq!(
termination_probability, &expect_final[state],
"p_final mismatch state {}",
state
);
}
snfa.check_consistency().unwrap();
}
#[test]
fn tau_removal_inter_scc_path() {
let mut snfa = StochasticNondeterministicFiniteAutomaton::new();
for _ in 0..3 {
snfa.add_state();
}
snfa.add_transition(a!(0), None, a!(0), frac!(1, 2))
.unwrap();
snfa.add_transition(a!(0), None, a!(1), frac!(1, 2))
.unwrap();
snfa.add_transition(a!(1), None, a!(2), frac!(1, 4))
.unwrap();
snfa.add_transition(a!(2), None, a!(1), frac!(1, 3))
.unwrap();
let actv = Some(snfa.activity_key.process_activity("v"));
let actw = Some(snfa.activity_key.process_activity("w"));
snfa.add_transition(a!(2), actv, a!(2), frac!(1, 2))
.unwrap();
snfa.add_transition(a!(1), actw, a!(1), frac!(1, 4))
.unwrap();
snfa.remove_tau_transitions().unwrap();
assert!(snfa.activities.iter().any(|label| *label == actv));
assert!(snfa.activities.iter().any(|label| *label == actv));
snfa.check_consistency().unwrap();
}
#[test]
fn tau_removal_empty_automaton() {
let mut snfa = StochasticNondeterministicFiniteAutomaton::new();
snfa.termination_probabilities.clear();
snfa.remove_tau_transitions().unwrap(); assert_eq!(snfa.termination_probabilities.len(), 0);
}
#[test]
fn tau_removal_no_tau_edges_is_identity() {
let mut snfa = StochasticNondeterministicFiniteAutomaton::new();
let actx = Some(snfa.activity_key.process_activity("x"));
snfa.add_transition(a!(0), actx, a!(0), frac!(2, 3))
.unwrap();
let before = collect(&snfa);
let before_final = snfa.termination_probabilities[0].clone();
snfa.remove_tau_transitions().unwrap();
assert_eq!(collect(&snfa), before);
assert_eq!(snfa.termination_probabilities[0], before_final);
}
#[test]
fn tau_removal_single_state_self_final() {
let mut snfa = StochasticNondeterministicFiniteAutomaton::new();
snfa.remove_tau_transitions().unwrap();
assert_eq!(snfa.termination_probabilities.len(), 1);
assert_eq!(snfa.termination_probabilities[0], f1!());
}
#[test]
fn tau_removal_tau_chain() {
let mut snfa = StochasticNondeterministicFiniteAutomaton::new();
for _ in 0..3 {
snfa.add_state();
}
let acta = Some(snfa.activity_key.process_activity("a"));
snfa.add_transition(a!(0), None, a!(1), frac!(1, 5))
.unwrap();
snfa.add_transition(a!(1), None, a!(2), frac!(1, 2))
.unwrap();
snfa.add_transition(a!(2), acta, a!(2), frac!(1, 2))
.unwrap();
snfa.remove_tau_transitions().unwrap();
snfa.check_consistency().unwrap();
for transition in 0..snfa.sources.len() {
if snfa.activities[transition] == acta && snfa.probabilities[transition] == frac!(1, 20)
{
return;
}
}
panic!("doesn't match");
}
#[test]
#[should_panic(expected = "singular")]
fn tau_removal_singular_matrix_panics() {
let mut a = StochasticNondeterministicFiniteAutomaton::new();
a.add_transition(a!(0), None, a!(0), f1!()).unwrap();
a.remove_tau_transitions().unwrap(); }
#[test]
fn tau_removal_only_tau_edges() {
let mut snfa = StochasticNondeterministicFiniteAutomaton::new();
snfa.add_state();
snfa.add_transition(a!(0), None, a!(1), frac!(1, 2))
.unwrap();
snfa.add_transition(a!(1), None, a!(0), frac!(1, 2))
.unwrap();
snfa.remove_tau_transitions().unwrap();
assert_eq!(snfa.sources.len(), 0);
snfa.check_consistency().unwrap();
}
#[test]
fn tau_removal_trivial_source_state() {
let mut snfa = StochasticNondeterministicFiniteAutomaton::new();
snfa.add_state();
let actx = Some(snfa.activity_key.process_activity("x"));
let acty = Some(snfa.activity_key.process_activity("y"));
snfa.add_transition(a!(0), actx, a!(0), frac!(1, 2))
.unwrap();
snfa.add_transition(a!(1), None, a!(1), frac!(1, 2))
.unwrap(); snfa.add_transition(a!(1), acty, a!(1), frac!(1, 2))
.unwrap();
snfa.remove_tau_transitions().unwrap();
snfa.check_consistency().unwrap();
}
}