use crate::arc::{Arc, ArcIterator};
use crate::fst::{Fst, Label, StateId};
use crate::semiring::Semiring;
use crate::Result;
#[allow(dead_code)]
const CACHE_LINE_SIZE: usize = 64;
#[allow(dead_code)]
#[inline]
fn align_to_cache_line(size: usize) -> usize {
(size + CACHE_LINE_SIZE - 1) & !(CACHE_LINE_SIZE - 1)
}
#[repr(C)]
struct AlignedVec<T: Clone> {
data: Box<[T]>,
}
impl<T: Clone> Clone for AlignedVec<T> {
fn clone(&self) -> Self {
Self {
data: self.data.clone(),
}
}
}
impl<T: Clone + Default> AlignedVec<T> {
#[allow(dead_code)]
fn with_capacity(len: usize) -> Self {
let data = vec![T::default(); len].into_boxed_slice();
Self { data }
}
}
impl<T: Clone> AlignedVec<T> {
fn from_vec(v: Vec<T>) -> Self {
Self {
data: v.into_boxed_slice(),
}
}
#[inline]
fn len(&self) -> usize {
self.data.len()
}
#[inline]
fn get(&self, index: usize) -> Option<&T> {
self.data.get(index)
}
#[inline]
fn as_slice(&self) -> &[T] {
&self.data
}
}
impl<T: Clone> std::ops::Index<usize> for AlignedVec<T> {
type Output = T;
#[inline]
fn index(&self, index: usize) -> &Self::Output {
&self.data[index]
}
}
impl<T: Clone> std::ops::IndexMut<usize> for AlignedVec<T> {
#[inline]
fn index_mut(&mut self, index: usize) -> &mut Self::Output {
&mut self.data[index]
}
}
#[derive(Clone)]
pub struct CsrFst<W: Semiring> {
start: Option<StateId>,
num_states: usize,
state_offsets: AlignedVec<u32>,
final_weights: AlignedVec<Option<W>>,
arc_ilabels: AlignedVec<Label>,
arc_olabels: AlignedVec<Label>,
arc_weights: AlignedVec<W>,
arc_nextstates: AlignedVec<StateId>,
properties: crate::properties::FstProperties,
}
impl<W: Semiring> CsrFst<W> {
pub fn new() -> Self {
Self {
start: None,
num_states: 0,
state_offsets: AlignedVec::from_vec(vec![0]),
final_weights: AlignedVec::from_vec(Vec::new()),
arc_ilabels: AlignedVec::from_vec(Vec::new()),
arc_olabels: AlignedVec::from_vec(Vec::new()),
arc_weights: AlignedVec::from_vec(Vec::new()),
arc_nextstates: AlignedVec::from_vec(Vec::new()),
properties: crate::properties::FstProperties::default(),
}
}
pub fn from_fst<F: Fst<W>>(fst: &F) -> Result<Self> {
let num_states = fst.num_states();
if num_states == 0 {
return Ok(Self::new());
}
let mut state_offsets = Vec::with_capacity(num_states + 1);
let mut total_arcs = 0u32;
state_offsets.push(0);
for state in 0..num_states as StateId {
total_arcs += fst.num_arcs(state) as u32;
state_offsets.push(total_arcs);
}
let total_arcs_usize = total_arcs as usize;
let mut arc_ilabels = Vec::with_capacity(total_arcs_usize);
let mut arc_olabels = Vec::with_capacity(total_arcs_usize);
let mut arc_weights = Vec::with_capacity(total_arcs_usize);
let mut arc_nextstates = Vec::with_capacity(total_arcs_usize);
let mut final_weights = Vec::with_capacity(num_states);
for state in 0..num_states as StateId {
final_weights.push(fst.final_weight(state).cloned());
for arc in fst.arcs(state) {
arc_ilabels.push(arc.ilabel);
arc_olabels.push(arc.olabel);
arc_weights.push(arc.weight.clone());
arc_nextstates.push(arc.nextstate);
}
}
let properties = crate::properties::compute_properties(fst);
Ok(Self {
start: fst.start(),
num_states,
state_offsets: AlignedVec::from_vec(state_offsets),
final_weights: AlignedVec::from_vec(final_weights),
arc_ilabels: AlignedVec::from_vec(arc_ilabels),
arc_olabels: AlignedVec::from_vec(arc_olabels),
arc_weights: AlignedVec::from_vec(arc_weights),
arc_nextstates: AlignedVec::from_vec(arc_nextstates),
properties,
})
}
#[inline]
fn arc_range(&self, state: StateId) -> std::ops::Range<usize> {
let start = self.state_offsets[state as usize] as usize;
let end = self.state_offsets[state as usize + 1] as usize;
start..end
}
#[inline]
pub fn ilabels_slice(&self) -> &[Label] {
self.arc_ilabels.as_slice()
}
#[inline]
pub fn olabels_slice(&self) -> &[Label] {
self.arc_olabels.as_slice()
}
#[inline]
pub fn weights_slice(&self) -> &[W] {
self.arc_weights.as_slice()
}
#[inline]
pub fn nextstates_slice(&self) -> &[StateId] {
self.arc_nextstates.as_slice()
}
#[inline]
pub fn state_ilabels(&self, state: StateId) -> &[Label] {
let range = self.arc_range(state);
&self.arc_ilabels.as_slice()[range]
}
#[inline]
pub fn state_olabels(&self, state: StateId) -> &[Label] {
let range = self.arc_range(state);
&self.arc_olabels.as_slice()[range]
}
#[inline]
pub fn state_weights(&self, state: StateId) -> &[W] {
let range = self.arc_range(state);
&self.arc_weights.as_slice()[range]
}
#[inline]
pub fn state_nextstates(&self, state: StateId) -> &[StateId] {
let range = self.arc_range(state);
&self.arc_nextstates.as_slice()[range]
}
#[inline]
pub fn total_arcs(&self) -> usize {
self.arc_ilabels.len()
}
#[inline]
pub fn prefetch_state(&self, state: StateId) {
if (state as usize) < self.num_states {
let range = self.arc_range(state);
if !range.is_empty() {
#[cfg(target_arch = "x86_64")]
unsafe {
use std::arch::x86_64::*;
let ilabel_ptr = self.arc_ilabels.as_slice().as_ptr().add(range.start);
let weight_ptr = self.arc_weights.as_slice().as_ptr().add(range.start);
let next_ptr = self.arc_nextstates.as_slice().as_ptr().add(range.start);
_mm_prefetch(ilabel_ptr as *const i8, _MM_HINT_T0);
_mm_prefetch(weight_ptr as *const i8, _MM_HINT_T0);
_mm_prefetch(next_ptr as *const i8, _MM_HINT_T0);
}
#[cfg(target_arch = "aarch64")]
unsafe {
let ilabel_ptr = self.arc_ilabels.as_slice().as_ptr().add(range.start);
let weight_ptr = self.arc_weights.as_slice().as_ptr().add(range.start);
let next_ptr = self.arc_nextstates.as_slice().as_ptr().add(range.start);
core::arch::asm!(
"prfm pldl1keep, [{0}]",
in(reg) ilabel_ptr,
options(nostack, preserves_flags)
);
core::arch::asm!(
"prfm pldl1keep, [{0}]",
in(reg) weight_ptr,
options(nostack, preserves_flags)
);
core::arch::asm!(
"prfm pldl1keep, [{0}]",
in(reg) next_ptr,
options(nostack, preserves_flags)
);
}
}
}
}
pub fn memory_usage(&self) -> usize {
let state_offsets_size = self.state_offsets.len() * std::mem::size_of::<u32>();
let final_weights_size = self.final_weights.len() * std::mem::size_of::<Option<W>>();
let ilabels_size = self.arc_ilabels.len() * std::mem::size_of::<Label>();
let olabels_size = self.arc_olabels.len() * std::mem::size_of::<Label>();
let weights_size = self.arc_weights.len() * std::mem::size_of::<W>();
let nextstates_size = self.arc_nextstates.len() * std::mem::size_of::<StateId>();
state_offsets_size
+ final_weights_size
+ ilabels_size
+ olabels_size
+ weights_size
+ nextstates_size
+ std::mem::size_of::<Self>()
}
}
impl<W: Semiring> Default for CsrFst<W> {
fn default() -> Self {
Self::new()
}
}
impl<W: Semiring> std::fmt::Debug for CsrFst<W> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CsrFst")
.field("num_states", &self.num_states)
.field("total_arcs", &self.total_arcs())
.field("start", &self.start)
.field("memory_bytes", &self.memory_usage())
.finish()
}
}
#[derive(Debug)]
pub struct CsrArcIterator<'a, W: Semiring> {
ilabels: &'a [Label],
olabels: &'a [Label],
weights: &'a [W],
nextstates: &'a [StateId],
index: usize,
}
impl<'a, W: Semiring> Iterator for CsrArcIterator<'a, W> {
type Item = Arc<W>;
#[inline]
fn next(&mut self) -> Option<Self::Item> {
if self.index < self.ilabels.len() {
let arc = Arc::new(
self.ilabels[self.index],
self.olabels[self.index],
self.weights[self.index].clone(),
self.nextstates[self.index],
);
self.index += 1;
Some(arc)
} else {
None
}
}
#[inline]
fn size_hint(&self) -> (usize, Option<usize>) {
let remaining = self.ilabels.len() - self.index;
(remaining, Some(remaining))
}
}
impl<'a, W: Semiring> ExactSizeIterator for CsrArcIterator<'a, W> {}
impl<'a, W: Semiring> ArcIterator<W> for CsrArcIterator<'a, W> {
fn reset(&mut self) {
self.index = 0;
}
}
impl<W: Semiring> Fst<W> for CsrFst<W> {
type ArcIter<'a>
= CsrArcIterator<'a, W>
where
W: 'a;
fn start(&self) -> Option<StateId> {
self.start
}
fn final_weight(&self, state: StateId) -> Option<&W> {
self.final_weights
.get(state as usize)
.and_then(|opt| opt.as_ref())
}
fn num_states(&self) -> usize {
self.num_states
}
fn num_arcs(&self, state: StateId) -> usize {
if (state as usize) >= self.num_states {
return 0;
}
let range = self.arc_range(state);
range.end - range.start
}
fn arcs(&self, state: StateId) -> Self::ArcIter<'_> {
if (state as usize) >= self.num_states {
return CsrArcIterator {
ilabels: &[],
olabels: &[],
weights: &[],
nextstates: &[],
index: 0,
};
}
let range = self.arc_range(state);
CsrArcIterator {
ilabels: &self.arc_ilabels.as_slice()[range.clone()],
olabels: &self.arc_olabels.as_slice()[range.clone()],
weights: &self.arc_weights.as_slice()[range.clone()],
nextstates: &self.arc_nextstates.as_slice()[range],
index: 0,
}
}
fn properties(&self) -> crate::properties::FstProperties {
self.properties
}
}
#[cfg(target_arch = "x86_64")]
pub mod simd {
use super::*;
#[allow(dead_code)]
#[target_feature(enable = "avx2")]
pub unsafe fn find_arcs_by_ilabel_avx2(ilabels: &[Label], target_label: Label) -> Vec<usize> {
use std::arch::x86_64::*;
let mut matches = Vec::new();
let target = _mm256_set1_epi32(target_label as i32);
let len = ilabels.len();
let mut i = 0;
while i + 8 <= len {
let labels = _mm256_loadu_si256(ilabels.as_ptr().add(i) as *const __m256i);
let cmp = _mm256_cmpeq_epi32(labels, target);
let mask = _mm256_movemask_ps(_mm256_castsi256_ps(cmp)) as u32;
if mask != 0 {
for j in 0..8 {
if (mask >> j) & 1 != 0 {
matches.push(i + j);
}
}
}
i += 8;
}
while i < len {
if ilabels[i] == target_label {
matches.push(i);
}
i += 1;
}
matches
}
#[allow(dead_code)]
#[target_feature(enable = "avx2")]
pub unsafe fn min_weight_avx2(weights: &[f32]) -> f32 {
use std::arch::x86_64::*;
if weights.is_empty() {
return f32::INFINITY;
}
let len = weights.len();
let mut min_vec = _mm256_set1_ps(f32::INFINITY);
let mut i = 0;
while i + 8 <= len {
let w = _mm256_loadu_ps(weights.as_ptr().add(i));
min_vec = _mm256_min_ps(min_vec, w);
i += 8;
}
let low = _mm256_castps256_ps128(min_vec);
let high = _mm256_extractf128_ps(min_vec, 1);
let min128 = _mm_min_ps(low, high);
let min64 = _mm_min_ps(min128, _mm_movehl_ps(min128, min128));
let min32 = _mm_min_ss(min64, _mm_shuffle_ps(min64, min64, 1));
let mut result = _mm_cvtss_f32(min32);
while i < len {
result = result.min(weights[i]);
i += 1;
}
result
}
#[allow(dead_code)]
#[target_feature(enable = "avx2")]
pub unsafe fn sum_weights_avx2(weights: &[f32]) -> f32 {
use std::arch::x86_64::*;
if weights.is_empty() {
return 0.0;
}
let len = weights.len();
let mut sum_vec = _mm256_setzero_ps();
let mut i = 0;
while i + 8 <= len {
let w = _mm256_loadu_ps(weights.as_ptr().add(i));
sum_vec = _mm256_add_ps(sum_vec, w);
i += 8;
}
let low = _mm256_castps256_ps128(sum_vec);
let high = _mm256_extractf128_ps(sum_vec, 1);
let sum128 = _mm_add_ps(low, high);
let sum64 = _mm_add_ps(sum128, _mm_movehl_ps(sum128, sum128));
let sum32 = _mm_add_ss(sum64, _mm_shuffle_ps(sum64, sum64, 1));
let mut result = _mm_cvtss_f32(sum32);
while i < len {
result += weights[i];
i += 1;
}
result
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
#[test]
fn test_csr_fst_empty() {
let csr: CsrFst<TropicalWeight> = CsrFst::new();
assert_eq!(csr.num_states(), 0);
assert!(csr.start().is_none());
assert_eq!(csr.total_arcs(), 0);
}
#[test]
fn test_csr_fst_from_vector_fst() {
let mut vector_fst = VectorFst::<TropicalWeight>::new();
let s0 = vector_fst.add_state();
let s1 = vector_fst.add_state();
let s2 = vector_fst.add_state();
vector_fst.set_start(s0);
vector_fst.set_final(s2, TropicalWeight::new(0.5));
vector_fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(1.0), s1));
vector_fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::new(2.0), s2));
vector_fst.add_arc(s1, Arc::new(3, 3, TropicalWeight::new(1.5), s2));
let csr = CsrFst::from_fst(&vector_fst).unwrap();
assert_eq!(csr.num_states(), 3);
assert_eq!(csr.start(), Some(0));
assert_eq!(csr.total_arcs(), 3);
assert_eq!(csr.num_arcs(s0), 2);
assert_eq!(csr.num_arcs(s1), 1);
assert_eq!(csr.num_arcs(s2), 0);
assert_eq!(csr.final_weight(s2), Some(&TropicalWeight::new(0.5)));
}
#[test]
fn test_csr_fst_arc_iteration() {
let mut vector_fst = VectorFst::<TropicalWeight>::new();
let s0 = vector_fst.add_state();
let s1 = vector_fst.add_state();
vector_fst.set_start(s0);
vector_fst.set_final(s1, TropicalWeight::one());
vector_fst.add_arc(s0, Arc::new(1, 2, TropicalWeight::new(0.5), s1));
vector_fst.add_arc(s0, Arc::new(3, 4, TropicalWeight::new(1.5), s1));
let csr = CsrFst::from_fst(&vector_fst).unwrap();
let arcs: Vec<Arc<TropicalWeight>> = csr.arcs(s0).collect();
assert_eq!(arcs.len(), 2);
assert_eq!(arcs[0].ilabel, 1);
assert_eq!(arcs[0].olabel, 2);
assert_eq!(*arcs[0].weight.value(), 0.5);
assert_eq!(arcs[1].ilabel, 3);
assert_eq!(arcs[1].olabel, 4);
assert_eq!(*arcs[1].weight.value(), 1.5);
}
#[test]
fn test_csr_fst_soa_access() {
let mut vector_fst = VectorFst::<TropicalWeight>::new();
let s0 = vector_fst.add_state();
let s1 = vector_fst.add_state();
vector_fst.set_start(s0);
vector_fst.set_final(s1, TropicalWeight::one());
for i in 1..=5 {
vector_fst.add_arc(s0, Arc::new(i, i * 10, TropicalWeight::new(i as f32), s1));
}
let csr = CsrFst::from_fst(&vector_fst).unwrap();
let ilabels = csr.state_ilabels(s0);
let olabels = csr.state_olabels(s0);
let weights = csr.state_weights(s0);
let nextstates = csr.state_nextstates(s0);
assert_eq!(ilabels, &[1, 2, 3, 4, 5]);
assert_eq!(olabels, &[10, 20, 30, 40, 50]);
assert_eq!(nextstates, &[1, 1, 1, 1, 1]);
assert_eq!(weights.len(), 5);
}
#[test]
fn test_csr_fst_large() {
let mut vector_fst = VectorFst::<TropicalWeight>::new();
let num_states = 1000;
let arcs_per_state = 10;
let mut states = Vec::with_capacity(num_states);
for _ in 0..num_states {
states.push(vector_fst.add_state());
}
vector_fst.set_start(states[0]);
vector_fst.set_final(states[num_states - 1], TropicalWeight::one());
for i in 0..num_states - 1 {
for j in 0..arcs_per_state {
vector_fst.add_arc(
states[i],
Arc::new(
(j + 1) as u32,
(j + 1) as u32,
TropicalWeight::new(j as f32 * 0.1),
states[i + 1],
),
);
}
}
let csr = CsrFst::from_fst(&vector_fst).unwrap();
assert_eq!(csr.num_states(), num_states);
assert_eq!(csr.total_arcs(), (num_states - 1) * arcs_per_state);
for state in states.iter().take(num_states - 1) {
assert_eq!(csr.num_arcs(*state), arcs_per_state);
}
assert_eq!(csr.num_arcs(states[num_states - 1]), 0);
}
#[test]
fn test_csr_fst_memory_usage() {
let mut vector_fst = VectorFst::<TropicalWeight>::new();
let s0 = vector_fst.add_state();
let s1 = vector_fst.add_state();
vector_fst.set_start(s0);
vector_fst.set_final(s1, TropicalWeight::one());
vector_fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::one(), s1));
let csr = CsrFst::from_fst(&vector_fst).unwrap();
let mem = csr.memory_usage();
assert!(mem > 0);
assert!(mem < 10000); }
#[test]
fn test_csr_fst_exact_size_iterator() {
let mut vector_fst = VectorFst::<TropicalWeight>::new();
let s0 = vector_fst.add_state();
let s1 = vector_fst.add_state();
vector_fst.set_start(s0);
vector_fst.set_final(s1, TropicalWeight::one());
for i in 1..=10 {
vector_fst.add_arc(s0, Arc::new(i, i, TropicalWeight::new(i as f32), s1));
}
let csr = CsrFst::from_fst(&vector_fst).unwrap();
let iter = csr.arcs(s0);
assert_eq!(iter.len(), 10);
assert_eq!(iter.size_hint(), (10, Some(10)));
}
#[cfg(target_arch = "x86_64")]
#[test]
fn test_simd_find_arcs() {
if !is_x86_feature_detected!("avx2") {
return;
}
let ilabels: Vec<u32> = (1..=100).collect();
let target = 50;
let matches = unsafe { simd::find_arcs_by_ilabel_avx2(&ilabels, target) };
assert_eq!(matches.len(), 1);
assert_eq!(matches[0], 49); }
#[cfg(target_arch = "x86_64")]
#[test]
fn test_simd_min_weight() {
if !is_x86_feature_detected!("avx2") {
return;
}
let weights: Vec<f32> = vec![5.0, 3.0, 7.0, 1.0, 9.0, 2.0, 8.0, 4.0, 6.0, 0.5];
let min = unsafe { simd::min_weight_avx2(&weights) };
assert_eq!(min, 0.5);
}
#[cfg(target_arch = "x86_64")]
#[test]
fn test_simd_sum_weights() {
if !is_x86_feature_detected!("avx2") {
return;
}
let weights: Vec<f32> = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
let sum = unsafe { simd::sum_weights_avx2(&weights) };
assert!((sum - 55.0).abs() < 0.001);
}
}