quant_test / strategies /runner.py
lucky-loster's picture
Upload folder using huggingface_hub
590a501 verified
Raw
History Blame Contribute Delete
6.28 kB
"""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