Skip to content

Commit 3e502cf

Browse files
Pigbibiclaudecursoragent
authored
feat(strategy_lifecycle): add walk-forward and sensitivity to BacktestOrchestrator (#192)
Extend BacktestOrchestrator with walk_forward() for multi-window runs and sensitivity() for basic param-grid sweeps, plus SensitivityReport contract and unit tests covering run, walk_forward, and sensitivity paths. Co-authored-by: Claude <noreply@anthropic.com> Co-authored-by: Cursor <cursoragent@cursor.com>
1 parent a52d78e commit 3e502cf

3 files changed

Lines changed: 302 additions & 2 deletions

File tree

src/quant_platform_kit/strategy_lifecycle/backtest_orchestrator.py

Lines changed: 120 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,11 +5,12 @@
55

66
from __future__ import annotations
77

8+
import itertools
89
import uuid
910
from datetime import date, datetime, timezone
10-
from typing import Any, Mapping, Protocol, runtime_checkable
11+
from typing import Any, Mapping, Protocol, Sequence, runtime_checkable
1112

12-
from quant_platform_kit.strategy_lifecycle.contracts import BacktestResult
13+
from quant_platform_kit.strategy_lifecycle.contracts import BacktestResult, SensitivityReport
1314
from quant_platform_kit.strategy_lifecycle.performance_store import PerformanceStore
1415

1516

@@ -140,3 +141,120 @@ def run(
140141
def run_latest(self, strategy_profile: str, *, domain: str) -> BacktestResult | None:
141142
"""Load the latest persisted backtest result for a strategy."""
142143
return self._store.load_latest_backtest(domain, strategy_profile)
144+
145+
def walk_forward(
146+
self,
147+
strategy_profile: str,
148+
*,
149+
domain: str,
150+
params: Mapping[str, Any],
151+
windows: Sequence[tuple[date | None, date | None]],
152+
param_set_id: str = "",
153+
param_version: int = 1,
154+
) -> list[BacktestResult]:
155+
"""Run backtests across multiple time windows.
156+
157+
Args:
158+
strategy_profile: Canonical strategy profile.
159+
domain: Market domain.
160+
params: Strategy parameters shared across windows.
161+
windows: Sequence of (start_date, end_date) pairs, one per fold.
162+
param_set_id: Base identifier for parameter set metadata.
163+
param_version: Version number for this parameter set.
164+
165+
Returns:
166+
One BacktestResult per window, in order.
167+
168+
Raises:
169+
ValueError: If windows is empty or no runner is registered.
170+
"""
171+
if not windows:
172+
raise ValueError("windows must contain at least one (start_date, end_date) pair")
173+
174+
base_id = param_set_id or _run_id()
175+
results: list[BacktestResult] = []
176+
for idx, (start_date, end_date) in enumerate(windows):
177+
results.append(
178+
self.run(
179+
strategy_profile,
180+
domain=domain,
181+
params=params,
182+
param_set_id=f"{base_id}_wf{idx}",
183+
param_version=param_version,
184+
start_date=start_date,
185+
end_date=end_date,
186+
)
187+
)
188+
return results
189+
190+
def sensitivity(
191+
self,
192+
strategy_profile: str,
193+
*,
194+
domain: str,
195+
base_params: Mapping[str, Any],
196+
param_ranges: Mapping[str, Sequence[Any]],
197+
start_date: date | None = None,
198+
end_date: date | None = None,
199+
max_combinations: int = 500,
200+
) -> SensitivityReport:
201+
"""Run a basic parameter-grid sensitivity sweep.
202+
203+
Expands ``param_ranges`` via Cartesian product, merging each combination
204+
onto ``base_params``, and runs one backtest per combination.
205+
206+
Args:
207+
strategy_profile: Canonical strategy profile.
208+
domain: Market domain.
209+
base_params: Fixed parameters applied to every combination.
210+
param_ranges: Per-parameter value lists to sweep.
211+
start_date: Backtest start date.
212+
end_date: Backtest end date.
213+
max_combinations: Upper bound on grid size (subsamples if exceeded).
214+
215+
Returns:
216+
SensitivityReport with one BacktestResult per combination tried.
217+
218+
Raises:
219+
ValueError: If param_ranges is empty or no runner is registered.
220+
"""
221+
if not param_ranges:
222+
raise ValueError("param_ranges must contain at least one parameter dimension")
223+
224+
keys = sorted(param_ranges.keys())
225+
value_lists = [list(param_ranges[k]) for k in keys]
226+
total = 1
227+
for values in value_lists:
228+
total *= len(values)
229+
230+
combos: list[dict[str, Any]] = []
231+
for idx, combo in enumerate(itertools.product(*value_lists)):
232+
if total > max_combinations and idx % max(1, total // max_combinations) != 0:
233+
continue
234+
merged = dict(base_params)
235+
merged.update(dict(zip(keys, combo)))
236+
combos.append(merged)
237+
if len(combos) >= max_combinations:
238+
break
239+
240+
results: list[BacktestResult] = []
241+
for idx, combo_params in enumerate(combos):
242+
results.append(
243+
self.run(
244+
strategy_profile,
245+
domain=domain,
246+
params=combo_params,
247+
param_set_id=f"{strategy_profile}_sens_{idx}",
248+
param_version=1,
249+
start_date=start_date,
250+
end_date=end_date,
251+
)
252+
)
253+
254+
return SensitivityReport(
255+
strategy_profile=strategy_profile,
256+
domain=domain,
257+
base_params=dict(base_params),
258+
results=tuple(results),
259+
combination_count=len(results),
260+
)

src/quant_platform_kit/strategy_lifecycle/contracts.py

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -282,6 +282,26 @@ def to_dict(self) -> dict[str, object]:
282282
}
283283

284284

285+
@dataclass(frozen=True)
286+
class SensitivityReport:
287+
"""Results from a parameter-grid sensitivity sweep."""
288+
289+
strategy_profile: str
290+
domain: str
291+
base_params: Mapping[str, Any]
292+
results: tuple[BacktestResult, ...] = ()
293+
combination_count: int = 0
294+
295+
def to_dict(self) -> dict[str, object]:
296+
return {
297+
"strategy_profile": self.strategy_profile,
298+
"domain": self.domain,
299+
"base_params": dict(self.base_params),
300+
"combination_count": self.combination_count,
301+
"results": [r.to_dict() for r in self.results],
302+
}
303+
304+
285305
@dataclass(frozen=True)
286306
class ParamSearchSpace:
287307
"""Definition of the search space for one strategy's parameters."""
Lines changed: 162 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,162 @@
1+
"""Tests for strategy_lifecycle.backtest_orchestrator."""
2+
3+
from __future__ import annotations
4+
5+
from datetime import date
6+
import tempfile
7+
import unittest
8+
from pathlib import Path
9+
from typing import Any, Mapping
10+
11+
from quant_platform_kit.strategy_lifecycle.backtest_orchestrator import BacktestOrchestrator
12+
from quant_platform_kit.strategy_lifecycle.contracts import BacktestResult
13+
from quant_platform_kit.strategy_lifecycle.performance_store import PerformanceStore
14+
15+
16+
class _RecordingRunner:
17+
"""Mock BacktestRunner that records calls and returns deterministic metrics."""
18+
19+
def __init__(self) -> None:
20+
self.calls: list[dict[str, Any]] = []
21+
22+
def run(
23+
self,
24+
strategy_profile: str,
25+
params: Mapping[str, Any],
26+
start_date: date | None = None,
27+
end_date: date | None = None,
28+
) -> BacktestResult:
29+
self.calls.append(
30+
{
31+
"strategy_profile": strategy_profile,
32+
"params": dict(params),
33+
"start_date": start_date,
34+
"end_date": end_date,
35+
}
36+
)
37+
lookback = int(params.get("lookback", 20))
38+
sharpe = 1.0 + lookback * 0.01
39+
return BacktestResult(
40+
strategy_profile=strategy_profile,
41+
domain="us_equity",
42+
param_set_id="mock",
43+
params=dict(params),
44+
sharpe_ratio=sharpe,
45+
cagr=0.12,
46+
max_drawdown=-0.08,
47+
start_date=start_date,
48+
end_date=end_date,
49+
observation_count=252,
50+
)
51+
52+
53+
class BacktestOrchestratorTests(unittest.TestCase):
54+
55+
def setUp(self) -> None:
56+
self.tmp = tempfile.TemporaryDirectory()
57+
self.store = PerformanceStore(local_root=Path(self.tmp.name))
58+
self.orchestrator = BacktestOrchestrator(store=self.store)
59+
self.runner = _RecordingRunner()
60+
self.orchestrator.register_runner("us_equity", self.runner)
61+
62+
def tearDown(self) -> None:
63+
self.tmp.cleanup()
64+
65+
def test_run_enriches_and_persists(self) -> None:
66+
result = self.orchestrator.run(
67+
"test_strat",
68+
domain="us_equity",
69+
params={"lookback": 30},
70+
start_date=date(2020, 1, 1),
71+
end_date=date(2024, 12, 31),
72+
)
73+
self.assertEqual(result.strategy_profile, "test_strat")
74+
self.assertEqual(result.domain, "us_equity")
75+
self.assertEqual(result.params, {"lookback": 30})
76+
self.assertAlmostEqual(result.sharpe_ratio, 1.3)
77+
self.assertTrue(result.run_id)
78+
self.assertTrue(result.computed_at)
79+
self.assertEqual(result.source_script, "backtest_orchestrator")
80+
81+
def test_run_raises_without_runner(self) -> None:
82+
with self.assertRaises(ValueError):
83+
self.orchestrator.run("test_strat", domain="cn_equity", params={})
84+
85+
def test_walk_forward_runs_each_window(self) -> None:
86+
windows = [
87+
(date(2020, 1, 1), date(2021, 12, 31)),
88+
(date(2022, 1, 1), date(2023, 12, 31)),
89+
(date(2024, 1, 1), date(2024, 12, 31)),
90+
]
91+
results = self.orchestrator.walk_forward(
92+
"test_strat",
93+
domain="us_equity",
94+
params={"lookback": 20},
95+
windows=windows,
96+
param_set_id="wf_test",
97+
)
98+
self.assertEqual(len(results), 3)
99+
self.assertEqual(len(self.runner.calls), 3)
100+
for idx, (result, window) in enumerate(zip(results, windows)):
101+
self.assertEqual(result.start_date, window[0])
102+
self.assertEqual(result.end_date, window[1])
103+
self.assertEqual(result.param_set_id, f"wf_test_wf{idx}")
104+
self.assertEqual(
105+
[self.runner.calls[i]["start_date"] for i in range(3)],
106+
[w[0] for w in windows],
107+
)
108+
109+
def test_walk_forward_empty_windows_raises(self) -> None:
110+
with self.assertRaises(ValueError):
111+
self.orchestrator.walk_forward(
112+
"test_strat",
113+
domain="us_equity",
114+
params={},
115+
windows=[],
116+
)
117+
118+
def test_sensitivity_runs_param_grid(self) -> None:
119+
report = self.orchestrator.sensitivity(
120+
"test_strat",
121+
domain="us_equity",
122+
base_params={"top_n": 2},
123+
param_ranges={"lookback": [10, 20, 30]},
124+
start_date=date(2020, 1, 1),
125+
end_date=date(2024, 12, 31),
126+
)
127+
self.assertEqual(report.combination_count, 3)
128+
self.assertEqual(len(report.results), 3)
129+
self.assertEqual(report.strategy_profile, "test_strat")
130+
self.assertEqual(report.base_params, {"top_n": 2})
131+
sharpes = [r.sharpe_ratio for r in report.results]
132+
self.assertEqual(sharpes, [1.1, 1.2, 1.3])
133+
for call, result in zip(self.runner.calls, report.results):
134+
self.assertEqual(call["params"]["top_n"], 2)
135+
self.assertIn("lookback", call["params"])
136+
self.assertEqual(result.params["top_n"], 2)
137+
138+
def test_sensitivity_two_dim_grid(self) -> None:
139+
report = self.orchestrator.sensitivity(
140+
"test_strat",
141+
domain="us_equity",
142+
base_params={},
143+
param_ranges={"lookback": [10, 20], "top_n": [2, 3]},
144+
)
145+
self.assertEqual(report.combination_count, 4)
146+
lookbacks = sorted({r.params["lookback"] for r in report.results})
147+
top_ns = sorted({r.params["top_n"] for r in report.results})
148+
self.assertEqual(lookbacks, [10, 20])
149+
self.assertEqual(top_ns, [2, 3])
150+
151+
def test_sensitivity_empty_ranges_raises(self) -> None:
152+
with self.assertRaises(ValueError):
153+
self.orchestrator.sensitivity(
154+
"test_strat",
155+
domain="us_equity",
156+
base_params={},
157+
param_ranges={},
158+
)
159+
160+
161+
if __name__ == "__main__":
162+
unittest.main()

0 commit comments

Comments
 (0)