use crate::fst::Fst;
use crate::semiring::Semiring;
use crate::utils::FstPath;
use crate::Result;
pub fn decode_linear_fst<W, F>(fst: &F) -> Result<FstPath<W>>
where
W: Semiring + Clone,
F: Fst<W>,
{
let start = fst
.start()
.ok_or_else(|| crate::Error::Algorithm("FST has no start state".to_string()))?;
let mut states = vec![start];
let mut arcs = Vec::new();
let mut current = start;
let mut weight = W::one();
let mut visited = std::collections::HashSet::new();
let mut just_followed_epsilon = false;
loop {
if !just_followed_epsilon {
if visited.contains(¤t) {
return Err(crate::Error::Algorithm(
"FST is not linear: contains cycles".to_string(),
));
}
visited.insert(current);
}
just_followed_epsilon = false;
if let Some(final_weight) = fst.final_weight(current) {
weight = weight * final_weight.clone();
return Ok(FstPath {
states,
arcs,
weight,
final_state: current,
});
}
let arc_iter = fst.arcs(current);
let mut non_epsilon_arcs: Vec<_> = arc_iter.collect();
let epsilon_arcs: Vec<_> = non_epsilon_arcs
.iter()
.filter(|arc| arc.is_epsilon())
.cloned()
.collect();
non_epsilon_arcs.retain(|arc| !arc.is_epsilon());
if non_epsilon_arcs.len() > 1 {
return Err(crate::Error::Algorithm(
"FST is not linear: multiple non-epsilon paths from state".to_string(),
));
}
if epsilon_arcs.len() > 1 {
return Err(crate::Error::Algorithm(
"FST is not linear: multiple epsilon paths from state".to_string(),
));
}
if let Some(epsilon_arc) = epsilon_arcs.first() {
weight = weight.clone() * epsilon_arc.weight.clone();
states.push(epsilon_arc.nextstate);
arcs.push(epsilon_arc.clone());
current = epsilon_arc.nextstate;
if visited.contains(¤t) {
return Err(crate::Error::Algorithm(
"FST is not linear: contains cycles".to_string(),
));
}
visited.insert(current);
just_followed_epsilon = true;
if let Some(final_weight) = fst.final_weight(current) {
weight = weight * final_weight.clone();
return Ok(FstPath {
states,
arcs,
weight,
final_state: current,
});
}
continue;
}
let arc = non_epsilon_arcs.first().ok_or_else(|| {
crate::Error::Algorithm("FST is not linear: no path to final state".to_string())
})?;
weight = weight * arc.weight.clone();
states.push(arc.nextstate);
arcs.push(arc.clone());
current = arc.nextstate;
}
}
pub fn decode_linear_fst_input<W, F>(fst: &F) -> Result<Vec<u32>>
where
W: Semiring + Clone,
F: Fst<W>,
{
let path = decode_linear_fst(fst)?;
Ok(path.input_labels())
}
pub fn decode_linear_fst_output<W, F>(fst: &F) -> Result<Vec<u32>>
where
W: Semiring + Clone,
F: Fst<W>,
{
let path = decode_linear_fst(fst)?;
Ok(path.output_labels())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
#[test]
fn test_decode_linear_fst() {
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 path = decode_linear_fst(&fst).unwrap();
assert_eq!(path.states, vec![0, 1, 2]);
assert_eq!(path.arcs.len(), 2);
assert_eq!(path.input_labels(), vec![1, 2]);
assert_eq!(path.output_labels(), vec![1, 2]);
}
#[test]
fn test_decode_linear_fst_no_start() {
let fst = VectorFst::<TropicalWeight>::new();
assert!(decode_linear_fst(&fst).is_err());
}
#[test]
fn test_decode_linear_fst_multiple_paths() {
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, Arc::new(1, 1, TropicalWeight::one(), s1));
fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::one(), s2)); assert!(decode_linear_fst(&fst).is_err());
}
}