use crate::semiring::{LogWeight, Semiring, TropicalWeight};
use crate::Error;
use crate::Result;
use rkyv::{rancor, Archive, Deserialize, Serialize};
#[derive(Archive, Deserialize, Serialize, Debug, Clone, PartialEq)]
pub struct RkyvArcF32 {
pub ilabel: u32,
pub olabel: u32,
pub weight: f32,
pub nextstate: u32,
}
#[derive(Archive, Deserialize, Serialize, Debug, Clone)]
pub struct RkyvStateF32 {
pub final_weight: Option<f32>,
pub arcs: Vec<RkyvArcF32>,
}
#[derive(Archive, Deserialize, Serialize, Debug, Clone)]
pub struct RkyvFstF32 {
pub states: Vec<RkyvStateF32>,
pub start: Option<u32>,
}
#[derive(Archive, Deserialize, Serialize, Debug, Clone, PartialEq)]
pub struct RkyvArcF64 {
pub ilabel: u32,
pub olabel: u32,
pub weight: f64,
pub nextstate: u32,
}
#[derive(Archive, Deserialize, Serialize, Debug, Clone)]
pub struct RkyvStateF64 {
pub final_weight: Option<f64>,
pub arcs: Vec<RkyvArcF64>,
}
#[derive(Archive, Deserialize, Serialize, Debug, Clone)]
pub struct RkyvFstF64 {
pub states: Vec<RkyvStateF64>,
pub start: Option<u32>,
}
pub fn serialize_tropical_fst<F>(fst: &F) -> Result<Vec<u8>>
where
F: crate::fst::Fst<TropicalWeight>,
{
let rkyv_fst = to_rkyv_fst_tropical(fst);
let bytes = rkyv::to_bytes::<rancor::Error>(&rkyv_fst)
.map_err(|e| Error::Serialization(format!("rkyv serialization failed: {}", e)))?;
Ok(bytes.to_vec())
}
pub fn deserialize_tropical_fst(bytes: &[u8]) -> Result<crate::fst::VectorFst<TropicalWeight>> {
let archived = access_archived_tropical_fst(bytes)?;
from_archived_fst_tropical(archived)
}
pub fn access_archived_tropical_fst(bytes: &[u8]) -> Result<&ArchivedRkyvFstF32> {
rkyv::access::<ArchivedRkyvFstF32, rancor::Error>(bytes)
.map_err(|e| Error::Serialization(format!("rkyv validation failed: {}", e)))
}
pub fn serialize_log_fst<F>(fst: &F) -> Result<Vec<u8>>
where
F: crate::fst::Fst<LogWeight>,
{
let rkyv_fst = to_rkyv_fst_log(fst);
let bytes = rkyv::to_bytes::<rancor::Error>(&rkyv_fst)
.map_err(|e| Error::Serialization(format!("rkyv serialization failed: {}", e)))?;
Ok(bytes.to_vec())
}
pub fn deserialize_log_fst(bytes: &[u8]) -> Result<crate::fst::VectorFst<LogWeight>> {
let archived = access_archived_log_fst(bytes)?;
from_archived_fst_log(archived)
}
pub fn access_archived_log_fst(bytes: &[u8]) -> Result<&ArchivedRkyvFstF64> {
rkyv::access::<ArchivedRkyvFstF64, rancor::Error>(bytes)
.map_err(|e| Error::Serialization(format!("rkyv validation failed: {}", e)))
}
fn to_rkyv_fst_tropical<F>(fst: &F) -> RkyvFstF32
where
F: crate::fst::Fst<TropicalWeight>,
{
let mut states = Vec::with_capacity(fst.num_states());
for state_id in 0..fst.num_states() as u32 {
let final_weight = fst.final_weight(state_id).map(|w| *w.value());
let arcs: Vec<RkyvArcF32> = fst
.arcs(state_id)
.map(|arc| RkyvArcF32 {
ilabel: arc.ilabel,
olabel: arc.olabel,
weight: *arc.weight.value(),
nextstate: arc.nextstate,
})
.collect();
states.push(RkyvStateF32 { final_weight, arcs });
}
RkyvFstF32 {
states,
start: fst.start(),
}
}
fn to_rkyv_fst_log<F>(fst: &F) -> RkyvFstF64
where
F: crate::fst::Fst<LogWeight>,
{
let mut states = Vec::with_capacity(fst.num_states());
for state_id in 0..fst.num_states() as u32 {
let final_weight = fst.final_weight(state_id).map(|w| *w.value());
let arcs: Vec<RkyvArcF64> = fst
.arcs(state_id)
.map(|arc| RkyvArcF64 {
ilabel: arc.ilabel,
olabel: arc.olabel,
weight: *arc.weight.value(),
nextstate: arc.nextstate,
})
.collect();
states.push(RkyvStateF64 { final_weight, arcs });
}
RkyvFstF64 {
states,
start: fst.start(),
}
}
fn from_archived_fst_tropical(
archived: &ArchivedRkyvFstF32,
) -> Result<crate::fst::VectorFst<TropicalWeight>> {
use crate::arc::Arc;
use crate::fst::MutableFst;
use rkyv::option::ArchivedOption;
let mut fst = crate::fst::VectorFst::new();
for _ in 0..archived.states.len() {
fst.add_state();
}
if let ArchivedOption::Some(start) = &archived.start {
let start_val: u32 = (*start).into();
fst.set_start(start_val);
}
for (state_id, archived_state) in archived.states.iter().enumerate() {
let state_id = state_id as u32;
if let ArchivedOption::Some(weight_val) = &archived_state.final_weight {
let val: f32 = (*weight_val).into();
fst.set_final(state_id, TropicalWeight::new(val));
}
for archived_arc in archived_state.arcs.iter() {
let ilabel: u32 = archived_arc.ilabel.into();
let olabel: u32 = archived_arc.olabel.into();
let weight: f32 = archived_arc.weight.into();
let nextstate: u32 = archived_arc.nextstate.into();
fst.add_arc(
state_id,
Arc::new(ilabel, olabel, TropicalWeight::new(weight), nextstate),
);
}
}
Ok(fst)
}
fn from_archived_fst_log(
archived: &ArchivedRkyvFstF64,
) -> Result<crate::fst::VectorFst<LogWeight>> {
use crate::arc::Arc;
use crate::fst::MutableFst;
use rkyv::option::ArchivedOption;
let mut fst = crate::fst::VectorFst::new();
for _ in 0..archived.states.len() {
fst.add_state();
}
if let ArchivedOption::Some(start) = &archived.start {
let start_val: u32 = (*start).into();
fst.set_start(start_val);
}
for (state_id, archived_state) in archived.states.iter().enumerate() {
let state_id = state_id as u32;
if let ArchivedOption::Some(weight_val) = &archived_state.final_weight {
let val: f64 = (*weight_val).into();
fst.set_final(state_id, LogWeight::new(val));
}
for archived_arc in archived_state.arcs.iter() {
let ilabel: u32 = archived_arc.ilabel.into();
let olabel: u32 = archived_arc.olabel.into();
let weight: f64 = archived_arc.weight.into();
let nextstate: u32 = archived_arc.nextstate.into();
fst.add_arc(
state_id,
Arc::new(ilabel, olabel, LogWeight::new(weight), nextstate),
);
}
}
Ok(fst)
}
pub fn write_tropical_rkyv<F>(fst: &F, path: &std::path::Path) -> Result<()>
where
F: crate::fst::Fst<TropicalWeight>,
{
use std::io::Write;
let bytes = serialize_tropical_fst(fst)?;
let mut file = std::fs::File::create(path)?;
file.write_all(&bytes)?;
Ok(())
}
pub fn read_tropical_rkyv(path: &std::path::Path) -> Result<crate::fst::VectorFst<TropicalWeight>> {
let bytes = std::fs::read(path)?;
deserialize_tropical_fst(&bytes)
}
pub fn write_log_rkyv<F>(fst: &F, path: &std::path::Path) -> Result<()>
where
F: crate::fst::Fst<LogWeight>,
{
use std::io::Write;
let bytes = serialize_log_fst(fst)?;
let mut file = std::fs::File::create(path)?;
file.write_all(&bytes)?;
Ok(())
}
pub fn read_log_rkyv(path: &std::path::Path) -> Result<crate::fst::VectorFst<LogWeight>> {
let bytes = std::fs::read(path)?;
deserialize_log_fst(&bytes)
}
#[cfg(feature = "memmap2")]
mod mmap_support {
use super::*;
use memmap2::Mmap;
use std::fs::File;
use std::path::Path;
pub struct MmapTropicalFst {
mmap: Mmap,
}
impl MmapTropicalFst {
pub fn open(path: &Path) -> Result<Self> {
let file = File::open(path)?;
let mmap = unsafe { Mmap::map(&file) }.map_err(Error::Io)?;
let _ = access_archived_tropical_fst(&mmap)?;
Ok(Self { mmap })
}
pub fn archived(&self) -> &ArchivedRkyvFstF32 {
rkyv::access::<ArchivedRkyvFstF32, rancor::Error>(&self.mmap)
.expect("Archive was validated during open")
}
pub fn num_states(&self) -> usize {
self.archived().states.len()
}
pub fn start(&self) -> Option<u32> {
use rkyv::option::ArchivedOption;
match &self.archived().start {
ArchivedOption::Some(s) => Some((*s).into()),
ArchivedOption::None => None,
}
}
pub fn as_bytes(&self) -> &[u8] {
&self.mmap
}
pub fn to_vector_fst(&self) -> Result<crate::fst::VectorFst<TropicalWeight>> {
from_archived_fst_tropical(self.archived())
}
}
impl std::fmt::Debug for MmapTropicalFst {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MmapTropicalFst")
.field("num_states", &self.num_states())
.field("start", &self.start())
.finish()
}
}
pub struct MmapLogFst {
mmap: Mmap,
}
impl MmapLogFst {
pub fn open(path: &Path) -> Result<Self> {
let file = File::open(path)?;
let mmap = unsafe { Mmap::map(&file) }.map_err(Error::Io)?;
let _ = access_archived_log_fst(&mmap)?;
Ok(Self { mmap })
}
pub fn archived(&self) -> &ArchivedRkyvFstF64 {
rkyv::access::<ArchivedRkyvFstF64, rancor::Error>(&self.mmap)
.expect("Archive was validated during open")
}
pub fn num_states(&self) -> usize {
self.archived().states.len()
}
pub fn start(&self) -> Option<u32> {
use rkyv::option::ArchivedOption;
match &self.archived().start {
ArchivedOption::Some(s) => Some((*s).into()),
ArchivedOption::None => None,
}
}
pub fn to_vector_fst(&self) -> Result<crate::fst::VectorFst<LogWeight>> {
from_archived_fst_log(self.archived())
}
}
impl std::fmt::Debug for MmapLogFst {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MmapLogFst")
.field("num_states", &self.num_states())
.field("start", &self.start())
.finish()
}
}
}
#[cfg(feature = "memmap2")]
pub use mmap_support::{MmapLogFst, MmapTropicalFst};
#[cfg(test)]
mod tests {
use super::*;
use crate::fst::MutableFst;
use crate::prelude::*;
fn create_test_fst() -> VectorFst<TropicalWeight> {
let mut fst = VectorFst::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(0.5));
fst.add_arc(s0, crate::arc::Arc::new(1, 2, TropicalWeight::new(1.0), s1));
fst.add_arc(s1, crate::arc::Arc::new(3, 4, TropicalWeight::new(2.0), s2));
fst.add_arc(s0, crate::arc::Arc::new(5, 5, TropicalWeight::new(0.5), s2));
fst
}
#[test]
fn test_tropical_rkyv_roundtrip() {
let original = create_test_fst();
let bytes = serialize_tropical_fst(&original).expect("Serialization failed");
let loaded = deserialize_tropical_fst(&bytes).expect("Deserialization failed");
assert_eq!(original.num_states(), loaded.num_states());
assert_eq!(original.start(), loaded.start());
for state in 0..original.num_states() as u32 {
assert_eq!(original.final_weight(state), loaded.final_weight(state));
let orig_arcs: Vec<_> = original.arcs(state).collect();
let loaded_arcs: Vec<_> = loaded.arcs(state).collect();
assert_eq!(orig_arcs, loaded_arcs);
}
}
#[test]
fn test_zero_copy_access() {
let fst = create_test_fst();
let bytes = serialize_tropical_fst(&fst).expect("Serialization failed");
let archived = access_archived_tropical_fst(&bytes).expect("Access failed");
assert_eq!(archived.states.len(), 3);
assert!(archived.start.is_some());
}
#[test]
fn test_empty_fst() {
let fst = VectorFst::<TropicalWeight>::new();
let bytes = serialize_tropical_fst(&fst).expect("Serialization failed");
let loaded = deserialize_tropical_fst(&bytes).expect("Deserialization failed");
assert_eq!(loaded.num_states(), 0);
assert_eq!(loaded.start(), None);
}
#[test]
fn test_single_state_fst() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
fst.set_start(s0);
fst.set_final(s0, TropicalWeight::one());
let bytes = serialize_tropical_fst(&fst).expect("Serialization failed");
let loaded = deserialize_tropical_fst(&bytes).expect("Deserialization failed");
assert_eq!(loaded.num_states(), 1);
assert_eq!(loaded.start(), Some(0));
assert_eq!(loaded.final_weight(0), Some(&TropicalWeight::one()));
}
#[test]
fn test_log_weight_roundtrip() {
let mut fst = VectorFst::<LogWeight>::new();
let s0 = fst.add_state();
let s1 = fst.add_state();
fst.set_start(s0);
fst.set_final(s1, LogWeight::new(0.5));
fst.add_arc(s0, crate::arc::Arc::new(1, 1, LogWeight::new(1.0), s1));
let bytes = serialize_log_fst(&fst).expect("Serialization failed");
let loaded = deserialize_log_fst(&bytes).expect("Deserialization failed");
assert_eq!(fst.num_states(), loaded.num_states());
assert_eq!(fst.start(), loaded.start());
}
}