1use std::collections::HashSet;
14
15use ordered_float::OrderedFloat;
16use rand::RngExt as _;
17use rand::SeedableRng as _;
18use rand::rngs::SmallRng;
19
20use super::error::EvalError;
21use super::generator::{MAX_RETRIES, VariationGenerator};
22use super::search_space::SearchSpace;
23use super::snapshot::ConfigSnapshot;
24use super::types::{Variation, VariationValue};
25
26pub struct Neighborhood {
55 search_space: SearchSpace,
56 radius: f64,
57 rng: SmallRng,
58}
59
60impl Neighborhood {
61 pub fn new(search_space: SearchSpace, radius: f64, seed: u64) -> Result<Self, EvalError> {
67 if !radius.is_finite() || radius <= 0.0 {
68 return Err(EvalError::InvalidRadius { radius });
69 }
70 Ok(Self {
71 search_space,
72 radius,
73 rng: SmallRng::seed_from_u64(seed),
74 })
75 }
76}
77
78impl VariationGenerator for Neighborhood {
79 fn next(
80 &mut self,
81 baseline: &ConfigSnapshot,
82 visited: &HashSet<Variation>,
83 ) -> Option<Variation> {
84 if self.search_space.parameters.is_empty() {
85 return None;
86 }
87 for _ in 0..MAX_RETRIES {
88 let idx = self.rng.random_range(0..self.search_space.parameters.len());
89 let range = &self.search_space.parameters[idx];
90 let current = baseline.get(range.kind());
91 let step = range.effective_step();
92 let delta = self.rng.random_range(-self.radius..=self.radius) * step;
93 if delta.abs() < f64::EPSILON {
95 continue;
96 }
97 let raw = current + delta;
98 let value = range.quantize(range.clamp(raw));
99 if (value - current).abs() < f64::EPSILON {
101 continue;
102 }
103 let variation = Variation {
104 parameter: range.kind(),
105 value: VariationValue::Float(OrderedFloat(value)),
106 };
107 if !visited.contains(&variation) {
108 return Some(variation);
109 }
110 }
111 None
112 }
113
114 fn name(&self) -> &'static str {
115 "neighborhood"
116 }
117}
118
119#[cfg(test)]
120mod tests {
121 #![allow(
122 clippy::collapsible_if,
123 clippy::field_reassign_with_default,
124 clippy::manual_midpoint,
125 clippy::manual_range_contains
126 )]
127
128 use std::collections::HashSet;
129
130 use super::super::search_space::ParameterRange;
131 use super::super::types::ParameterKind;
132 use super::*;
133
134 fn make_space(kind: ParameterKind, min: f64, max: f64, step: f64) -> SearchSpace {
135 SearchSpace {
136 parameters: vec![
137 ParameterRange::new(kind, min, max, Some(step), f64::midpoint(min, max)).unwrap(),
138 ],
139 }
140 }
141
142 #[test]
143 fn neighborhood_produces_values_in_range() {
144 let space = make_space(ParameterKind::Temperature, 0.0, 2.0, 0.1);
145 let mut generator = Neighborhood::new(space, 1.0, 42).unwrap();
146 let baseline = ConfigSnapshot::default();
147 let visited = HashSet::new();
148 for _ in 0..20 {
149 if let Some(v) = generator.next(&baseline, &visited) {
150 let val = v.value.as_f64();
151 assert!((0.0..=2.0).contains(&val), "out of range: {val}");
152 }
153 }
154 }
155
156 #[test]
157 fn neighborhood_is_deterministic_with_same_seed() {
158 let space = SearchSpace::default();
159 let baseline = ConfigSnapshot::default();
160 let visited = HashSet::new();
161 let mut gen1 = Neighborhood::new(space.clone(), 1.0, 99).unwrap();
162 let mut gen2 = Neighborhood::new(space, 1.0, 99).unwrap();
163 let v1 = gen1.next(&baseline, &visited);
164 let v2 = gen2.next(&baseline, &visited);
165 assert_eq!(v1, v2, "same seed must produce same first variation");
166 }
167
168 #[test]
169 fn neighborhood_skips_visited() {
170 let space = make_space(ParameterKind::Temperature, 0.5, 0.6, 0.1);
173 let mut generator = Neighborhood::new(space, 1.0, 0).unwrap();
174 let baseline = ConfigSnapshot::default();
175 let mut visited = HashSet::new();
176 visited.insert(Variation {
177 parameter: ParameterKind::Temperature,
178 value: VariationValue::Float(OrderedFloat(0.5)),
179 });
180 visited.insert(Variation {
181 parameter: ParameterKind::Temperature,
182 value: VariationValue::Float(OrderedFloat(0.6)),
183 });
184 assert!(generator.next(&baseline, &visited).is_none());
185 }
186
187 #[test]
188 fn neighborhood_empty_space_returns_none() {
189 let mut generator = Neighborhood::new(SearchSpace { parameters: vec![] }, 1.0, 0).unwrap();
190 let baseline = ConfigSnapshot::default();
191 let visited = HashSet::new();
192 assert!(generator.next(&baseline, &visited).is_none());
193 }
194
195 #[test]
196 fn neighborhood_zero_radius_returns_error() {
197 let result = Neighborhood::new(SearchSpace::default(), 0.0, 0);
198 assert!(result.is_err(), "zero radius must be rejected");
199 }
200
201 #[test]
202 fn neighborhood_negative_radius_returns_error() {
203 let result = Neighborhood::new(SearchSpace::default(), -1.0, 0);
204 assert!(result.is_err(), "negative radius must be rejected");
205 }
206
207 #[test]
208 fn neighborhood_nan_radius_returns_error() {
209 let result = Neighborhood::new(SearchSpace::default(), f64::NAN, 0);
210 assert!(result.is_err(), "NaN radius must be rejected");
211 }
212
213 #[test]
214 fn neighborhood_step_none_uses_default_steps() {
215 let space = SearchSpace {
217 parameters: vec![
218 super::super::search_space::ParameterRange::new(
219 ParameterKind::Temperature,
220 0.0,
221 2.0,
222 None,
223 1.0,
224 )
225 .unwrap(),
226 ],
227 };
228 let mut generator = Neighborhood::new(space, 1.0, 77).unwrap();
229 let baseline = ConfigSnapshot::default();
230 let visited = HashSet::new();
231 let mut got_any = false;
233 for _ in 0..50 {
234 if generator.next(&baseline, &visited).is_some() {
235 got_any = true;
236 break;
237 }
238 }
239 assert!(
240 got_any,
241 "should produce at least one variation for continuous parameter"
242 );
243 }
244
245 #[test]
246 fn neighborhood_quantizes_perturbed_values() {
247 let space = make_space(ParameterKind::TopP, 0.1, 1.0, 0.05);
248 let mut generator = Neighborhood::new(space, 2.0, 11).unwrap();
249 let mut baseline = ConfigSnapshot::default();
250 baseline.top_p = 0.5;
251 let visited = HashSet::new();
252 for _ in 0..30 {
253 if let Some(v) = generator.next(&baseline, &visited) {
254 let val = v.value.as_f64();
255 let steps = (val - 0.1) / 0.05;
258 assert!(
259 (steps - steps.round()).abs() < 1e-10,
260 "value {val} is not on the 0.05-step grid anchored at 0.1"
261 );
262 }
263 }
264 }
265
266 #[test]
267 fn neighborhood_name() {
268 let generator = Neighborhood::new(SearchSpace::default(), 1.0, 0).unwrap();
269 assert_eq!(generator.name(), "neighborhood");
270 }
271
272 #[test]
273 fn neighborhood_perturbs_around_baseline() {
274 let space = make_space(ParameterKind::Temperature, 0.0, 2.0, 0.1);
277 let mut generator = Neighborhood::new(space, 1.0, 55).unwrap();
278 let baseline = ConfigSnapshot::default(); let visited = HashSet::new();
280 let mut temp_values = vec![];
281 for _ in 0..50 {
282 if let Some(v) = generator.next(&baseline, &visited)
283 && v.parameter == ParameterKind::Temperature
284 {
285 temp_values.push(v.value.as_f64());
286 }
287 }
288 assert!(
289 !temp_values.is_empty(),
290 "should produce temperature variations"
291 );
292 for val in &temp_values {
294 assert!(
295 *val >= 0.6 - 1e-10 && *val <= 0.8 + 1e-10,
296 "value {val} not within ±1 step of 0.7"
297 );
298 }
299 }
300}