use crate::fst::{Fst, Label, StateId};
use crate::semiring::Semiring;
use crate::Result;
use std::collections::{HashMap, HashSet};
pub fn partition<W, F>(fst: &F) -> Result<Vec<StateId>>
where
W: Semiring + Eq + std::hash::Hash,
F: Fst<W>,
{
let n = fst.num_states();
if n == 0 {
return Ok(Vec::new());
}
let mut class_map = initialize_partition(fst);
let mut changed = true;
while changed {
changed = false;
let old_class_map = class_map.clone();
let mut classes: HashMap<StateId, Vec<StateId>> = HashMap::new();
for (state_idx, class_id) in old_class_map.iter().enumerate() {
let state = state_idx as StateId;
classes.entry(*class_id).or_default().push(state);
}
for states_in_class in classes.values() {
if states_in_class.len() <= 1 {
continue; }
let signatures: HashMap<StateId, Signature> = states_in_class
.iter()
.map(|&state| (state, compute_signature(fst, state, &old_class_map)))
.collect();
let mut sig_groups: HashMap<Signature, Vec<StateId>> = HashMap::new();
for (&state, sig) in &signatures {
sig_groups.entry(sig.clone()).or_default().push(state);
}
if sig_groups.len() > 1 {
changed = true;
let max_class = *class_map.iter().max().unwrap_or(&0);
let mut new_class_id = max_class + 1;
for (idx, group) in sig_groups.values().enumerate() {
let target_class = if idx == 0 {
old_class_map[group[0] as usize]
} else {
let class_id = new_class_id;
new_class_id += 1;
class_id
};
for &state in group {
class_map[state as usize] = target_class;
}
}
}
}
}
renumber_classes(&mut class_map);
Ok(class_map)
}
fn initialize_partition<W, F>(fst: &F) -> Vec<StateId>
where
W: Semiring + Eq + std::hash::Hash,
F: Fst<W>,
{
let n = fst.num_states();
let mut class_map = vec![0; n];
let mut weight_to_class: HashMap<Option<W>, StateId> = HashMap::new();
let mut next_class = 0;
for (state_idx, class_entry) in class_map.iter_mut().enumerate().take(n) {
let state = state_idx as StateId;
let final_weight = fst.final_weight(state).cloned();
let class_id = weight_to_class.entry(final_weight).or_insert_with(|| {
let id = next_class;
next_class += 1;
id
});
*class_entry = *class_id;
}
class_map
}
type Signature = Vec<(Label, Label, String, StateId)>;
fn compute_signature<W, F>(fst: &F, state: StateId, class_map: &[StateId]) -> Signature
where
W: Semiring,
F: Fst<W>,
{
let mut sig: Signature = fst
.arcs(state)
.map(|arc| {
(
arc.ilabel,
arc.olabel,
format!("{:?}", arc.weight), class_map[arc.nextstate as usize],
)
})
.collect();
sig.sort();
sig
}
fn renumber_classes(class_map: &mut [StateId]) {
let unique_classes: HashSet<StateId> = class_map.iter().copied().collect();
let mut sorted_classes: Vec<StateId> = unique_classes.into_iter().collect();
sorted_classes.sort();
let renumbering: HashMap<StateId, StateId> = sorted_classes
.into_iter()
.enumerate()
.map(|(new_id, old_id)| (old_id, new_id as StateId))
.collect();
for class_id in class_map.iter_mut() {
*class_id = renumbering[class_id];
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
#[test]
fn test_partition_empty_fst() {
let fst = VectorFst::<TropicalWeight>::new();
let classes = partition(&fst).unwrap();
assert_eq!(classes.len(), 0);
}
#[test]
fn test_partition_single_state() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
fst.set_start(s0);
fst.set_final(s0, TropicalWeight::one());
let classes = partition(&fst).unwrap();
assert_eq!(classes.len(), 1);
assert_eq!(classes[0], 0);
}
#[test]
fn test_partition_two_equivalent_states() {
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.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));
let classes = partition(&fst).unwrap();
assert_eq!(classes[s1 as usize], classes[s2 as usize]);
assert_ne!(classes[s0 as usize], classes[s1 as usize]);
}
#[test]
fn test_partition_distinct_states() {
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::new(1.0));
fst.set_final(s2, TropicalWeight::new(2.0));
let classes = partition(&fst).unwrap();
assert_eq!(classes.len(), 3);
assert_ne!(classes[s0 as usize], classes[s1 as usize]);
assert_ne!(classes[s1 as usize], classes[s2 as usize]);
assert_ne!(classes[s0 as usize], classes[s2 as usize]);
}
#[test]
fn test_partition_by_arc_structure() {
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.set_final(s2, TropicalWeight::one());
fst.set_final(s3, TropicalWeight::one());
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s2));
fst.add_arc(s1, Arc::new(1, 1, TropicalWeight::one(), s3));
let classes = partition(&fst).unwrap();
assert_eq!(classes[s2 as usize], classes[s3 as usize]);
assert_eq!(classes[s0 as usize], classes[s1 as usize]);
}
#[test]
fn test_partition_different_arc_labels() {
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.set_final(s2, TropicalWeight::one());
fst.set_final(s3, TropicalWeight::one());
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s2));
fst.add_arc(s1, Arc::new(2, 2, TropicalWeight::one(), s3));
let classes = partition(&fst).unwrap();
assert_eq!(classes[s2 as usize], classes[s3 as usize]);
assert_ne!(classes[s0 as usize], classes[s1 as usize]);
}
#[test]
fn test_partition_minimal_fst() {
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));
let classes = partition(&fst).unwrap();
assert_eq!(classes.len(), 2);
assert_ne!(classes[s0 as usize], classes[s1 as usize]);
}
#[test]
fn test_partition_self_loop() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
fst.set_start(s0);
fst.set_final(s0, TropicalWeight::one());
fst.set_final(s1, TropicalWeight::one());
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s0));
let classes = partition(&fst).unwrap();
assert_ne!(classes[s0 as usize], classes[s1 as usize]);
}
#[test]
fn test_partition_complex_equivalence() {
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();
let s4 = fst.add_state();
fst.set_start(s0);
fst.set_final(s2, TropicalWeight::one());
fst.set_final(s4, TropicalWeight::one());
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s1));
fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::one(), s3));
fst.add_arc(s1, Arc::new(3, 3, TropicalWeight::one(), s2));
fst.add_arc(s3, Arc::new(3, 3, TropicalWeight::one(), s4));
let classes = partition(&fst).unwrap();
assert_eq!(classes[s2 as usize], classes[s4 as usize]);
assert_eq!(classes[s1 as usize], classes[s3 as usize]);
}
#[test]
fn test_partition_with_boolean_weight() {
let mut fst = VectorFst::<BooleanWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
fst.set_start(s0);
fst.set_final(s1, BooleanWeight::one());
let classes = partition(&fst).unwrap();
assert_eq!(classes.len(), 2);
}
#[test]
fn test_partition_renumbering() {
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(s0, TropicalWeight::new(1.0));
fst.set_final(s1, TropicalWeight::new(2.0));
fst.set_final(s2, TropicalWeight::new(3.0));
let classes = partition(&fst).unwrap();
let mut sorted_classes = classes.clone();
sorted_classes.sort();
sorted_classes.dedup();
assert_eq!(sorted_classes, vec![0, 1, 2]);
}
#[test]
fn test_partition_no_final_states() {
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));
let classes = partition(&fst).unwrap();
assert_eq!(classes.len(), 2);
}
#[test]
fn test_partition_multiple_arcs_same_dest() {
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(), s2));
fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::one(), s2));
fst.add_arc(s1, Arc::new(1, 1, TropicalWeight::one(), s2));
fst.add_arc(s1, Arc::new(2, 2, TropicalWeight::one(), s2));
let classes = partition(&fst).unwrap();
assert_eq!(classes[s0 as usize], classes[s1 as usize]);
}
}