Skip to main content

zeph_experiments/
random.rs

1// SPDX-FileCopyrightText: 2026 Andrei G <bug-ops>
2// SPDX-License-Identifier: MIT OR Apache-2.0
3
4//! Uniform random sampling strategy for parameter variation.
5//!
6//! [`Random`] selects a parameter uniformly at random on each call, then samples
7//! its value uniformly from `[min, max]`, quantizing to the nearest step. It
8//! provides broad coverage without systematic ordering, which can be useful when
9//! the search space is large and a full [`GridStep`] sweep is too expensive.
10//!
11//! [`GridStep`]: crate::GridStep
12
13use std::collections::HashSet;
14use std::sync::Mutex;
15
16use ordered_float::OrderedFloat;
17use rand::RngExt as _;
18use rand::SeedableRng as _;
19use rand::rngs::SmallRng;
20
21use super::generator::{MAX_RETRIES, VariationGenerator};
22use super::search_space::SearchSpace;
23use super::snapshot::ConfigSnapshot;
24use super::types::{Variation, VariationValue};
25
26/// Uniform random sampling within parameter bounds.
27///
28/// At each call, a parameter is chosen uniformly at random, then a value is
29/// sampled uniformly from its `[min, max]` range and quantized to the nearest
30/// step (if configured). The sample is rejected if it was already visited.
31/// Returns `None` after 1000 consecutive rejections (the space is considered
32/// effectively exhausted for this seed).
33///
34/// The internal RNG is wrapped in a [`Mutex`] so that `Random` implements [`Sync`],
35/// which is required by [`VariationGenerator`] to allow [`ExperimentEngine`] to be
36/// used in an async context. The experiment loop is sequential, so the mutex is
37/// never contended.
38///
39/// # Examples
40///
41/// ```rust
42/// use std::collections::HashSet;
43/// use zeph_experiments::{ConfigSnapshot, Random, SearchSpace, VariationGenerator};
44///
45/// let mut generator = Random::new(SearchSpace::default(), 42);
46/// let baseline = ConfigSnapshot::default();
47/// let visited = HashSet::new();
48///
49/// // Two generators with the same seed produce the same first variation.
50/// let mut gen2 = Random::new(SearchSpace::default(), 42);
51/// let v1 = generator.next(&baseline, &visited);
52/// let v2 = gen2.next(&baseline, &visited);
53/// assert_eq!(v1, v2);
54/// ```
55///
56/// [`ExperimentEngine`]: crate::ExperimentEngine
57pub struct Random {
58    search_space: SearchSpace,
59    rng: Mutex<SmallRng>,
60}
61
62impl Random {
63    /// Create a new [`Random`] generator with a deterministic seed.
64    ///
65    /// Generators with the same `seed` and `search_space` will produce identical
66    /// variation sequences, making experiments reproducible.
67    ///
68    /// # Examples
69    ///
70    /// ```rust
71    /// use zeph_experiments::{Random, SearchSpace, VariationGenerator};
72    ///
73    /// let generator = Random::new(SearchSpace::default(), 1234);
74    /// assert_eq!(generator.name(), "random");
75    /// ```
76    #[must_use]
77    pub fn new(search_space: SearchSpace, seed: u64) -> Self {
78        Self {
79            search_space,
80            rng: Mutex::new(SmallRng::seed_from_u64(seed)),
81        }
82    }
83}
84
85impl VariationGenerator for Random {
86    fn next(
87        &mut self,
88        _baseline: &ConfigSnapshot,
89        visited: &HashSet<Variation>,
90    ) -> Option<Variation> {
91        if self.search_space.parameters.is_empty() {
92            return None;
93        }
94        let mut rng = self.rng.lock().expect("rng mutex poisoned");
95        for _ in 0..MAX_RETRIES {
96            let idx = rng.random_range(0..self.search_space.parameters.len());
97            let range = &self.search_space.parameters[idx];
98            let raw: f64 = rng.random_range(range.min()..=range.max());
99            let value = range.quantize(raw);
100            let variation = Variation {
101                parameter: range.kind(),
102                value: VariationValue::Float(OrderedFloat(value)),
103            };
104            if !visited.contains(&variation) {
105                return Some(variation);
106            }
107        }
108        None
109    }
110
111    fn name(&self) -> &'static str {
112        "random"
113    }
114}
115
116#[cfg(test)]
117mod tests {
118    #![allow(clippy::manual_range_contains)]
119
120    use std::collections::HashSet;
121
122    use super::super::search_space::ParameterRange;
123    use super::super::types::ParameterKind;
124    use super::*;
125
126    #[test]
127    fn random_produces_values_in_range() {
128        let space = SearchSpace {
129            parameters: vec![
130                ParameterRange::new(ParameterKind::Temperature, 0.0, 1.0, Some(0.1), 0.5).unwrap(),
131            ],
132        };
133        let mut generator = Random::new(space, 42);
134        let baseline = ConfigSnapshot::default();
135        let visited = HashSet::new();
136        for _ in 0..20 {
137            if let Some(v) = generator.next(&baseline, &visited) {
138                let val = v.value.as_f64();
139                assert!((0.0..=1.0).contains(&val), "out of range: {val}");
140            }
141        }
142    }
143
144    #[test]
145    fn random_skips_visited() {
146        // Use a very small range with a large step so only one grid point exists (0.5).
147        let space = SearchSpace {
148            parameters: vec![
149                ParameterRange::new(ParameterKind::Temperature, 0.5, 0.6, Some(1.0), 0.55).unwrap(),
150            ],
151        };
152        let mut generator = Random::new(space, 0);
153        let baseline = ConfigSnapshot::default();
154        let mut visited = HashSet::new();
155        visited.insert(Variation {
156            parameter: ParameterKind::Temperature,
157            value: VariationValue::Float(OrderedFloat(0.5)),
158        });
159        // Only one grid point (0.5 clamped from min), after visiting it must return None.
160        let result = generator.next(&baseline, &visited);
161        assert!(
162            result.is_none(),
163            "expected None when only option is already visited"
164        );
165    }
166
167    #[test]
168    fn random_empty_space_returns_none() {
169        let mut generator = Random::new(SearchSpace { parameters: vec![] }, 0);
170        let baseline = ConfigSnapshot::default();
171        let visited = HashSet::new();
172        assert!(generator.next(&baseline, &visited).is_none());
173    }
174
175    #[test]
176    fn random_is_deterministic_with_same_seed() {
177        let space = SearchSpace::default();
178        let baseline = ConfigSnapshot::default();
179        let visited = HashSet::new();
180        let mut gen1 = Random::new(space.clone(), 123);
181        let mut gen2 = Random::new(space, 123);
182        let v1 = gen1.next(&baseline, &visited);
183        let v2 = gen2.next(&baseline, &visited);
184        assert_eq!(v1, v2, "same seed must produce same first variation");
185    }
186
187    #[test]
188    fn random_quantizes_sampled_values() {
189        let space = SearchSpace {
190            parameters: vec![
191                ParameterRange::new(ParameterKind::TopP, 0.1, 1.0, Some(0.05), 0.9).unwrap(),
192            ],
193        };
194        let mut generator = Random::new(space, 7);
195        let baseline = ConfigSnapshot::default();
196        let visited = HashSet::new();
197        for _ in 0..30 {
198            if let Some(v) = generator.next(&baseline, &visited) {
199                let val = v.value.as_f64();
200                // Quantized values must be on the 0.05-step grid anchored at min=0.1:
201                // i.e. (val - 0.1) / 0.05 must be an integer.
202                let steps = (val - 0.1) / 0.05;
203                assert!(
204                    (steps - steps.round()).abs() < 1e-10,
205                    "value {val} is not on the 0.05-step grid anchored at 0.1"
206                );
207            }
208        }
209    }
210
211    #[test]
212    fn random_name() {
213        let generator = Random::new(SearchSpace::default(), 0);
214        assert_eq!(generator.name(), "random");
215    }
216
217    #[test]
218    fn random_is_sync() {
219        fn assert_sync<T: Sync>() {}
220        assert_sync::<Random>();
221    }
222}