use crate::arc::Arc;
use crate::fst::{Label, MutableFst, StateId};
use crate::semiring::Semiring;
use crate::Result;
use std::collections::HashMap;
type ArcKey = (Label, Label, StateId);
pub fn arc_sum<W, F>(fst: &mut F) -> Result<()>
where
W: Semiring + Clone,
F: MutableFst<W>,
{
let num_states = fst.num_states();
for state in 0..num_states as StateId {
let arcs: Vec<Arc<W>> = fst.arcs(state).collect();
if arcs.is_empty() {
continue;
}
let mut arc_groups: HashMap<ArcKey, W> = HashMap::new();
for arc in arcs {
let key = (arc.ilabel, arc.olabel, arc.nextstate);
arc_groups
.entry(key)
.and_modify(|w| *w = w.plus(&arc.weight))
.or_insert(arc.weight);
}
fst.delete_arcs(state);
for ((ilabel, olabel, nextstate), weight) in arc_groups {
fst.add_arc(state, Arc::new(ilabel, olabel, weight, nextstate));
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
#[test]
fn test_exact_duplicates() {
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, 2, TropicalWeight::new(1.0), s1));
fst.add_arc(s0, Arc::new(1, 2, TropicalWeight::new(2.0), s1));
arc_sum(&mut fst).unwrap();
assert_eq!(fst.num_arcs(s0), 1);
let arc = fst.arcs(s0).next().unwrap();
assert_eq!(arc.ilabel, 1);
assert_eq!(arc.olabel, 2);
assert_eq!(arc.weight, TropicalWeight::new(1.0)); }
#[test]
fn test_multiple_duplicates() {
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(1, 1, TropicalWeight::new(1.0), s1));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(2.0), s1));
arc_sum(&mut fst).unwrap();
assert_eq!(fst.num_arcs(s0), 1);
let arc = fst.arcs(s0).next().unwrap();
assert_eq!(arc.weight, TropicalWeight::new(1.0)); }
#[test]
fn test_no_duplicates() {
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::new(1, 1, TropicalWeight::new(1.0), s1));
fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::new(2.0), s2));
arc_sum(&mut fst).unwrap();
assert_eq!(fst.num_arcs(s0), 2);
}
#[test]
fn test_partial_duplicates() {
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::new(1, 1, TropicalWeight::new(1.0), s1));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(2.0), s1)); fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::new(3.0), s2));
arc_sum(&mut fst).unwrap();
assert_eq!(fst.num_arcs(s0), 2);
}
#[test]
fn test_different_weights() {
let mut fst = VectorFst::<ProbabilityWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
fst.set_start(s0);
fst.add_arc(s0, Arc::new(1, 1, ProbabilityWeight::new(0.3), s1));
fst.add_arc(s0, Arc::new(1, 1, ProbabilityWeight::new(0.4), s1));
arc_sum(&mut fst).unwrap();
assert_eq!(fst.num_arcs(s0), 1);
let arc = fst.arcs(s0).next().unwrap();
assert_eq!(arc.weight, ProbabilityWeight::new(0.7));
}
#[test]
fn test_empty_fst() {
let mut fst = VectorFst::<TropicalWeight>::new();
arc_sum(&mut fst).unwrap();
assert_eq!(fst.num_states(), 0);
}
#[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(5.0), s1));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(3.0), s1));
arc_sum(&mut fst).unwrap();
let arc = fst.arcs(s0).next().unwrap();
assert_eq!(arc.weight, 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(1, 1, LogWeight::new(2.0), s1));
arc_sum(&mut fst).unwrap();
assert_eq!(fst.num_arcs(s0), 1);
}
#[test]
fn test_multiple_states_with_duplicates() {
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::new(1, 1, TropicalWeight::new(1.0), s1));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(2.0), s1));
fst.add_arc(s1, Arc::new(2, 2, TropicalWeight::new(3.0), s2));
fst.add_arc(s1, Arc::new(2, 2, TropicalWeight::new(4.0), s2));
arc_sum(&mut fst).unwrap();
assert_eq!(fst.num_arcs(s0), 1);
assert_eq!(fst.num_arcs(s1), 1);
}
#[test]
fn test_same_labels_different_nextstates() {
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::new(1, 1, TropicalWeight::new(1.0), s1));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(2.0), s2));
arc_sum(&mut fst).unwrap();
assert_eq!(fst.num_arcs(s0), 2);
}
}