mirror of
https://ghfast.top/https://github.com/aeroxw/easy_tdx_max.git
synced 2026-09-12 16:54:20 +08:00
- screen() now calls run_combination() internally, eliminating duplicated signal extraction/combination logic - _run_combo_screen creates one CombinationRunner before the size loop, so signal cache is reused across 2-factor and 3-factor screens - Add MAJORITY(2)=AND note to screen() docstring
551 lines
19 KiB
Python
551 lines
19 KiB
Python
"""批量回测脚本:依次运行 strategies/ 目录下所有策略并比较结果。
|
||
|
||
用法::
|
||
|
||
python run_all_strategies.py SZ 300308 --count 2000 --cash 1000000 --adjust QFQ
|
||
|
||
输出每个策略的绩效指标,并按总收益率排名。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import sys
|
||
import time
|
||
from pathlib import Path
|
||
from typing import Any
|
||
|
||
# 确保 easy_tdx 可导入
|
||
sys.path.insert(0, str(Path(__file__).parent / "src"))
|
||
|
||
import click
|
||
|
||
|
||
def _setup_chinese_font() -> None:
|
||
"""配置 matplotlib 中文字体,按平台自动选择。"""
|
||
import platform
|
||
|
||
import matplotlib
|
||
|
||
system = platform.system()
|
||
if system == "Windows":
|
||
candidates = ["Microsoft YaHei", "SimHei", "KaiTi", "FangSong"]
|
||
elif system == "Darwin":
|
||
candidates = ["PingFang SC", "Heiti SC", "STHeiti"]
|
||
else:
|
||
candidates = ["WenQuanYi Micro Hei", "Noto Sans CJK SC", "Droid Sans Fallback"]
|
||
|
||
import matplotlib.font_manager as fm
|
||
|
||
available = {f.name for f in fm.fontManager.ttflist}
|
||
for font in candidates:
|
||
if font in available:
|
||
matplotlib.rcParams["font.sans-serif"] = [font, "DejaVu Sans"]
|
||
break
|
||
matplotlib.rcParams["axes.unicode_minus"] = False
|
||
|
||
|
||
def _map_trade_values(trades_df: Any, equity: Any, initial_cash: float) -> list[float]:
|
||
"""将交易的 datetime 映射到 equity_curve 对应的归一化值。"""
|
||
eq_dt = equity["datetime"].values
|
||
eq_norm = equity["total"].values / initial_cash
|
||
result_vals: list[float] = []
|
||
for dt in trades_df["datetime"].values:
|
||
idx = eq_dt.searchsorted(dt, side="right") - 1
|
||
if idx < 0:
|
||
idx = 0
|
||
if idx >= len(eq_norm):
|
||
idx = len(eq_norm) - 1
|
||
result_vals.append(float(eq_norm[idx]))
|
||
return result_vals
|
||
|
||
|
||
def _run_combo_screen(
|
||
strategy_files: list[Path],
|
||
df: Any,
|
||
cash: float,
|
||
commission: float,
|
||
combo_sizes: tuple[int, ...],
|
||
combo_mode: str,
|
||
) -> None:
|
||
"""运行多因子组合回测并输出排名。"""
|
||
import importlib.util
|
||
|
||
from easy_tdx.backtest.combo import CombinationRunner
|
||
from easy_tdx.backtest.strategy import Strategy
|
||
|
||
# 加载所有策略类
|
||
strategy_classes: list[type[Strategy]] = []
|
||
for sf in strategy_files:
|
||
spec = importlib.util.spec_from_file_location("strategy_module", sf)
|
||
if spec is None or spec.loader is None:
|
||
continue
|
||
module = importlib.util.module_from_spec(spec)
|
||
spec.loader.exec_module(module)
|
||
|
||
for attr_name in dir(module):
|
||
obj = getattr(module, attr_name)
|
||
try:
|
||
if isinstance(obj, type) and issubclass(obj, Strategy) and obj is not Strategy:
|
||
strategy_classes.append(obj)
|
||
break
|
||
except TypeError:
|
||
pass
|
||
|
||
if len(strategy_classes) < 2:
|
||
click.echo("[!] 策略数量不足 2 个,跳过组合回测")
|
||
return
|
||
|
||
from math import comb
|
||
|
||
# 单个 Runner 跨 size 复用信号缓存,避免重复提取
|
||
runner = CombinationRunner(
|
||
strategy_classes=strategy_classes,
|
||
df=df,
|
||
cash=cash,
|
||
commission=commission,
|
||
)
|
||
|
||
for size in combo_sizes:
|
||
total = comb(len(strategy_classes), size)
|
||
click.echo("\n" + "=" * 80)
|
||
click.echo(f"[*] {size}因子组合回测 (共{total}组, 模式={combo_mode})")
|
||
click.echo("=" * 80)
|
||
|
||
results = runner.screen(combo_sizes=(size,), mode=combo_mode.upper())
|
||
|
||
if not results:
|
||
click.echo(" 无有效交易组合(所有组合均为零交易)")
|
||
continue
|
||
|
||
# 表头
|
||
click.echo(
|
||
f"{'排名':>4} {'因子组合':<50} {'总收益率':>10} {'年化收益':>10} "
|
||
f"{'最大回撤':>10} {'夏普':>8} {'胜率':>8} {'交易':>6}"
|
||
)
|
||
click.echo("-" * 120)
|
||
|
||
for i, r in enumerate(results[:20], 1): # 只显示 top 20
|
||
medal = " *1*" if i == 1 else " *2*" if i == 2 else " *3*" if i == 3 else " "
|
||
perf = r.result.performance
|
||
click.echo(
|
||
f"{medal}{i:>2} {r.name:<50} "
|
||
f"{perf.get('total_return', 0):>9.2%} "
|
||
f"{perf.get('annual_return', 0):>9.2%} "
|
||
f"{perf.get('max_drawdown', 0):>9.2%} "
|
||
f"{perf.get('sharpe', 0):>8.2f} "
|
||
f"{perf.get('win_rate', 0):>7.1%} "
|
||
f"{perf.get('total_trades', 0):>6}"
|
||
)
|
||
|
||
if len(results) > 20:
|
||
click.echo(f" ... 共 {len(results)} 个有效组合,仅显示前 20")
|
||
|
||
# 最佳组合详细报告
|
||
best = results[0]
|
||
bp = best.result.performance
|
||
click.echo(f"\n[BEST {size}因子] {best.name}")
|
||
click.echo(f" 总收益率: {bp.get('total_return', 0):.2%}")
|
||
click.echo(f" 年化收益: {bp.get('annual_return', 0):.2%}")
|
||
click.echo(f" 最大回撤: {bp.get('max_drawdown', 0):.2%}")
|
||
click.echo(f" 夏普比率: {bp.get('sharpe', 0):.2f}")
|
||
click.echo(f" 胜率: {bp.get('win_rate', 0):.1%}")
|
||
|
||
|
||
def _show_best_chart(
|
||
df: Any,
|
||
result: Any,
|
||
strategy_name: str,
|
||
stock_label: str,
|
||
stock_name: str,
|
||
initial_cash: float,
|
||
) -> None:
|
||
"""展示最佳策略资金曲线与股价归一化对比图。"""
|
||
try:
|
||
import matplotlib.pyplot as plt
|
||
except ImportError:
|
||
click.echo("[!] 需要 matplotlib 才能展示图表: pip install matplotlib")
|
||
return
|
||
|
||
_setup_chinese_font()
|
||
|
||
equity = result.equity_curve
|
||
if equity.empty:
|
||
click.echo("[!] 最佳策略无资金曲线数据,跳过绘图")
|
||
return
|
||
|
||
fig, ax1 = plt.subplots(figsize=(14, 7))
|
||
|
||
# 归一化股价(以第一天收盘价为基准)
|
||
close_prices = df["close"].values
|
||
norm_price = close_prices / close_prices[0]
|
||
dates = df["datetime"] if "datetime" in df.columns else df.index
|
||
ax1.plot(dates, norm_price, color="steelblue", linewidth=1.2, label="股价 (归一化)")
|
||
ax1.set_ylabel("股价归一化", color="steelblue", fontsize=11)
|
||
ax1.tick_params(axis="y", labelcolor="steelblue")
|
||
|
||
# 归一化资金曲线(以初始资金为基准)
|
||
eq_dates = equity["datetime"]
|
||
eq_values = equity["total"].values / initial_cash
|
||
ax2 = ax1.twinx()
|
||
ax2.plot(eq_dates, eq_values, color="crimson", linewidth=1.5, label=f"策略: {strategy_name}")
|
||
ax2.set_ylabel("资金曲线 (归一化)", color="crimson", fontsize=11)
|
||
ax2.tick_params(axis="y", labelcolor="crimson")
|
||
|
||
# 标记买卖点
|
||
trades = result.trades
|
||
if not trades.empty:
|
||
buy_trades = trades[trades["direction"] == "BUY"]
|
||
sell_trades = trades[trades["direction"] == "SELL"]
|
||
if not buy_trades.empty:
|
||
ax2.scatter(
|
||
buy_trades["datetime"].values,
|
||
_map_trade_values(buy_trades, equity, initial_cash),
|
||
marker="^",
|
||
color="green",
|
||
s=30,
|
||
alpha=0.7,
|
||
zorder=5,
|
||
label="买入",
|
||
)
|
||
if not sell_trades.empty:
|
||
ax2.scatter(
|
||
sell_trades["datetime"].values,
|
||
_map_trade_values(sell_trades, equity, initial_cash),
|
||
marker="v",
|
||
color="orange",
|
||
s=30,
|
||
alpha=0.7,
|
||
zorder=5,
|
||
label="卖出",
|
||
)
|
||
|
||
# 标题:股票代码 + 名称 + 策略绩效
|
||
title = f"{stock_label}"
|
||
if stock_name:
|
||
title += f" {stock_name}"
|
||
perf = result.performance
|
||
ret_str = f"{perf.get('total_return', 0):.1%}"
|
||
dd_str = f"{perf.get('max_drawdown', 0):.1%}"
|
||
sharpe_str = f"{perf.get('sharpe', 0):.2f}"
|
||
title += f" | 最佳策略: {strategy_name} | 收益 {ret_str} 回撤 {dd_str} 夏普 {sharpe_str}"
|
||
|
||
ax1.set_title(title, fontsize=12, pad=15)
|
||
ax1.set_xlabel("日期", fontsize=11)
|
||
|
||
# 合并两个轴的图例
|
||
lines1, labels1 = ax1.get_legend_handles_labels()
|
||
lines2, labels2 = ax2.get_legend_handles_labels()
|
||
ax1.legend(lines1 + lines2, labels1 + labels2, loc="upper left", fontsize=9)
|
||
|
||
fig.autofmt_xdate()
|
||
plt.tight_layout()
|
||
click.echo("\n正在显示图表,关闭窗口后继续...")
|
||
plt.show()
|
||
|
||
|
||
@click.command()
|
||
@click.argument("market")
|
||
@click.argument("code")
|
||
@click.option("--count", default=2000, type=int, help="K线数量")
|
||
@click.option("--cash", default=1000000.0, type=float, help="初始资金")
|
||
@click.option("--commission", default=0.0003, type=float, help="佣金率")
|
||
@click.option("--adjust", default="QFQ", help="复权: NONE/QFQ/HFQ")
|
||
@click.option("--period", default="DAILY", help="K线周期")
|
||
@click.option(
|
||
"--combo",
|
||
"combo_sizes",
|
||
multiple=True,
|
||
type=int,
|
||
help="多因子组合回测(可多次指定,如 --combo 2 --combo 3)",
|
||
)
|
||
@click.option(
|
||
"--combo-mode",
|
||
"combo_mode",
|
||
default="MAJORITY",
|
||
type=click.Choice(["AND", "OR", "MAJORITY"], case_sensitive=False),
|
||
help="多因子信号合并模式(默认 MAJORITY)",
|
||
)
|
||
@click.option("--show", "show_chart", is_flag=True, help="显示最佳策略资金曲线 vs 股价对比图")
|
||
def run_all(
|
||
market: str,
|
||
code: str,
|
||
count: int,
|
||
cash: float,
|
||
commission: float,
|
||
adjust: str,
|
||
period: str,
|
||
combo_sizes: tuple[int, ...],
|
||
combo_mode: str,
|
||
show_chart: bool,
|
||
) -> None:
|
||
"""批量运行 strategies/ 目录下所有策略并比较结果。"""
|
||
from easy_tdx.backtest.engine import BacktestEngine
|
||
from easy_tdx.backtest.strategy import Strategy
|
||
from easy_tdx.cli.parsers import parse_adjust, parse_market, parse_period
|
||
from easy_tdx.mac.client import MacClient
|
||
|
||
# 1. 发现策略文件
|
||
strategies_dir = Path(__file__).parent / "strategies"
|
||
strategy_files = sorted(strategies_dir.glob("*.py"))
|
||
if not strategy_files:
|
||
click.echo("未找到策略文件 (strategies/*.py)", err=True)
|
||
raise SystemExit(1)
|
||
|
||
click.echo(f"发现 {len(strategy_files)} 个策略文件")
|
||
click.echo(f"标的: {market} {code} | K线: {count} | 资金: {cash:,.0f} | 复权: {adjust}")
|
||
click.echo("=" * 80)
|
||
|
||
# 2. 获取数据(所有策略共享同一份数据)
|
||
mkt = parse_market(market)
|
||
click.echo("正在获取行情数据...")
|
||
client = MacClient.from_best_host()
|
||
client.connect()
|
||
stock_name = ""
|
||
try:
|
||
df = client.get_stock_kline(
|
||
mkt,
|
||
code,
|
||
period=parse_period(period),
|
||
start=0,
|
||
count=count,
|
||
adjust=parse_adjust(adjust),
|
||
)
|
||
# 获取股票名称
|
||
if show_chart:
|
||
try:
|
||
quotes_df = client.get_stock_quotes([(mkt, code)])
|
||
if not quotes_df.empty and "name" in quotes_df.columns:
|
||
stock_name = str(quotes_df.iloc[0]["name"])
|
||
except Exception:
|
||
pass
|
||
finally:
|
||
client.close()
|
||
click.echo(f"获取到 {len(df)} 条K线数据")
|
||
click.echo("=" * 80)
|
||
|
||
# 3. 逐个运行策略
|
||
results: list[dict] = []
|
||
backtest_results: dict[str, Any] = {} # strategy_name -> BacktestResult
|
||
|
||
for sf in strategy_files:
|
||
strategy_name = sf.stem
|
||
click.echo(f"\n>> 运行策略: {strategy_name} ...", nl=False)
|
||
|
||
# 加载策略类
|
||
import importlib.util
|
||
|
||
spec = importlib.util.spec_from_file_location("strategy_module", sf)
|
||
if spec is None or spec.loader is None:
|
||
click.echo(" [加载失败]")
|
||
continue
|
||
|
||
module = importlib.util.module_from_spec(spec)
|
||
spec.loader.exec_module(module)
|
||
|
||
# 查找 Strategy 子类
|
||
strategy_cls = None
|
||
for attr_name in dir(module):
|
||
obj = getattr(module, attr_name)
|
||
try:
|
||
if isinstance(obj, type) and issubclass(obj, Strategy) and obj is not Strategy:
|
||
strategy_cls = obj
|
||
break
|
||
except TypeError:
|
||
pass
|
||
|
||
if strategy_cls is None:
|
||
click.echo(" [未找到 Strategy 子类]")
|
||
continue
|
||
|
||
# 运行回测
|
||
t0 = time.perf_counter()
|
||
try:
|
||
engine = BacktestEngine(
|
||
strategy=strategy_cls,
|
||
cash=cash,
|
||
commission=commission,
|
||
)
|
||
result = engine.run(df)
|
||
elapsed = time.perf_counter() - t0
|
||
perf = result.performance
|
||
click.echo(f" 完成 ({elapsed:.1f}s)")
|
||
|
||
results.append(
|
||
{
|
||
"strategy": strategy_name,
|
||
"total_return": perf.get("total_return", 0),
|
||
"annual_return": perf.get("annual_return", 0),
|
||
"max_drawdown": perf.get("max_drawdown", 0),
|
||
"sharpe": perf.get("sharpe", 0),
|
||
"sortino": perf.get("sortino", 0),
|
||
"calmar": perf.get("calmar", 0),
|
||
"win_rate": perf.get("win_rate", 0),
|
||
"total_trades": perf.get("total_trades", 0),
|
||
"profit_factor": perf.get("profit_factor", 0),
|
||
"volatility": perf.get("volatility", 0),
|
||
}
|
||
)
|
||
backtest_results[strategy_name] = result
|
||
except Exception as e:
|
||
elapsed = time.perf_counter() - t0
|
||
click.echo(f" 错误 ({elapsed:.1f}s): {e}")
|
||
results.append(
|
||
{
|
||
"strategy": strategy_name,
|
||
"error": str(e),
|
||
}
|
||
)
|
||
|
||
# 4. 输出排名
|
||
click.echo("\n" + "=" * 80)
|
||
click.echo("[*] 策略绩效排名 (按总收益率降序)")
|
||
click.echo("=" * 80)
|
||
|
||
# 过滤掉有错误的策略
|
||
valid = [r for r in results if "error" not in r]
|
||
errored = [r for r in results if "error" in r]
|
||
|
||
if not valid:
|
||
click.echo("所有策略均运行失败!")
|
||
for r in errored:
|
||
click.echo(f" {r['strategy']}: {r['error']}")
|
||
raise SystemExit(1)
|
||
|
||
# 按总收益率排序
|
||
valid.sort(key=lambda x: x["total_return"], reverse=True)
|
||
|
||
# 表头
|
||
click.echo(
|
||
f"{'排名':>4} {'策略':<22} {'总收益率':>10} {'年化收益':>10} "
|
||
f"{'最大回撤':>10} {'夏普':>8} {'胜率':>8} {'交易次数':>8} {'盈亏比':>8}"
|
||
)
|
||
click.echo("-" * 100)
|
||
|
||
for i, r in enumerate(valid, 1):
|
||
medal = " *1*" if i == 1 else " *2*" if i == 2 else " *3*" if i == 3 else " "
|
||
click.echo(
|
||
f"{medal}{i:>2} {r['strategy']:<22} "
|
||
f"{r['total_return']:>9.2%} "
|
||
f"{r['annual_return']:>9.2%} "
|
||
f"{r['max_drawdown']:>9.2%} "
|
||
f"{r['sharpe']:>8.2f} "
|
||
f"{r['win_rate']:>7.1%} "
|
||
f"{r['total_trades']:>8} "
|
||
f"{r['profit_factor']:>8.2f}"
|
||
)
|
||
|
||
# 最佳策略详细报告
|
||
best = valid[0]
|
||
click.echo("\n" + "=" * 80)
|
||
click.echo(f"[BEST] 最佳策略: {best['strategy']}")
|
||
click.echo("=" * 80)
|
||
click.echo(f" 总收益率: {best['total_return']:.2%}")
|
||
click.echo(f" 年化收益: {best['annual_return']:.2%}")
|
||
click.echo(f" 最大回撤: {best['max_drawdown']:.2%}")
|
||
click.echo(f" 夏普比率: {best['sharpe']:.2f}")
|
||
click.echo(f" 索提诺: {best['sortino']:.2f}")
|
||
click.echo(f" 卡玛比率: {best['calmar']:.2f}")
|
||
click.echo(f" 胜率: {best['win_rate']:.1%}")
|
||
click.echo(f" 交易次数: {best['total_trades']}")
|
||
click.echo(f" 盈亏比: {best['profit_factor']:.2f}")
|
||
click.echo(f" 年化波动: {best['volatility']:.4f}")
|
||
|
||
# 综合评分(综合夏普、收益率、回撤)
|
||
click.echo("\n" + "=" * 80)
|
||
click.echo("[*] 综合评分排名 (Sharpe*0.4 + Ret/DD*0.3 + WinRate*0.3)")
|
||
click.echo("=" * 80)
|
||
|
||
scored = []
|
||
for r in valid:
|
||
# 避免除以零
|
||
ret_dd_ratio = r["annual_return"] / r["max_drawdown"] if r["max_drawdown"] > 1e-6 else 999.0
|
||
score = r["sharpe"] * 0.4 + ret_dd_ratio * 0.3 + r["win_rate"] * 100 * 0.3
|
||
scored.append((r, score))
|
||
|
||
scored.sort(key=lambda x: x[1], reverse=True)
|
||
|
||
click.echo(
|
||
f"{'排名':>4} {'策略':<22} {'综合评分':>10} {'夏普':>8} {'收益/回撤':>10} {'胜率':>8}"
|
||
)
|
||
click.echo("-" * 70)
|
||
for i, (r, score) in enumerate(scored, 1):
|
||
ret_dd_ratio = r["annual_return"] / r["max_drawdown"] if r["max_drawdown"] > 1e-6 else 999.0
|
||
medal = " *1*" if i == 1 else " *2*" if i == 2 else " *3*" if i == 3 else " "
|
||
click.echo(
|
||
f"{medal}{i:>2} {r['strategy']:<22} {score:>10.2f} "
|
||
f"{r['sharpe']:>8.2f} {ret_dd_ratio:>10.2f} {r['win_rate']:>7.1%}"
|
||
)
|
||
|
||
# 最佳策略完整交易明细
|
||
best_name = valid[0]["strategy"]
|
||
if best_name in backtest_results:
|
||
bt = backtest_results[best_name]
|
||
bp = bt.performance
|
||
bc = bt.config
|
||
|
||
click.echo("\n" + "=" * 80)
|
||
click.echo(f"[DETAIL] 最佳策略交易明细: {best_name}")
|
||
click.echo("=" * 80)
|
||
|
||
click.echo("=== 回测绩效概要 ===")
|
||
click.echo(f"总收益率: {bp.get('total_return', 0):.2%}")
|
||
click.echo(f"年化收益: {bp.get('annual_return', 0):.2%}")
|
||
click.echo(f"最大回撤: {bp.get('max_drawdown', 0):.2%}")
|
||
click.echo(f"夏普比率: {bp.get('sharpe', 0):.2f}")
|
||
click.echo(f"胜率: {bp.get('win_rate', 0):.2%}")
|
||
click.echo(f"交易次数: {bp.get('total_trades', 0)}")
|
||
click.echo()
|
||
click.echo("=== 配置参数 ===")
|
||
click.echo(f"初始资金: {bc.get('cash', 0):.2f}")
|
||
click.echo(f"佣金率: {bc.get('commission', 0):.4f}")
|
||
click.echo(f"成交规则: {bc.get('execution', 'next_open')}")
|
||
click.echo()
|
||
|
||
if not bt.trades.empty:
|
||
click.echo("=== 最近交易记录 ===")
|
||
recent_trades = bt.trades.tail(10)
|
||
for _, trade in recent_trades.iterrows():
|
||
direction = "买入" if trade["direction"] == "BUY" else "卖出"
|
||
status = "拒绝" if trade["rejected"] else "成交"
|
||
click.echo(
|
||
f" [{trade['datetime']}] {direction} "
|
||
f"数量={trade['size']:.0f} 价格={trade['price']:.2f} "
|
||
f"盈亏={trade['pnl']:.2f} [{status}]"
|
||
)
|
||
else:
|
||
click.echo("无交易记录")
|
||
|
||
# 报告错误
|
||
if errored:
|
||
click.echo("\n[!] 以下策略运行失败:")
|
||
for r in errored:
|
||
click.echo(f" {r['strategy']}: {r['error']}")
|
||
|
||
# 5. 多因子组合回测
|
||
if combo_sizes:
|
||
_run_combo_screen(
|
||
strategy_files=strategy_files,
|
||
df=df,
|
||
cash=cash,
|
||
commission=commission,
|
||
combo_sizes=combo_sizes,
|
||
combo_mode=combo_mode,
|
||
)
|
||
|
||
# 6. 展示最佳策略曲线图
|
||
if show_chart and valid:
|
||
best_name = valid[0]["strategy"]
|
||
if best_name in backtest_results:
|
||
_show_best_chart(
|
||
df=df,
|
||
result=backtest_results[best_name],
|
||
strategy_name=best_name,
|
||
stock_label=f"{market}{code}",
|
||
stock_name=stock_name,
|
||
initial_cash=cash,
|
||
)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
run_all()
|