1use 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
23pub struct GridStep {
59 search_space: SearchSpace,
60 current_param: usize,
61 current_step: usize,
62}
63
64impl GridStep {
65 #[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 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 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 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 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 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 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 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 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 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}