Skip to main content

zeph_experiments/
neighborhood.rs

1// SPDX-FileCopyrightText: 2026 Andrei G <bug-ops>
2// SPDX-License-Identifier: MIT OR Apache-2.0
3
4//! Neighborhood perturbation strategy for parameter variation.
5//!
6//! [`Neighborhood`] is a local-search strategy that generates variations by
7//! perturbing the current baseline value of a randomly chosen parameter by a
8//! small random amount proportional to the configured `radius`. It is most
9//! effective after a coarse [`GridStep`] sweep has identified a promising region.
10//!
11//! [`GridStep`]: crate::GridStep
12
13use 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
26/// Perturbation strategy that explores the neighborhood of the current baseline.
27///
28/// At each call, a parameter is chosen uniformly at random. The new value is
29/// computed as `baseline_value ± U(-radius, radius) * step`, then clamped and
30/// quantized to the nearest grid step. Useful after a [`GridStep`] sweep has
31/// narrowed the search to a promising region.
32///
33/// The generator is seeded deterministically via `seed`, making experiments
34/// reproducible. `radius` must be finite and positive (enforced in [`Neighborhood::new`]).
35///
36/// # Examples
37///
38/// ```rust
39/// use std::collections::HashSet;
40/// use zeph_experiments::{ConfigSnapshot, Neighborhood, SearchSpace, VariationGenerator};
41///
42/// let mut generator = Neighborhood::new(SearchSpace::default(), 1.0, 42).unwrap();
43/// let baseline = ConfigSnapshot::default();
44/// let visited = HashSet::new();
45///
46/// // Each call perturbs a random parameter by a small amount.
47/// if let Some(v) = generator.next(&baseline, &visited) {
48///     let val = v.value.as_f64();
49///     assert!(val.is_finite());
50/// }
51/// ```
52///
53/// [`GridStep`]: crate::GridStep
54pub struct Neighborhood {
55    search_space: SearchSpace,
56    radius: f64,
57    rng: SmallRng,
58}
59
60impl Neighborhood {
61    /// Create a new `Neighborhood` generator.
62    ///
63    /// # Errors
64    ///
65    /// Returns [`EvalError::InvalidRadius`] if `radius` is not finite and positive.
66    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            // Skip zero perturbations — they produce the baseline value, wasting an attempt.
94            if delta.abs() < f64::EPSILON {
95                continue;
96            }
97            let raw = current + delta;
98            let value = range.quantize(range.clamp(raw));
99            // Skip if the quantized value equals the baseline (no effective change).
100            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        // Narrow range [0.5, 0.6] with large step: perturbations always clamp to 0.5 or 0.6.
171        // After visiting both, must return None.
172        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        // Continuous parameter (step=None) — neighborhood must still produce values.
216        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        // effective_step() falls back to 20 divisions: 2.0/20.0 = 0.1; must get at least one result.
232        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                // Quantized values must be multiples of 0.05 anchored at min=0.1:
256                // i.e. (val - 0.1) / 0.05 must be an integer.
257                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        // Baseline temperature 0.7, radius 1.0, step 0.1 => perturbation in [-0.1, +0.1]
275        // All values should be in [0.6, 0.8] within [0.0, 2.0]
276        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(); // temperature = 0.7
279        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        // All values must be within ±1 step of 0.7 (i.e., ±0.1, so [0.6, 0.8])
293        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}