mirror of
https://ghfast.top/https://github.com/aeroxw/easy-tdx.git
synced 2026-09-12 15:44:15 +08:00
feat(backtest): add --show flag for best strategy equity curve chart
- Add --show/--show-chart flag to run_all_strategies.py - Display dual-axis chart: normalized stock price vs strategy equity curve - Mark buy/sell points with green/orange triangles - Auto-detect Chinese fonts (SimHei/YaHei on Windows, PingFang on macOS) - Fetch stock name via get_stock_quotes for chart title - Add 3 demo screenshots to README with disclaimer - Update README with --show usage and visual examples
This commit is contained in:
@@ -295,8 +295,13 @@ easy-tdx backtest SZ 300308 --strategy-file strategies/expma_cross.py --count 20
|
|||||||
|
|
||||||
```bash
|
```bash
|
||||||
python -X utf8 run_all_strategies.py SZ 300308 --count 2000 --cash 1000000 --adjust QFQ
|
python -X utf8 run_all_strategies.py SZ 300308 --count 2000 --cash 1000000 --adjust QFQ
|
||||||
|
|
||||||
|
# 加 --show 自动弹出最佳策略的资金曲线 vs 股价对比图
|
||||||
|
python -X utf8 run_all_strategies.py SZ 300308 --count 2000 --cash 1000000 --adjust QFQ --show
|
||||||
```
|
```
|
||||||
|
|
||||||
|
`--show` 会用 matplotlib 弹出一个双轴对比窗口:左轴蓝色线是归一化股价,右轴红色线是最佳策略的资金曲线,绿三角=买入、黄三角=卖出,标题显示股票名称和关键绩效指标。需要 `pip install matplotlib`。
|
||||||
|
|
||||||
输出示例(以 SZ 300308 为例):
|
输出示例(以 SZ 300308 为例):
|
||||||
|
|
||||||
```
|
```
|
||||||
@@ -345,6 +350,25 @@ python -X utf8 run_all_strategies.py SZ 300308 --count 2000 --cash 1000000 --adj
|
|||||||
python -X utf8 run_all_strategies.py SH 600519 --count 2000 --cash 1000000 --adjust QFQ
|
python -X utf8 run_all_strategies.py SH 600519 --count 2000 --cash 1000000 --adjust QFQ
|
||||||
```
|
```
|
||||||
|
|
||||||
|
#### `--show` 可视化效果
|
||||||
|
|
||||||
|
<p align="center">
|
||||||
|
<img src="strategies/demo/1.png" width="700"><br>
|
||||||
|
<sub>SH601088 中国神华 — bollinger_breakout 策略 | 收益 1281.8%</sub>
|
||||||
|
</p>
|
||||||
|
|
||||||
|
<p align="center">
|
||||||
|
<img src="strategies/demo/2.png" width="700"><br>
|
||||||
|
<sub>SH600522 中天科技 — kdj_golden 策略 | 收益 568.6%</sub>
|
||||||
|
</p>
|
||||||
|
|
||||||
|
<p align="center">
|
||||||
|
<img src="strategies/demo/3.png" width="700"><br>
|
||||||
|
<sub>SH601179 中国西电 — expma_cross 策略 | 收益 168.0%</sub>
|
||||||
|
</p>
|
||||||
|
|
||||||
|
> **⚠️ Demo 展示,不作为操作依据。** 历史回测收益不代表未来表现,策略参数未经过样本外验证。
|
||||||
|
|
||||||
#### 自带策略示例
|
#### 自带策略示例
|
||||||
|
|
||||||
`strategies/` 目录下有 9 个开箱即用的策略文件,可直接用于 `--strategy-file`:
|
`strategies/` 目录下有 9 个开箱即用的策略文件,可直接用于 `--strategy-file`:
|
||||||
|
|||||||
+146
-1
@@ -9,7 +9,6 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import json
|
|
||||||
import sys
|
import sys
|
||||||
import time
|
import time
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -21,6 +20,128 @@ sys.path.insert(0, str(Path(__file__).parent / "src"))
|
|||||||
import click
|
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 _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.command()
|
||||||
@click.argument("market")
|
@click.argument("market")
|
||||||
@click.argument("code")
|
@click.argument("code")
|
||||||
@@ -29,6 +150,7 @@ import click
|
|||||||
@click.option("--commission", default=0.0003, type=float, help="佣金率")
|
@click.option("--commission", default=0.0003, type=float, help="佣金率")
|
||||||
@click.option("--adjust", default="QFQ", help="复权: NONE/QFQ/HFQ")
|
@click.option("--adjust", default="QFQ", help="复权: NONE/QFQ/HFQ")
|
||||||
@click.option("--period", default="DAILY", help="K线周期")
|
@click.option("--period", default="DAILY", help="K线周期")
|
||||||
|
@click.option("--show", "show_chart", is_flag=True, help="显示最佳策略资金曲线 vs 股价对比图")
|
||||||
def run_all(
|
def run_all(
|
||||||
market: str,
|
market: str,
|
||||||
code: str,
|
code: str,
|
||||||
@@ -37,6 +159,7 @@ def run_all(
|
|||||||
commission: float,
|
commission: float,
|
||||||
adjust: str,
|
adjust: str,
|
||||||
period: str,
|
period: str,
|
||||||
|
show_chart: bool,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""批量运行 strategies/ 目录下所有策略并比较结果。"""
|
"""批量运行 strategies/ 目录下所有策略并比较结果。"""
|
||||||
from easy_tdx.backtest.engine import BacktestEngine
|
from easy_tdx.backtest.engine import BacktestEngine
|
||||||
@@ -60,6 +183,7 @@ def run_all(
|
|||||||
click.echo("正在获取行情数据...")
|
click.echo("正在获取行情数据...")
|
||||||
client = MacClient.from_best_host()
|
client = MacClient.from_best_host()
|
||||||
client.connect()
|
client.connect()
|
||||||
|
stock_name = ""
|
||||||
try:
|
try:
|
||||||
df = client.get_stock_kline(
|
df = client.get_stock_kline(
|
||||||
mkt,
|
mkt,
|
||||||
@@ -69,6 +193,14 @@ def run_all(
|
|||||||
count=count,
|
count=count,
|
||||||
adjust=parse_adjust(adjust),
|
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:
|
finally:
|
||||||
client.close()
|
client.close()
|
||||||
click.echo(f"获取到 {len(df)} 条K线数据")
|
click.echo(f"获取到 {len(df)} 条K线数据")
|
||||||
@@ -268,6 +400,19 @@ def run_all(
|
|||||||
for r in errored:
|
for r in errored:
|
||||||
click.echo(f" {r['strategy']}: {r['error']}")
|
click.echo(f" {r['strategy']}: {r['error']}")
|
||||||
|
|
||||||
|
# 5. 展示最佳策略曲线图
|
||||||
|
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__":
|
if __name__ == "__main__":
|
||||||
run_all()
|
run_all()
|
||||||
|
|||||||
Binary file not shown.
|
After Width: | Height: | Size: 165 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 144 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 174 KiB |
Reference in New Issue
Block a user