use crate::fst::*;
use crate::semiring::Semiring;
use std::collections::HashMap;
#[derive(Debug, Clone)]
pub struct CacheMetadata {
pub avg_arcs_per_state: f64,
pub arc_count_distribution: Vec<(usize, usize)>,
pub cache_line_utilization: f64,
pub prefetch_distance: usize,
pub access_pattern: AccessPattern,
}
#[derive(Debug, Clone, PartialEq)]
pub enum AccessPattern {
Sequential,
Random,
Clustered,
Sparse,
}
impl CacheMetadata {
pub fn analyze<W: Semiring>(fst: &VectorFst<W>) -> Self {
let num_states = fst.num_states();
if num_states == 0 {
return Self::default();
}
let mut arc_counts = Vec::new();
let mut total_arcs = 0;
for state in 0..num_states as StateId {
let num_arcs = fst.num_arcs(state);
arc_counts.push(num_arcs);
total_arcs += num_arcs;
}
let avg_arcs_per_state = total_arcs as f64 / num_states as f64;
let mut distribution_map = HashMap::new();
for &count in &arc_counts {
*distribution_map.entry(count).or_insert(0) += 1;
}
let mut arc_count_distribution: Vec<_> = distribution_map.into_iter().collect();
arc_count_distribution.sort_by_key(|&(count, _)| count);
let cache_line_utilization = Self::estimate_cache_utilization(&arc_counts);
let prefetch_distance = Self::calculate_prefetch_distance(avg_arcs_per_state);
let access_pattern = Self::classify_access_pattern(&arc_counts);
Self {
avg_arcs_per_state,
arc_count_distribution,
cache_line_utilization,
prefetch_distance,
access_pattern,
}
}
pub fn optimal_prefetch_distance(&self) -> usize {
self.prefetch_distance
}
pub fn is_cache_friendly(&self) -> bool {
self.cache_line_utilization > 0.7
&& matches!(
self.access_pattern,
AccessPattern::Sequential | AccessPattern::Clustered
)
}
pub fn recommendations(&self) -> Vec<OptimizationRecommendation> {
let mut recommendations = Vec::new();
if self.cache_line_utilization < 0.5 {
recommendations.push(OptimizationRecommendation::ImproveDataLayout);
}
if self.avg_arcs_per_state > 10.0 {
recommendations.push(OptimizationRecommendation::UsePrefetching);
}
if matches!(
self.access_pattern,
AccessPattern::Random | AccessPattern::Sparse
) {
recommendations.push(OptimizationRecommendation::UseBlocking);
}
if self.avg_arcs_per_state < 2.0 {
recommendations.push(OptimizationRecommendation::ConsiderCompression);
}
recommendations
}
fn estimate_cache_utilization(arc_counts: &[usize]) -> f64 {
const CACHE_LINE_SIZE: usize = 64;
const ESTIMATED_ARC_SIZE: usize = 16; const ARCS_PER_CACHE_LINE: usize = CACHE_LINE_SIZE / ESTIMATED_ARC_SIZE;
let total_states = arc_counts.len();
let well_utilized_states = arc_counts
.iter()
.filter(|&&count| count >= ARCS_PER_CACHE_LINE / 2)
.count();
well_utilized_states as f64 / total_states as f64
}
fn calculate_prefetch_distance(avg_arcs: f64) -> usize {
if avg_arcs < 2.0 {
1
} else if avg_arcs < 5.0 {
2
} else if avg_arcs < 10.0 {
3
} else {
4
}
}
fn classify_access_pattern(arc_counts: &[usize]) -> AccessPattern {
if arc_counts.is_empty() {
return AccessPattern::Sequential;
}
let mean = arc_counts.iter().sum::<usize>() as f64 / arc_counts.len() as f64;
let variance = arc_counts
.iter()
.map(|&x| (x as f64 - mean).powi(2))
.sum::<f64>()
/ arc_counts.len() as f64;
let std_dev = variance.sqrt();
let coefficient_of_variation = std_dev / mean;
if coefficient_of_variation < 0.3 {
AccessPattern::Sequential
} else if coefficient_of_variation < 0.7 {
AccessPattern::Clustered
} else if mean < 2.0 {
AccessPattern::Sparse
} else {
AccessPattern::Random
}
}
}
impl Default for CacheMetadata {
fn default() -> Self {
Self {
avg_arcs_per_state: 0.0,
arc_count_distribution: Vec::new(),
cache_line_utilization: 0.0,
prefetch_distance: 1,
access_pattern: AccessPattern::Sequential,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum OptimizationRecommendation {
ImproveDataLayout,
UsePrefetching,
UseBlocking,
ConsiderCompression,
UseMemoryPooling,
}
pub mod layout_optimization {
use super::*;
use crate::arc::Arc;
use crate::semiring::Semiring;
pub fn reorder_for_locality<W: Semiring>(fst: &VectorFst<W>) -> VectorFst<W> {
use std::collections::{HashMap, VecDeque};
let num_states = fst.num_states();
if num_states == 0 {
return fst.clone();
}
let mut adjacency: HashMap<StateId, Vec<StateId>> = HashMap::new();
let mut in_degree: HashMap<StateId, usize> = HashMap::new();
for state in 0..num_states as StateId {
let neighbors: Vec<StateId> = fst.arcs(state).map(|arc| arc.nextstate).collect();
for &next in &neighbors {
*in_degree.entry(next).or_insert(0) += 1;
}
adjacency.insert(state, neighbors);
}
let mut new_order = Vec::new();
let mut visited = vec![false; num_states];
let mut queue = VecDeque::new();
if let Some(start) = fst.start() {
queue.push_back(start);
visited[start as usize] = true;
}
while let Some(state) = queue.pop_front() {
new_order.push(state);
let mut neighbors = adjacency.get(&state).cloned().unwrap_or_default();
neighbors.sort_by_key(|&s| std::cmp::Reverse(in_degree.get(&s).copied().unwrap_or(0)));
for next in neighbors {
if !visited[next as usize] {
visited[next as usize] = true;
queue.push_back(next);
}
}
}
for state in 0..num_states as StateId {
if !visited[state as usize] {
new_order.push(state);
}
}
let mut old_to_new: HashMap<StateId, StateId> = HashMap::new();
for (new_id, &old_id) in new_order.iter().enumerate() {
old_to_new.insert(old_id, new_id as StateId);
}
let mut new_fst = VectorFst::new();
for _ in 0..num_states {
new_fst.add_state();
}
if let Some(start) = fst.start() {
new_fst.set_start(*old_to_new.get(&start).unwrap());
}
for (old_state, &new_state) in &old_to_new {
for arc in fst.arcs(*old_state) {
let new_arc = Arc::new(
arc.ilabel,
arc.olabel,
arc.weight.clone(),
*old_to_new.get(&arc.nextstate).unwrap(),
);
new_fst.add_arc(new_state, new_arc);
}
if let Some(weight) = fst.final_weight(*old_state) {
new_fst.set_final(new_state, weight.clone());
}
}
new_fst
}
pub fn optimize_arc_layout<W: Semiring>(fst: &mut VectorFst<W>) {
use std::collections::HashMap;
let mut target_frequency: HashMap<StateId, usize> = HashMap::new();
for state in 0..fst.num_states() as StateId {
for arc in fst.arcs(state) {
*target_frequency.entry(arc.nextstate).or_insert(0) += 1;
}
}
for state in 0..fst.num_states() as StateId {
let mut arcs: Vec<_> = fst.arcs(state).collect();
arcs.sort_by(|a, b| {
const CACHE_GROUP_SIZE: StateId = 8; let a_group = a.nextstate / CACHE_GROUP_SIZE;
let b_group = b.nextstate / CACHE_GROUP_SIZE;
match a_group.cmp(&b_group) {
std::cmp::Ordering::Equal => {
let a_freq = target_frequency.get(&a.nextstate).copied().unwrap_or(0);
let b_freq = target_frequency.get(&b.nextstate).copied().unwrap_or(0);
match b_freq.cmp(&a_freq) {
std::cmp::Ordering::Equal => {
match a.ilabel.cmp(&b.ilabel) {
std::cmp::Ordering::Equal => a.olabel.cmp(&b.olabel),
other => other,
}
}
other => other,
}
}
other => other,
}
});
fst.delete_arcs(state);
for arc in arcs {
fst.add_arc(state, arc);
}
}
}
pub fn analyze_packing<T>(data: &[T]) -> PackingAnalysis {
let size_of_t = std::mem::size_of::<T>();
let cache_line_size = 64; let items_per_line = cache_line_size / size_of_t;
let waste_per_line = cache_line_size % size_of_t;
let total_lines = data.len().div_ceil(items_per_line);
let total_waste = total_lines * waste_per_line;
PackingAnalysis {
items_per_cache_line: items_per_line,
waste_bytes_per_line: waste_per_line,
total_cache_lines: total_lines,
total_waste_bytes: total_waste,
efficiency: 1.0 - (total_waste as f64 / (total_lines * cache_line_size) as f64),
}
}
}
#[derive(Debug, Clone)]
pub struct PackingAnalysis {
pub items_per_cache_line: usize,
pub waste_bytes_per_line: usize,
pub total_cache_lines: usize,
pub total_waste_bytes: usize,
pub efficiency: f64,
}
pub mod cache_aware_iteration {
use super::*;
use crate::arc::Arc;
use crate::semiring::Semiring;
#[derive(Debug)]
pub struct CacheAwareStateIterator {
states: Vec<StateId>,
current: usize,
}
impl CacheAwareStateIterator {
pub fn new(num_states: usize, metadata: &CacheMetadata) -> Self {
let mut states: Vec<StateId> = (0..num_states as StateId).collect();
match metadata.access_pattern {
AccessPattern::Sequential => {
}
AccessPattern::Random => {
Self::apply_cache_blocking(&mut states);
}
AccessPattern::Clustered => {
Self::apply_clustering(&mut states);
}
AccessPattern::Sparse => {
Self::apply_space_filling_order(&mut states);
}
}
Self { states, current: 0 }
}
fn apply_cache_blocking(states: &mut [StateId]) {
const BLOCK_SIZE: usize = 64;
for _chunk in states.chunks_mut(BLOCK_SIZE) {
}
}
fn apply_clustering(states: &mut [StateId]) {
states.sort();
}
fn apply_space_filling_order(states: &mut [StateId]) {
states.reverse();
}
}
impl Iterator for CacheAwareStateIterator {
type Item = StateId;
fn next(&mut self) -> Option<Self::Item> {
if self.current < self.states.len() {
let state = self.states[self.current];
self.current += 1;
Some(state)
} else {
None
}
}
}
pub fn process_cache_aware<W, F, R>(fst: &VectorFst<W>, mut processor: F) -> Vec<R>
where
W: Semiring,
F: FnMut(StateId, Vec<Arc<W>>) -> R,
{
let metadata = CacheMetadata::analyze(fst);
let state_iter = CacheAwareStateIterator::new(fst.num_states(), &metadata);
state_iter
.map(|state| {
let arcs: Vec<_> = fst.arcs(state).collect();
processor(state, arcs)
})
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::prelude::*;
#[test]
fn test_cache_metadata_analysis() {
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(1, 1, TropicalWeight::new(0.5), s1));
fst.add_arc(s0, Arc::new(2, 2, TropicalWeight::new(1.0), s1));
fst.add_arc(s1, Arc::new(3, 3, TropicalWeight::new(1.5), s2));
let metadata = CacheMetadata::analyze(&fst);
assert!(metadata.avg_arcs_per_state > 0.0);
assert!(!metadata.arc_count_distribution.is_empty());
assert!(metadata.prefetch_distance > 0);
}
#[test]
fn test_optimization_recommendations() {
let mut fst = VectorFst::<TropicalWeight>::new();
let s0 = fst.add_state();
fst.set_start(s0);
let metadata = CacheMetadata::analyze(&fst);
let recommendations = metadata.recommendations();
assert!(!recommendations.is_empty());
}
#[test]
fn test_access_pattern_classification() {
let uniform_counts = vec![2, 2, 2, 2, 2];
let pattern = CacheMetadata::classify_access_pattern(&uniform_counts);
assert_eq!(pattern, AccessPattern::Sequential);
let sparse_counts = vec![0, 1, 0, 1, 0];
let pattern = CacheMetadata::classify_access_pattern(&sparse_counts);
assert_eq!(pattern, AccessPattern::Sparse);
}
#[test]
fn test_cache_aware_iteration() {
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(0.5), s1));
let results =
cache_aware_iteration::process_cache_aware(&fst, |state, arcs| (state, arcs.len()));
assert_eq!(results.len(), 2);
assert!(results.iter().any(|&(state, _)| state == s0));
assert!(results.iter().any(|&(state, _)| state == s1));
}
#[test]
fn test_packing_analysis() {
let data = vec![1u32; 100];
let analysis = layout_optimization::analyze_packing(&data);
assert!(analysis.items_per_cache_line > 0);
assert!(analysis.efficiency >= 0.0 && analysis.efficiency <= 1.0);
}
}