Files
sanguo_vnpy_v2/sanguo_backtest/metrics.py
T
claude_dev 8d55e414fa 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.
2026-07-12 23:39:45 +08:00

97 lines
4.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""回测相对/绝对指标计算(empyrical,聚宽同源口径)。纯函数。"""
from dataclasses import dataclass, field
from typing import Dict, Literal
import math
import numpy as np
import pandas as pd
import empyrical
BenchmarkCode = Literal["hs300", "zz500"]
BENCHMARK_SYMBOL: Dict[str, str] = {"hs300": "sh000300", "zz500": "sz000905"}
@dataclass
class MetricsResult:
scalars: Dict[str, float] = field(default_factory=dict)
series: Dict[str, pd.Series] = field(default_factory=dict)
def compute_metrics(
daily_df: pd.DataFrame,
benchmark_returns: pd.Series,
period: int = 252,
) -> MetricsResult:
"""对 vnpy daily_df + 基准日收益计算聚宽级指标。
H4: 从 daily_df["balance"] 自算 simple returnvnpy 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: 基准日收益率 Seriesindex 为日期(不必与 daily_df 对齐)。
period: 年化周期(整数,默认252交易日),empyrical 内部使用 'daily'
"""
# H4: simple return from balancepct_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 = {
"total_return": float(empyrical.cum_returns_final(s)),
"annual_return": float(empyrical.annual_return(s, period='daily')),
"alpha": float(empyrical.alpha(s, b, period='daily')),
"beta": float(empyrical.beta(s, b)),
"sharpe_ratio": float(empyrical.sharpe_ratio(s, period='daily')),
"sortino_ratio": float(empyrical.sortino_ratio(s, period='daily')),
"information_ratio": float(empyrical.excess_sharpe(s, b)),
"annual_volatility": float(empyrical.annual_volatility(s, period='daily')),
"max_drawdown": float(empyrical.max_drawdown(s)),
"benchmark_return": float(empyrical.cum_returns_final(b)),
"benchmark_volatility": float(empyrical.annual_volatility(b, period='daily')),
}
# Sanitize non-finite floats (NaN/Inf from degenerate inputs) → None for JSON safety
scalars = {k: (None if isinstance(v, float) and not math.isfinite(v) else v) for k, v in scalars.items()}
equity = empyrical.cum_returns(s)
bench_curve = empyrical.cum_returns(b)
# rolling alpha/beta (expanding window 用于画图,口径由 scalars 保证)
roll_beta = pd.Series(index=s.index, dtype=float)
roll_alpha = pd.Series(index=s.index, dtype=float)
for i in range(len(s)):
sub = aligned.iloc[: i + 1]
if len(sub) >= 2 and sub["b"].var() > 0:
beta = sub["s"].cov(sub["b"]) / sub["b"].var()
alpha = sub["s"].mean() - beta * sub["b"].mean()
roll_beta.iloc[i] = beta
roll_alpha.iloc[i] = alpha * period
# drawdown: 从峰值回落 (值 <= 0)
cummax = equity.cummax()
drawdown = (equity - cummax) / cummax
# Rolling annualized volatility (quarterly window) for the volatility chart
_vol_window = min(63, len(s))
if _vol_window >= 2:
vol_strategy = s.rolling(_vol_window, min_periods=2).std() * np.sqrt(period)
vol_benchmark = b.rolling(_vol_window, min_periods=2).std() * np.sqrt(period)
else:
vol_strategy = pd.Series([np.nan] * len(s), index=s.index)
vol_benchmark = pd.Series([np.nan] * len(s), index=s.index)
series = {
"equity_curve": equity,
"benchmark_curve": bench_curve,
"alpha": roll_alpha,
"beta": roll_beta,
"drawdown": drawdown,
"volatility_strategy": vol_strategy,
"volatility_benchmark": vol_benchmark,
}
return MetricsResult(scalars=scalars, series=series)