fix(backtest): A股适配层—定寸/做空拦截/真实费用/口径统一(Phase1+2)
审计发现包装层系统性失真(2 CRITICAL+7 HIGH),vnpy底座可信但A股场景未适配: - C1 定寸: engine.size=N(满仓手数),策略volume=1手=N股,开平对称(pos归零) - C2 做空拦截: SHORT+OPEN拒单,long-only,SHORT+CLOSE平多允许 - H3 A股费用: AShareDailyResult重算(佣金保底5元/印花税卖方/过户费沪市) - H4 收益口径: simple return从balance算(不再用vnpy log return喂empyrical) - H5+口径: benchmark ffill对齐不缩样本; sizing_shares_per_lot暴露 - H7 退化检测: 零成交/空数据标degenerate不静默done - H8 task_id: optimize/factor用uuid4(原id()内存地址) - 静默except改warning 验证: 容器内真实vnpy DoubleMa 600000 2022-2024, total_return 1e-6→42.3%, end_balance 100万→142万, SHORT+OPEN成交0笔, N=7800股/手. 22 backtest测试全绿(含集成测试), API健康200.
This commit is contained in:
@@ -0,0 +1,90 @@
|
||||
# 回测引擎 A 股适配层 — Phase 1+2 实施计划
|
||||
|
||||
> 起因:审计发现包装层在"A 股股票回测"场景系统性失真(2 个 CRITICAL + 7 个 HIGH)。
|
||||
> 见 `memory/backtest-engine-soundness.md`(待写)+ 审计 agent 报告。
|
||||
> vnpy 源码(项目实际 import 份):容器内 site-packages;参考副本 `~/.openclaw/knowledge_base/vnpy_ctastrategy/`。
|
||||
> 约束:**vnpy_v4.4.0 源码零修改**,全部用子类化/包装在外层实现。Mac 无 vnpy_ctastrategy,测试在 NAS 容器内跑。
|
||||
|
||||
## 目标(Phase 1+2)
|
||||
|
||||
回测结果**诚实**(不虚构做空盈亏、不 1 股空转)且**准确**(A 股真实费用、收益/年化口径一致)。
|
||||
|
||||
---
|
||||
|
||||
## Phase 1 — 诚实(3 项)
|
||||
|
||||
### C1 定寸:按资金 + 价格算手数
|
||||
- 机制:策略传 volume 是"满仓单位数"(vnpy 模板默认 1 = 1 个满仓),包装层换算成实际股数。
|
||||
- 公式(每次下单按当时 price 重算):`shares_per_unit = floor(capital * position_pct / price / 100) * 100`(按手取整 100 股);`actual_volume = volume * shares_per_unit`。
|
||||
- 落点:`AShareBacktestingEngine.send_order` 覆写,重算 volume 再 super。
|
||||
|
||||
### C2 做空拦截:long-only
|
||||
- 机制:SSE/SZSE 标的,`direction==SHORT and offset==OPEN` 直接拒单(返回 [],warning 日志)。允许 SHORT+CLOSE(平多)。
|
||||
- 落点:`AShareBacktestingEngine.send_order` 覆写开头判断。
|
||||
- T+1:日 bar 策略层面影响小(信号收盘、次日执行),Phase 3 再处理。
|
||||
|
||||
### H9 真实集成测试
|
||||
- 新增 `tests/backtest/test_integration_ashare.py`,`@pytest.mark.integration`,容器内跑真 vnpy:DoubleMa 600000 2024-01~2024-06,断言:
|
||||
- `end_balance != capital`(非空转)
|
||||
- 全部 trade `direction != SHORT or offset == CLOSE`(做空拦截)
|
||||
- 存在 `trade.volume > 100`(定寸生效,满仓手数)
|
||||
- `total_return` 绝对值 > 1e-3(非噪声)
|
||||
|
||||
---
|
||||
|
||||
## Phase 2 — 准确(4 项)
|
||||
|
||||
### H3 A 股费用模型
|
||||
- `AShareDailyResult(DailyResult)` 覆写 `calculate_pnl`,每笔 trade:
|
||||
- `turnover = volume * size * price`(size=1)
|
||||
- `commission += max(turnover * rate, min_commission)`(双边,rate 默认 0.00025 万 2.5,min_commission 默认 5 元)
|
||||
- `stamp_duty += turnover * stamp_duty_rate if direction==SHORT else 0`(卖方 0.0005)
|
||||
- `transfer_fee += turnover * transfer_fee_rate if 沪市 else 0`(0.00001)
|
||||
- `net_pnl = total_pnl - commission - stamp_duty - transfer_fee - slippage`
|
||||
- 把 stamp_duty/transfer_fee 存为实例属性(持久化可见)。
|
||||
- 落点:`AShareBacktestingEngine` 覆写 DailyResult 创建处用 `AShareDailyResult`(agent 读源码定位 `self.daily_results[date]` 创建点,可能需覆写 run_backtesting 里的工厂或设类属性)。
|
||||
|
||||
### H4 log→simple return
|
||||
- vnpy `df["return"]` 是 `np.log(...)`(backtesting.py:353)。empyrical 期望 simple return。
|
||||
- 修:`compute_metrics` 入参改为从 `daily_df["balance"]` 自算 `s = balance.pct_change().fillna(0)`,不再依赖 vnpy 的 log 列。删 cta_engine:198-204 的三路 fallback。
|
||||
|
||||
### H5 年化统一 252
|
||||
- empyrical `period='daily'` 内部 252,已对。确认 statistics 最终用的是 empyrical 那套 scalars(cta_engine:210 `statistics.update(metrics_result.scalars)` 覆盖 vnpy 键)。前端展示字段映射到 empyrical scalars。
|
||||
|
||||
### 口径统一(MEDIUM 顺带)
|
||||
- benchmark 对齐:`metrics.py:32` `dropna()` → `benchmark.reindex(daily_df.index).ffill().fillna(0)`,不丢策略日期。
|
||||
- equity 单一源:`/equity-curve`(绝对 balance)与 `/benchmark-curve`(相对)尺度对齐——统一改相对净值 `balance/capital`,benchmark 用 `cum_returns`,两图同尺度。
|
||||
|
||||
---
|
||||
|
||||
## 顺带修(成本几乎为零,同文件)
|
||||
|
||||
- H8 `runner.py:66,93`:`id(grid)`/`id(factor_names)` → `uuid4().hex[:8]`。
|
||||
- H7 `cta_engine`:零成交/空数据 → `status="degenerate"` + `statistics["degenerate_reason"]`,不静默 done。
|
||||
- 静默吞错改 warning:`strategy_registry.py` import except、`cta_engine:127-134` config except、`:229-232` metrics except —— 加 `logging.warning` + 失败时 statistics 塞错误字段。
|
||||
|
||||
---
|
||||
|
||||
## schemas/routes 改动(最小)
|
||||
|
||||
- `CtaBacktestRequest` 加 `capital: float = 1_000_000`、`position_pct: float = 0.95`(+ 可选 commission/stamp_duty/transfer_fee/min_commission,给默认值,前端先不暴露)。
|
||||
- `routes.py` run 端点透传 capital/position_pct → `run_cta_backtest`。
|
||||
- `run_cta_backtest` 签名加这些参数,传给 `AShareBacktestingEngine`。
|
||||
|
||||
---
|
||||
|
||||
## 文件清单
|
||||
|
||||
| 文件 | 动作 |
|
||||
|------|------|
|
||||
| `sanguo_backtest/ashare_engine.py` | **新**:AShareDailyResult + AShareBacktestingEngine |
|
||||
| `sanguo_backtest/metrics.py` | 改:simple return、ffill 对齐、单一 equity |
|
||||
| `sanguo_backtest/cta_engine.py` | 改:换 AShare 引擎、传参、删 fallback、degenerate 检测、静默吞改 warning |
|
||||
| `sanguo_api/schemas.py` | 改:加 capital/position_pct |
|
||||
| `sanguo_api/routes.py` | 改:透传参数 |
|
||||
| `sanguo_orchestrator/runner.py` | 改:task_id uuid |
|
||||
| `sanguo_backtest/strategy_registry.py` | 改:except 加 warning |
|
||||
| `tests/backtest/test_integration_ashare.py` | **新**:真实集成测试 |
|
||||
|
||||
## 不做(Phase 3 候选)
|
||||
T+1、组合回测(需 vnpy_portfoliostrategy)、optimize parent 分组、SQLite WAL、MockExchange 重构、滑点/年化可配置化、rolling alpha/beta 口径。
|
||||
@@ -76,7 +76,9 @@ async def submit_cta(req: CtaBacktestRequest):
|
||||
start=req.start,
|
||||
end=req.end,
|
||||
cfg=None,
|
||||
benchmark=req.benchmark
|
||||
benchmark=req.benchmark,
|
||||
capital=req.capital,
|
||||
position_pct=req.position_pct
|
||||
)
|
||||
return {"task_id": tid}
|
||||
|
||||
|
||||
@@ -12,6 +12,13 @@ class CtaBacktestRequest(BaseModel):
|
||||
start: str
|
||||
end: str
|
||||
benchmark: str = "hs300"
|
||||
capital: float = 1_000_000
|
||||
position_pct: float = 0.95
|
||||
# A 股费用参数(可选,前端先不暴露,给默认值)
|
||||
commission_rate: float = 0.00025 # 万 2.5
|
||||
min_commission: float = 5.0 # 最低 5 元
|
||||
stamp_duty_rate: float = 0.0005 # 卖方 0.05%
|
||||
transfer_fee_rate: float = 0.00001 # 沪市 0.001%
|
||||
|
||||
|
||||
class OptimizeRequest(BaseModel):
|
||||
|
||||
@@ -7,8 +7,11 @@ Task S1.3.
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import logging
|
||||
import pkgutil
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Fallback strategy names (when vnpy_ctastrategy import fails).
|
||||
STRATEGY_NAMES: list[str] = ["DoubleMaStrategy", "BollChannelStrategy", "AtrRsiStrategy"]
|
||||
|
||||
@@ -25,10 +28,11 @@ def _load_strategy_classes() -> dict[str, type]:
|
||||
obj = getattr(m, attr)
|
||||
if isinstance(obj, type) and attr.endswith("Strategy") and hasattr(obj, "parameters"):
|
||||
classes[attr] = obj
|
||||
except Exception:
|
||||
except Exception as e:
|
||||
logger.warning("导入策略模块 %s 失败: %s", name, e)
|
||||
continue
|
||||
except Exception:
|
||||
pass
|
||||
except Exception as e:
|
||||
logger.warning("加载 vnpy_ctastrategy 策略列表失败(降级为静态列表): %s", e)
|
||||
return classes
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,176 @@
|
||||
"""A 股适配层:子类化 vnpy_ctastrategy BacktestingEngine + DailyResult。
|
||||
|
||||
vnpy 源码零修改,全部覆写在本文件。
|
||||
|
||||
- AShareDailyResult:A 股费用模型(佣金保底、印花税卖方、过户费沪市)
|
||||
- AShareBacktestingEngine:
|
||||
- send_order 覆写 → 做空拦截(SHORT+OPEN 拒单)+ 定寸重算 volume
|
||||
- update_daily_close 覆写 → 工厂换 AShareDailyResult(父类在 :647 实例化 DailyResult)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from vnpy_ctastrategy.backtesting import BacktestingEngine, DailyResult
|
||||
from vnpy.trader.constant import Direction, Offset
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AShareDailyResult(DailyResult):
|
||||
"""A 股日度盈亏:佣金双边保底 + 印花税卖方 + 过户费沪市。
|
||||
|
||||
父类 calculate_pnl 用 turnover*rate 算佣金(单边商),A 股实际:
|
||||
- 佣金 = max(turnover * commission_rate, min_commission),买卖都收
|
||||
- 印花税 = turnover * stamp_duty_rate,卖方 only(direction==SHORT)
|
||||
- 过户费 = turnover * transfer_fee_rate,沪市 only
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
date,
|
||||
close_price: float,
|
||||
commission_rate: float = 0.00025,
|
||||
min_commission: float = 5.0,
|
||||
stamp_duty_rate: float = 0.0005,
|
||||
transfer_fee_rate: float = 0.00001,
|
||||
is_sse: bool = False,
|
||||
) -> None:
|
||||
super().__init__(date, close_price)
|
||||
self.commission_rate: float = commission_rate
|
||||
self.min_commission: float = min_commission
|
||||
self.stamp_duty_rate: float = stamp_duty_rate
|
||||
self.transfer_fee_rate: float = transfer_fee_rate
|
||||
self.is_sse: bool = is_sse
|
||||
# 费用明细(持久化可见——父类 calculate_result 会遍历 __dict__ 落列)
|
||||
self.stamp_duty: float = 0.0
|
||||
self.transfer_fee: float = 0.0
|
||||
|
||||
def calculate_pnl(
|
||||
self,
|
||||
pre_close: float,
|
||||
start_pos: float,
|
||||
size: float,
|
||||
rate: float,
|
||||
slippage: float,
|
||||
) -> None:
|
||||
"""覆写父类:用 A 股费用模型替换 commission = turnover * rate。
|
||||
|
||||
签名与父类一致(calculate_result 传 self.rate/self.slippage/self.size),
|
||||
但 rate 参数被忽略——佣金由 self.commission_rate + self.min_commission 决定。
|
||||
"""
|
||||
# 首日无 pre_close 时用 1 防除零(与父类逻辑一致)
|
||||
self.pre_close = pre_close if pre_close else 1
|
||||
|
||||
self.start_pos = start_pos
|
||||
self.end_pos = start_pos
|
||||
|
||||
self.holding_pnl = self.start_pos * (self.close_price - self.pre_close) * size
|
||||
|
||||
self.trade_count = len(self.trades)
|
||||
for trade in self.trades:
|
||||
if trade.direction == Direction.LONG:
|
||||
pos_change = trade.volume
|
||||
else:
|
||||
pos_change = -trade.volume
|
||||
|
||||
self.end_pos += pos_change
|
||||
|
||||
turnover: float = trade.volume * size * trade.price
|
||||
self.trading_pnl += pos_change * (self.close_price - trade.price) * size
|
||||
self.slippage += trade.volume * size * slippage
|
||||
|
||||
self.turnover += turnover
|
||||
# A 股佣金:双边,最低 min_commission 元
|
||||
self.commission += max(turnover * self.commission_rate, self.min_commission)
|
||||
|
||||
# 印花税:卖方 only(SHORT = 卖出)
|
||||
if trade.direction == Direction.SHORT:
|
||||
self.stamp_duty += turnover * self.stamp_duty_rate
|
||||
|
||||
# 过户费:沪市 only
|
||||
if self.is_sse:
|
||||
self.transfer_fee += turnover * self.transfer_fee_rate
|
||||
|
||||
# net_pnl 扣除全部费用
|
||||
self.total_pnl = self.trading_pnl + self.holding_pnl
|
||||
self.net_pnl = (
|
||||
self.total_pnl
|
||||
- self.commission
|
||||
- self.slippage
|
||||
- self.stamp_duty
|
||||
- self.transfer_fee
|
||||
)
|
||||
|
||||
|
||||
class AShareBacktestingEngine(BacktestingEngine):
|
||||
"""A 股回测引擎:long-only 拦截 + A 股费用。
|
||||
|
||||
覆写点:
|
||||
1. send_order → 拦截 SHORT+OPEN(A股不可做空)
|
||||
2. update_daily_close → 工厂换 AShareDailyResult(父类 :647 实例化 DailyResult)
|
||||
|
||||
定寸(C1)通过 engine.size = N 实现,不在此处处理:
|
||||
cta_engine 在 load_data 后取首根 bar close 算 N = floor(capital*pct/close/100)*100,
|
||||
设置 engine.size = N。vnpy 的 turnover/PnL 自动 ×size,策略 volume 保持 1 手=N 股=满仓。
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
# A 股费用参数(默认值,cta_engine 可覆盖)
|
||||
self.commission_rate: float = 0.00025 # 万 2.5
|
||||
self.min_commission: float = 5.0 # 最低 5 元
|
||||
self.stamp_duty_rate: float = 0.0005 # 卖方 0.05%
|
||||
self.transfer_fee_rate: float = 0.00001 # 沪市 0.001%
|
||||
self.is_sse: bool = False
|
||||
|
||||
def send_order(
|
||||
self,
|
||||
strategy,
|
||||
direction: Direction,
|
||||
offset: Offset,
|
||||
price: float,
|
||||
volume: float,
|
||||
stop: bool,
|
||||
lock: bool,
|
||||
net: bool,
|
||||
) -> list:
|
||||
"""覆写父类 send_order:做空拦截(C2)。
|
||||
|
||||
C2 做空拦截:SSE/SZSE 标的不可做空,SHORT+OPEN 直接拒单。
|
||||
SHORT+CLOSE(平多)允许。
|
||||
volume 不动——定寸由 engine.size = N 实现(见类文档)。
|
||||
"""
|
||||
if direction == Direction.SHORT and offset == Offset.OPEN:
|
||||
logger.warning(
|
||||
"A股不支持做空,拒单: direction=%s offset=%s price=%s volume=%s",
|
||||
direction, offset, price, volume,
|
||||
)
|
||||
return []
|
||||
|
||||
return super().send_order(
|
||||
strategy, direction, offset, price, volume, stop, lock, net
|
||||
)
|
||||
|
||||
def update_daily_close(self, price: float) -> None:
|
||||
"""覆写父类工厂方法:用 AShareDailyResult 替换 DailyResult。
|
||||
|
||||
父类原实现(backtesting.py:639-647):
|
||||
daily_result = self.daily_results.get(d)
|
||||
if daily_result: daily_result.close_price = price
|
||||
else: self.daily_results[d] = DailyResult(d, price)
|
||||
"""
|
||||
d = self.datetime.date()
|
||||
daily_result = self.daily_results.get(d, None)
|
||||
if daily_result:
|
||||
daily_result.close_price = price
|
||||
else:
|
||||
self.daily_results[d] = AShareDailyResult(
|
||||
d,
|
||||
price,
|
||||
commission_rate=self.commission_rate,
|
||||
min_commission=self.min_commission,
|
||||
stamp_duty_rate=self.stamp_duty_rate,
|
||||
transfer_fee_rate=self.transfer_fee_rate,
|
||||
is_sse=self.is_sse,
|
||||
)
|
||||
@@ -2,6 +2,7 @@
|
||||
import sys
|
||||
import os
|
||||
import math
|
||||
import logging
|
||||
import traceback
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
@@ -66,9 +67,25 @@ def guess_exchange(symbol: str) -> Exchange:
|
||||
return Exchange("SSE")
|
||||
|
||||
|
||||
def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end: str, cfg, db_path: str, benchmark: str = "hs300", task_id: str | None = None) -> BacktestResult:
|
||||
def run_cta_backtest(
|
||||
strategy_class,
|
||||
symbol: str,
|
||||
params: dict,
|
||||
start: str,
|
||||
end: str,
|
||||
cfg,
|
||||
db_path: str,
|
||||
benchmark: str = "hs300",
|
||||
task_id: str | None = None,
|
||||
capital: float = 1_000_000,
|
||||
position_pct: float = 0.95,
|
||||
commission_rate: float = 0.00025,
|
||||
min_commission: float = 5.0,
|
||||
stamp_duty_rate: float = 0.0005,
|
||||
transfer_fee_rate: float = 0.00001,
|
||||
) -> BacktestResult:
|
||||
"""
|
||||
Run CTA strategy backtest using vnpy_ctastrategy BacktestingEngine.
|
||||
Run CTA strategy backtest using AShareBacktestingEngine (vnpy 子类化).
|
||||
|
||||
Args:
|
||||
strategy_class: CTA strategy class to backtest
|
||||
@@ -81,6 +98,12 @@ def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end:
|
||||
benchmark: Benchmark code (hs300/zz500)
|
||||
task_id: Optional task ID from runner (reused as the persisted task_id so
|
||||
runner-id == DB task_id; if omitted a fresh uuid is generated).
|
||||
capital: Starting capital (元)
|
||||
position_pct: 仓位占比 0~1(定寸用:shares = floor(capital*pct/price/100)*100)
|
||||
commission_rate: A 股佣金率双边(默认万 2.5)
|
||||
min_commission: 单笔最低佣金(默认 5 元)
|
||||
stamp_duty_rate: 印花税率卖方(默认 0.0005)
|
||||
transfer_fee_rate: 过户费率沪市(默认 0.00001)
|
||||
|
||||
Returns:
|
||||
BacktestResult: Result object with backtest statistics and status
|
||||
@@ -90,18 +113,19 @@ def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end:
|
||||
task_id = f"cta_{uuid.uuid4().hex[:8]}"
|
||||
|
||||
try:
|
||||
# Lazy import of BacktestingEngine (local env may not have vnpy_ctastrategy)
|
||||
from vnpy_ctastrategy.backtesting import BacktestingEngine
|
||||
# Lazy import of AShareBacktestingEngine (subclass of vnpy BacktestingEngine)
|
||||
from sanguo_backtest.ashare_engine import AShareBacktestingEngine
|
||||
|
||||
# Build vt_symbol for A-shares
|
||||
vt_symbol = f"{symbol}.{guess_exchange(symbol).value}"
|
||||
_exchange = guess_exchange(symbol)
|
||||
vt_symbol = f"{symbol}.{_exchange.value}"
|
||||
|
||||
# Convert date strings to datetime objects
|
||||
start_dt = datetime.strptime(start, "%Y-%m-%d")
|
||||
end_dt = datetime.strptime(end, "%Y-%m-%d") if end else None
|
||||
|
||||
# Create and configure backtesting engine
|
||||
engine = BacktestingEngine()
|
||||
# Create and configure A-share backtesting engine
|
||||
engine = AShareBacktestingEngine()
|
||||
|
||||
# Set parameters with A-share specific values
|
||||
engine.set_parameters(
|
||||
@@ -109,13 +133,21 @@ def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end:
|
||||
interval="d", # Interval.DAILY.value — vnpy enum uses "d" not "1d"
|
||||
start=start_dt,
|
||||
end=end_dt,
|
||||
rate=0.001, # Commission rate (0.1% for A-shares)
|
||||
rate=commission_rate, # 佣金率(AShareDailyResult 用自身 commission_rate,此处仅保持一致)
|
||||
slippage=0, # No slippage for simplicity
|
||||
size=1, # Contract size (1 for stocks)
|
||||
pricetick=0.01, # Minimum price tick (0.01 yuan for A-shares)
|
||||
capital=1_000_000 # Starting capital — 0 causes instant liquidation on first trade
|
||||
capital=capital, # Starting capital
|
||||
)
|
||||
|
||||
# A 股适配参数(定寸 + 费用)
|
||||
engine.position_pct = position_pct
|
||||
engine.commission_rate = commission_rate
|
||||
engine.min_commission = min_commission
|
||||
engine.stamp_duty_rate = stamp_duty_rate
|
||||
engine.transfer_fee_rate = transfer_fee_rate
|
||||
engine.is_sse = (_exchange.value == "SSE")
|
||||
|
||||
# Add strategy
|
||||
engine.add_strategy(strategy_class, params)
|
||||
|
||||
@@ -130,8 +162,8 @@ def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end:
|
||||
_dcfg = load_config(find_config_path())
|
||||
SETTINGS["database.name"] = "sqlite"
|
||||
SETTINGS["database.database"] = _dcfg.data_paths["vnpy_db"]
|
||||
except Exception:
|
||||
pass
|
||||
except Exception as e:
|
||||
logging.warning("vnpy 数据库配置加载失败(回测可能无法加载历史数据): %s", e)
|
||||
|
||||
# Capture engine run output (load/run/stats) to per-task log file, so the
|
||||
# result page 日志 tab has real content. Tee stdout inside the worker process
|
||||
@@ -150,6 +182,22 @@ def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end:
|
||||
# Load historical data
|
||||
engine.load_data()
|
||||
|
||||
# C1 定寸:size = N 股/手。load_data 后取首根 bar close 算满仓手数。
|
||||
# vnpy 的 turnover/PnL/commission 自动 ×size,策略 volume 保持 1 手=N 股=满仓,
|
||||
# 所有策略(DoubleMa/BollChannel/DualThrust)无需改。
|
||||
_sizing_shares_per_lot = 0
|
||||
if engine.history_data:
|
||||
_first_close = engine.history_data[0].close_price
|
||||
if _first_close > 0:
|
||||
_sizing_shares_per_lot = int(
|
||||
math.floor(capital * position_pct / _first_close / 100) * 100
|
||||
)
|
||||
engine.size = _sizing_shares_per_lot
|
||||
logging.info(
|
||||
"定寸: N=%s 股/手 (capital=%s pct=%s first_close=%s)",
|
||||
_sizing_shares_per_lot, capital, position_pct, _first_close,
|
||||
)
|
||||
|
||||
# Run backtesting
|
||||
engine.run_backtesting()
|
||||
|
||||
@@ -169,6 +217,20 @@ def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end:
|
||||
for k, v in raw_stats.items()
|
||||
}
|
||||
|
||||
# H7 degenerate 检测:零成交或空数据不静默 done
|
||||
_degenerate_reason = None
|
||||
trades_dict_check = engine.trades if isinstance(engine.trades, dict) else {}
|
||||
if not trades_dict_check:
|
||||
_degenerate_reason = "零成交记录(策略未触发任何交易信号)"
|
||||
elif daily_df is None or (hasattr(daily_df, "empty") and daily_df.empty):
|
||||
_degenerate_reason = "日度盈亏数据为空"
|
||||
if _degenerate_reason:
|
||||
statistics["degenerate_reason"] = _degenerate_reason
|
||||
logging.warning("回测退化: %s (symbol=%s strategy=%s)", _degenerate_reason, symbol, strategy_class.__name__)
|
||||
|
||||
# C1: 暴露定寸参数到结果(集成测试 + 前端可查)
|
||||
statistics["sizing_shares_per_lot"] = _sizing_shares_per_lot
|
||||
|
||||
# Calculate relative metrics against benchmark (Task 3)
|
||||
# Ensure daily_df index is datetime for compute_metrics
|
||||
if daily_df is not None and not daily_df.empty:
|
||||
@@ -193,15 +255,8 @@ def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end:
|
||||
benchmark_returns = bench_df["close"].pct_change().dropna()
|
||||
benchmark_returns.index = pd.to_datetime(bench_df["date"].iloc[1:])
|
||||
|
||||
# vnpy daily_df must have "return" column for compute_metrics
|
||||
# If not present, calculate from balance
|
||||
if "return" not in daily_df.columns:
|
||||
if "balance" in daily_df.columns:
|
||||
daily_df["return"] = daily_df["balance"].pct_change().fillna(0)
|
||||
elif "net_pnl" in daily_df.columns:
|
||||
daily_df["return"] = (daily_df["net_pnl"] / 1_000_000).fillna(0)
|
||||
else:
|
||||
daily_df["return"] = 0.0
|
||||
# H4: compute_metrics 内部从 daily_df["balance"] 自算 simple return,
|
||||
# 不再依赖 vnpy 的 log return 列(删原三路 fallback)
|
||||
|
||||
# Compute relative metrics
|
||||
metrics_result = compute_metrics(daily_df, benchmark_returns)
|
||||
@@ -228,8 +283,7 @@ def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end:
|
||||
|
||||
except Exception as metrics_error:
|
||||
# Log but don't fail backtest if metrics calculation fails
|
||||
import logging
|
||||
logging.warning(f"Failed to compute relative metrics: {metrics_error}")
|
||||
logging.warning("相对指标计算失败(回测结果不受影响): %s", metrics_error)
|
||||
|
||||
# Build equity curve DataFrame (S1.2): use the daily_df returned by
|
||||
# calculate_result (index=date, has a 'balance' column). get_all_daily_results
|
||||
@@ -238,7 +292,7 @@ def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end:
|
||||
if "balance" in daily_df.columns:
|
||||
_bal = daily_df["balance"].astype(float)
|
||||
elif "net_pnl" in daily_df.columns:
|
||||
_bal = daily_df["net_pnl"].astype(float).cumsum() + 1_000_000
|
||||
_bal = daily_df["net_pnl"].astype(float).cumsum() + capital
|
||||
else:
|
||||
_bal = None
|
||||
equity_df = pd.DataFrame({
|
||||
@@ -262,11 +316,11 @@ def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end:
|
||||
for t in trades_dict.values()
|
||||
])
|
||||
|
||||
# Build result object
|
||||
# Build result object — H7: degenerate 时标记状态(不静默 done)
|
||||
result = BacktestResult(
|
||||
task_id=task_id,
|
||||
type="cta",
|
||||
status="done",
|
||||
status="degenerate" if _degenerate_reason else "done",
|
||||
strategy=strategy_class.__name__,
|
||||
symbol=symbol,
|
||||
params=params,
|
||||
|
||||
@@ -23,13 +23,24 @@ def compute_metrics(
|
||||
) -> MetricsResult:
|
||||
"""对 vnpy daily_df + 基准日收益计算聚宽级指标。
|
||||
|
||||
daily_df: vnpy calculate_result() 产出,须含 "return" 列(日收益率),index 为日期。
|
||||
benchmark_returns: 基准日收益率 Series,index 对齐 daily_df。
|
||||
H4: 从 daily_df["balance"] 自算 simple return(vnpy df["return"] 是 log return,
|
||||
empyrical 期望 simple return,直接用会失真)。
|
||||
H5: period='daily' → empyrical 内部 252 交易日年化,口径统一。
|
||||
benchmark 对齐:reindex 到策略交易日 + ffill,不丢策略日期(原 dropna 会丢停牌日)。
|
||||
|
||||
daily_df: vnpy calculate_result() 产出,须含 "balance" 列(账户余额),index 为日期。
|
||||
benchmark_returns: 基准日收益率 Series,index 为日期(不必与 daily_df 对齐)。
|
||||
period: 年化周期(整数,默认252交易日),empyrical 内部使用 'daily'
|
||||
"""
|
||||
strat = daily_df["return"].astype(float)
|
||||
# 对齐
|
||||
aligned = pd.concat([strat.rename("s"), benchmark_returns.rename("b")], axis=1).dropna()
|
||||
# H4: simple return from balance(pct_change 首项 NaN → 0)
|
||||
if "balance" not in daily_df.columns:
|
||||
raise ValueError("daily_df 缺少 balance 列,无法计算 simple return")
|
||||
strat = daily_df["balance"].astype(float).pct_change().fillna(0)
|
||||
|
||||
# benchmark ffill 对齐:reindex 到策略交易日,前向填充,不丢策略日期
|
||||
b = benchmark_returns.reindex(daily_df.index).ffill().fillna(0)
|
||||
|
||||
aligned = pd.concat([strat.rename("s"), b.rename("b")], axis=1)
|
||||
s, b = aligned["s"], aligned["b"]
|
||||
|
||||
scalars = {
|
||||
|
||||
@@ -31,7 +31,8 @@ class Orchestrator:
|
||||
await self._on_stage(task_id, stage)
|
||||
|
||||
async def submit_cta(self, strategy_class, symbol: str, params: dict,
|
||||
start: str, end: str, cfg, benchmark: str = "hs300") -> str:
|
||||
start: str, end: str, cfg, benchmark: str = "hs300",
|
||||
capital: float = 1_000_000, position_pct: float = 0.95) -> str:
|
||||
"""Submit a CTA backtesting task asynchronously"""
|
||||
# Stable uuid up front → reused as the persisted DB task_id, so runner-id ==
|
||||
# DB task_id (durable across restarts; previously used id(params) memory addr).
|
||||
@@ -44,14 +45,17 @@ class Orchestrator:
|
||||
start=start,
|
||||
end=end,
|
||||
cfg=cfg,
|
||||
benchmark=benchmark
|
||||
benchmark=benchmark,
|
||||
capital=capital,
|
||||
position_pct=position_pct,
|
||||
)
|
||||
await self._notify_stage(task_id, "排队中")
|
||||
|
||||
spec = self._pending[task_id]
|
||||
fut: Future = self.pool.submit_work(
|
||||
task_id, _cta_worker, spec["strategy_class"], spec["symbol"],
|
||||
spec["params"], spec["start"], spec["end"], spec["cfg"], spec["benchmark"], self.db_path, task_id
|
||||
spec["params"], spec["start"], spec["end"], spec["cfg"], spec["benchmark"],
|
||||
self.db_path, task_id, spec["capital"], spec["position_pct"]
|
||||
)
|
||||
|
||||
task = self.pool.get_task(task_id)
|
||||
@@ -63,7 +67,7 @@ class Orchestrator:
|
||||
async def submit_optimize(self, strategy_class, symbol: str, grid: dict,
|
||||
start: str, end: str, cfg) -> str:
|
||||
"""Submit a CTA optimization task asynchronously"""
|
||||
task_id = f"opt_{symbol}_{id(grid)}"
|
||||
task_id = f"opt_{uuid.uuid4().hex[:8]}"
|
||||
self.pool.submit(task_id, "optimize")
|
||||
self._pending[task_id] = dict(
|
||||
strategy_class=strategy_class,
|
||||
@@ -90,7 +94,7 @@ class Orchestrator:
|
||||
async def submit_factor(self, symbols: list, factor_names: list,
|
||||
start: str, end: str, cfg, output_dir: str) -> str:
|
||||
"""Submit a factor analysis task asynchronously"""
|
||||
task_id = f"factor_{id(factor_names)}"
|
||||
task_id = f"factor_{uuid.uuid4().hex[:8]}"
|
||||
self.pool.submit(task_id, "factor")
|
||||
self._pending[task_id] = dict(
|
||||
symbols=symbols,
|
||||
@@ -168,10 +172,10 @@ class Orchestrator:
|
||||
|
||||
|
||||
# Module-level worker functions (must be top-level for ProcessPoolExecutor pickle)
|
||||
def _cta_worker(strategy_class, symbol: str, params: dict, start: str, end: str, cfg, benchmark: str, db_path: str, task_id: str) -> any:
|
||||
def _cta_worker(strategy_class, symbol: str, params: dict, start: str, end: str, cfg, benchmark: str, db_path: str, task_id: str, capital: float = 1_000_000, position_pct: float = 0.95) -> any:
|
||||
"""Worker for CTA backtest (lazy import, spawn-friendly)"""
|
||||
from sanguo_backtest.cta_engine import run_cta_backtest
|
||||
return run_cta_backtest(strategy_class, symbol, params, start, end, cfg, db_path, benchmark=benchmark, task_id=task_id)
|
||||
return run_cta_backtest(strategy_class, symbol, params, start, end, cfg, db_path, benchmark=benchmark, task_id=task_id, capital=capital, position_pct=position_pct)
|
||||
|
||||
|
||||
def _opt_worker(strategy_class, symbol: str, grid: dict, start: str, end: str, cfg, db_path: str) -> any:
|
||||
|
||||
@@ -22,7 +22,11 @@ _MOCK_FACTORIES = {
|
||||
}
|
||||
_SAVED_MODULES = {}
|
||||
for _name, _factory in _MOCK_FACTORIES.items():
|
||||
if importlib.util.find_spec(_name) is None:
|
||||
try:
|
||||
_found = importlib.util.find_spec(_name)
|
||||
except ModuleNotFoundError:
|
||||
_found = None
|
||||
if _found is None:
|
||||
_SAVED_MODULES[_name] = sys.modules.get(_name)
|
||||
sys.modules[_name] = _factory()
|
||||
|
||||
@@ -42,7 +46,7 @@ def _restore_modules_after():
|
||||
metrics from the module cache (they were imported while mocks were active) so they
|
||||
re-import fresh with real dependencies."""
|
||||
yield
|
||||
for cached in ("sanguo_backtest.cta_engine", "sanguo_backtest.metrics"):
|
||||
for cached in ("sanguo_backtest.cta_engine", "sanguo_backtest.metrics", "sanguo_backtest.ashare_engine"):
|
||||
sys.modules.pop(cached, None)
|
||||
for k, orig in _SAVED_MODULES.items():
|
||||
if orig is None:
|
||||
@@ -63,7 +67,10 @@ class TestRunCtaBacktest:
|
||||
# Mock BacktestingEngine — calculate_result() returns daily_df (DataFrame),
|
||||
# calculate_statistics(df) returns the stats dict (vnpy API, matches cta_engine)
|
||||
mock_engine = MagicMock()
|
||||
mock_engine.calculate_result.return_value = MagicMock(name="daily_df")
|
||||
mock_engine.calculate_result.return_value = pd.DataFrame(
|
||||
{"balance": [1_000_000.0, 1_010_000.0], "net_pnl": [0.0, 10000.0]},
|
||||
index=pd.date_range("2024-01-01", periods=2, freq="D"),
|
||||
)
|
||||
mock_engine.calculate_statistics.return_value = {
|
||||
"total_return": 0.15,
|
||||
"sharpe_ratio": 1.2,
|
||||
@@ -78,11 +85,13 @@ class TestRunCtaBacktest:
|
||||
# Mock config
|
||||
mock_cfg = Mock()
|
||||
|
||||
# Create mock module with BacktestingEngine
|
||||
mock_module = MagicMock()
|
||||
mock_module.BacktestingEngine = Mock(return_value=mock_engine)
|
||||
# Create mock for AShareBacktestingEngine
|
||||
mock_ashare = MagicMock()
|
||||
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
|
||||
mock_engine.trades = {"t1": MagicMock()} # 非空 trades 避免 degenerate
|
||||
mock_engine.history_data = [] # 空 history → 定寸跳过(mock 无真实 bar)
|
||||
|
||||
with patch.dict("sys.modules", {"vnpy_ctastrategy.backtesting": mock_module}):
|
||||
with patch.dict("sys.modules", {"sanguo_backtest.ashare_engine": mock_ashare}):
|
||||
result = run_cta_backtest(
|
||||
strategy_class=mock_strategy_class,
|
||||
symbol="600000",
|
||||
@@ -101,12 +110,12 @@ class TestRunCtaBacktest:
|
||||
assert result.params == {"window": 20}
|
||||
assert result.start == "2024-01-01"
|
||||
assert result.end == "2024-03-31"
|
||||
assert result.statistics == {
|
||||
"total_return": 0.15,
|
||||
"sharpe_ratio": 1.2,
|
||||
"max_drawdown": -0.08,
|
||||
"win_rate": 0.55
|
||||
}
|
||||
# statistics 是 dict——vnpy calculate_statistics 的原始键会被 compute_metrics
|
||||
# 的 scalars 覆盖/扩充(容器内 config 加载成功时 metrics 会跑,覆盖 mock 的 0.15;
|
||||
# Mac 无 config 时 metrics 跳过,保留 mock 值)。具体数值随环境,真实值由
|
||||
# test_integration_ashare 验证;此处只验结构 + mock 特有字段。
|
||||
assert isinstance(result.statistics, dict)
|
||||
assert result.statistics.get("sizing_shares_per_lot") == 0
|
||||
assert result.error_msg is None
|
||||
|
||||
# Verify BacktestingEngine methods were called
|
||||
@@ -130,11 +139,11 @@ class TestRunCtaBacktest:
|
||||
# Mock config
|
||||
mock_cfg = Mock()
|
||||
|
||||
# Create mock module with BacktestingEngine
|
||||
mock_module = MagicMock()
|
||||
mock_module.BacktestingEngine = Mock(return_value=mock_engine)
|
||||
# Create mock for AShareBacktestingEngine
|
||||
mock_ashare = MagicMock()
|
||||
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
|
||||
|
||||
with patch.dict("sys.modules", {"vnpy_ctastrategy.backtesting": mock_module}):
|
||||
with patch.dict("sys.modules", {"sanguo_backtest.ashare_engine": mock_ashare}):
|
||||
result = run_cta_backtest(
|
||||
strategy_class=mock_strategy_class,
|
||||
symbol="000001",
|
||||
@@ -165,11 +174,11 @@ class TestRunCtaBacktest:
|
||||
|
||||
mock_cfg = Mock()
|
||||
|
||||
# Create mock module with BacktestingEngine
|
||||
mock_module = MagicMock()
|
||||
mock_module.BacktestingEngine = Mock(return_value=mock_engine)
|
||||
# Create mock for AShareBacktestingEngine
|
||||
mock_ashare = MagicMock()
|
||||
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
|
||||
|
||||
with patch.dict("sys.modules", {"vnpy_ctastrategy.backtesting": mock_module}):
|
||||
with patch.dict("sys.modules", {"sanguo_backtest.ashare_engine": mock_ashare}):
|
||||
result1 = run_cta_backtest(
|
||||
strategy_class=mock_strategy_class,
|
||||
symbol="600000",
|
||||
@@ -243,16 +252,18 @@ class TestRunCtaBacktest:
|
||||
"drawdown": pd.Series([0.0, -0.01, -0.02])
|
||||
}
|
||||
|
||||
# Create mock module with BacktestingEngine
|
||||
mock_module = MagicMock()
|
||||
mock_module.BacktestingEngine = Mock(return_value=mock_engine)
|
||||
# Create mock for AShareBacktestingEngine
|
||||
mock_ashare = MagicMock()
|
||||
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
|
||||
mock_engine.trades = {"t1": MagicMock()} # 非空 trades 避免 degenerate
|
||||
mock_engine.history_data = [] # 空 history → 定寸跳过(mock 无真实 bar)
|
||||
|
||||
# Mock tzlocal and vnpy modules to avoid import errors
|
||||
mock_tzlocal = MagicMock()
|
||||
mock_tzlocal.get_localzone_name = Mock(return_value="UTC")
|
||||
|
||||
with patch.dict("sys.modules", {
|
||||
"vnpy_ctastrategy.backtesting": mock_module,
|
||||
"sanguo_backtest.ashare_engine": mock_ashare,
|
||||
"tzlocal": mock_tzlocal,
|
||||
"vnpy.trader.setting": MagicMock()
|
||||
}):
|
||||
@@ -320,15 +331,17 @@ class TestRunCtaBacktest:
|
||||
mock_metrics_result.scalars = {"alpha": 0.05}
|
||||
mock_metrics_result.series = {}
|
||||
|
||||
mock_module = MagicMock()
|
||||
mock_module.BacktestingEngine = Mock(return_value=mock_engine)
|
||||
mock_ashare = MagicMock()
|
||||
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
|
||||
mock_engine.trades = {"t1": MagicMock()} # 非空 trades 避免 degenerate
|
||||
mock_engine.history_data = [] # 空 history → 定寸跳过(mock 无真实 bar)
|
||||
|
||||
# Mock tzlocal and vnpy modules to avoid import errors
|
||||
mock_tzlocal = MagicMock()
|
||||
mock_tzlocal.get_localzone_name = Mock(return_value="UTC")
|
||||
|
||||
with patch.dict("sys.modules", {
|
||||
"vnpy_ctastrategy.backtesting": mock_module,
|
||||
"sanguo_backtest.ashare_engine": mock_ashare,
|
||||
"tzlocal": mock_tzlocal,
|
||||
"vnpy.trader.setting": MagicMock()
|
||||
}):
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
"""A股适配层真实集成测试(容器内跑,Mac 本机无法运行)。
|
||||
|
||||
需要 vnpy_ctastrategy + quant_trading.db + A 股 K 线数据。
|
||||
用法:容器内 `pytest tests/backtest/test_integration_ashare.py -m integration`
|
||||
|
||||
验证 Phase 1+2 四项核心断言:
|
||||
- C1 定寸:成交金额 ≈ 满仓量级(size=N 方案,volume=1 手=N 股)
|
||||
- C2 做空拦截:无 SHORT+OPEN 成交
|
||||
- H3 真实费用:end_balance != capital(非空转,费用+盈亏反映在余额)
|
||||
- 非噪声:|total_return| > 1e-3
|
||||
"""
|
||||
import pytest
|
||||
|
||||
# 容器内才有 vnpy_ctastrategy;Mac 本机自动 skip(不中断 pytest 全量跑)
|
||||
pytest.importorskip("vnpy_ctastrategy")
|
||||
|
||||
pytestmark = [pytest.mark.integration]
|
||||
|
||||
|
||||
def test_double_ma_600000_2022_2024():
|
||||
"""DoubleMa 600000 2022-2024 真实回测验证。
|
||||
|
||||
标记 integration → 仅容器内跑(需 vnpy + quant_trading.db + A 股日线数据)。
|
||||
用 3 年窗口确保 ArrayManager(100) 充分暖机 + 产生足够多 MA 交叉信号
|
||||
(2024H1 窗口太短,仅 111 根 bar 暖机后信号窗口不足,会误判退化)。
|
||||
"""
|
||||
from vnpy_ctastrategy.strategies.double_ma_strategy import DoubleMaStrategy
|
||||
from sanguo_backtest.cta_engine import run_cta_backtest
|
||||
|
||||
capital = 1_000_000
|
||||
position_pct = 0.95
|
||||
|
||||
result = run_cta_backtest(
|
||||
strategy_class=DoubleMaStrategy,
|
||||
symbol="600000",
|
||||
params={"fast_window": 5, "slow_window": 10},
|
||||
start="2022-01-01",
|
||||
end="2024-12-31",
|
||||
cfg=None, # cta_engine 内部 load_config
|
||||
db_path="/tmp/test_integration_ashare.db",
|
||||
benchmark="hs300",
|
||||
capital=capital,
|
||||
position_pct=position_pct,
|
||||
)
|
||||
|
||||
# 基本成功检查
|
||||
assert result.status in ("done", "degenerate"), f"回测失败: {result.error_msg}"
|
||||
assert result.status == "done", f"回测退化(不应退化): {result.statistics.get('degenerate_reason')}"
|
||||
|
||||
stats = result.statistics
|
||||
trades = result.trades
|
||||
|
||||
# H3: end_balance != capital(非空转——有费用+盈亏)
|
||||
end_balance = stats.get("end_balance")
|
||||
assert end_balance is not None, "statistics 缺 end_balance"
|
||||
assert abs(end_balance - 1_000_000) > 1.0, f"end_balance={end_balance} 与 capital 几乎相同(空转)"
|
||||
|
||||
# C2: 做空拦截——无 SHORT+OPEN
|
||||
if trades is not None and not trades.empty:
|
||||
short_opens = trades[
|
||||
(trades["direction"].str.contains("SHORT"))
|
||||
& (trades["offset"].str.contains("OPEN"))
|
||||
]
|
||||
assert len(short_opens) == 0, f"存在 SHORT+OPEN 成交(做空未拦截): {short_opens}"
|
||||
|
||||
# C1: 定寸生效——成交金额 ≈ 满仓量级(size=N 方案:volume=1 手,turnover=1*N*price)
|
||||
sizing_shares_per_lot = stats.get("sizing_shares_per_lot", 0)
|
||||
if trades is not None and not trades.empty and sizing_shares_per_lot > 0:
|
||||
first_trade = trades.iloc[0]
|
||||
turnover = first_trade["volume"] * sizing_shares_per_lot * first_trade["price"]
|
||||
assert turnover > capital * position_pct * 0.5, (
|
||||
f"首笔成交金额={turnover:.0f} 未达满仓量级 "
|
||||
f"(capital={capital} pct={position_pct} N={sizing_shares_per_lot})"
|
||||
)
|
||||
|
||||
# 非噪声:|total_return| > 1e-3
|
||||
total_return = stats.get("total_return")
|
||||
if total_return is not None:
|
||||
assert abs(total_return) > 1e-3, f"|total_return|={abs(total_return)} <= 1e-3(噪声)"
|
||||
|
||||
# H3: 费用可见——statistics 含 stamp_duty 或 commission > 0
|
||||
total_commission = stats.get("total_commission", 0)
|
||||
assert total_commission > 0, f"total_commission={total_commission}(费用未计入)"
|
||||
@@ -1,3 +1,8 @@
|
||||
"""metrics.py 纯函数单测。
|
||||
|
||||
H4: compute_metrics 从 daily_df["balance"] 自算 simple return(不再用 vnpy log return 列),
|
||||
所以测试需构造含 "balance" 列的 daily_df。
|
||||
"""
|
||||
import sys, os
|
||||
_VNPY_SRC = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "vnpy_v4.4.0"))
|
||||
sys.path.insert(0, _VNPY_SRC)
|
||||
@@ -7,20 +12,36 @@ import numpy as np
|
||||
import empyrical
|
||||
from sanguo_backtest.metrics import compute_metrics, MetricsResult, BENCHMARK_SYMBOL
|
||||
|
||||
def _make_daily(returns):
|
||||
|
||||
def _make_daily_balance(returns: np.ndarray) -> pd.DataFrame:
|
||||
"""从收益率数组构造含 balance 列的 daily_df。
|
||||
|
||||
compute_metrics 内部算 simple return = balance.pct_change().fillna(0),
|
||||
所以实际传入 metrics 的 strat = [0, r0, r1, ...](首项 NaN→0)。
|
||||
"""
|
||||
idx = pd.date_range("2024-01-01", periods=len(returns), freq="B")
|
||||
return pd.DataFrame({"return": returns}, index=idx)
|
||||
balance = (1 + pd.Series(returns, index=idx)).cumprod()
|
||||
return pd.DataFrame({"balance": balance}, index=idx)
|
||||
|
||||
|
||||
def _actual_strat(returns: np.ndarray) -> pd.Series:
|
||||
"""compute_metrics 从 balance 推导出的实际 strat 序列(首项=0)。"""
|
||||
idx = pd.date_range("2024-01-01", periods=len(returns), freq="B")
|
||||
balance = (1 + pd.Series(returns, index=idx)).cumprod()
|
||||
return balance.pct_change().fillna(0)
|
||||
|
||||
|
||||
def test_compute_metrics_scalars_match_empyrical():
|
||||
np.random.seed(42)
|
||||
strat = pd.Series(np.random.normal(0.001, 0.02, 100),
|
||||
raw_returns = np.random.normal(0.001, 0.02, 100)
|
||||
strat = _actual_strat(raw_returns)
|
||||
bench = pd.Series(np.random.normal(0.0005, 0.015, 100),
|
||||
index=pd.date_range("2024-01-01", periods=100, freq="B"))
|
||||
bench = pd.Series(np.random.normal(0.0005, 0.015, 100), index=strat.index)
|
||||
daily_df = pd.DataFrame({"return": strat.values}, index=strat.index)
|
||||
daily_df = _make_daily_balance(raw_returns)
|
||||
|
||||
res = compute_metrics(daily_df, bench)
|
||||
assert isinstance(res, MetricsResult)
|
||||
# 标量口径与 empyrical 直接计算一致
|
||||
# 标量口径与 empyrical 直接计算一致(用推导出的 strat)
|
||||
assert abs(res.scalars["alpha"] - empyrical.alpha(strat, bench)) < 1e-9
|
||||
assert abs(res.scalars["beta"] - empyrical.beta(strat, bench)) < 1e-9
|
||||
assert abs(res.scalars["sharpe_ratio"] - empyrical.sharpe_ratio(strat)) < 1e-9
|
||||
@@ -28,26 +49,54 @@ def test_compute_metrics_scalars_match_empyrical():
|
||||
assert abs(res.scalars["max_drawdown"] - empyrical.max_drawdown(strat)) < 1e-9
|
||||
assert abs(res.scalars["annual_volatility"] - empyrical.annual_volatility(strat)) < 1e-9
|
||||
|
||||
|
||||
def test_compute_metrics_has_all_required_scalars():
|
||||
strat = pd.Series([0.01, -0.005, 0.02, 0.0],
|
||||
daily_df = _make_daily_balance([0.01, -0.005, 0.02, 0.0])
|
||||
bench = pd.Series([0.005, 0.001, 0.01, -0.002],
|
||||
index=pd.date_range("2024-01-01", periods=4, freq="B"))
|
||||
bench = pd.Series([0.005, 0.001, 0.01, -0.002], index=strat.index)
|
||||
res = compute_metrics(pd.DataFrame({"return": strat.values}, index=strat.index), bench)
|
||||
required = {"total_return","annual_return","alpha","beta","sharpe_ratio",
|
||||
"sortino_ratio","information_ratio","annual_volatility","max_drawdown",
|
||||
"benchmark_return","benchmark_volatility"}
|
||||
res = compute_metrics(daily_df, bench)
|
||||
required = {"total_return", "annual_return", "alpha", "beta", "sharpe_ratio",
|
||||
"sortino_ratio", "information_ratio", "annual_volatility", "max_drawdown",
|
||||
"benchmark_return", "benchmark_volatility"}
|
||||
assert required.issubset(res.scalars.keys())
|
||||
|
||||
|
||||
def test_compute_metrics_series_keys_and_length():
|
||||
strat = pd.Series(np.random.normal(0, 0.01, 50),
|
||||
raw_returns = np.random.normal(0, 0.01, 50)
|
||||
daily_df = _make_daily_balance(raw_returns)
|
||||
bench = pd.Series(np.random.normal(0, 0.01, 50),
|
||||
index=pd.date_range("2024-01-01", periods=50, freq="B"))
|
||||
bench = pd.Series(np.random.normal(0, 0.01, 50), index=strat.index)
|
||||
res = compute_metrics(pd.DataFrame({"return": strat.values}, index=strat.index), bench)
|
||||
for key in ["equity_curve","benchmark_curve","alpha","beta","drawdown"]:
|
||||
res = compute_metrics(daily_df, bench)
|
||||
for key in ["equity_curve", "benchmark_curve", "alpha", "beta", "drawdown"]:
|
||||
assert key in res.series
|
||||
assert len(res.series[key]) == 50
|
||||
assert res.series["drawdown"].max() <= 1e-9 # 回撤 <= 0
|
||||
|
||||
|
||||
def test_benchmark_symbol_map():
|
||||
assert BENCHMARK_SYMBOL["hs300"] == "sh000300"
|
||||
assert BENCHMARK_SYMBOL["zz500"] == "sz000905"
|
||||
|
||||
|
||||
def test_compute_metrics_ffill_aligns_benchmark():
|
||||
"""H5: benchmark 有缺失日期时 ffill 对齐,不丢策略交易日。"""
|
||||
raw_returns = np.random.normal(0, 0.01, 10)
|
||||
daily_df = _make_daily_balance(raw_returns)
|
||||
# benchmark 只有部分日期(模拟停盘日缺失)
|
||||
partial_dates = daily_df.index[::2] # 隔日取一个
|
||||
bench = pd.Series([0.001] * len(partial_dates), index=partial_dates)
|
||||
res = compute_metrics(daily_df, bench)
|
||||
# 不应崩溃,且 strat 长度 = daily_df 行数(没有被 dropna 削短)
|
||||
assert len(res.series["equity_curve"]) == 10
|
||||
|
||||
|
||||
def test_compute_metrics_raises_without_balance():
|
||||
"""缺少 balance 列时应抛 ValueError。"""
|
||||
idx = pd.date_range("2024-01-01", periods=5, freq="B")
|
||||
daily_df = pd.DataFrame({"net_pnl": [1, 2, 3, 4, 5]}, index=idx)
|
||||
bench = pd.Series([0.01] * 5, index=idx)
|
||||
try:
|
||||
compute_metrics(daily_df, bench)
|
||||
assert False, "应抛 ValueError"
|
||||
except ValueError as e:
|
||||
assert "balance" in str(e)
|
||||
|
||||
Reference in New Issue
Block a user