use super::traits::*;
use crate::arc::{Arc, ArcIterator};
use crate::properties::FstProperties;
use crate::semiring::Semiring;
use crate::Result;
use std::collections::HashMap;
use std::sync::{Arc as StdArc, RwLock};
pub struct LazyFstImpl<W: Semiring, F> {
compute_fn: F,
state_cache: RwLock<HashMap<StateId, LazyState<W>>>,
final_weight_cache: RwLock<HashMap<StateId, W>>,
cache_config: CacheConfig,
start_state: Option<StateId>,
properties: FstProperties,
estimated_states: usize,
}
impl<W: Semiring, F> std::fmt::Debug for LazyFstImpl<W, F> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("LazyFstImpl")
.field("cache_config", &self.cache_config)
.field("start_state", &self.start_state)
.field("properties", &self.properties)
.field("estimated_states", &self.estimated_states)
.finish()
}
}
#[derive(Debug, Clone)]
pub struct CacheConfig {
pub max_cached_states: usize,
pub memory_limit: Option<usize>,
pub eviction_policy: EvictionPolicy,
pub enable_memory_mapping: bool,
pub enable_prefetching: bool,
pub streaming_config: LazyStreamingConfig,
}
impl Default for CacheConfig {
fn default() -> Self {
Self {
max_cached_states: 10000,
memory_limit: Some(100_000_000), eviction_policy: EvictionPolicy::LRU,
enable_memory_mapping: false,
enable_prefetching: false,
streaming_config: LazyStreamingConfig::default(),
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum EvictionPolicy {
LRU,
LFU,
Random,
None,
}
#[derive(Debug, Clone)]
pub struct LazyStreamingConfig {
pub enable_streaming: bool,
pub stream_buffer_size: usize,
pub checkpoint_interval: usize,
pub memory_checkpoint_threshold: usize,
pub enable_state_compression: bool,
}
impl Default for LazyStreamingConfig {
fn default() -> Self {
Self {
enable_streaming: false,
stream_buffer_size: 1000,
checkpoint_interval: 10000,
memory_checkpoint_threshold: 100_000_000, enable_state_compression: true,
}
}
}
pub trait StateGenerator<W: Semiring>: Send + Sync + std::fmt::Debug {
fn generate_batch(
&self,
start_state: StateId,
batch_size: usize,
) -> Vec<(StateId, LazyState<W>)>;
fn has_more_states(&self, state: StateId) -> bool;
fn estimated_size(&self) -> Option<usize>;
fn reset(&self);
}
pub trait MemoryMappedProvider<W: Semiring>: Send + Sync {
fn load_state(&self, state: StateId) -> Option<LazyState<W>>;
fn store_state(
&self,
state: StateId,
lazy_state: &LazyState<W>,
) -> std::result::Result<(), Box<dyn std::error::Error>>;
fn contains_state(&self, state: StateId) -> bool;
fn state_count(&self) -> usize;
fn flush(&self) -> std::result::Result<(), Box<dyn std::error::Error>>;
}
#[derive(Debug, Clone)]
pub struct LazyState<W: Semiring> {
pub arcs: Vec<Arc<W>>,
#[allow(dead_code)]
pub final_weight: Option<W>,
}
impl<W: Semiring, F> LazyFstImpl<W, F>
where
F: Fn(StateId) -> Option<LazyState<W>> + Send + Sync,
{
pub fn new_with_config(
compute_fn: F,
estimated_states: usize,
cache_config: CacheConfig,
) -> Self {
Self {
compute_fn,
state_cache: RwLock::new(HashMap::new()),
final_weight_cache: RwLock::new(HashMap::new()),
cache_config,
start_state: Some(0), properties: FstProperties::default(),
estimated_states,
}
}
pub fn new(compute_fn: F, estimated_states: usize) -> Self {
Self::new_with_config(compute_fn, estimated_states, CacheConfig::default())
}
fn get_or_compute_state(&self, state: StateId) -> Option<LazyState<W>> {
{
let cache = self.state_cache.read().unwrap();
if let Some(cached_state) = cache.get(&state) {
return Some(cached_state.clone());
}
}
let computed = (self.compute_fn)(state)?;
{
let mut state_cache = self.state_cache.write().unwrap();
let mut final_weight_cache = self.final_weight_cache.write().unwrap();
state_cache.insert(state, computed.clone());
if let Some(ref final_weight) = computed.final_weight {
final_weight_cache.insert(state, final_weight.clone());
}
self.apply_eviction_policy(&mut state_cache, &mut final_weight_cache);
}
Some(computed)
}
fn apply_eviction_policy(
&self,
state_cache: &mut HashMap<StateId, LazyState<W>>,
final_weight_cache: &mut HashMap<StateId, W>,
) {
if self.cache_config.eviction_policy == EvictionPolicy::None {
return;
}
let needs_eviction = state_cache.len() > self.cache_config.max_cached_states
|| self.cache_config.memory_limit.is_some_and(|limit| {
let current_usage = self.estimate_cache_memory_usage(state_cache);
current_usage > limit
});
if !needs_eviction {
return;
}
if let Some(oldest_state) = state_cache.keys().next().copied() {
state_cache.remove(&oldest_state);
final_weight_cache.remove(&oldest_state);
}
}
fn estimate_cache_memory_usage(&self, state_cache: &HashMap<StateId, LazyState<W>>) -> usize {
state_cache
.values()
.map(|lazy_state| {
std::mem::size_of::<LazyState<W>>()
+ lazy_state.arcs.len() * std::mem::size_of::<Arc<W>>()
})
.sum()
}
pub fn final_weight_ref(&self, state: StateId) -> Option<StdArc<W>> {
self.get_or_compute_state(state)?;
let final_weight_cache = self.final_weight_cache.read().unwrap();
final_weight_cache
.get(&state)
.map(|w| StdArc::new(w.clone()))
}
pub fn final_weight_owned(&self, state: StateId) -> Option<W> {
self.get_or_compute_state(state)
.and_then(|s| s.final_weight)
}
pub fn update_cache_config(&mut self, new_config: CacheConfig) {
self.cache_config = new_config;
let mut state_cache = self.state_cache.write().unwrap();
let mut final_weight_cache = self.final_weight_cache.write().unwrap();
self.apply_eviction_policy(&mut state_cache, &mut final_weight_cache);
}
pub fn cache_stats(&self) -> CacheStats {
let state_cache = self.state_cache.read().unwrap();
let final_weight_cache = self.final_weight_cache.read().unwrap();
CacheStats {
cached_states: state_cache.len(),
cached_final_weights: final_weight_cache.len(),
memory_usage: self.estimate_cache_memory_usage(&state_cache),
estimated_total_states: self.estimated_states,
hit_rate: 0.0, }
}
pub fn clear_cache(&self) {
let mut state_cache = self.state_cache.write().unwrap();
let mut final_weight_cache = self.final_weight_cache.write().unwrap();
state_cache.clear();
final_weight_cache.clear();
}
pub fn set_start_state(&mut self, state: StateId) {
self.start_state = Some(state);
}
pub fn enable_streaming(&mut self, streaming_config: LazyStreamingConfig) {
self.cache_config.streaming_config = streaming_config;
self.cache_config.enable_memory_mapping = true;
if self.cache_config.streaming_config.enable_streaming {
self.cache_config.max_cached_states =
self.cache_config.streaming_config.stream_buffer_size;
}
}
pub fn new_streaming<G>(
generator: G,
streaming_config: LazyStreamingConfig,
cache_config: CacheConfig,
) -> StreamingLazyFst<W, G>
where
G: StateGenerator<W>,
{
StreamingLazyFst::new(generator, streaming_config, cache_config)
}
pub fn checkpoint(&self) -> std::result::Result<(), Box<dyn std::error::Error>> {
if !self.cache_config.streaming_config.enable_streaming {
return Ok(());
}
self.clear_cache();
Ok(())
}
pub fn estimated_memory_usage(&self) -> usize {
let state_cache = self.state_cache.read().unwrap();
self.estimate_cache_memory_usage(&state_cache)
}
pub fn needs_checkpoint(&self) -> bool {
if !self.cache_config.streaming_config.enable_streaming {
return false;
}
let current_memory = self.estimated_memory_usage();
current_memory
> self
.cache_config
.streaming_config
.memory_checkpoint_threshold
}
}
#[derive(Debug, Clone)]
pub struct CacheStats {
pub cached_states: usize,
pub cached_final_weights: usize,
pub memory_usage: usize,
pub estimated_total_states: usize,
pub hit_rate: f64,
}
#[derive(Debug)]
pub struct LazyArcIterator<W: Semiring> {
arcs: Vec<Arc<W>>,
pos: usize,
}
impl<W: Semiring> Iterator for LazyArcIterator<W> {
type Item = Arc<W>;
fn next(&mut self) -> Option<Self::Item> {
if self.pos < self.arcs.len() {
let arc = self.arcs[self.pos].clone();
self.pos += 1;
Some(arc)
} else {
None
}
}
}
impl<W: Semiring> ArcIterator<W> for LazyArcIterator<W> {
fn reset(&mut self) {
self.pos = 0;
}
}
impl<W: Semiring, F> Fst<W> for LazyFstImpl<W, F>
where
F: Fn(StateId) -> Option<LazyState<W>> + Send + Sync,
{
type ArcIter<'a>
= LazyArcIterator<W>
where
Self: 'a;
fn start(&self) -> Option<StateId> {
self.start_state
}
fn final_weight(&self, state: StateId) -> Option<&W> {
self.get_or_compute_state(state);
None
}
fn num_arcs(&self, state: StateId) -> usize {
self.get_or_compute_state(state)
.map(|s| s.arcs.len())
.unwrap_or(0)
}
fn num_states(&self) -> usize {
self.estimated_states
}
fn properties(&self) -> FstProperties {
self.properties
}
fn arcs(&self, state: StateId) -> Self::ArcIter<'_> {
let arcs = self
.get_or_compute_state(state)
.map(|s| s.arcs)
.unwrap_or_default();
LazyArcIterator { arcs, pos: 0 }
}
}
impl<W: Semiring, F> LazyFst<W> for LazyFstImpl<W, F>
where
F: Fn(StateId) -> Option<LazyState<W>> + Send + Sync,
{
fn expand(&self, state: StateId) -> Result<()> {
self.get_or_compute_state(state);
Ok(())
}
}
#[derive(Debug)]
pub struct StreamingLazyFst<W: Semiring, G: StateGenerator<W>> {
generator: G,
state_cache: RwLock<HashMap<StateId, LazyState<W>>>,
final_weight_cache: RwLock<HashMap<StateId, W>>,
streaming_config: LazyStreamingConfig,
#[allow(dead_code)]
cache_config: CacheConfig,
generation_offset: RwLock<StateId>,
generation_buffer: RwLock<Vec<(StateId, LazyState<W>)>>,
start_state: Option<StateId>,
properties: FstProperties,
}
impl<W: Semiring, G: StateGenerator<W>> StreamingLazyFst<W, G> {
pub fn new(
generator: G,
streaming_config: LazyStreamingConfig,
cache_config: CacheConfig,
) -> Self {
Self {
generator,
state_cache: RwLock::new(HashMap::new()),
final_weight_cache: RwLock::new(HashMap::new()),
streaming_config,
cache_config,
generation_offset: RwLock::new(0),
generation_buffer: RwLock::new(Vec::new()),
start_state: Some(0),
properties: FstProperties::default(),
}
}
fn generate_states_on_demand(&self, target_state: StateId) -> Option<LazyState<W>> {
let current_offset = *self.generation_offset.read().unwrap();
if target_state >= current_offset {
let batch_size = self.streaming_config.stream_buffer_size;
let new_states = self.generator.generate_batch(current_offset, batch_size);
{
let mut buffer = self.generation_buffer.write().unwrap();
buffer.extend(new_states.clone());
*self.generation_offset.write().unwrap() = current_offset + batch_size as StateId;
}
self.cache_generated_states(&new_states);
new_states
.into_iter()
.find(|(state, _)| *state == target_state)
.map(|(_, lazy_state)| lazy_state)
} else {
None
}
}
fn cache_generated_states(&self, states: &[(StateId, LazyState<W>)]) {
let mut state_cache = self.state_cache.write().unwrap();
let mut final_weight_cache = self.final_weight_cache.write().unwrap();
for (state_id, lazy_state) in states {
state_cache.insert(*state_id, lazy_state.clone());
if let Some(ref final_weight) = lazy_state.final_weight {
final_weight_cache.insert(*state_id, final_weight.clone());
}
}
self.apply_streaming_memory_management(&mut state_cache, &mut final_weight_cache);
}
fn apply_streaming_memory_management(
&self,
state_cache: &mut HashMap<StateId, LazyState<W>>,
final_weight_cache: &mut HashMap<StateId, W>,
) {
let current_memory = self.estimate_memory_usage(state_cache);
if current_memory > self.streaming_config.memory_checkpoint_threshold {
let checkpoint_interval = self.streaming_config.checkpoint_interval;
let current_offset = *self.generation_offset.read().unwrap();
if current_offset > checkpoint_interval as StateId {
let cutoff = current_offset - checkpoint_interval as StateId;
state_cache.retain(|&state, _| state >= cutoff);
final_weight_cache.retain(|&state, _| state >= cutoff);
}
}
}
fn estimate_memory_usage(&self, state_cache: &HashMap<StateId, LazyState<W>>) -> usize {
state_cache
.values()
.map(|lazy_state| {
std::mem::size_of::<LazyState<W>>()
+ lazy_state.arcs.len() * std::mem::size_of::<Arc<W>>()
})
.sum()
}
fn get_or_generate_state(&self, state: StateId) -> Option<LazyState<W>> {
{
let cache = self.state_cache.read().unwrap();
if let Some(cached_state) = cache.get(&state) {
return Some(cached_state.clone());
}
}
{
let buffer = self.generation_buffer.read().unwrap();
if let Some((_, lazy_state)) = buffer.iter().find(|(s, _)| *s == state) {
return Some(lazy_state.clone());
}
}
self.generate_states_on_demand(state)
}
pub fn reset_streaming(&self) {
self.generator.reset();
*self.generation_offset.write().unwrap() = 0;
self.generation_buffer.write().unwrap().clear();
self.state_cache.write().unwrap().clear();
self.final_weight_cache.write().unwrap().clear();
}
pub fn streaming_stats(&self) -> StreamingStats {
let state_cache = self.state_cache.read().unwrap();
let buffer = self.generation_buffer.read().unwrap();
let current_offset = *self.generation_offset.read().unwrap();
StreamingStats {
cached_states: state_cache.len(),
buffered_states: buffer.len(),
generation_offset: current_offset,
memory_usage: self.estimate_memory_usage(&state_cache),
estimated_total_states: self.generator.estimated_size(),
}
}
}
#[derive(Debug, Clone)]
pub struct StreamingStats {
pub cached_states: usize,
pub buffered_states: usize,
pub generation_offset: StateId,
pub memory_usage: usize,
pub estimated_total_states: Option<usize>,
}
impl<W: Semiring, G: StateGenerator<W>> Fst<W> for StreamingLazyFst<W, G> {
type ArcIter<'a>
= LazyArcIterator<W>
where
Self: 'a;
fn start(&self) -> Option<StateId> {
self.start_state
}
fn final_weight(&self, _state: StateId) -> Option<&W> {
None
}
fn num_arcs(&self, state: StateId) -> usize {
self.get_or_generate_state(state)
.map(|s| s.arcs.len())
.unwrap_or(0)
}
fn num_states(&self) -> usize {
self.generator.estimated_size().unwrap_or(usize::MAX)
}
fn properties(&self) -> FstProperties {
self.properties
}
fn arcs(&self, state: StateId) -> Self::ArcIter<'_> {
let arcs = self
.get_or_generate_state(state)
.map(|s| s.arcs)
.unwrap_or_default();
LazyArcIterator { arcs, pos: 0 }
}
}
#[derive(Debug)]
pub struct EvictingCacheFst<F: Fst<W>, W: Semiring> {
inner: F,
arc_cache: RwLock<HashMap<StateId, Vec<Arc<W>>>>,
final_weight_cache: RwLock<HashMap<StateId, Option<W>>>,
metadata_cache: RwLock<HashMap<StateId, StateMetadata>>,
cache_config: CacheConfig,
access_tracker: RwLock<AccessTracker>,
stats: RwLock<CachePerformanceStats>,
}
#[derive(Debug, Clone)]
struct StateMetadata {
num_arcs: usize,
#[allow(dead_code)]
is_computed: bool,
memory_size: usize,
}
#[derive(Debug)]
struct AccessTracker {
lru_timestamps: HashMap<StateId, u64>,
lfu_counts: HashMap<StateId, u64>,
access_counter: u64,
rng_state: u64,
}
impl AccessTracker {
fn new() -> Self {
Self {
lru_timestamps: HashMap::new(),
lfu_counts: HashMap::new(),
access_counter: 0,
rng_state: 0x9E3779B97F4A7C15, }
}
fn record_access(&mut self, state: StateId) {
self.access_counter += 1;
self.lru_timestamps.insert(state, self.access_counter);
*self.lfu_counts.entry(state).or_insert(0) += 1;
}
fn choose_eviction_candidate(
&mut self,
policy: &EvictionPolicy,
states: &[StateId],
) -> Option<StateId> {
if states.is_empty() {
return None;
}
match policy {
EvictionPolicy::LRU => states
.iter()
.min_by_key(|&&state| self.lru_timestamps.get(&state).unwrap_or(&0))
.copied(),
EvictionPolicy::LFU => states
.iter()
.min_by_key(|&&state| self.lfu_counts.get(&state).unwrap_or(&0))
.copied(),
EvictionPolicy::Random => {
self.rng_state = self.rng_state.wrapping_mul(0x9E3779B97F4A7C15);
let index = (self.rng_state as usize) % states.len();
Some(states[index])
}
EvictionPolicy::None => None,
}
}
fn remove_state(&mut self, state: StateId) {
self.lru_timestamps.remove(&state);
self.lfu_counts.remove(&state);
}
}
#[derive(Debug, Clone, Default)]
pub struct CachePerformanceStats {
pub total_accesses: u64,
pub cache_hits: u64,
pub cache_misses: u64,
pub evictions: u64,
}
impl CachePerformanceStats {
fn hit_rate(&self) -> f64 {
if self.total_accesses == 0 {
0.0
} else {
self.cache_hits as f64 / self.total_accesses as f64
}
}
}
impl<F: Fst<W>, W: Semiring> EvictingCacheFst<F, W> {
pub fn new(inner: F, cache_config: CacheConfig) -> Self {
Self {
inner,
arc_cache: RwLock::new(HashMap::new()),
final_weight_cache: RwLock::new(HashMap::new()),
metadata_cache: RwLock::new(HashMap::new()),
cache_config,
access_tracker: RwLock::new(AccessTracker::new()),
stats: RwLock::new(CachePerformanceStats::default()),
}
}
pub fn with_default_config(inner: F) -> Self {
Self::new(inner, CacheConfig::default())
}
fn get_cached_arcs(&self, state: StateId) -> Vec<Arc<W>> {
{
let mut tracker = self.access_tracker.write().unwrap();
tracker.record_access(state);
}
{
let cache = self.arc_cache.read().unwrap();
if let Some(arcs) = cache.get(&state) {
{
let mut stats = self.stats.write().unwrap();
stats.total_accesses += 1;
stats.cache_hits += 1;
}
return arcs.clone();
}
}
let arcs: Vec<Arc<W>> = self.inner.arcs(state).collect();
{
let mut stats = self.stats.write().unwrap();
stats.total_accesses += 1;
stats.cache_misses += 1;
}
{
let mut arc_cache = self.arc_cache.write().unwrap();
let mut metadata_cache = self.metadata_cache.write().unwrap();
arc_cache.insert(state, arcs.clone());
let memory_size = arcs.len() * std::mem::size_of::<Arc<W>>();
metadata_cache.insert(
state,
StateMetadata {
num_arcs: arcs.len(),
is_computed: true,
memory_size,
},
);
self.apply_eviction_policy(&mut arc_cache, &mut metadata_cache);
}
arcs
}
#[allow(dead_code)]
fn get_cached_final_weight(&self, state: StateId) -> Option<W> {
{
let mut tracker = self.access_tracker.write().unwrap();
tracker.record_access(state);
}
{
let cache = self.final_weight_cache.read().unwrap();
if let Some(weight_opt) = cache.get(&state) {
{
let mut stats = self.stats.write().unwrap();
stats.total_accesses += 1;
stats.cache_hits += 1;
}
return weight_opt.clone();
}
}
let weight = self.inner.final_weight(state).cloned();
{
let mut stats = self.stats.write().unwrap();
stats.total_accesses += 1;
stats.cache_misses += 1;
}
{
let mut cache = self.final_weight_cache.write().unwrap();
cache.insert(state, weight.clone());
}
weight
}
fn apply_eviction_policy(
&self,
arc_cache: &mut HashMap<StateId, Vec<Arc<W>>>,
metadata_cache: &mut HashMap<StateId, StateMetadata>,
) {
if self.cache_config.eviction_policy == EvictionPolicy::None {
return;
}
let needs_eviction = arc_cache.len() > self.cache_config.max_cached_states
|| self.cache_config.memory_limit.is_some_and(|limit| {
let current_usage: usize =
metadata_cache.values().map(|meta| meta.memory_size).sum();
current_usage > limit
});
if !needs_eviction {
return;
}
let states: Vec<StateId> = arc_cache.keys().copied().collect();
let evict_state = {
let mut tracker = self.access_tracker.write().unwrap();
tracker.choose_eviction_candidate(&self.cache_config.eviction_policy, &states)
};
if let Some(state) = evict_state {
arc_cache.remove(&state);
metadata_cache.remove(&state);
{
let mut final_weight_cache = self.final_weight_cache.write().unwrap();
final_weight_cache.remove(&state);
}
{
let mut tracker = self.access_tracker.write().unwrap();
tracker.remove_state(state);
}
{
let mut stats = self.stats.write().unwrap();
stats.evictions += 1;
}
}
}
pub fn cache_stats(&self) -> CacheStats {
let arc_cache = self.arc_cache.read().unwrap();
let final_weight_cache = self.final_weight_cache.read().unwrap();
let metadata_cache = self.metadata_cache.read().unwrap();
let stats = self.stats.read().unwrap();
let memory_usage: usize = metadata_cache.values().map(|meta| meta.memory_size).sum();
CacheStats {
cached_states: arc_cache.len(),
cached_final_weights: final_weight_cache.len(),
memory_usage,
estimated_total_states: self.inner.num_states(),
hit_rate: stats.hit_rate(),
}
}
pub fn clear_cache(&self) {
let mut arc_cache = self.arc_cache.write().unwrap();
let mut final_weight_cache = self.final_weight_cache.write().unwrap();
let mut metadata_cache = self.metadata_cache.write().unwrap();
let mut tracker = self.access_tracker.write().unwrap();
arc_cache.clear();
final_weight_cache.clear();
metadata_cache.clear();
*tracker = AccessTracker::new();
}
pub fn update_cache_config(&mut self, new_config: CacheConfig) {
self.cache_config = new_config;
let mut arc_cache = self.arc_cache.write().unwrap();
let mut metadata_cache = self.metadata_cache.write().unwrap();
self.apply_eviction_policy(&mut arc_cache, &mut metadata_cache);
}
pub fn inner(&self) -> &F {
&self.inner
}
pub fn performance_stats(&self) -> CachePerformanceStats {
self.stats.read().unwrap().clone()
}
}
#[derive(Debug)]
pub struct EvictingCacheArcIterator<W: Semiring> {
arcs: Vec<Arc<W>>,
pos: usize,
}
impl<W: Semiring> Iterator for EvictingCacheArcIterator<W> {
type Item = Arc<W>;
fn next(&mut self) -> Option<Self::Item> {
if self.pos < self.arcs.len() {
let arc = self.arcs[self.pos].clone();
self.pos += 1;
Some(arc)
} else {
None
}
}
}
impl<W: Semiring> ArcIterator<W> for EvictingCacheArcIterator<W> {
fn reset(&mut self) {
self.pos = 0;
}
}
impl<F: Fst<W>, W: Semiring> Fst<W> for EvictingCacheFst<F, W> {
type ArcIter<'a>
= EvictingCacheArcIterator<W>
where
Self: 'a;
fn start(&self) -> Option<StateId> {
self.inner.start()
}
fn final_weight(&self, state: StateId) -> Option<&W> {
self.inner.final_weight(state)
}
fn num_arcs(&self, state: StateId) -> usize {
{
let metadata_cache = self.metadata_cache.read().unwrap();
if let Some(metadata) = metadata_cache.get(&state) {
return metadata.num_arcs;
}
}
self.get_cached_arcs(state).len()
}
fn num_states(&self) -> usize {
self.inner.num_states()
}
fn properties(&self) -> FstProperties {
self.inner.properties()
}
fn arcs(&self, state: StateId) -> Self::ArcIter<'_> {
let arcs = self.get_cached_arcs(state);
EvictingCacheArcIterator { arcs, pos: 0 }
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
use num_traits::One;
#[test]
fn test_lazy_fst_creation() {
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::new(0.5), s1));
assert_eq!(fst.num_states(), 2);
assert_eq!(fst.start(), Some(s0));
assert!(fst.is_final(s1));
}
#[test]
fn test_lazy_fst_impl_new() {
let compute_fn = |state: StateId| -> Option<LazyState<TropicalWeight>> {
if state > 100 {
return None;
}
let arcs = if state == 0 {
vec![Arc::new(1, 1, TropicalWeight::one(), 1)]
} else {
vec![]
};
let final_weight = if state == 1 {
Some(TropicalWeight::one())
} else {
None
};
Some(LazyState { arcs, final_weight })
};
let lazy_fst = LazyFstImpl::new(compute_fn, 101);
assert_eq!(lazy_fst.start(), Some(0));
assert_eq!(lazy_fst.num_states(), 101);
assert_eq!(lazy_fst.num_arcs(0), 1);
assert_eq!(lazy_fst.num_arcs(1), 0);
assert!(lazy_fst.final_weight_ref(1).is_some());
assert!(lazy_fst.final_weight_ref(0).is_none());
}
#[test]
fn test_evicting_cache_fst_wrapper() {
let mut inner_fst = VectorFst::<TropicalWeight>::new();
let s0 = inner_fst.add_state();
let s1 = inner_fst.add_state();
let s2 = inner_fst.add_state();
inner_fst.set_start(s0);
inner_fst.set_final(s2, TropicalWeight::one());
inner_fst.add_arc(s0, Arc::new(1, 1, TropicalWeight::new(0.5), s1));
inner_fst.add_arc(s1, Arc::new(2, 2, TropicalWeight::new(1.0), s2));
let config = CacheConfig {
max_cached_states: 10,
memory_limit: Some(1024),
eviction_policy: EvictionPolicy::LRU,
..Default::default()
};
let cached_fst = EvictingCacheFst::new(inner_fst, config);
assert_eq!(cached_fst.start(), Some(s0));
assert_eq!(cached_fst.num_states(), 3);
assert_eq!(cached_fst.num_arcs(s0), 1);
assert_eq!(cached_fst.num_arcs(s1), 1);
assert_eq!(cached_fst.num_arcs(s2), 0);
let arcs1: Vec<_> = cached_fst.arcs(s0).collect();
let arcs2: Vec<_> = cached_fst.arcs(s0).collect();
assert_eq!(arcs1.len(), arcs2.len());
assert_eq!(arcs1[0].ilabel, arcs2[0].ilabel);
let stats = cached_fst.cache_stats();
assert!(stats.cached_states > 0);
assert!(stats.hit_rate >= 0.0 && stats.hit_rate <= 1.0);
}
#[test]
fn test_cache_eviction_policies() {
let mut inner_fst = VectorFst::<TropicalWeight>::new();
for i in 0..10 {
let state = inner_fst.add_state();
if i == 0 {
inner_fst.set_start(state);
}
if i < 9 {
inner_fst.add_arc(state, Arc::new(1, 1, TropicalWeight::one(), state + 1));
} else {
inner_fst.set_final(state, TropicalWeight::one());
}
}
let config = CacheConfig {
max_cached_states: 3, eviction_policy: EvictionPolicy::LRU,
..Default::default()
};
let cached_fst = EvictingCacheFst::new(inner_fst, config);
for i in 0..5 {
let _arcs: Vec<_> = cached_fst.arcs(i).collect();
}
let stats = cached_fst.cache_stats();
assert!(stats.cached_states <= 3); assert!(stats.hit_rate >= 0.0);
}
#[test]
fn test_cache_config() {
let config = CacheConfig::default();
assert_eq!(config.max_cached_states, 10000);
assert_eq!(config.eviction_policy, EvictionPolicy::LRU);
assert!(!config.enable_memory_mapping);
assert!(!config.enable_prefetching);
let custom_config = CacheConfig {
max_cached_states: 1000,
memory_limit: Some(5_000_000),
eviction_policy: EvictionPolicy::Random,
enable_memory_mapping: true,
enable_prefetching: true,
streaming_config: LazyStreamingConfig::default(),
};
assert_eq!(custom_config.max_cached_states, 1000);
assert_eq!(custom_config.memory_limit, Some(5_000_000));
assert_eq!(custom_config.eviction_policy, EvictionPolicy::Random);
assert!(custom_config.enable_memory_mapping);
assert!(custom_config.enable_prefetching);
}
#[test]
fn test_eviction_policy_enum() {
assert_eq!(EvictionPolicy::LRU, EvictionPolicy::LRU);
assert_ne!(EvictionPolicy::LRU, EvictionPolicy::LFU);
assert_ne!(EvictionPolicy::Random, EvictionPolicy::None);
}
#[test]
fn test_streaming_config() {
let config = LazyStreamingConfig::default();
assert!(!config.enable_streaming);
assert_eq!(config.stream_buffer_size, 1000);
assert_eq!(config.checkpoint_interval, 10000);
assert_eq!(config.memory_checkpoint_threshold, 100_000_000);
assert!(config.enable_state_compression);
let custom_config = LazyStreamingConfig {
enable_streaming: true,
stream_buffer_size: 500,
checkpoint_interval: 5000,
memory_checkpoint_threshold: 50_000_000,
enable_state_compression: false,
};
assert!(custom_config.enable_streaming);
assert_eq!(custom_config.stream_buffer_size, 500);
assert_eq!(custom_config.checkpoint_interval, 5000);
assert_eq!(custom_config.memory_checkpoint_threshold, 50_000_000);
assert!(!custom_config.enable_state_compression);
}
#[derive(Debug)]
struct TestGenerator {
max_state: StateId,
}
impl TestGenerator {
fn new(max_state: StateId) -> Self {
Self { max_state }
}
}
impl StateGenerator<TropicalWeight> for TestGenerator {
fn generate_batch(
&self,
start_state: StateId,
batch_size: usize,
) -> Vec<(StateId, LazyState<TropicalWeight>)> {
(start_state..std::cmp::min(start_state + batch_size as StateId, self.max_state + 1))
.map(|state| {
let arcs = if state < self.max_state {
vec![Arc::new(1, 1, TropicalWeight::one(), state + 1)]
} else {
vec![]
};
let final_weight = if state == self.max_state {
Some(TropicalWeight::one())
} else {
None
};
(state, LazyState { arcs, final_weight })
})
.collect()
}
fn has_more_states(&self, state: StateId) -> bool {
state <= self.max_state
}
fn estimated_size(&self) -> Option<usize> {
Some((self.max_state + 1) as usize)
}
fn reset(&self) {
}
}
#[test]
fn test_streaming_lazy_fst() {
let generator = TestGenerator::new(10);
let streaming_config = LazyStreamingConfig {
enable_streaming: true,
stream_buffer_size: 5,
checkpoint_interval: 8,
memory_checkpoint_threshold: 1000,
enable_state_compression: true,
};
let cache_config = CacheConfig::default();
let streaming_fst = StreamingLazyFst::new(generator, streaming_config, cache_config);
assert_eq!(streaming_fst.start(), Some(0));
assert_eq!(streaming_fst.num_states(), 11);
assert_eq!(streaming_fst.num_arcs(0), 1);
assert_eq!(streaming_fst.num_arcs(10), 0);
let arcs: Vec<_> = streaming_fst.arcs(0).collect();
assert_eq!(arcs.len(), 1);
assert_eq!(arcs[0].ilabel, 1);
assert_eq!(arcs[0].nextstate, 1);
let stats = streaming_fst.streaming_stats();
assert!(stats.cached_states > 0);
assert_eq!(stats.estimated_total_states, Some(11));
}
#[test]
fn test_state_generator_trait() {
let generator = TestGenerator::new(5);
assert!(generator.has_more_states(3));
assert!(generator.has_more_states(5));
assert!(!generator.has_more_states(6));
assert_eq!(generator.estimated_size(), Some(6));
let batch = generator.generate_batch(0, 3);
assert_eq!(batch.len(), 3);
assert_eq!(batch[0].0, 0);
assert_eq!(batch[1].0, 1);
assert_eq!(batch[2].0, 2);
let final_batch = generator.generate_batch(5, 2);
assert_eq!(final_batch.len(), 1);
assert_eq!(final_batch[0].0, 5);
assert!(final_batch[0].1.final_weight.is_some());
}
#[test]
fn test_lazy_fst_streaming_methods() {
let compute_fn = |state: StateId| -> Option<LazyState<TropicalWeight>> {
if state > 100 {
return None;
}
let arcs = if state == 0 {
vec![Arc::new(1, 1, TropicalWeight::one(), 1)]
} else {
vec![]
};
let final_weight = if state == 1 {
Some(TropicalWeight::one())
} else {
None
};
Some(LazyState { arcs, final_weight })
};
let mut lazy_fst = LazyFstImpl::new(compute_fn, 101);
let streaming_config = LazyStreamingConfig {
enable_streaming: true,
stream_buffer_size: 50,
checkpoint_interval: 100,
memory_checkpoint_threshold: 1_000_000,
enable_state_compression: true,
};
lazy_fst.enable_streaming(streaming_config);
let _memory_usage = lazy_fst.estimated_memory_usage();
assert!(!lazy_fst.needs_checkpoint()); assert!(lazy_fst.checkpoint().is_ok()); }
}