use crate::arc::Arc;
use crate::fst::{Fst, MutableFst};
use crate::semiring::Semiring;
use crate::Result;
pub fn weight_convert<W1, W2, F, M, C>(fst: &F, converter: C) -> Result<M>
where
W1: Semiring,
W2: Semiring,
F: Fst<W1>,
M: MutableFst<W2> + Default,
C: Fn(&W1) -> W2,
{
let mut result = M::default();
for _ in 0..fst.num_states() {
result.add_state();
}
if let Some(start) = fst.start() {
result.set_start(start);
}
for state in fst.states() {
if let Some(weight) = fst.final_weight(state) {
result.set_final(state, converter(weight));
}
for arc in fst.arcs(state) {
result.add_arc(
state,
Arc::new(
arc.ilabel,
arc.olabel,
converter(&arc.weight),
arc.nextstate,
),
);
}
}
Ok(result)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
#[test]
fn test_weight_convert_tropical_to_boolean_threshold() {
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(1.5));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(0.5), s1)); fst.add_arc(s1, Arc::new(2, 2, TropicalWeight::new(2.5), s2));
let converted: VectorFst<BooleanWeight> = weight_convert(&fst, |w| {
if *w.value() < 2.0 {
BooleanWeight::one()
} else {
BooleanWeight::zero()
}
})
.unwrap();
assert_eq!(converted.num_states(), fst.num_states());
assert_eq!(converted.start(), fst.start());
assert!(converted.is_final(s2));
let arcs: Vec<_> = converted.arcs(s0).collect();
assert_eq!(arcs.len(), 1);
assert_eq!(arcs[0].weight, BooleanWeight::one());
let arcs: Vec<_> = converted.arcs(s1).collect();
assert_eq!(arcs.len(), 1);
assert_eq!(arcs[0].weight, BooleanWeight::zero()); }
#[test]
fn test_weight_convert_tropical_to_log() {
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::new(1.0));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(2.0), s1));
let converted: VectorFst<LogWeight> =
weight_convert(&fst, |w| LogWeight::new((*w.value()) as f64)).unwrap();
assert_eq!(converted.num_states(), fst.num_states());
assert_eq!(converted.start(), fst.start());
assert_eq!(converted.final_weight(s1), Some(&LogWeight::new(1.0)));
let arcs: Vec<_> = converted.arcs(s0).collect();
assert_eq!(arcs.len(), 1);
assert_eq!(arcs[0].weight, LogWeight::new(2.0));
}
#[test]
fn test_weight_convert_empty_fst() {
let fst = VectorFst::<TropicalWeight>::new();
let converted: VectorFst<BooleanWeight> =
weight_convert(&fst, |_| BooleanWeight::one()).unwrap();
assert_eq!(converted.num_states(), 0);
assert!(converted.is_empty());
assert!(converted.start().is_none());
}
#[test]
fn test_weight_convert_single_state() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
fst.set_start(s0);
fst.set_final(s0, TropicalWeight::new(3.5));
let converted: VectorFst<BooleanWeight> =
weight_convert(&fst, |_| BooleanWeight::one()).unwrap();
assert_eq!(converted.num_states(), 1);
assert_eq!(converted.start(), Some(s0));
assert!(converted.is_final(s0));
assert_eq!(converted.final_weight(s0), Some(&BooleanWeight::one()));
assert_eq!(converted.num_arcs_total(), 0);
}
#[test]
fn test_weight_convert_no_start_state() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
fst.set_final(s0, TropicalWeight::one());
let converted: VectorFst<BooleanWeight> =
weight_convert(&fst, |_| BooleanWeight::one()).unwrap();
assert_eq!(converted.num_states(), 1);
assert!(converted.start().is_none());
assert!(converted.is_final(s0));
}
#[test]
fn test_weight_convert_multiple_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.set_final(s1, TropicalWeight::one());
fst.set_final(s2, TropicalWeight::one());
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));
fst.add_arc(s0, Arc::new(3, 3, TropicalWeight::new(3.0), s1));
let converted: VectorFst<TropicalWeight> =
weight_convert(&fst, |w| TropicalWeight::new(*w.value() * 2.0)).unwrap();
assert_eq!(converted.num_states(), fst.num_states());
assert_eq!(converted.num_arcs_total(), fst.num_arcs_total());
let arcs: Vec<_> = converted.arcs(s0).collect();
assert_eq!(arcs.len(), 3);
let weights: Vec<f32> = arcs.iter().map(|arc| *arc.weight.value()).collect();
assert!(weights.contains(&2.0)); assert!(weights.contains(&4.0)); assert!(weights.contains(&6.0)); }
#[test]
fn test_weight_convert_self_loops() {
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.5), s0));
fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::new(2.5), s1));
let converted: VectorFst<BooleanWeight> = weight_convert(&fst, |w| {
if *w.value() < 2.0 {
BooleanWeight::one()
} else {
BooleanWeight::zero()
}
})
.unwrap();
assert_eq!(converted.num_states(), 2);
assert_eq!(converted.num_arcs_total(), 2);
let arcs: Vec<_> = converted.arcs(s0).collect();
assert_eq!(arcs.len(), 2);
let self_loop = arcs.iter().find(|arc| arc.nextstate == s0).unwrap();
let regular_arc = arcs.iter().find(|arc| arc.nextstate == s1).unwrap();
assert_eq!(self_loop.weight, BooleanWeight::one()); assert_eq!(regular_arc.weight, BooleanWeight::zero()); }
#[test]
fn test_weight_convert_linear_chain() {
let mut fst = VectorFst::<TropicalWeight>::new();
let states: Vec<_> = (0..4).map(|_| fst.add_state()).collect();
fst.set_start(states[0]);
fst.set_final(states[3], TropicalWeight::new(4.0));
for i in 0..3 {
fst.add_arc(
states[i],
Arc::new(
(i + 1) as u32,
(i + 1) as u32,
TropicalWeight::new((i + 1) as f32),
states[i + 1],
),
);
}
let converted: VectorFst<LogWeight> =
weight_convert(&fst, |w| LogWeight::new((*w.value()) as f64)).unwrap();
assert_eq!(converted.num_states(), 4);
assert_eq!(converted.num_arcs_total(), 3);
assert_eq!(
converted.final_weight(states[3]),
Some(&LogWeight::new(4.0))
);
for (i, &state) in states[..3].iter().enumerate() {
let arcs: Vec<_> = converted.arcs(state).collect();
assert_eq!(arcs.len(), 1);
assert_eq!(*arcs[0].weight.value(), (i + 1) as f64);
}
}
#[test]
fn test_weight_convert_complex_converter() {
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::new(5.0));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(3.0), s1));
let converted: VectorFst<LogWeight> = weight_convert(&fst, |w| {
let value = *w.value();
if value > 0.0 {
LogWeight::new((value + 1.0).ln() as f64)
} else {
LogWeight::zero()
}
})
.unwrap();
assert_eq!(converted.num_states(), 2);
let arcs: Vec<_> = converted.arcs(s0).collect();
assert_eq!(arcs.len(), 1);
assert!((*arcs[0].weight.value() - 4.0_f64.ln()).abs() < 1e-6);
let final_weight = converted.final_weight(s1).unwrap();
assert!((*final_weight.value() - 6.0_f64.ln()).abs() < 1e-6);
}
#[test]
fn test_weight_convert_preserves_labels() {
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(100, 200, TropicalWeight::new(1.5), s1));
let converted: VectorFst<BooleanWeight> =
weight_convert(&fst, |_| BooleanWeight::zero()).unwrap();
let arcs: Vec<_> = converted.arcs(s0).collect();
assert_eq!(arcs.len(), 1);
assert_eq!(arcs[0].ilabel, 100);
assert_eq!(arcs[0].olabel, 200);
assert_eq!(arcs[0].nextstate, s1);
assert_eq!(arcs[0].weight, BooleanWeight::zero());
}
#[test]
fn test_weight_convert_epsilon_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.set_final(s2, TropicalWeight::one());
fst.add_arc(s0, Arc::new(0, 0, TropicalWeight::new(0.1), s1));
fst.add_arc(s1, Arc::new(1, 1, TropicalWeight::new(0.2), s2));
let converted: VectorFst<BooleanWeight> =
weight_convert(&fst, |_| BooleanWeight::one()).unwrap();
assert_eq!(converted.num_states(), 3);
assert_eq!(converted.num_arcs_total(), 2);
let epsilon_arcs: Vec<_> = converted.arcs(s0).collect();
assert_eq!(epsilon_arcs.len(), 1);
assert_eq!(epsilon_arcs[0].ilabel, 0);
assert_eq!(epsilon_arcs[0].olabel, 0);
assert_eq!(epsilon_arcs[0].weight, BooleanWeight::one());
}
#[test]
fn test_weight_convert_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::new(1.0), s1));
let converted: VectorFst<BooleanWeight> =
weight_convert(&fst, |_| BooleanWeight::one()).unwrap();
assert_eq!(converted.num_states(), 2);
assert!(converted.start().is_some());
let final_count = converted
.states()
.filter(|&s| converted.is_final(s))
.count();
assert_eq!(final_count, 0);
assert_eq!(converted.num_arcs_total(), 1);
}
#[test]
fn test_weight_convert_identity_transformation() {
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::new(2.5));
fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(1.5), s1));
let converted: VectorFst<TropicalWeight> = weight_convert(&fst, |w| *w).unwrap();
assert_eq!(converted.num_states(), fst.num_states());
assert_eq!(converted.start(), fst.start());
assert_eq!(converted.num_arcs_total(), fst.num_arcs_total());
assert_eq!(converted.final_weight(s1), fst.final_weight(s1));
let orig_arcs: Vec<_> = fst.arcs(s0).collect();
let conv_arcs: Vec<_> = converted.arcs(s0).collect();
assert_eq!(orig_arcs.len(), conv_arcs.len());
assert_eq!(orig_arcs[0].weight, conv_arcs[0].weight);
}
}