use arcweight::prelude::*;
use criterion::{criterion_group, criterion_main, Criterion};
use std::hint::black_box;
fn create_fst_with_unreachable(size: usize, unreachable_count: usize) -> VectorFst<TropicalWeight> {
let mut fst = VectorFst::new();
let mut states = Vec::new();
for _ in 0..size {
states.push(fst.add_state());
}
let unreachable_start = states.len();
for _ in 0..unreachable_count {
states.push(fst.add_state());
}
fst.set_start(states[0]);
fst.set_final(states[size - 1], TropicalWeight::one());
for i in 0..size - 1 {
fst.add_arc(
states[i],
Arc::new(
(i % 26 + 97) as u32,
(i % 26 + 97) as u32,
TropicalWeight::new(1.0),
states[i + 1],
),
);
}
for i in unreachable_start..states.len() {
if i < states.len() - 1 {
fst.add_arc(
states[i],
Arc::new(
((i - unreachable_start) % 26 + 97) as u32,
((i - unreachable_start) % 26 + 97) as u32,
TropicalWeight::new(1.0),
states[i + 1],
),
);
}
fst.add_arc(
states[i],
Arc::new(
((i - unreachable_start + 10) % 26 + 97) as u32,
((i - unreachable_start + 10) % 26 + 97) as u32,
TropicalWeight::new(0.5),
states[i],
),
);
}
fst
}
fn bench_connect_with_unreachable(c: &mut Criterion) {
let mut group = c.benchmark_group("connect_with_unreachable");
for size in [50, 200, 500, 1000].iter() {
let unreachable_count = size / 4;
let fst = create_fst_with_unreachable(*size, unreachable_count);
group.bench_function(format!("size_{size}_unreachable_{unreachable_count}"), |b| {
b.iter(|| {
let result: VectorFst<TropicalWeight> = connect(black_box(&fst)).unwrap();
black_box(result)
})
});
}
group.finish();
}
fn bench_connect_fully_connected(c: &mut Criterion) {
let mut group = c.benchmark_group("connect_fully_connected");
for size in [50, 200, 500, 1000].iter() {
let mut fst = VectorFst::new();
let states: Vec<_> = (0..*size).map(|_| fst.add_state()).collect();
fst.set_start(states[0]);
fst.set_final(states[size - 1], TropicalWeight::one());
for i in 0..size - 1 {
fst.add_arc(
states[i],
Arc::new(
(i % 26 + 97) as u32,
(i % 26 + 97) as u32,
TropicalWeight::new(1.0),
states[i + 1],
),
);
}
group.bench_function(format!("fully_connected_{size}"), |b| {
b.iter(|| {
let result: VectorFst<TropicalWeight> = connect(black_box(&fst)).unwrap();
black_box(result)
})
});
}
group.finish();
}
criterion_group!(
benches,
bench_connect_with_unreachable,
bench_connect_fully_connected
);
criterion_main!(benches);