Skip to main content

zeph_experiments/
grid.rs

1// SPDX-FileCopyrightText: 2026 Andrei G <bug-ops>
2// SPDX-License-Identifier: MIT OR Apache-2.0
3
4//! Systematic grid sweep strategy for parameter variation.
5//!
6//! [`GridStep`] iterates each parameter through its discrete steps in order,
7//! skipping variations that have already been visited. This gives exhaustive
8//! coverage of the search space and is well-suited as a first-pass exploration
9//! before switching to a [`Neighborhood`] or [`Random`] strategy.
10//!
11//! [`Neighborhood`]: crate::Neighborhood
12//! [`Random`]: crate::Random
13
14use std::collections::HashSet;
15
16use ordered_float::OrderedFloat;
17
18use super::generator::VariationGenerator;
19use super::search_space::SearchSpace;
20use super::snapshot::ConfigSnapshot;
21use super::types::{Variation, VariationValue};
22
23/// Systematic grid sweep: iterate each parameter through its discrete steps, skip visited.
24///
25/// Parameters are swept one at a time. For each parameter, all grid points from
26/// `min` to `max` (with the configured `step`) are enumerated in order. Already-visited
27/// variations are skipped. When all steps for a parameter are exhausted, the next
28/// parameter is tried. Returns `None` when the full grid has been visited.
29///
30/// When a parameter has no discrete `step`, [`GridStep`] falls back to
31/// `(max - min) / 20` as the step size.
32///
33/// # Examples
34///
35/// ```rust
36/// use std::collections::HashSet;
37/// use zeph_experiments::{
38///     ConfigSnapshot, GridStep, ParameterKind, ParameterRange, SearchSpace, VariationGenerator,
39/// };
40///
41/// let space = SearchSpace {
42///     parameters: vec![
43///         ParameterRange::new(ParameterKind::Temperature, 0.0, 1.0, Some(0.5), 0.5).unwrap(),
44///     ],
45/// };
46/// let mut generator = GridStep::new(space);
47/// let baseline = ConfigSnapshot::default();
48/// let mut visited = HashSet::new();
49///
50/// // Produces 0.0, 0.5, 1.0 in order.
51/// let mut count = 0;
52/// while let Some(v) = generator.next(&baseline, &visited) {
53///     visited.insert(v);
54///     count += 1;
55/// }
56/// assert_eq!(count, 3);
57/// ```
58pub struct GridStep {
59    search_space: SearchSpace,
60    current_param: usize,
61    current_step: usize,
62}
63
64impl GridStep {
65    /// Create a new [`GridStep`] generator starting at the first grid point.
66    ///
67    /// # Examples
68    ///
69    /// ```rust
70    /// use zeph_experiments::{GridStep, SearchSpace, VariationGenerator};
71    ///
72    /// let generator = GridStep::new(SearchSpace::default());
73    /// assert_eq!(generator.name(), "grid");
74    /// ```
75    #[must_use]
76    pub fn new(search_space: SearchSpace) -> Self {
77        Self {
78            search_space,
79            current_param: 0,
80            current_step: 0,
81        }
82    }
83}
84
85impl VariationGenerator for GridStep {
86    fn next(
87        &mut self,
88        _baseline: &ConfigSnapshot,
89        visited: &HashSet<Variation>,
90    ) -> Option<Variation> {
91        while self.current_param < self.search_space.parameters.len() {
92            let range = &self.search_space.parameters[self.current_param];
93            let step = range.effective_step();
94            if step <= 0.0 {
95                self.current_param += 1;
96                self.current_step = 0;
97                continue;
98            }
99
100            #[allow(clippy::cast_precision_loss)]
101            let raw = range.min() + step * self.current_step as f64;
102
103            if raw > range.max() + f64::EPSILON {
104                self.current_param += 1;
105                self.current_step = 0;
106                continue;
107            }
108
109            self.current_step += 1;
110
111            // Quantize to avoid floating-point accumulation before deduplication.
112            let value = range.quantize(raw);
113
114            let variation = Variation {
115                parameter: range.kind(),
116                value: VariationValue::Float(OrderedFloat(value)),
117            };
118
119            if !visited.contains(&variation) {
120                return Some(variation);
121            }
122        }
123        None
124    }
125
126    fn name(&self) -> &'static str {
127        "grid"
128    }
129}
130
131#[cfg(test)]
132mod tests {
133    use std::collections::HashSet;
134
135    use super::super::search_space::ParameterRange;
136    use super::super::types::ParameterKind;
137    use super::*;
138
139    fn single_param_space(min: f64, max: f64, step: f64) -> SearchSpace {
140        // default = midpoint so it satisfies min <= default <= max
141        let default = f64::midpoint(min, max);
142        SearchSpace {
143            parameters: vec![
144                ParameterRange::new(ParameterKind::Temperature, min, max, Some(step), default)
145                    .unwrap(),
146            ],
147        }
148    }
149
150    #[test]
151    fn grid_step_produces_values_in_range() {
152        let mut generator = GridStep::new(single_param_space(0.0, 1.0, 0.5));
153        let baseline = ConfigSnapshot::default();
154        let mut visited = HashSet::new();
155        let mut values = vec![];
156        while let Some(v) = generator.next(&baseline, &visited) {
157            visited.insert(v.clone());
158            values.push(v.value.as_f64());
159        }
160        assert_eq!(values.len(), 3, "0.0, 0.5, 1.0");
161        for v in &values {
162            assert!(*v >= 0.0 && *v <= 1.0);
163        }
164    }
165
166    #[test]
167    fn grid_step_skips_visited() {
168        let mut generator = GridStep::new(single_param_space(0.0, 1.0, 0.5));
169        let baseline = ConfigSnapshot::default();
170        let mut visited = HashSet::new();
171        visited.insert(Variation {
172            parameter: ParameterKind::Temperature,
173            value: VariationValue::Float(OrderedFloat(0.0)),
174        });
175        let first = generator.next(&baseline, &visited).unwrap();
176        assert!(
177            (first.value.as_f64() - 0.5).abs() < 1e-10,
178            "expected 0.5, got {}",
179            first.value.as_f64()
180        );
181    }
182
183    #[test]
184    fn grid_step_returns_none_when_exhausted() {
185        // step=1.0 over [0.0, 0.5]: only one grid point (0.0), then exhausted.
186        let mut generator = GridStep::new(single_param_space(0.0, 0.5, 1.0));
187        let baseline = ConfigSnapshot::default();
188        let mut visited = HashSet::new();
189        // Only one point: 0.0
190        generator.next(&baseline, &visited).unwrap();
191        visited.insert(Variation {
192            parameter: ParameterKind::Temperature,
193            value: VariationValue::Float(OrderedFloat(0.0)),
194        });
195        assert!(generator.next(&baseline, &visited).is_none());
196    }
197
198    #[test]
199    fn grid_step_multiple_params() {
200        let space = SearchSpace {
201            parameters: vec![
202                ParameterRange::new(ParameterKind::Temperature, 0.0, 0.5, Some(0.5), 0.0).unwrap(),
203                ParameterRange::new(ParameterKind::TopP, 0.5, 1.0, Some(0.5), 0.5).unwrap(),
204            ],
205        };
206        let mut generator = GridStep::new(space);
207        let baseline = ConfigSnapshot::default();
208        let mut visited = HashSet::new();
209        let mut results = vec![];
210        while let Some(v) = generator.next(&baseline, &visited) {
211            visited.insert(v.clone());
212            results.push(v);
213        }
214        // Temperature: 0.0, 0.5 — TopP: 0.5, 1.0
215        assert_eq!(results.len(), 4);
216        let temp_count = results
217            .iter()
218            .filter(|v| v.parameter == ParameterKind::Temperature)
219            .count();
220        let top_p_count = results
221            .iter()
222            .filter(|v| v.parameter == ParameterKind::TopP)
223            .count();
224        assert_eq!(temp_count, 2);
225        assert_eq!(top_p_count, 2);
226    }
227
228    #[test]
229    fn grid_step_quantizes_to_avoid_fp_drift() {
230        // 0.1 * 7 via accumulation = 0.7000000000000001
231        // quantize must snap to 0.7
232        let mut generator = GridStep::new(single_param_space(0.0, 1.0, 0.1));
233        let baseline = ConfigSnapshot::default();
234        let mut visited = HashSet::new();
235        let mut values = vec![];
236        while let Some(v) = generator.next(&baseline, &visited) {
237            visited.insert(v.clone());
238            values.push(v.value.as_f64());
239        }
240        // All values should be clean multiples of 0.1
241        for v in &values {
242            let rounded = (v * 10.0).round() / 10.0;
243            assert!(
244                (v - rounded).abs() < 1e-10,
245                "value {v} is not a clean multiple of 0.1"
246            );
247        }
248    }
249
250    #[test]
251    fn grid_step_empty_space_returns_none() {
252        let mut generator = GridStep::new(SearchSpace { parameters: vec![] });
253        let baseline = ConfigSnapshot::default();
254        let visited = HashSet::new();
255        assert!(generator.next(&baseline, &visited).is_none());
256    }
257
258    #[test]
259    fn grid_step_none_step_uses_fallback() {
260        // Parameter with step=None — GridStep falls back to (max-min)/20.0 as step size.
261        let space = SearchSpace {
262            parameters: vec![
263                ParameterRange::new(ParameterKind::Temperature, 0.0, 1.0, None, 0.5).unwrap(),
264            ],
265        };
266        let mut generator = GridStep::new(space);
267        let baseline = ConfigSnapshot::default();
268        let mut visited = HashSet::new();
269        let mut count = 0;
270        while let Some(v) = generator.next(&baseline, &visited) {
271            visited.insert(v.clone());
272            count += 1;
273        }
274        // With step = 1.0/20.0, there should be 21 steps (0.0, 0.05, ..., 1.0)
275        assert_eq!(
276            count, 21,
277            "expected 21 steps for step=None with effective_step() falling back to 20 divisions"
278        );
279    }
280
281    #[test]
282    fn grid_step_name() {
283        let generator = GridStep::new(SearchSpace::default());
284        assert_eq!(generator.name(), "grid");
285    }
286}