1use crate::{
7 BacktestConfig, BacktestEngine, BacktestError, PerformanceMetrics, internal_invariant,
8};
9use polars::prelude::*;
10use std::collections::HashMap;
11
12#[derive(Debug, Clone, PartialEq)]
14pub struct SweepVariant {
15 pub params: HashMap<String, f64>,
16 pub signal_col: String,
17}
18
19pub 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
49pub 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 ¶m_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 ¶m_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}