use crate::arc::Arc;
use crate::fst::{Fst, Label, MutableFst, StateId};
use crate::semiring::Semiring;
use crate::utils::SymbolTable;
use crate::{Error, Result};
use std::collections::HashMap;
use std::io::{BufRead, Write};
use std::str::FromStr;
pub fn write_text<W, F, Writer>(
fst: &F,
writer: &mut Writer,
isyms: Option<&SymbolTable>,
osyms: Option<&SymbolTable>,
) -> Result<()>
where
W: Semiring,
F: Fst<W>,
Writer: Write,
{
if let Some(start) = fst.start() {
writeln!(writer, "START\t{start}")?;
}
for state in fst.states() {
writeln!(writer, "STATE\t{state}")?;
}
for state in fst.states() {
for arc in fst.arcs(state) {
let nextstate = arc.nextstate;
write!(writer, "{state}\t{nextstate}\t")?;
if let Some(syms) = isyms {
let symbol = syms.find(arc.ilabel).unwrap_or("?");
write!(writer, "{symbol}\t")?;
} else {
let ilabel = arc.ilabel;
write!(writer, "{ilabel}\t")?;
}
if let Some(syms) = osyms {
let symbol = syms.find(arc.olabel).unwrap_or("?");
write!(writer, "{symbol}\t")?;
} else {
let olabel = arc.olabel;
write!(writer, "{olabel}\t")?;
}
let weight = &arc.weight;
writeln!(writer, "{weight}")?;
}
if let Some(weight) = fst.final_weight(state) {
writeln!(writer, "FINAL\t{state}\t{weight}")?;
}
}
Ok(())
}
pub fn read_text<W, M, Reader>(
reader: &mut Reader,
isyms: Option<&SymbolTable>,
osyms: Option<&SymbolTable>,
) -> Result<M>
where
W: Semiring + FromStr,
W::Err: std::error::Error + Send + Sync + 'static,
M: MutableFst<W> + Default,
Reader: BufRead,
{
let buf_reader = reader;
let mut fst = M::default();
let mut state_map = HashMap::new();
let get_state = |state_map: &mut HashMap<StateId, StateId>,
fst: &mut M,
id: StateId|
-> StateId { *state_map.entry(id).or_insert_with(|| fst.add_state()) };
for line in buf_reader.lines() {
let line = line?;
let parts: Vec<&str> = line.split_whitespace().collect();
match parts.len() {
2 => {
if parts[0] == "START" {
let state = parts[1]
.parse::<StateId>()
.map_err(|e| Error::Serialization(e.to_string()))?;
let state = get_state(&mut state_map, &mut fst, state);
fst.set_start(state);
} else if parts[0] == "STATE" {
let state = parts[1]
.parse::<StateId>()
.map_err(|e| Error::Serialization(e.to_string()))?;
get_state(&mut state_map, &mut fst, state);
} else {
let state = parts[0]
.parse::<StateId>()
.map_err(|e| Error::Serialization(e.to_string()))?;
let weight = parts[1]
.parse::<W>()
.map_err(|e| Error::Serialization(e.to_string()))?;
let state = get_state(&mut state_map, &mut fst, state);
fst.set_final(state, weight);
}
}
3 => {
if parts[0] == "FINAL" {
let state = parts[1]
.parse::<StateId>()
.map_err(|e| Error::Serialization(e.to_string()))?;
let weight = parts[2]
.parse::<W>()
.map_err(|e| Error::Serialization(e.to_string()))?;
let state = get_state(&mut state_map, &mut fst, state);
fst.set_final(state, weight);
}
}
5 => {
let from = parts[0]
.parse::<StateId>()
.map_err(|e| Error::Serialization(e.to_string()))?;
let to = parts[1]
.parse::<StateId>()
.map_err(|e| Error::Serialization(e.to_string()))?;
let ilabel = if let Some(syms) = isyms {
syms.find_id(parts[2]).unwrap_or(0)
} else {
parts[2]
.parse::<Label>()
.map_err(|e| Error::Serialization(e.to_string()))?
};
let olabel = if let Some(syms) = osyms {
syms.find_id(parts[3]).unwrap_or(0)
} else {
parts[3]
.parse::<Label>()
.map_err(|e| Error::Serialization(e.to_string()))?
};
let weight = parts[4]
.parse::<W>()
.map_err(|e| Error::Serialization(e.to_string()))?;
let from = get_state(&mut state_map, &mut fst, from);
let to = get_state(&mut state_map, &mut fst, to);
fst.add_arc(from, Arc::new(ilabel, olabel, weight, to));
}
_ => continue, }
}
Ok(fst)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
use num_traits::identities::One;
use std::io::{BufReader, Cursor};
#[test]
fn test_write_read_text_roundtrip() {
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(2.5));
fst.add_arc(s0, Arc::new(1, 2, TropicalWeight::new(1.0), s1));
fst.add_arc(s1, Arc::new(3, 4, TropicalWeight::new(1.5), s2));
fst.add_arc(s0, Arc::epsilon(TropicalWeight::new(0.5), s2));
let mut buffer = Vec::new();
write_text(&fst, &mut buffer, None, None).unwrap();
let cursor = Cursor::new(buffer);
let mut buf_reader = BufReader::new(cursor);
let read_fst: VectorFst<TropicalWeight> =
read_text::<TropicalWeight, VectorFst<TropicalWeight>, _>(&mut buf_reader, None, None)
.unwrap();
assert_eq!(read_fst.num_states(), fst.num_states());
assert_eq!(read_fst.start(), fst.start());
assert_eq!(read_fst.num_arcs_total(), fst.num_arcs_total());
for state in fst.states() {
let original_final = fst.final_weight(state);
let read_final = read_fst.final_weight(state);
match (original_final, read_final) {
(Some(w1), Some(w2)) => assert_eq!(w1, w2),
(None, None) => {}
_ => panic!("Final weight mismatch for state {state}"),
}
}
for state in fst.states() {
let original_arcs: Vec<_> = fst.arcs(state).collect();
let read_arcs: Vec<_> = read_fst.arcs(state).collect();
assert_eq!(original_arcs.len(), read_arcs.len());
for (orig, read) in original_arcs.iter().zip(read_arcs.iter()) {
assert_eq!(orig.ilabel, read.ilabel);
assert_eq!(orig.olabel, read.olabel);
assert_eq!(orig.weight, read.weight);
assert_eq!(orig.nextstate, read.nextstate);
}
}
}
#[test]
fn test_write_read_empty_fst() {
let fst = VectorFst::<TropicalWeight>::new();
let mut buffer = Vec::new();
write_text(&fst, &mut buffer, None, None).unwrap();
let mut cursor = Cursor::new(buffer);
let read_fst: VectorFst<TropicalWeight> =
read_text::<TropicalWeight, VectorFst<TropicalWeight>, _>(&mut cursor, None, None)
.unwrap();
assert!(read_fst.is_empty());
assert_eq!(read_fst.num_states(), 0);
assert_eq!(read_fst.start(), None);
}
#[test]
fn test_write_read_single_state() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
fst.set_start(s0);
fst.set_final(s0, TropicalWeight::one());
let mut buffer = Vec::new();
write_text(&fst, &mut buffer, None, None).unwrap();
let mut cursor = Cursor::new(buffer);
let read_fst: VectorFst<TropicalWeight> =
read_text::<TropicalWeight, VectorFst<TropicalWeight>, _>(&mut cursor, None, None)
.unwrap();
assert_eq!(read_fst.num_states(), 1);
assert_eq!(read_fst.start(), Some(0));
assert!(read_fst.is_final(0));
}
#[test]
fn test_text_format_different_weights() {
let mut bool_fst = VectorFst::<BooleanWeight>::new();
let s0 = bool_fst.add_state();
let s1 = bool_fst.add_state();
bool_fst.set_start(s0);
bool_fst.set_final(s1, BooleanWeight::new(true));
bool_fst.add_arc(s0, Arc::new(1, 1, BooleanWeight::new(false), s1));
let mut buffer = Vec::new();
write_text(&bool_fst, &mut buffer, None, None).unwrap();
let mut cursor = Cursor::new(buffer);
let read_fst: VectorFst<BooleanWeight> =
read_text::<BooleanWeight, VectorFst<BooleanWeight>, _>(&mut cursor, None, None)
.unwrap();
assert_eq!(read_fst.num_states(), bool_fst.num_states());
assert_eq!(read_fst.start(), bool_fst.start());
let mut prob_fst = VectorFst::<ProbabilityWeight>::new();
let s0 = prob_fst.add_state();
let s1 = prob_fst.add_state();
prob_fst.set_start(s0);
prob_fst.set_final(s1, ProbabilityWeight::new(0.8));
prob_fst.add_arc(s0, Arc::new(1, 1, ProbabilityWeight::new(0.3), s1));
let mut buffer = Vec::new();
write_text(&prob_fst, &mut buffer, None, None).unwrap();
let mut cursor = Cursor::new(buffer);
let read_fst: VectorFst<ProbabilityWeight> = read_text::<
ProbabilityWeight,
VectorFst<ProbabilityWeight>,
_,
>(&mut cursor, None, None)
.unwrap();
assert_eq!(read_fst.num_states(), prob_fst.num_states());
assert_eq!(read_fst.start(), prob_fst.start());
}
}