Skip to main content

simsam/
lib.rs

1//! simsam — sample from custom discrete and continuous distributions.
2//!
3//! Build a distribution from a PDF or CDF (closures, histograms, location-scale transforms,
4//! or [simsym](https://docs.rs/simsym) symbolic expressions), then draw samples via inverse
5//! transform sampling — similar to [SciPy `rv_continuous`](https://docs.scipy.org/doc/scipy/reference/generated/scipy.stats.rv_continuous.html).
6//!
7//! ## Features
8//!
9//! - `pdf`, `cdf`, `ppf`, `sf`, `isf`, `logpdf`, `logcdf`, `logsf`
10//! - `mean`, `var`, `std`, `median`, `entropy`, `expect`, `interval`
11//! - Fast sampling via [`HermitePpfTable`](continuous::HermitePpfTable) ([`BuildOptions::with_hermite`])
12//!   or TDR ([`BuildOptions::with_tdr`])
13//! - [`LocScale`](continuous::LocScale), [`Truncated`](continuous::Truncated), [`HistogramPdf`](continuous::HistogramPdf)
14//!
15//! ## Limitations
16//!
17//! - Continuous distributions require **finite support** `[lo, hi]` (use [`Truncated`] on a wide interval).
18//! - Inverse transform assumes a **unimodal** CDF on that interval.
19//!
20//! ## Example
21//!
22//! ```
23//! use simsam::{from_pdf_fn, Interval};
24//!
25//! let support = Interval::new(0.0, 1.0).unwrap();
26//! let sampler = from_pdf_fn(|x| 3.0 * x * x, support).unwrap();
27//! let x = sampler.sample().unwrap();
28//! assert!((0.0..=1.0).contains(&x));
29//! assert!((sampler.mean().unwrap() - 0.75).abs() < 1e-2);
30//! ```
31
32mod continuous;
33mod discrete;
34mod error;
35mod multivar;
36mod sample;
37mod support;
38
39pub use continuous::{
40    from_cdf_fn, from_histogram, from_pdf_dpdf_fn, from_pdf_fn, from_pdf_fn_with_options,
41    from_pdf_loc_scale, AffineCdf,
42    BuildOptions, Cdf, CdfFn, CdfSource, ContinuousSampler, HasSupport, HermitePpfTable,
43    HistogramPdf, IntegratedPdf, InvertOptions, LocScale, Pdf, PdfFn, PpfMethod, SampleMethod,
44    Truncated, TdrBuildConfig, default_quad_tol, tdr_from_fns, Dpdf, DpdfFn, TdrHat, TdrOptions,
45    TdrSampler, TdrTransform,
46};
47#[cfg(feature = "symbolic")]
48pub use continuous::SymbolicPdfDpdf1d;
49#[cfg(feature = "symbolic")]
50pub use continuous::{SymbolicContinuous, SymbolicPdfAdapter};
51pub use discrete::{CdfDiscrete, DiscreteSampler, Pmf};
52pub use error::{BuildError, SampleError};
53pub use multivar::{HasSupportNd, HyperRect, PdfNd, PdfNdFn, RejectionOptions, RejectionSamplerNd};
54pub use multivar::{
55    CdfMcEstimator, CdfMcOptions, GaussianCopula, MhOptions, MetropolisHastingsNd,
56};
57pub use multivar::{
58    ConditionalFactorSampler, ConditionalFactorization, ConditionalSampler, GibbsOptions,
59    GibbsSamplerNd, HmcOptions, HmcSamplerNd, LogPdfNd,
60};
61#[cfg(feature = "symbolic")]
62pub use multivar::GradientLogPdfNd;
63#[cfg(feature = "symbolic")]
64pub use multivar::SymbolicPdfNd;
65pub use support::Interval;
66
67#[cfg(test)]
68mod tests {
69    use super::*;
70
71    #[test]
72    fn uniform_pdf_mean() {
73        let support = Interval::new(0.0, 1.0).unwrap();
74        let sampler = from_pdf_fn(|_| 1.0, support).unwrap();
75        let mut sum = 0.0;
76        let n = 4000;
77        for _ in 0..n {
78            sum += sampler.sample().unwrap();
79        }
80        let mean = sum / n as f64;
81        assert!((mean - 0.5).abs() < 0.05, "mean={mean}");
82        assert!((sampler.mean().unwrap() - 0.5).abs() < 1e-3);
83    }
84
85    #[test]
86    fn triangular_ppf() {
87        let support = Interval::new(0.0, 1.0).unwrap();
88        let sampler = from_pdf_fn(|x| 2.0 * x, support).unwrap();
89        let x = sampler.ppf(0.25).unwrap();
90        assert!((x - 0.5).abs() < 1e-3, "ppf(0.25)={x}");
91        assert!((sampler.median().unwrap() - 2.0_f64.sqrt().recip()).abs() < 0.02);
92    }
93
94    #[test]
95    fn cdf_only_quadratic() {
96        let support = Interval::new(0.0, 1.0).unwrap();
97        let sampler = from_cdf_fn(|x| x * x, support, BuildOptions::default()).unwrap();
98        let x = sampler.ppf(0.25).unwrap();
99        assert!((x - 0.5).abs() < 1e-3);
100        assert!((sampler.sf(0.5) - 0.75).abs() < 1e-3);
101    }
102
103    #[test]
104    fn discrete_bernoulli() {
105        let dist = DiscreteSampler::from_pmf(vec![0.0, 1.0], vec![0.3, 0.7]).unwrap();
106        assert!((dist.mean() - 0.7).abs() < 1e-10);
107        let mut ones = 0;
108        let n = 5000;
109        for _ in 0..n {
110            if dist.sample().unwrap() > 0.5 {
111                ones += 1;
112            }
113        }
114        let p = ones as f64 / n as f64;
115        assert!((p - 0.7).abs() < 0.05, "p={p}");
116    }
117
118    #[test]
119    #[cfg(feature = "symbolic")]
120    fn symbolic_triangular() {
121        use simsym::prelude::*;
122
123        let x = symbol("x");
124        let pdf = rational(2, 1) * x;
125        let support = Interval::new(0.0, 1.0).unwrap();
126        let sym = SymbolicContinuous::with_defaults(pdf, x, support).unwrap();
127        let sampler = sym.sampler(BuildOptions::default()).unwrap();
128        assert!((sampler.mean().unwrap() - 2.0 / 3.0).abs() < 1e-2);
129    }
130
131    #[test]
132    fn hermite_matches_bisection() {
133        let support = Interval::new(0.0, 1.0).unwrap();
134        let opts = BuildOptions::default().with_hermite(64);
135        let fast = from_pdf_fn_with_options(|x| 2.0 * x, support, opts).unwrap();
136        assert!(fast.uses_hermite_table());
137        let u = 0.37;
138        let mut slow = from_pdf_fn(|x| 2.0 * x, support).unwrap();
139        slow.clear_hermite_table();
140        let xh = fast.ppf(u).unwrap();
141        let xb = slow.ppf(u).unwrap();
142        assert!((xh - xb).abs() < 0.02);
143    }
144
145    #[test]
146    fn histogram_and_loc_scale() {
147        let edges = vec![0.0, 0.5, 1.0];
148        let counts = vec![1.0, 1.0];
149        let h = from_histogram(edges, counts, false, BuildOptions::default()).unwrap();
150        assert!((h.mean().unwrap() - 0.5).abs() < 1e-2);
151
152        let inner = Interval::new(0.0, 1.0).unwrap();
153        let base = from_pdf_fn(|_| 1.0, inner).unwrap();
154        let _ = base.mean().unwrap();
155
156        let scaled = from_pdf_loc_scale(PdfFn::new(|_| 1.0, inner), 10.0, 2.0, BuildOptions::default())
157            .unwrap();
158        let (lo, hi) = scaled.interval(0.9).unwrap();
159        assert!(lo > 9.0 && hi < 12.0);
160    }
161
162    #[test]
163    fn truncated_uniform() {
164        let inner = Interval::new(0.0, 1.0).unwrap();
165        let pdf = PdfFn::new(|_| 1.0, inner);
166        let trunc =
167            Truncated::new(pdf, Interval::new(0.25, 0.75).unwrap(), crate::default_quad_tol())
168                .unwrap();
169        let s = ContinuousSampler::from_pdf(trunc, BuildOptions::default()).unwrap();
170        assert!((s.mean().unwrap() - 0.5).abs() < 1e-2);
171    }
172
173    #[test]
174    fn rand_distribution_trait() {
175        use rand::distr::Distribution as RandDist;
176
177        let support = Interval::new(0.0, 1.0).unwrap();
178        let sampler = from_pdf_fn(|_| 1.0, support).unwrap();
179        let mut rng = rand::rng();
180        let x: f64 = RandDist::sample(&sampler, &mut rng);
181        assert!((0.0..=1.0).contains(&x));
182    }
183
184    #[test]
185    fn multivar_rejection_uniform_2d() {
186        let support = HyperRect::new(vec![0.0, 0.0], vec![1.0, 1.0]).unwrap();
187        let pdf = PdfNdFn::new(|_| 1.0, support);
188        let sampler = RejectionSamplerNd::new(pdf, 1.0, RejectionOptions::default()).unwrap();
189        let n = 2000;
190        let mut sx = 0.0;
191        let mut sy = 0.0;
192        for _ in 0..n {
193            let v = sampler.sample().unwrap();
194            sx += v[0];
195            sy += v[1];
196        }
197        let mx = sx / n as f64;
198        let my = sy / n as f64;
199        assert!((mx - 0.5).abs() < 0.06, "mx={mx}");
200        assert!((my - 0.5).abs() < 0.06, "my={my}");
201    }
202
203    #[test]
204    fn multivar_mh_uniform_2d_smoke() {
205        let support = HyperRect::new(vec![0.0, 0.0], vec![1.0, 1.0]).unwrap();
206        let log_pdf = PdfNdFn::new(|_| 1.0, support);
207        let mut mh = MetropolisHastingsNd::new(log_pdf, MhOptions::default()).unwrap();
208        mh.init().unwrap();
209        let samples = mh.sample_n(2000).unwrap();
210        let mut sx = 0.0;
211        let mut sy = 0.0;
212        for v in &samples {
213            sx += v[0];
214            sy += v[1];
215        }
216        let mx = sx / samples.len() as f64;
217        let my = sy / samples.len() as f64;
218        assert!((mx - 0.5).abs() < 0.07, "mx={mx}");
219        assert!((my - 0.5).abs() < 0.07, "my={my}");
220        assert!(mh.accept_rate() > 0.01, "accept_rate={}", mh.accept_rate());
221    }
222
223    #[test]
224    fn multivar_cdf_mc_uniform_2d() {
225        let support = HyperRect::new(vec![0.0, 0.0], vec![1.0, 1.0]).unwrap();
226        let pdf = PdfNdFn::new(|_| 1.0, support);
227        let mut est = CdfMcEstimator::new(
228            pdf,
229            CdfMcOptions {
230                normalization_samples: 20_000,
231                cdf_samples: 20_000,
232            },
233        )
234        .unwrap();
235        let p = est.cdf(&[0.25, 0.4]).unwrap();
236        assert!((p - 0.1).abs() < 0.05, "p={p}");
237    }
238
239    #[test]
240    fn multivar_copula_independent_uniforms() {
241        // Independent copula with uniform marginals on [0,1].
242        let corr = vec![vec![1.0, 0.0], vec![0.0, 1.0]];
243        let cop = GaussianCopula::new(corr).unwrap();
244
245        let s0 =
246            from_cdf_fn(|x| x, Interval::new(0.0, 1.0).unwrap(), BuildOptions::default()).unwrap();
247        let s1 =
248            from_cdf_fn(|x| x, Interval::new(0.0, 1.0).unwrap(), BuildOptions::default()).unwrap();
249
250        let mut rng = rand::rng();
251        let n = 3000;
252        let mut sx = 0.0;
253        let mut sy = 0.0;
254        let mut sxy = 0.0;
255        for _ in 0..n {
256            let v = cop
257                .sample_with_ppfs(&mut rng, &[&|u| s0.ppf(u), &|u| s1.ppf(u)])
258                .unwrap();
259            sx += v[0];
260            sy += v[1];
261            sxy += v[0] * v[1];
262        }
263        let mx = sx / n as f64;
264        let my = sy / n as f64;
265        let cov = sxy / n as f64 - mx * my;
266        assert!(cov.abs() < 0.03, "cov={cov}");
267    }
268
269    #[test]
270    #[cfg(feature = "symbolic")]
271    fn multivar_symbolic_pdf_smoke() {
272        use simsym::prelude::*;
273
274        let x = symbol("x");
275        let y = symbol("y");
276        let expr = simsym::expr::const_(rational(1, 1)) - (x.pow(2) + y.pow(2));
277        let support = HyperRect::new(vec![-1.0, -1.0], vec![1.0, 1.0]).unwrap();
278        let pdf = SymbolicPdfNd::new(expr, vec![x, y], support).unwrap();
279        let v = pdf.pdf(&[0.0, 0.0]);
280        assert!((v - 1.0).abs() < 1e-12);
281    }
282
283}