diff --git a/Cargo.toml b/Cargo.toml index 7d55127..25580ee 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -65,3 +65,7 @@ harness = false [[bench]] name = "f_perf" harness = false + +[[bench]] +name = "sequence_vec" +harness = false diff --git a/benches/sequence_vec.rs b/benches/sequence_vec.rs new file mode 100644 index 0000000..268f9b2 --- /dev/null +++ b/benches/sequence_vec.rs @@ -0,0 +1,29 @@ +//! Benchmark for large sequences of simple observations. + +use criterion::{black_box, criterion_group, criterion_main, Criterion}; +use fugue::*; +use rand::rngs::StdRng; +use rand::SeedableRng; + +fn benchmark_stack_overflow_model(c: &mut Criterion) { + let n = 1_000usize; + + c.bench_function("stack_overflow_model/sequence_vec", |b| { + b.iter(|| { + let models: Vec> = (0..n) + .map(|i| sample(addr!("coin", i), Bernoulli::new(0.5).unwrap())) + .collect(); + let model = sequence_vec(models); + let mut rng = StdRng::seed_from_u64(42); + let handler = runtime::interpreters::PriorHandler { + rng: &mut rng, + trace: runtime::trace::Trace::default(), + }; + let (result, _trace) = runtime::handler::run(handler, model); + black_box(result.len()); + }) + }); +} + +criterion_group!(stack_overflow_benches, benchmark_stack_overflow_model); +criterion_main!(stack_overflow_benches); diff --git a/src/runtime/handler.rs b/src/runtime/handler.rs index 252fba9..611d2a9 100644 --- a/src/runtime/handler.rs +++ b/src/runtime/handler.rs @@ -287,4 +287,30 @@ mod tests { .join() .expect("deep model interpretation overflowed the stack"); } + + // Companion regression from PR #34: the same stack-safety guarantee through + // the `sequence_vec` + `observe` path (the deep test above covers + // sample+bind). 100k observations interpreted through `sequence_vec` must + // complete without overflowing the stack. + #[test] + fn run_handles_large_observe_sequence() { + use crate::core::model::{observe, sequence_vec}; + + let n = 100_000usize; + let models: Vec> = (0..n) + .map(|i| observe(addr!("obs", i), Bernoulli::new(0.5).unwrap(), true)) + .collect(); + let model = sequence_vec(models).map(|_| ()); + + let mut rng = StdRng::seed_from_u64(123); + let (_a, trace) = crate::runtime::handler::run( + PriorHandler { + rng: &mut rng, + trace: Trace::default(), + }, + model, + ); + + assert!(trace.log_likelihood.is_finite()); + } }