Skip to main content

quantwave_backtest/
sweep.rs

1//! Parameter sweep helper (quantwave-cr6v.12).
2//!
3//! Runs one backtest per variant and returns a Polars DataFrame with param columns
4//! plus [`PerformanceMetrics`] columns (vectorbt / RaptorBT `batch_spread` pattern).
5
6use crate::{
7    BacktestConfig, BacktestEngine, BacktestError, PerformanceMetrics, internal_invariant,
8};
9use polars::prelude::*;
10use std::collections::HashMap;
11
12/// One grid point: parameter values and the signal column to backtest.
13#[derive(Debug, Clone, PartialEq)]
14pub struct SweepVariant {
15    pub params: HashMap<String, f64>,
16    pub signal_col: String,
17}
18
19/// Build variants for a single-parameter grid (e.g. `hurst_period` → signal columns).
20pub fn single_param_variants(
21    param_name: impl Into<String>,
22    param_values: &[f64],
23    signal_cols: &[impl AsRef<str>],
24) -> Result<Vec<SweepVariant>, BacktestError> {
25    if param_values.len() != signal_cols.len() {
26        return Err(BacktestError::InvalidInput(format!(
27            "param_values len {} != signal_cols len {}",
28            param_values.len(),
29            signal_cols.len()
30        )));
31    }
32    if param_values.is_empty() {
33        return Err(BacktestError::InvalidInput(
34            "sweep requires at least one variant".into(),
35        ));
36    }
37
38    let name = param_name.into();
39    Ok(param_values
40        .iter()
41        .zip(signal_cols.iter())
42        .map(|(&value, col)| SweepVariant {
43            params: HashMap::from([(name.clone(), value)]),
44            signal_col: col.as_ref().to_string(),
45        })
46        .collect())
47}
48
49/// Run backtests for each variant and return a param × metrics DataFrame.
50pub fn run_param_sweep(
51    lf: LazyFrame,
52    variants: &[SweepVariant],
53    base_config: &BacktestConfig,
54) -> Result<DataFrame, BacktestError> {
55    if variants.is_empty() {
56        return Err(BacktestError::InvalidInput(
57            "sweep requires at least one variant".into(),
58        ));
59    }
60
61    let param_keys = sorted_param_keys(variants);
62    let mut param_cols: HashMap<String, Vec<f64>> =
63        param_keys.iter().map(|k| (k.clone(), Vec::new())).collect();
64    let mut metric_cols: HashMap<&'static str, Vec<f64>> = PerformanceMetrics::column_names()
65        .iter()
66        .map(|&name| (name, Vec::new()))
67        .collect();
68
69    for variant in variants {
70        for key in &param_keys {
71            let value = variant.params.get(key).copied().ok_or_else(|| {
72                BacktestError::InvalidInput(format!(
73                    "variant missing param key '{key}' (expected keys: {param_keys:?})"
74                ))
75            })?;
76            param_cols
77                .get_mut(key)
78                .ok_or_else(|| {
79                    internal_invariant(format!(
80                        "param column '{key}' missing from sweep accumulator"
81                    ))
82                })?
83                .push(value);
84        }
85
86        let mut config = base_config.clone();
87        config.signal_col = variant.signal_col.clone();
88        let report = BacktestEngine::new(config).backtest_with_report(lf.clone())?;
89        for (name, value) in report.metrics.row_iter() {
90            metric_cols
91                .get_mut(name)
92                .ok_or_else(|| {
93                    internal_invariant(format!(
94                        "metric column '{name}' missing from sweep accumulator"
95                    ))
96                })?
97                .push(value);
98        }
99    }
100
101    let mut columns: Vec<Column> = Vec::new();
102    for key in &param_keys {
103        columns.push(Column::new(
104            PlSmallStr::from_str(key),
105            param_cols.remove(key).ok_or_else(|| {
106                internal_invariant(format!(
107                    "param column '{key}' missing when building sweep df"
108                ))
109            })?,
110        ));
111    }
112    for name in PerformanceMetrics::column_names() {
113        columns.push(Column::new(
114            PlSmallStr::from_str(name),
115            metric_cols.remove(name).ok_or_else(|| {
116                internal_invariant(format!(
117                    "metric column '{name}' missing when building sweep df"
118                ))
119            })?,
120        ));
121    }
122
123    DataFrame::new(columns).map_err(BacktestError::from)
124}
125
126pub(crate) fn sorted_param_keys(variants: &[SweepVariant]) -> Vec<String> {
127    let mut keys: Vec<String> = variants[0].params.keys().cloned().collect();
128    keys.sort();
129    keys
130}
131
132#[cfg(test)]
133mod tests {
134    use super::*;
135    use approx::assert_relative_eq;
136
137    fn sweep_base_df() -> DataFrame {
138        DataFrame::new(vec![
139            Column::new(
140                "timestamp".into(),
141                (0..6)
142                    .map(|i| 1_700_000_000i64 + (i as i64) * 3600)
143                    .collect::<Vec<_>>(),
144            ),
145            Column::new(
146                "close".into(),
147                vec![100.0, 101.0, 102.5, 103.0, 102.0, 101.0],
148            ),
149            Column::new("signal_early".into(), vec![0.0, 1.0, 1.0, 1.0, 0.0, 0.0]),
150            Column::new("signal_late".into(), vec![0.0, 0.0, 1.0, 1.0, 0.0, 0.0]),
151            Column::new("signal_flat".into(), vec![0.0, 0.0, 0.0, 0.0, 0.0, 0.0]),
152        ])
153        .unwrap()
154    }
155
156    fn zero_cost_config() -> BacktestConfig {
157        BacktestConfig {
158            cost_model: crate::CostModel {
159                commission_bps: 0.0,
160                slippage_bps: 0.0,
161                initial_cash: 100_000.0,
162            },
163            ..Default::default()
164        }
165    }
166
167    #[test]
168    fn test_sweep_single_param_returns_metrics_df() {
169        let variants = single_param_variants(
170            "threshold",
171            &[0.5, 1.0, 2.0],
172            &["signal_early", "signal_late", "signal_flat"],
173        )
174        .unwrap();
175
176        let df = run_param_sweep(sweep_base_df().lazy(), &variants, &zero_cost_config()).unwrap();
177
178        assert_eq!(df.height(), 3);
179        assert!(df.column("threshold").is_ok());
180        assert!(df.column("num_trades").is_ok());
181        assert!(df.column("final_equity").is_ok());
182        assert!(df.column("total_return").is_ok());
183
184        let thresholds = df.column("threshold").unwrap().f64().unwrap();
185        assert_relative_eq!(thresholds.get(0).unwrap(), 0.5, epsilon = 1e-9);
186        assert_relative_eq!(thresholds.get(1).unwrap(), 1.0, epsilon = 1e-9);
187        assert_relative_eq!(thresholds.get(2).unwrap(), 2.0, epsilon = 1e-9);
188
189        let trades = df.column("num_trades").unwrap().f64().unwrap();
190        assert_relative_eq!(trades.get(0).unwrap(), 1.0, epsilon = 1e-9);
191        assert_relative_eq!(trades.get(1).unwrap(), 1.0, epsilon = 1e-9);
192        assert_relative_eq!(trades.get(2).unwrap(), 0.0, epsilon = 1e-9);
193    }
194
195    #[test]
196    fn test_sweep_variants_produce_different_final_equity() {
197        let variants =
198            single_param_variants("entry_bar", &[1.0, 2.0], &["signal_early", "signal_late"])
199                .unwrap();
200
201        let df = run_param_sweep(sweep_base_df().lazy(), &variants, &zero_cost_config()).unwrap();
202        assert_eq!(df.height(), 2);
203
204        let equity = df.column("final_equity").unwrap().f64().unwrap();
205        let e0 = equity.get(0).unwrap();
206        let e1 = equity.get(1).unwrap();
207        assert!(
208            (e0 - e1).abs() > 1.0,
209            "early vs late entry should differ: {e0} vs {e1}"
210        );
211    }
212
213    #[test]
214    fn test_sweep_multi_param_explicit_variants() {
215        let variants = vec![
216            SweepVariant {
217                params: HashMap::from([("stop_pct".to_string(), 0.05), ("mode".to_string(), 1.0)]),
218                signal_col: "signal_early".into(),
219            },
220            SweepVariant {
221                params: HashMap::from([("stop_pct".to_string(), 0.10), ("mode".to_string(), 1.0)]),
222                signal_col: "signal_late".into(),
223            },
224            SweepVariant {
225                params: HashMap::from([("stop_pct".to_string(), 0.05), ("mode".to_string(), 2.0)]),
226                signal_col: "signal_flat".into(),
227            },
228        ];
229
230        let df = run_param_sweep(sweep_base_df().lazy(), &variants, &zero_cost_config()).unwrap();
231        assert_eq!(df.height(), 3);
232        assert!(df.column("mode").is_ok());
233        assert!(df.column("stop_pct").is_ok());
234        assert_eq!(df.column("mode").unwrap().f64().unwrap().len(), 3);
235    }
236}