(
build_classifier: F,
param_sets: Vec<ParamSet>,
data: SearchData<'_>,
n_splits: usize,
pct_embargo: f64,
scoring: SearchScoring,
)
| 318 | } |
| 319 | |
| 320 | fn search_over_params<C, F>( |
| 321 | build_classifier: F, |
| 322 | param_sets: Vec<ParamSet>, |
| 323 | data: SearchData<'_>, |
| 324 | n_splits: usize, |
| 325 | pct_embargo: f64, |
| 326 | scoring: SearchScoring, |
| 327 | ) -> Result<SearchResult, String> |
| 328 | where |
| 329 | C: SimpleClassifier, |
| 330 | F: Fn(&ParamSet) -> C, |
| 331 | { |
| 332 | validate_search_data(&data, n_splits)?; |
| 333 | let cv = PurgedKFold::new(n_splits, data.samples_info_sets.to_vec(), pct_embargo)?; |
| 334 | let splits = cv.split(data.x.len())?; |
| 335 | |
| 336 | let mut trials = Vec::with_capacity(param_sets.len()); |
| 337 | for params in param_sets { |
| 338 | let fold_scores = evaluate_params( |
| 339 | &build_classifier, |
| 340 | ¶ms, |
| 341 | &splits, |
| 342 | data.x, |
| 343 | data.y, |
| 344 | data.sample_weight, |
| 345 | scoring, |
| 346 | )?; |
| 347 | let mean_score = fold_scores.iter().sum::<f64>() / fold_scores.len() as f64; |
| 348 | trials.push(SearchTrial { params, fold_scores, mean_score }); |
| 349 | } |
| 350 | |
| 351 | let best = trials |
| 352 | .iter() |
| 353 | .max_by(|a, b| a.mean_score.partial_cmp(&b.mean_score).unwrap_or(std::cmp::Ordering::Equal)) |
| 354 | .cloned() |
| 355 | .ok_or_else(|| "no trials produced".to_string())?; |
| 356 | |
| 357 | Ok(SearchResult { best_params: best.params, best_score: best.mean_score, trials }) |
| 358 | } |
| 359 | |
| 360 | fn validate_search_data(data: &SearchData<'_>, n_splits: usize) -> Result<(), String> { |
| 361 | if data.x.is_empty() { |
no test coverage detected