zeph_experiments/
random.rs1use 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
26pub struct Random {
58 search_space: SearchSpace,
59 rng: Mutex<SmallRng>,
60}
61
62impl Random {
63 #[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 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 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 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}