use crate::fst::{Fst, StateId};
use crate::semiring::Semiring;
use crate::{Error, Result};
use std::collections::HashSet;
const MAX_ITERATIONS: usize = 1000;
pub fn shortest_distance<W, F>(fst: &F) -> Result<Vec<W>>
where
W: Semiring + Clone + PartialEq,
F: Fst<W>,
{
let start = fst
.start()
.ok_or_else(|| Error::Algorithm("FST has no start state".into()))?;
let num_states = fst.num_states();
if let Ok(topo_order) = compute_topo_order(fst) {
shortest_distance_acyclic(fst, start, &topo_order)
} else {
shortest_distance_cyclic(fst, start, num_states)
}
}
fn compute_topo_order<W: Semiring, F: Fst<W>>(fst: &F) -> Result<Vec<StateId>> {
let mut visited = HashSet::new();
let mut finished = HashSet::new();
let mut order = Vec::new();
fn dfs<W: Semiring, F: Fst<W>>(
fst: &F,
state: StateId,
visited: &mut HashSet<StateId>,
finished: &mut HashSet<StateId>,
order: &mut Vec<StateId>,
) -> Result<()> {
visited.insert(state);
for arc in fst.arcs(state) {
if !visited.contains(&arc.nextstate) {
dfs(fst, arc.nextstate, visited, finished, order)?;
} else if !finished.contains(&arc.nextstate) {
return Err(Error::Algorithm("FST has cycles".into()));
}
}
finished.insert(state);
order.push(state);
Ok(())
}
if let Some(start) = fst.start() {
if !visited.contains(&start) {
dfs(fst, start, &mut visited, &mut finished, &mut order)?;
}
}
for state in fst.states() {
if !visited.contains(&state) {
dfs(fst, state, &mut visited, &mut finished, &mut order)?;
}
}
order.reverse();
Ok(order)
}
fn shortest_distance_acyclic<W, F>(
fst: &F,
start: StateId,
topo_order: &[StateId],
) -> Result<Vec<W>>
where
W: Semiring + Clone,
F: Fst<W>,
{
let num_states = fst.num_states();
let mut distance = vec![W::zero(); num_states];
distance[start as usize] = W::one();
for &state in topo_order {
let dist = distance[state as usize].clone();
if Semiring::is_zero(&dist) {
continue;
}
for arc in fst.arcs(state) {
let new_dist = dist.times(&arc.weight);
distance[arc.nextstate as usize] = distance[arc.nextstate as usize].plus(&new_dist);
}
}
Ok(distance)
}
fn shortest_distance_cyclic<W, F>(fst: &F, start: StateId, num_states: usize) -> Result<Vec<W>>
where
W: Semiring + Clone + PartialEq,
F: Fst<W>,
{
let mut distance = vec![W::zero(); num_states];
distance[start as usize] = W::one();
for iteration in 0..MAX_ITERATIONS {
let mut changed = false;
let old_distance = distance.clone();
for state in 0..num_states as StateId {
let dist = distance[state as usize].clone();
if Semiring::is_zero(&dist) {
continue;
}
for arc in fst.arcs(state) {
let new_dist = dist.times(&arc.weight);
let next_idx = arc.nextstate as usize;
let updated = distance[next_idx].plus(&new_dist);
if updated != distance[next_idx] {
distance[next_idx] = updated;
changed = true;
}
}
}
if !changed {
return Ok(distance);
}
if iteration > 0 && distances_converged(&distance, &old_distance) {
return Ok(distance);
}
}
Err(Error::Algorithm(
format!(
"Shortest distance failed to converge after {} iterations (FST may be cyclic with non-k-closed semiring)",
MAX_ITERATIONS
)
))
}
fn distances_converged<W: Semiring + Clone + PartialEq>(current: &[W], previous: &[W]) -> bool {
current.iter().zip(previous.iter()).all(|(c, p)| c == p)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
use num_traits::One;
#[test]
fn test_acyclic_linear_chain() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
let s2 = fst.add_state();
fst.set_start(s0);
fst.set_final(s2, TropicalWeight::new(0.5));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(1.0), s1));
fst.add_arc(s1, Arc::new(2, 2, TropicalWeight::new(2.0), s2));
let distances = shortest_distance(&fst).unwrap();
assert_eq!(distances[s0 as usize], TropicalWeight::one());
assert_eq!(distances[s1 as usize], TropicalWeight::new(1.0));
assert_eq!(distances[s2 as usize], TropicalWeight::new(3.0));
}
#[test]
fn test_acyclic_branching() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
fst.set_start(s0);
fst.set_final(s1, TropicalWeight::one());
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(1.0), s1));
fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::new(2.0), s1));
let distances = shortest_distance(&fst).unwrap();
assert_eq!(distances[s0 as usize], TropicalWeight::one());
assert_eq!(distances[s1 as usize], TropicalWeight::new(1.0)); }
#[test]
fn test_cyclic_simple_loop() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
fst.set_start(s0);
fst.set_final(s1, TropicalWeight::one());
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(1.0), s1));
fst.add_arc(s1, Arc::new(2, 2, TropicalWeight::new(10.0), s0));
let distances = shortest_distance(&fst).unwrap();
assert_eq!(distances[s0 as usize], TropicalWeight::one());
assert_eq!(distances[s1 as usize], TropicalWeight::new(1.0));
}
#[test]
fn test_empty_fst() {
let fst = VectorFst::<TropicalWeight>::new();
let result = shortest_distance(&fst);
assert!(result.is_err()); }
#[test]
fn test_no_start_state() {
let mut fst = VectorFst::<TropicalWeight>::new();
fst.add_state();
let result = shortest_distance(&fst);
assert!(result.is_err());
}
#[test]
fn test_single_state() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
fst.set_start(s0);
fst.set_final(s0, TropicalWeight::new(2.0));
let distances = shortest_distance(&fst).unwrap();
assert_eq!(distances[s0 as usize], TropicalWeight::one());
assert_eq!(distances.len(), 1);
}
#[test]
fn test_tropical_semiring() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
fst.set_start(s0);
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(3.0), s1));
fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::new(5.0), s1));
let distances = shortest_distance(&fst).unwrap();
assert_eq!(distances[s1 as usize], TropicalWeight::new(3.0));
}
#[test]
fn test_log_semiring() {
let mut fst = VectorFst::<LogWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
fst.set_start(s0);
fst.add_arc(s0, Arc::new(1, 1, LogWeight::new(1.0), s1));
fst.add_arc(s0, Arc::new(2, 2, LogWeight::new(2.0), s1));
let distances = shortest_distance(&fst).unwrap();
assert_eq!(distances[s0 as usize], LogWeight::one());
assert!(distances[s1 as usize].value() < &2.0);
}
#[test]
fn test_boolean_semiring() {
let mut fst = VectorFst::<BooleanWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
fst.set_start(s0);
fst.add_arc(s0, Arc::new(1, 1, BooleanWeight::one(), s1));
let distances = shortest_distance(&fst).unwrap();
assert_eq!(distances[s0 as usize], BooleanWeight::one());
assert_eq!(distances[s1 as usize], BooleanWeight::one());
}
#[test]
fn test_with_epsilon_transitions() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
let s2 = fst.add_state();
fst.set_start(s0);
fst.add_arc(s0, Arc::epsilon(TropicalWeight::new(0.5), s1));
fst.add_arc(s1, Arc::new(1, 1, TropicalWeight::new(1.0), s2));
let distances = shortest_distance(&fst).unwrap();
assert_eq!(distances[s0 as usize], TropicalWeight::one());
assert_eq!(distances[s1 as usize], TropicalWeight::new(0.5));
assert_eq!(distances[s2 as usize], TropicalWeight::new(1.5));
}
#[test]
fn test_convergence_cyclic() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
fst.set_start(s0);
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(5.0), s0));
let distances = shortest_distance(&fst).unwrap();
assert_eq!(distances[s0 as usize], TropicalWeight::one());
}
#[test]
fn test_multiple_paths_diamond() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
let s2 = fst.add_state();
let s3 = fst.add_state();
fst.set_start(s0);
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(1.0), s1));
fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::new(3.0), s2));
fst.add_arc(s1, Arc::new(3, 3, TropicalWeight::new(2.0), s3));
fst.add_arc(s2, Arc::new(4, 4, TropicalWeight::new(1.0), s3));
let distances = shortest_distance(&fst).unwrap();
assert_eq!(distances[s0 as usize], TropicalWeight::one());
assert_eq!(distances[s1 as usize], TropicalWeight::new(1.0));
assert_eq!(distances[s2 as usize], TropicalWeight::new(3.0));
assert_eq!(distances[s3 as usize], TropicalWeight::new(3.0));
}
}