use crate::fst::{Fst, StateId};
use crate::semiring::Semiring;
use crate::{Error, Result};
pub fn shortest_distance_acyclic<W, F>(fst: &F) -> Result<Vec<W>>
where
W: Semiring,
F: Fst<W>,
{
let num_states = fst.num_states();
if num_states == 0 {
return Ok(Vec::new());
}
let start = match fst.start() {
Some(s) => s,
None => return Ok(vec![W::zero(); num_states]),
};
let topo_order = topological_sort(fst)?;
let mut distances = vec![W::zero(); num_states];
distances[start as usize] = W::one();
for &state in &topo_order {
let dist = distances[state as usize].clone();
if Semiring::is_zero(&dist) {
continue;
}
for arc in fst.arcs(state) {
let new_dist = dist.times(&arc.weight);
distances[arc.nextstate as usize] = distances[arc.nextstate as usize].plus(&new_dist);
}
}
Ok(distances)
}
pub fn shortest_distance_acyclic_reverse<W, F>(fst: &F) -> Result<Vec<W>>
where
W: Semiring,
F: Fst<W>,
{
let num_states = fst.num_states();
if num_states == 0 {
return Ok(Vec::new());
}
let mut topo_order = topological_sort(fst)?;
topo_order.reverse();
let mut distances = vec![W::zero(); num_states];
for state in fst.states() {
if let Some(w) = fst.final_weight(state) {
distances[state as usize] = w.clone();
}
}
let mut incoming: Vec<Vec<(StateId, W)>> = vec![Vec::new(); num_states];
for state in fst.states() {
for arc in fst.arcs(state) {
incoming[arc.nextstate as usize].push((state, arc.weight.clone()));
}
}
for &state in &topo_order {
let dist = distances[state as usize].clone();
if Semiring::is_zero(&dist) {
continue;
}
for (prev_state, weight) in &incoming[state as usize] {
let new_dist = weight.times(&dist);
distances[*prev_state as usize] = distances[*prev_state as usize].plus(&new_dist);
}
}
Ok(distances)
}
fn topological_sort<W, F>(fst: &F) -> Result<Vec<StateId>>
where
W: Semiring,
F: Fst<W>,
{
let num_states = fst.num_states();
if num_states == 0 {
return Ok(Vec::new());
}
let mut in_degree = vec![0usize; num_states];
for state in fst.states() {
for arc in fst.arcs(state) {
in_degree[arc.nextstate as usize] += 1;
}
}
let mut queue: Vec<StateId> = Vec::new();
for state in fst.states() {
if in_degree[state as usize] == 0 {
queue.push(state);
}
}
let mut result = Vec::with_capacity(num_states);
while let Some(state) = queue.pop() {
result.push(state);
for arc in fst.arcs(state) {
in_degree[arc.nextstate as usize] -= 1;
if in_degree[arc.nextstate as usize] == 0 {
queue.push(arc.nextstate);
}
}
}
if result.len() != num_states {
return Err(Error::InvalidOperation(
"FST contains cycles - use general shortest_distance instead".to_string(),
));
}
Ok(result)
}
pub fn is_acyclic<W, F>(fst: &F) -> bool
where
W: Semiring,
F: Fst<W>,
{
topological_sort(fst).is_ok()
}
pub fn shortest_distance_auto<W, F>(fst: &F) -> Result<Vec<W>>
where
W: Semiring,
F: Fst<W>,
{
match shortest_distance_acyclic(fst) {
Ok(distances) => Ok(distances),
Err(_) => {
crate::algorithms::shortest_distance(fst)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::fst::{MutableFst, VectorFst};
use crate::semiring::{Semiring, TropicalWeight};
use num_traits::One;
#[test]
fn test_acyclic_shortest_distance_simple() {
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::one());
fst.add_arc(s0, crate::arc::Arc::new(1, 1, TropicalWeight::new(3.0), s2));
fst.add_arc(s0, crate::arc::Arc::new(2, 2, TropicalWeight::new(1.0), s1));
fst.add_arc(s1, crate::arc::Arc::new(3, 3, TropicalWeight::new(0.5), s2));
let distances = shortest_distance_acyclic(&fst).unwrap();
assert_eq!(*distances[0].value(), 0.0); assert_eq!(*distances[1].value(), 1.0); assert_eq!(*distances[2].value(), 1.5); }
#[test]
fn test_acyclic_shortest_distance_chain() {
let mut fst = VectorFst::<TropicalWeight>::new();
let mut states = Vec::new();
for _ in 0..5 {
states.push(fst.add_state());
}
fst.set_start(states[0]);
fst.set_final(states[4], TropicalWeight::one());
for i in 0..4 {
fst.add_arc(
states[i],
crate::arc::Arc::new(1, 1, TropicalWeight::new(1.0), states[i + 1]),
);
}
let distances = shortest_distance_acyclic(&fst).unwrap();
for (i, dist) in distances.iter().enumerate() {
assert_eq!(*dist.value(), i as f32);
}
}
#[test]
fn test_acyclic_shortest_distance_unreachable() {
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(s1, TropicalWeight::one());
fst.add_arc(s0, crate::arc::Arc::new(1, 1, TropicalWeight::new(1.0), s1));
let distances = shortest_distance_acyclic(&fst).unwrap();
assert_eq!(*distances[0].value(), 0.0);
assert_eq!(*distances[1].value(), 1.0);
assert!(Semiring::is_zero(&distances[2])); }
#[test]
fn test_cyclic_fst_detection() {
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, crate::arc::Arc::new(1, 1, TropicalWeight::new(1.0), s1));
fst.add_arc(s1, crate::arc::Arc::new(2, 2, TropicalWeight::new(1.0), s0));
assert!(shortest_distance_acyclic::<TropicalWeight, _>(&fst).is_err());
assert!(!is_acyclic::<TropicalWeight, _>(&fst));
}
#[test]
fn test_is_acyclic() {
let mut fst1 = VectorFst::<TropicalWeight>::new();
let s0 = fst1.add_state();
let s1 = fst1.add_state();
fst1.set_start(s0);
fst1.set_final(s1, TropicalWeight::one());
fst1.add_arc(s0, crate::arc::Arc::new(1, 1, TropicalWeight::one(), s1));
assert!(is_acyclic(&fst1));
let mut fst2 = VectorFst::<TropicalWeight>::new();
let t0 = fst2.add_state();
let t1 = fst2.add_state();
fst2.set_start(t0);
fst2.add_arc(t0, crate::arc::Arc::new(1, 1, TropicalWeight::one(), t1));
fst2.add_arc(t1, crate::arc::Arc::new(2, 2, TropicalWeight::one(), t0));
assert!(!is_acyclic(&fst2));
}
#[test]
fn test_reverse_shortest_distance() {
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, crate::arc::Arc::new(1, 1, TropicalWeight::new(1.0), s1));
fst.add_arc(s1, crate::arc::Arc::new(2, 2, TropicalWeight::new(2.0), s2));
let distances = shortest_distance_acyclic_reverse(&fst).unwrap();
assert_eq!(*distances[2].value(), 0.5);
assert_eq!(*distances[1].value(), 2.5);
assert_eq!(*distances[0].value(), 3.5);
}
#[test]
fn test_empty_fst() {
let fst = VectorFst::<TropicalWeight>::new();
let distances = shortest_distance_acyclic(&fst).unwrap();
assert!(distances.is_empty());
}
#[test]
fn test_auto_selects_acyclic() {
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, crate::arc::Arc::new(1, 1, TropicalWeight::new(1.0), s1));
let distances = shortest_distance_auto(&fst).unwrap();
assert_eq!(*distances[1].value(), 1.0);
}
}