use crate::algorithms::{determinize, intersect};
use crate::arc::Arc;
use crate::fst::{Fst, Label, MutableFst};
use crate::semiring::{DivisibleSemiring, Semiring};
use crate::{Error, Result};
use core::hash::Hash;
use std::collections::{HashMap, HashSet};
pub fn difference<W, F1, F2, M>(fst1: &F1, fst2: &F2) -> Result<M>
where
W: DivisibleSemiring + Hash + Clone + Ord + Eq,
F1: Fst<W>,
F2: Fst<W>,
M: MutableFst<W> + Default,
{
validate_acceptor(fst1)?;
validate_acceptor(fst2)?;
let alphabet = collect_alphabet(fst1, fst2)?;
let complement_fst2: M = build_complement(fst2, &alphabet)?;
intersect(fst1, &complement_fst2)
}
fn validate_acceptor<W: Semiring, F: Fst<W>>(fst: &F) -> Result<()> {
for state in fst.states() {
for arc in fst.arcs(state) {
if arc.ilabel != arc.olabel {
return Err(Error::Algorithm(
"FST must be an acceptor (input = output labels)".into(),
));
}
}
}
Ok(())
}
fn collect_alphabet<W: Semiring, F1: Fst<W>, F2: Fst<W>>(
fst1: &F1,
fst2: &F2,
) -> Result<HashSet<Label>> {
let mut alphabet = HashSet::new();
for state in fst1.states() {
for arc in fst1.arcs(state) {
if arc.ilabel != 0 {
alphabet.insert(arc.ilabel);
}
}
}
for state in fst2.states() {
for arc in fst2.arcs(state) {
if arc.ilabel != 0 {
alphabet.insert(arc.ilabel);
}
}
}
if alphabet.is_empty() {
return Err(Error::Algorithm(
"Empty alphabet in difference operation".into(),
));
}
Ok(alphabet)
}
fn build_complement<
W: DivisibleSemiring + Clone + Ord + Hash + Eq,
F: Fst<W>,
M: MutableFst<W> + Default,
>(
fst: &F,
alphabet: &HashSet<Label>,
) -> Result<M> {
let det_fst: M = determinize(fst)?;
let complete_fst: M = make_complete(&det_fst, alphabet)?;
let mut complement = M::default();
let mut state_map = HashMap::new();
for state in complete_fst.states() {
let new_state = complement.add_state();
state_map.insert(state, new_state);
}
if let Some(start) = complete_fst.start() {
if let Some(&new_start) = state_map.get(&start) {
complement.set_start(new_start);
}
}
for state in complete_fst.states() {
if let Some(&new_state) = state_map.get(&state) {
for arc in complete_fst.arcs(state) {
if let Some(&new_nextstate) = state_map.get(&arc.nextstate) {
complement.add_arc(
new_state,
Arc::new(arc.ilabel, arc.olabel, arc.weight.clone(), new_nextstate),
);
}
}
}
}
for state in complete_fst.states() {
if let Some(&new_state) = state_map.get(&state) {
if complete_fst.final_weight(state).is_none() {
complement.set_final(new_state, W::one());
}
}
}
Ok(complement)
}
fn make_complete<
W: DivisibleSemiring + Clone + Ord + Hash + Eq,
F: Fst<W>,
M: MutableFst<W> + Default,
>(
fst: &F,
alphabet: &HashSet<Label>,
) -> Result<M> {
let mut complete = M::default();
let mut state_map = HashMap::new();
for state in fst.states() {
let new_state = complete.add_state();
state_map.insert(state, new_state);
}
let sink_state = complete.add_state();
if let Some(start) = fst.start() {
if let Some(&new_start) = state_map.get(&start) {
complete.set_start(new_start);
}
}
for state in fst.states() {
if let Some(&new_state) = state_map.get(&state) {
if let Some(weight) = fst.final_weight(state) {
complete.set_final(new_state, weight.clone());
}
for arc in fst.arcs(state) {
if let Some(&new_nextstate) = state_map.get(&arc.nextstate) {
complete.add_arc(
new_state,
Arc::new(arc.ilabel, arc.olabel, arc.weight.clone(), new_nextstate),
);
}
}
}
}
for state in fst.states() {
if let Some(&new_state) = state_map.get(&state) {
let mut existing_labels = HashSet::new();
for arc in fst.arcs(state) {
if arc.ilabel != 0 {
existing_labels.insert(arc.ilabel);
}
}
for &label in alphabet {
if !existing_labels.contains(&label) {
complete.add_arc(new_state, Arc::new(label, label, W::one(), sink_state));
}
}
}
}
for &label in alphabet {
complete.add_arc(sink_state, Arc::new(label, label, W::one(), sink_state));
}
Ok(complete)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
#[test]
fn test_validate_acceptor_valid() {
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::one(), s1));
assert!(validate_acceptor(&fst).is_ok());
}
#[test]
fn test_validate_acceptor_invalid() {
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, 2, TropicalWeight::one(), s1));
assert!(validate_acceptor(&fst).is_err());
}
#[test]
fn test_collect_alphabet_basic() {
let mut fst1 = VectorFst::<TropicalWeight>::new();
let mut fst2 = VectorFst::<TropicalWeight>::new();
let s0 = fst1.add_state();
let s1 = fst1.add_state();
fst1.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s1));
fst1.add_arc(s1, Arc::new(2, 2, TropicalWeight::one(), s0));
let s0 = fst2.add_state();
let s1 = fst2.add_state();
fst2.add_arc(s0, Arc::new(2, 2, TropicalWeight::one(), s1));
fst2.add_arc(s1, Arc::new(3, 3, TropicalWeight::one(), s0));
let alphabet = collect_alphabet(&fst1, &fst2).unwrap();
assert_eq!(alphabet.len(), 3);
assert!(alphabet.contains(&1));
assert!(alphabet.contains(&2));
assert!(alphabet.contains(&3));
}
#[test]
fn test_collect_alphabet_with_epsilon() {
let mut fst1 = VectorFst::<TropicalWeight>::new();
let mut fst2 = VectorFst::<TropicalWeight>::new();
let s0 = fst1.add_state();
let s1 = fst1.add_state();
fst1.add_arc(s0, Arc::new(0, 0, TropicalWeight::one(), s1)); fst1.add_arc(s1, Arc::new(1, 1, TropicalWeight::one(), s0));
let s0 = fst2.add_state();
let s1 = fst2.add_state();
fst2.add_arc(s0, Arc::new(2, 2, TropicalWeight::one(), s1));
let alphabet = collect_alphabet(&fst1, &fst2).unwrap();
assert_eq!(alphabet.len(), 2);
assert!(alphabet.contains(&1));
assert!(alphabet.contains(&2));
assert!(!alphabet.contains(&0)); }
#[test]
fn test_collect_alphabet_empty() {
let mut fst1 = VectorFst::<TropicalWeight>::new();
let mut fst2 = VectorFst::<TropicalWeight>::new();
let s0 = fst1.add_state();
let s1 = fst1.add_state();
fst1.add_arc(s0, Arc::new(0, 0, TropicalWeight::one(), s1));
let _s0 = fst2.add_state();
let result = collect_alphabet(&fst1, &fst2);
assert!(result.is_err()); }
#[test]
fn test_make_complete_basic() {
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::one(), s1));
let mut alphabet = HashSet::new();
alphabet.insert(1);
alphabet.insert(2);
let complete: VectorFst<TropicalWeight> = make_complete(&fst, &alphabet).unwrap();
assert_eq!(complete.num_states(), 3);
assert!(complete.start().is_some());
assert!(complete.num_arcs_total() > fst.num_arcs_total());
}
#[test]
fn test_build_complement_simple() {
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::one(), s1));
let mut alphabet = HashSet::new();
alphabet.insert(1);
let complement: VectorFst<TropicalWeight> = build_complement(&fst, &alphabet).unwrap();
assert!(complement.num_states() >= fst.num_states());
assert!(complement.start().is_some());
}
#[test]
fn test_difference_basic() {
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, Arc::new(1, 1, TropicalWeight::one(), s1));
let mut fst2 = VectorFst::<TropicalWeight>::new();
let s0 = fst2.add_state();
let s1 = fst2.add_state();
fst2.set_start(s0);
fst2.set_final(s1, TropicalWeight::one());
fst2.add_arc(s0, Arc::new(2, 2, TropicalWeight::one(), s1));
let diff: VectorFst<TropicalWeight> = difference(&fst1, &fst2).unwrap();
assert!(diff.start().is_some());
assert!(diff.num_states() > 0);
}
#[test]
fn test_difference_self() {
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::one(), s1));
let diff: VectorFst<TropicalWeight> = difference(&fst, &fst).unwrap();
let final_count = diff.states().filter(|&s| diff.is_final(s)).count();
assert_eq!(final_count, 0);
}
#[test]
fn test_difference_empty_fst() {
let fst1 = VectorFst::<TropicalWeight>::new();
let mut fst2 = VectorFst::<TropicalWeight>::new();
let s0 = fst2.add_state();
fst2.set_start(s0);
fst2.set_final(s0, TropicalWeight::one());
let result = difference::<TropicalWeight, _, _, VectorFst<TropicalWeight>>(&fst1, &fst2);
assert!(result.is_err());
}
#[test]
fn test_difference_non_acceptor() {
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, Arc::new(1, 2, TropicalWeight::one(), s1));
let mut fst2 = VectorFst::<TropicalWeight>::new();
let s0 = fst2.add_state();
let s1 = fst2.add_state();
fst2.set_start(s0);
fst2.set_final(s1, TropicalWeight::one());
fst2.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s1));
let result = difference::<TropicalWeight, _, _, VectorFst<TropicalWeight>>(&fst1, &fst2);
assert!(result.is_err());
}
#[test]
fn test_difference_overlapping_languages() {
let mut fst1 = VectorFst::<TropicalWeight>::new();
let s0 = fst1.add_state();
let s1 = fst1.add_state();
let s2 = fst1.add_state();
fst1.set_start(s0);
fst1.set_final(s1, TropicalWeight::one()); fst1.set_final(s2, TropicalWeight::one()); fst1.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s1)); fst1.add_arc(s1, Arc::new(2, 2, TropicalWeight::one(), s2));
let mut fst2 = VectorFst::<TropicalWeight>::new();
let s0 = fst2.add_state();
let s1 = fst2.add_state();
let s2 = fst2.add_state();
fst2.set_start(s0);
fst2.set_final(s2, TropicalWeight::one());
fst2.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s1)); fst2.add_arc(s1, Arc::new(2, 2, TropicalWeight::one(), s2));
let diff: VectorFst<TropicalWeight> = difference(&fst1, &fst2).unwrap();
assert!(diff.start().is_some());
assert!(diff.num_states() > 0);
}
#[test]
fn test_difference_disjoint_languages() {
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, Arc::new(1, 1, TropicalWeight::one(), s1));
let mut fst2 = VectorFst::<TropicalWeight>::new();
let s0 = fst2.add_state();
let s1 = fst2.add_state();
fst2.set_start(s0);
fst2.set_final(s1, TropicalWeight::one());
fst2.add_arc(s0, Arc::new(2, 2, TropicalWeight::one(), s1));
let diff: VectorFst<TropicalWeight> = difference(&fst1, &fst2).unwrap();
assert!(diff.start().is_some());
assert!(diff.num_states() > 0);
let final_count = diff.states().filter(|&s| diff.is_final(s)).count();
assert!(final_count > 0);
}
#[test]
fn test_difference_complex_automata() {
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, Arc::new(1, 1, TropicalWeight::one(), s1)); fst1.add_arc(s0, Arc::new(2, 2, TropicalWeight::one(), s0)); fst1.add_arc(s1, Arc::new(1, 1, TropicalWeight::one(), s1)); fst1.add_arc(s1, Arc::new(2, 2, TropicalWeight::one(), s0));
let mut fst2 = VectorFst::<TropicalWeight>::new();
let s0 = fst2.add_state();
let s1 = fst2.add_state();
fst2.set_start(s0);
fst2.set_final(s1, TropicalWeight::one());
fst2.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s1));
let diff: VectorFst<TropicalWeight> = difference(&fst1, &fst2).unwrap();
assert!(diff.start().is_some());
assert!(diff.num_states() > 0);
}
#[test]
fn test_difference_preserves_weights() {
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::new(2.0));
fst1.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(1.5), s1));
let mut fst2 = VectorFst::<TropicalWeight>::new();
let s0 = fst2.add_state();
fst2.set_start(s0);
fst2.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s0));
let diff: VectorFst<TropicalWeight> = difference(&fst1, &fst2).unwrap();
assert!(diff.start().is_some());
let has_final = diff.states().any(|s| diff.is_final(s));
assert!(has_final);
}
#[test]
fn test_difference_single_state_acceptors() {
let mut fst1 = VectorFst::<TropicalWeight>::new();
let s0 = fst1.add_state();
fst1.set_start(s0);
fst1.set_final(s0, TropicalWeight::one());
fst1.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s0));
let mut fst2 = VectorFst::<TropicalWeight>::new();
let s0 = fst2.add_state();
fst2.set_start(s0);
fst2.set_final(s0, TropicalWeight::one());
fst2.add_arc(s0, Arc::new(2, 2, TropicalWeight::one(), s0));
let diff: VectorFst<TropicalWeight> = difference(&fst1, &fst2).unwrap();
assert!(diff.start().is_some());
}
#[test]
fn test_difference_epsilon_handling() {
let mut fst1 = VectorFst::<TropicalWeight>::new();
let s0 = fst1.add_state();
let s1 = fst1.add_state();
let s2 = fst1.add_state();
fst1.set_start(s0);
fst1.set_final(s2, TropicalWeight::one());
fst1.add_arc(s0, Arc::new(0, 0, TropicalWeight::one(), s1)); fst1.add_arc(s1, Arc::new(1, 1, TropicalWeight::one(), s2));
let mut fst2 = VectorFst::<TropicalWeight>::new();
let s0 = fst2.add_state();
let s1 = fst2.add_state();
fst2.set_start(s0);
fst2.set_final(s1, TropicalWeight::one());
fst2.add_arc(s0, Arc::new(2, 2, TropicalWeight::one(), s1));
let diff: VectorFst<TropicalWeight> = difference(&fst1, &fst2).unwrap();
assert!(diff.start().is_some());
}
}