Skip to main content

sklears_svm/hyperparameter_optimization/
grid_search.rs

1//! Grid Search Cross-Validation for hyperparameter optimization
2
3use std::time::Instant;
4
5#[cfg(feature = "parallel")]
6use rayon::prelude::*;
7use scirs2_core::ndarray::{Array1, Array2};
8use scirs2_core::random::Random;
9
10use crate::kernels::KernelType;
11use crate::svc::SVC;
12use sklears_core::error::{Result, SklearsError};
13use sklears_core::traits::{Fit, Predict};
14
15use super::{
16    OptimizationConfig, OptimizationResult, ParameterSet, ParameterSpec, ScoringMetric, SearchSpace,
17};
18
19/// Grid Search hyperparameter optimizer
20pub struct GridSearchCV {
21    config: OptimizationConfig,
22    search_space: SearchSpace,
23    #[allow(dead_code)] // intentionally deferred: randomized grid search not yet called
24    rng: Random<scirs2_core::random::rngs::StdRng>,
25}
26
27impl GridSearchCV {
28    /// Create a new grid search optimizer
29    pub fn new(config: OptimizationConfig, search_space: SearchSpace) -> Self {
30        let rng = if let Some(seed) = config.random_state {
31            Random::seed(seed)
32        } else {
33            Random::seed(42) // Default seed for reproducibility
34        };
35
36        Self {
37            config,
38            search_space,
39            rng,
40        }
41    }
42
43    /// Run grid search optimization
44    pub fn fit(&mut self, x: &Array2<f64>, y: &Array1<f64>) -> Result<OptimizationResult> {
45        let start_time = Instant::now();
46
47        // Generate parameter grid
48        let param_grid = self.generate_parameter_grid()?;
49
50        if self.config.verbose {
51            println!(
52                "Grid search with {} parameter combinations",
53                param_grid.len()
54            );
55        }
56
57        // Evaluate all parameter combinations
58        let cv_results: Vec<(ParameterSet, f64)> = {
59            #[cfg(feature = "parallel")]
60            if self.config.n_jobs.is_some() {
61                // Parallel evaluation
62                param_grid
63                    .into_par_iter()
64                    .map(|params| {
65                        let score = self
66                            .evaluate_params(&params, x, y)
67                            .unwrap_or(-f64::INFINITY);
68                        (params, score)
69                    })
70                    .collect()
71            } else {
72                // Sequential evaluation
73                param_grid
74                    .into_iter()
75                    .map(|params| {
76                        let score = self
77                            .evaluate_params(&params, x, y)
78                            .unwrap_or(-f64::INFINITY);
79                        if self.config.verbose {
80                            println!("Params: {:?}, Score: {:.6}", params, score);
81                        }
82                        (params, score)
83                    })
84                    .collect()
85            }
86
87            #[cfg(not(feature = "parallel"))]
88            {
89                // Sequential evaluation (parallel feature disabled)
90                param_grid
91                    .into_iter()
92                    .map(|params| {
93                        let score = self
94                            .evaluate_params(&params, x, y)
95                            .unwrap_or(-f64::INFINITY);
96                        if self.config.verbose {
97                            println!("Params: {:?}, Score: {:.6}", params, score);
98                        }
99                        (params, score)
100                    })
101                    .collect()
102            }
103        };
104
105        // Find best parameters
106        let (best_params, best_score) = cv_results
107            .iter()
108            .max_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal))
109            .map(|(p, s)| (p.clone(), *s))
110            .ok_or_else(|| {
111                SklearsError::Other("No valid parameter combinations found".to_string())
112            })?;
113
114        let score_history: Vec<f64> = cv_results.iter().map(|(_, score)| *score).collect();
115        let n_iterations = cv_results.len();
116
117        Ok(OptimizationResult {
118            best_params,
119            best_score,
120            cv_results,
121            n_iterations,
122            optimization_time: start_time.elapsed().as_secs_f64(),
123            score_history,
124        })
125    }
126
127    /// Generate parameter grid for grid search
128    fn generate_parameter_grid(&mut self) -> Result<Vec<ParameterSet>> {
129        let mut param_grid = Vec::new();
130
131        // Clone specs to avoid borrowing conflicts
132        let c_spec = self.search_space.c.clone();
133        let kernel_spec = self.search_space.kernel.clone();
134        let tol_spec = self.search_space.tol.clone();
135        let max_iter_spec = self.search_space.max_iter.clone();
136
137        // Generate C values
138        let c_values = self.generate_values(&c_spec, 10)?;
139
140        // Generate kernel values
141        let kernel_values = if let Some(kernel_spec) = kernel_spec {
142            self.generate_kernel_values(&kernel_spec)?
143        } else {
144            vec![KernelType::Rbf { gamma: 1.0 }]
145        };
146
147        // Generate tolerance values
148        let tol_values = if let Some(tol_spec) = tol_spec {
149            self.generate_values(&tol_spec, 5)?
150        } else {
151            vec![1e-3]
152        };
153
154        // Generate max_iter values
155        let max_iter_values = if let Some(max_iter_spec) = max_iter_spec {
156            self.generate_values(&max_iter_spec, 3)?
157                .into_iter()
158                .map(|v| v as usize)
159                .collect()
160        } else {
161            vec![1000]
162        };
163
164        // Generate all combinations
165        for &c in &c_values {
166            for kernel in &kernel_values {
167                for &tol in &tol_values {
168                    for &max_iter in &max_iter_values {
169                        param_grid.push(ParameterSet {
170                            c,
171                            kernel: kernel.clone(),
172                            tol,
173                            max_iter,
174                        });
175                    }
176                }
177            }
178        }
179
180        Ok(param_grid)
181    }
182
183    /// Generate values from parameter specification
184    fn generate_values(&mut self, spec: &ParameterSpec, n_values: usize) -> Result<Vec<f64>> {
185        match spec {
186            ParameterSpec::Fixed(value) => Ok(vec![*value]),
187            ParameterSpec::Uniform { min, max } => Ok((0..n_values)
188                .map(|i| min + (max - min) * i as f64 / (n_values - 1) as f64)
189                .collect()),
190            ParameterSpec::LogUniform { min, max } => {
191                let log_min = min.ln();
192                let log_max = max.ln();
193                Ok((0..n_values)
194                    .map(|i| {
195                        let log_val =
196                            log_min + (log_max - log_min) * i as f64 / (n_values - 1) as f64;
197                        log_val.exp()
198                    })
199                    .collect())
200            }
201            ParameterSpec::Choice(choices) => Ok(choices.clone()),
202            ParameterSpec::KernelChoice(_) => Err(SklearsError::InvalidInput(
203                "Use generate_kernel_values for kernel specs".to_string(),
204            )),
205        }
206    }
207
208    /// Generate kernel values from kernel specification
209    fn generate_kernel_values(&mut self, spec: &ParameterSpec) -> Result<Vec<KernelType>> {
210        match spec {
211            ParameterSpec::KernelChoice(kernels) => Ok(kernels.clone()),
212            _ => Err(SklearsError::InvalidInput(
213                "Invalid kernel specification".to_string(),
214            )),
215        }
216    }
217
218    /// Evaluate parameter set using cross-validation
219    fn evaluate_params(
220        &self,
221        params: &ParameterSet,
222        x: &Array2<f64>,
223        y: &Array1<f64>,
224    ) -> Result<f64> {
225        let scores = self.cross_validate(params, x, y)?;
226        Ok(scores.iter().sum::<f64>() / scores.len() as f64)
227    }
228
229    /// Perform cross-validation
230    fn cross_validate(
231        &self,
232        params: &ParameterSet,
233        x: &Array2<f64>,
234        y: &Array1<f64>,
235    ) -> Result<Vec<f64>> {
236        let n_samples = x.nrows();
237        let fold_size = n_samples / self.config.cv_folds;
238        let mut scores = Vec::new();
239
240        for fold in 0..self.config.cv_folds {
241            let start_idx = fold * fold_size;
242            let end_idx = if fold == self.config.cv_folds - 1 {
243                n_samples
244            } else {
245                (fold + 1) * fold_size
246            };
247
248            // Create train/test splits
249            let mut x_train_data = Vec::new();
250            let mut y_train_vals = Vec::new();
251            let mut x_test_data = Vec::new();
252            let mut y_test_vals = Vec::new();
253
254            for i in 0..n_samples {
255                if i >= start_idx && i < end_idx {
256                    // Test set
257                    for j in 0..x.ncols() {
258                        x_test_data.push(x[[i, j]]);
259                    }
260                    y_test_vals.push(y[i]);
261                } else {
262                    // Training set
263                    for j in 0..x.ncols() {
264                        x_train_data.push(x[[i, j]]);
265                    }
266                    y_train_vals.push(y[i]);
267                }
268            }
269
270            let n_train = y_train_vals.len();
271            let n_test = y_test_vals.len();
272            let n_features = x.ncols();
273
274            let x_train = Array2::from_shape_vec((n_train, n_features), x_train_data)?;
275            let y_train = Array1::from_vec(y_train_vals);
276            let x_test = Array2::from_shape_vec((n_test, n_features), x_test_data)?;
277            let y_test = Array1::from_vec(y_test_vals);
278
279            // Train and evaluate model
280            let svm = SVC::new()
281                .c(params.c)
282                .kernel(params.kernel.clone())
283                .tol(params.tol)
284                .max_iter(params.max_iter);
285
286            let fitted_svm = svm.fit(&x_train, &y_train)?;
287            let y_pred = fitted_svm.predict(&x_test)?;
288
289            let score = self.calculate_score(&y_test, &y_pred)?;
290            scores.push(score);
291        }
292
293        Ok(scores)
294    }
295
296    /// Calculate score based on scoring metric
297    fn calculate_score(&self, y_true: &Array1<f64>, y_pred: &Array1<f64>) -> Result<f64> {
298        match self.config.scoring {
299            ScoringMetric::Accuracy => {
300                let correct = y_true
301                    .iter()
302                    .zip(y_pred.iter())
303                    .map(|(&t, &p)| if (t - p).abs() < 0.5 { 1.0 } else { 0.0 })
304                    .sum::<f64>();
305                Ok(correct / y_true.len() as f64)
306            }
307            ScoringMetric::MeanSquaredError => {
308                let mse = y_true
309                    .iter()
310                    .zip(y_pred.iter())
311                    .map(|(&t, &p)| (t - p).powi(2))
312                    .sum::<f64>()
313                    / y_true.len() as f64;
314                Ok(-mse) // Negative because we want to maximize
315            }
316            ScoringMetric::MeanAbsoluteError => {
317                let mae = y_true
318                    .iter()
319                    .zip(y_pred.iter())
320                    .map(|(&t, &p)| (t - p).abs())
321                    .sum::<f64>()
322                    / y_true.len() as f64;
323                Ok(-mae) // Negative because we want to maximize
324            }
325            _ => {
326                // For now, default to accuracy for other metrics
327                let correct = y_true
328                    .iter()
329                    .zip(y_pred.iter())
330                    .map(|(&t, &p)| if (t - p).abs() < 0.5 { 1.0 } else { 0.0 })
331                    .sum::<f64>();
332                Ok(correct / y_true.len() as f64)
333            }
334        }
335    }
336}