| """Unified strategy backtest runner.""" |
|
|
| from __future__ import annotations |
|
|
| import json |
| from dataclasses import dataclass, field |
| from pathlib import Path |
| from typing import Any |
|
|
| from contextlib import contextmanager |
|
|
| import pandas as pd |
|
|
| from config.settings import load_settings |
| from data_pipeline.factor_loader import build_signal_from_source |
| from data_pipeline.init_qlib import init_qlib |
| from qlib.backtest import backtest_loop, get_strategy_executor |
| from qlib.contrib.evaluate import risk_analysis |
| from strategies.registry import resolve_strategy |
|
|
|
|
| @dataclass |
| class StrategyRunResult: |
| strategy_name: str |
| report: pd.DataFrame |
| positions: dict |
| risk: pd.DataFrame |
| signal_stats: dict[str, Any] = field(default_factory=dict) |
| meta: dict[str, Any] = field(default_factory=dict) |
|
|
|
|
| def _build_exchange_kwargs(backtest_cfg: dict[str, Any], freq: str = "day") -> dict[str, Any]: |
| return { |
| "freq": freq, |
| "limit_threshold": backtest_cfg.get("limit_threshold", 0.095), |
| "deal_price": backtest_cfg.get("deal_price", "close"), |
| "open_cost": backtest_cfg.get("open_cost", 0.0005), |
| "close_cost": backtest_cfg.get("close_cost", 0.0015), |
| "min_cost": backtest_cfg.get("min_cost", 5), |
| } |
|
|
|
|
| @contextmanager |
| def _disable_qlib_benchmark_default(): |
| """ |
| qlib passes benchmark=None as {} and PortfolioMetrics then defaults to CSI300, |
| which breaks 30min-only datasets. Treat empty config as no benchmark. |
| """ |
| from qlib.backtest.report import PortfolioMetrics |
|
|
| original = PortfolioMetrics._cal_benchmark |
|
|
| @staticmethod |
| def _cal_benchmark_no_default(benchmark_config, freq): |
| if not benchmark_config or benchmark_config.get("benchmark") is None: |
| return None |
| return original(benchmark_config, freq) |
|
|
| PortfolioMetrics._cal_benchmark = _cal_benchmark_no_default |
| try: |
| yield |
| finally: |
| PortfolioMetrics._cal_benchmark = original |
|
|
|
|
| def _run_qlib_backtest( |
| *, |
| start_time: str, |
| end_time: str, |
| strategy, |
| executor_config: dict[str, Any], |
| account: float | int, |
| benchmark: str | None, |
| exchange_kwargs: dict[str, Any], |
| ): |
| """Run qlib backtest; disable benchmark when None (qlib defaults empty config to CSI300).""" |
| with _disable_qlib_benchmark_default(): |
| trade_strategy, trade_executor = get_strategy_executor( |
| start_time=start_time, |
| end_time=end_time, |
| strategy=strategy, |
| executor=executor_config, |
| benchmark=benchmark, |
| account=account, |
| exchange_kwargs=exchange_kwargs, |
| ) |
| return backtest_loop(start_time, end_time, trade_strategy, trade_executor) |
|
|
|
|
| def run_strategy_backtest( |
| strategy_name: str, |
| signal_source: dict[str, Any], |
| strategy_kwargs: dict[str, Any] | None = None, |
| start_time: str | None = None, |
| end_time: str | None = None, |
| config_path: str | None = None, |
| ) -> StrategyRunResult: |
| settings = load_settings(config_path) |
| init_qlib(config_path) |
| bt = settings.backtest_config |
| data_freq = settings.backtest_freq |
|
|
| signal = build_signal_from_source(signal_source) |
| if signal.index.names != ["instrument", "datetime"]: |
| signal = signal.swaplevel().sort_index() |
| signal.index.names = ["instrument", "datetime"] |
|
|
| strategy = resolve_strategy(strategy_name, signal, strategy_kwargs) |
|
|
| start_time = start_time or settings.segments.get("test", settings.segments["valid"])[0] |
| end_time = end_time or settings.segments.get("test", settings.segments["valid"])[1] |
|
|
| executor_config = { |
| "class": "SimulatorExecutor", |
| "module_path": "qlib.backtest.executor", |
| "kwargs": { |
| "time_per_step": data_freq, |
| "generate_portfolio_metrics": True, |
| }, |
| } |
|
|
| benchmark = bt.get("benchmark", settings.benchmark) |
| if benchmark in (None, "null", "none", ""): |
| benchmark = None |
|
|
| portfolio_metric_dict, indicator_dict = _run_qlib_backtest( |
| start_time=start_time, |
| end_time=end_time, |
| strategy=strategy, |
| executor_config=executor_config, |
| account=bt.get("account", 100_000_000), |
| benchmark=benchmark, |
| exchange_kwargs=_build_exchange_kwargs(bt, freq=data_freq), |
| ) |
|
|
| freq_key = next(iter(portfolio_metric_dict.keys())) |
| report, positions = portfolio_metric_dict[freq_key] |
|
|
| risk = risk_analysis(report["return"]) if "return" in report.columns else pd.DataFrame() |
| signal_stats = { |
| "n_obs": int(signal.notna().sum()), |
| "n_instruments": int(signal.index.get_level_values("instrument").nunique()), |
| "date_range": [str(signal.index.get_level_values("datetime").min()), str(signal.index.get_level_values("datetime").max())], |
| } |
|
|
| return StrategyRunResult( |
| strategy_name=strategy_name, |
| report=report, |
| positions=positions, |
| risk=risk, |
| signal_stats=signal_stats, |
| meta={"start_time": start_time, "end_time": end_time, "benchmark": benchmark, "freq": freq_key}, |
| ) |
|
|
|
|
| def run_strategy_suite( |
| signal_source: dict[str, Any], |
| strategy_names: list[str], |
| start_time: str | None = None, |
| end_time: str | None = None, |
| config_path: str | None = None, |
| ) -> dict[str, StrategyRunResult]: |
| results = {} |
| for name in strategy_names: |
| results[name] = run_strategy_backtest( |
| strategy_name=name, |
| signal_source=signal_source, |
| start_time=start_time, |
| end_time=end_time, |
| config_path=config_path, |
| ) |
| return results |
|
|
|
|
| def save_backtest_result(result: StrategyRunResult, output_dir: str | Path) -> Path: |
| output_dir = Path(output_dir) |
| output_dir.mkdir(parents=True, exist_ok=True) |
| result.report.to_csv(output_dir / f"{result.strategy_name}_report.csv") |
| if not result.risk.empty: |
| result.risk.to_csv(output_dir / f"{result.strategy_name}_risk.csv") |
| summary = { |
| "strategy": result.strategy_name, |
| "signal_stats": result.signal_stats, |
| "meta": result.meta, |
| } |
| with open(output_dir / f"{result.strategy_name}_summary.json", "w", encoding="utf-8") as f: |
| json.dump(summary, f, indent=2, ensure_ascii=False) |
| return output_dir |
|
|