use crate::arc::Arc;
use crate::fst::{MutableFst, StateId};
use crate::semiring::Semiring;
use crate::Result;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ArcSortType {
ByInput,
ByOutput,
ByInputOutput,
ByOutputInput,
}
pub fn arc_sort<W, F>(fst: &mut F, sort_type: ArcSortType) -> Result<()>
where
W: Semiring + Clone,
F: MutableFst<W>,
{
let num_states = fst.num_states();
for state in 0..num_states as StateId {
let mut arcs: Vec<Arc<W>> = fst.arcs(state).collect();
if arcs.is_empty() {
continue;
}
match sort_type {
ArcSortType::ByInput => {
arcs.sort_by_key(|arc| arc.ilabel);
}
ArcSortType::ByOutput => {
arcs.sort_by_key(|arc| arc.olabel);
}
ArcSortType::ByInputOutput => {
arcs.sort_by_key(|arc| (arc.ilabel, arc.olabel));
}
ArcSortType::ByOutputInput => {
arcs.sort_by_key(|arc| (arc.olabel, arc.ilabel));
}
}
fst.delete_arcs(state);
for arc in arcs {
fst.add_arc(state, arc);
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
#[test]
fn test_sort_by_input() {
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(3, 10, TropicalWeight::one(), s1));
fst.add_arc(s0, Arc::new(1, 20, TropicalWeight::one(), s1));
fst.add_arc(s0, Arc::new(2, 30, TropicalWeight::one(), s1));
arc_sort(&mut fst, ArcSortType::ByInput).unwrap();
let labels: Vec<_> = fst.arcs(s0).map(|a| a.ilabel).collect();
assert_eq!(labels, vec![1, 2, 3]);
}
#[test]
fn test_sort_by_output() {
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(10, 3, TropicalWeight::one(), s1));
fst.add_arc(s0, Arc::new(20, 1, TropicalWeight::one(), s1));
fst.add_arc(s0, Arc::new(30, 2, TropicalWeight::one(), s1));
arc_sort(&mut fst, ArcSortType::ByOutput).unwrap();
let labels: Vec<_> = fst.arcs(s0).map(|a| a.olabel).collect();
assert_eq!(labels, vec![1, 2, 3]);
}
#[test]
fn test_sort_by_input_output() {
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, 3, TropicalWeight::one(), s1));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s1));
fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::one(), s1));
arc_sort(&mut fst, ArcSortType::ByInputOutput).unwrap();
let pairs: Vec<_> = fst.arcs(s0).map(|a| (a.ilabel, a.olabel)).collect();
assert_eq!(pairs, vec![(1, 1), (1, 3), (2, 2)]);
}
#[test]
fn test_sort_by_output_input() {
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(3, 1, TropicalWeight::one(), s1));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s1));
fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::one(), s1));
arc_sort(&mut fst, ArcSortType::ByOutputInput).unwrap();
let pairs: Vec<_> = fst.arcs(s0).map(|a| (a.ilabel, a.olabel)).collect();
assert_eq!(pairs, vec![(1, 1), (3, 1), (2, 2)]);
}
#[test]
fn test_empty_fst() {
let mut fst = VectorFst::<TropicalWeight>::new();
arc_sort(&mut fst, ArcSortType::ByInput).unwrap();
assert_eq!(fst.num_states(), 0);
}
#[test]
fn test_single_state_no_arcs() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
fst.set_start(s0);
arc_sort(&mut fst, ArcSortType::ByInput).unwrap();
assert_eq!(fst.num_arcs(s0), 0);
}
#[test]
fn test_single_arc() {
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::one(), s1));
arc_sort(&mut fst, ArcSortType::ByInput).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);
}
#[test]
fn test_already_sorted() {
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::one(), s1));
fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::one(), s1));
fst.add_arc(s0, Arc::new(3, 3, TropicalWeight::one(), s1));
arc_sort(&mut fst, ArcSortType::ByInput).unwrap();
let labels: Vec<_> = fst.arcs(s0).map(|a| a.ilabel).collect();
assert_eq!(labels, vec![1, 2, 3]);
}
#[test]
fn test_reverse_sorted() {
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(3, 1, TropicalWeight::one(), s1));
fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::one(), s1));
fst.add_arc(s0, Arc::new(1, 3, TropicalWeight::one(), s1));
arc_sort(&mut fst, ArcSortType::ByInput).unwrap();
let labels: Vec<_> = fst.arcs(s0).map(|a| a.ilabel).collect();
assert_eq!(labels, vec![1, 2, 3]);
}
#[test]
fn test_multiple_states_with_arcs() {
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(3, 1, TropicalWeight::one(), s1));
fst.add_arc(s0, Arc::new(1, 2, TropicalWeight::one(), s1));
fst.add_arc(s1, Arc::new(5, 3, TropicalWeight::one(), s2));
fst.add_arc(s1, Arc::new(4, 4, TropicalWeight::one(), s2));
arc_sort(&mut fst, ArcSortType::ByInput).unwrap();
let s0_labels: Vec<_> = fst.arcs(s0).map(|a| a.ilabel).collect();
let s1_labels: Vec<_> = fst.arcs(s1).map(|a| a.ilabel).collect();
assert_eq!(s0_labels, vec![1, 3]);
assert_eq!(s1_labels, vec![4, 5]);
}
#[test]
fn test_different_out_degrees() {
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);
for i in (1..=5).rev() {
fst.add_arc(s0, Arc::new(i, i, TropicalWeight::one(), s1));
}
fst.add_arc(s1, Arc::new(2, 2, TropicalWeight::one(), s2));
fst.add_arc(s1, Arc::new(1, 1, TropicalWeight::one(), s2));
arc_sort(&mut fst, ArcSortType::ByInput).unwrap();
let s0_labels: Vec<_> = fst.arcs(s0).map(|a| a.ilabel).collect();
let s1_labels: Vec<_> = fst.arcs(s1).map(|a| a.ilabel).collect();
assert_eq!(s0_labels, vec![1, 2, 3, 4, 5]);
assert_eq!(s1_labels, vec![1, 2]);
}
#[test]
fn test_stable_sort() {
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::one(), s1));
fst.add_arc(s0, Arc::new(1, 2, TropicalWeight::one(), s2));
arc_sort(&mut fst, ArcSortType::ByInput).unwrap();
let arcs: Vec<_> = fst.arcs(s0).collect();
assert_eq!(arcs[0].olabel, 1);
assert_eq!(arcs[0].nextstate, s1);
assert_eq!(arcs[1].olabel, 2);
assert_eq!(arcs[1].nextstate, s2);
}
#[test]
fn test_duplicate_labels() {
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(1.0), s1));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(2.0), s1));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(3.0), s1));
arc_sort(&mut fst, ArcSortType::ByInput).unwrap();
assert_eq!(fst.num_arcs(s0), 3);
let labels: Vec<_> = fst.arcs(s0).map(|a| a.ilabel).collect();
assert_eq!(labels, vec![1, 1, 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(3, 1, TropicalWeight::new(5.0), s1));
fst.add_arc(s0, Arc::new(1, 2, TropicalWeight::new(3.0), s1));
fst.add_arc(s0, Arc::new(2, 3, TropicalWeight::new(4.0), s1));
arc_sort(&mut fst, ArcSortType::ByInput).unwrap();
let labels: Vec<_> = fst.arcs(s0).map(|a| a.ilabel).collect();
assert_eq!(labels, vec![1, 2, 3]);
}
#[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(3, 1, LogWeight::new(1.0), s1));
fst.add_arc(s0, Arc::new(1, 2, LogWeight::new(2.0), s1));
fst.add_arc(s0, Arc::new(2, 3, LogWeight::new(3.0), s1));
arc_sort(&mut fst, ArcSortType::ByOutput).unwrap();
let labels: Vec<_> = fst.arcs(s0).map(|a| a.olabel).collect();
assert_eq!(labels, vec![1, 2, 3]);
}
#[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(3, 1, BooleanWeight::one(), s1));
fst.add_arc(s0, Arc::new(1, 2, BooleanWeight::one(), s1));
arc_sort(&mut fst, ArcSortType::ByInput).unwrap();
let labels: Vec<_> = fst.arcs(s0).map(|a| a.ilabel).collect();
assert_eq!(labels, vec![1, 3]);
}
#[test]
fn test_probability_semiring() {
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(2, 1, ProbabilityWeight::new(0.5), s1));
fst.add_arc(s0, Arc::new(1, 2, ProbabilityWeight::new(0.3), s1));
arc_sort(&mut fst, ArcSortType::ByInputOutput).unwrap();
let pairs: Vec<_> = fst.arcs(s0).map(|a| (a.ilabel, a.olabel)).collect();
assert_eq!(pairs, vec![(1, 2), (2, 1)]);
}
#[test]
fn test_language_preserved() {
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(3, 3, 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), s1));
fst.add_arc(s1, Arc::new(4, 4, TropicalWeight::new(4.0), s2));
let arcs_before_count = fst.num_arcs(s0);
arc_sort(&mut fst, ArcSortType::ByInput).unwrap();
let arcs_after_count = fst.num_arcs(s0);
assert_eq!(arcs_after_count, arcs_before_count);
let labels_after: Vec<_> = fst.arcs(s0).map(|a| a.ilabel).collect();
assert_eq!(labels_after, vec![1, 2, 3]); }
#[test]
fn test_weights_preserved() {
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(3, 1, TropicalWeight::new(5.0), s1));
fst.add_arc(s0, Arc::new(1, 2, TropicalWeight::new(3.0), s1));
fst.add_arc(s0, Arc::new(2, 3, TropicalWeight::new(4.0), s1));
arc_sort(&mut fst, ArcSortType::ByInput).unwrap();
let arcs: Vec<_> = fst.arcs(s0).collect();
assert_eq!(arcs[0].ilabel, 1);
assert_eq!(arcs[0].weight, TropicalWeight::new(3.0));
assert_eq!(arcs[1].ilabel, 2);
assert_eq!(arcs[1].weight, TropicalWeight::new(4.0));
assert_eq!(arcs[2].ilabel, 3);
assert_eq!(arcs[2].weight, TropicalWeight::new(5.0));
}
#[test]
fn test_many_arcs_per_state() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
fst.set_start(s0);
for i in (1..=100).rev() {
fst.add_arc(s0, Arc::new(i, i, TropicalWeight::one(), s1));
}
arc_sort(&mut fst, ArcSortType::ByInput).unwrap();
let labels: Vec<_> = fst.arcs(s0).map(|a| a.ilabel).collect();
let expected: Vec<_> = (1..=100).collect();
assert_eq!(labels, expected);
}
}