fix(backtest): 结果页垃圾值/无图表端到端修复(empyrical×numpy2.0根因)

根因: empyrical 0.5.5 引用 numpy2.0 已移除的 np.NINF → compute_metrics 静默崩 → _metrics.json 不生成 → 结果页回退 vnpy 原始字段(单位混乱: total_return当百分数、max_drawdown当元 → 前端×100显 3305%/-50M%)。
- metrics.py: 导入 empyrical 前补回 np 别名(NINF/Inf/PINF/NaN/NAN/infty)
- routes.py: benchmark-curve/risk-series 缺 metrics 文件时返空200(不再404拖垮整页); get_result 从 statistics 抽 relative_metrics
- cta_engine.py: bench_df 日期 strip tz 防 pct_change 崩; metrics 块加 traceback 日志
- Result.vue: onMounted 用 Promise.allSettled 隔离7端点, 单接口失败不拖垮整页
- result_store.py: _safe_read_json 容错迁移后残留 NAS 绝对路径, stale path 不崩 list_results
- datareader.py: read_index_daily 改从 vnpy DB 读 + 前缀解析交易所(sh→SSE, 避免 000300 被 guess_exchange 误判 SZSE)
This commit is contained in:
2026-07-17 08:24:01 +08:00
parent 37c850d0c5
commit 7fe3fb0844
6 changed files with 159 additions and 66 deletions
+67 -32
View File
@@ -110,40 +110,75 @@ const filteredBenchmarkCurve = computed(() => filterDataByTimeRange(benchmarkCur
const filteredRiskSeries = computed(() => filterDataByTimeRange(riskSeries.value))
onMounted(async () => {
try {
const [info, relMetrics, benchCurve, riskSer, eq, p, tr] = await Promise.all([
getResult(taskId),
getRelativeMetrics(taskId),
getBenchmarkCurve(taskId),
getRiskSeries(taskId),
getEquityCurve(taskId),
getDailyPnl(taskId),
getTrades(taskId),
])
// 请求隔离:任一接口失败(如 benchmark-curve/risk-series 数据缺失)不得阻塞
// 其余请求。statistics/equity/trades 有数据时必须正常渲染。用 Promise.allSettled
// 保留并发,逐个取值,失败项保留默认空值 + console.warn。
const settled = await Promise.allSettled([
getResult(taskId),
getRelativeMetrics(taskId),
getBenchmarkCurve(taskId),
getRiskSeries(taskId),
getEquityCurve(taskId),
getDailyPnl(taskId),
getTrades(taskId),
])
const [rInfo, rRel, rBench, rRisk, rEq, rPnl, rTr] = settled
statistics.value = info.statistics || {}
relativeMetrics.value = relMetrics
benchmarkCurve.value = benchCurve
riskSeries.value = riskSer
equity.value = eq
pnl.value = p
trades.value = tr
if (info.symbol && info.start && info.end) {
try {
kline.value = await getKline(info.symbol, info.start, info.end)
} catch {
kline.value = []
}
}
try {
logText.value = await getLog(taskId)
} catch {
logText.value = ''
}
} finally {
loading.value = false
if (rInfo.status === 'fulfilled') {
statistics.value = rInfo.value.statistics || {}
} else {
console.warn('[Result] getResult failed:', rInfo.reason)
}
if (rRel.status === 'fulfilled') {
relativeMetrics.value = rRel.value
} else {
console.warn('[Result] getRelativeMetrics failed:', rRel.reason)
}
if (rBench.status === 'fulfilled') {
benchmarkCurve.value = rBench.value
} else {
console.warn('[Result] getBenchmarkCurve failed:', rBench.reason)
}
if (rRisk.status === 'fulfilled') {
riskSeries.value = rRisk.value
} else {
console.warn('[Result] getRiskSeries failed:', rRisk.reason)
}
if (rEq.status === 'fulfilled') {
equity.value = rEq.value
} else {
console.warn('[Result] getEquityCurve failed:', rEq.reason)
}
if (rPnl.status === 'fulfilled') {
pnl.value = rPnl.value
} else {
console.warn('[Result] getDailyPnl failed:', rPnl.reason)
}
if (rTr.status === 'fulfilled') {
trades.value = rTr.value
} else {
console.warn('[Result] getTrades failed:', rTr.reason)
}
// kline 依赖 getResult 返回的 symbol/start/end,单独隔离
const info = rInfo.status === 'fulfilled' ? rInfo.value : null
if (info?.symbol && info?.start && info?.end) {
try {
kline.value = await getKline(info.symbol, info.start, info.end)
} catch (e) {
console.warn('[Result] getKline failed:', e)
kline.value = []
}
}
try {
logText.value = await getLog(taskId)
} catch (e) {
console.warn('[Result] getLog failed:', e)
logText.value = ''
}
loading.value = false
})
// 每日收益格式化:浮点精度 → 2 位;日期去 00:00:00
+4 -2
View File
@@ -358,7 +358,8 @@ def benchmark_curve(task_id: str):
"""Get benchmark curve data (strategy vs benchmark)."""
metrics_file = _get_metrics_file_path(task_id)
if not metrics_file:
raise HTTPException(status_code=404, detail="metrics file not found")
# metrics 文件缺失时返回空 200(图表优雅降级),不再 404 触发前端整页空白
return {"dates": [], "strategy": [], "benchmark": []}
import json
with open(metrics_file, 'r') as f:
@@ -380,7 +381,8 @@ def risk_series(task_id: str):
"""Get risk series data (alpha, beta, drawdown)."""
metrics_file = _get_metrics_file_path(task_id)
if not metrics_file:
raise HTTPException(status_code=404, detail="metrics file not found")
# metrics 文件缺失时返回空 200(图表优雅降级),不再 404 触发前端整页空白
return {"dates": [], "alpha": [], "beta": [], "drawdown": [], "strategy_vol": [], "benchmark_vol": []}
import json
with open(metrics_file, 'r') as f:
+11 -2
View File
@@ -254,6 +254,11 @@ def run_cta_backtest(
if bench_df is not None and not bench_df.empty and "close" in bench_df.columns:
# Calculate benchmark daily returns
bench_df["date"] = pd.to_datetime(bench_df["date"])
# 去时区:daily_df.index 是 tz-naivebenchmark 若 tz-aware 会让
# compute_metrics 内部 reindex 抛 TypeError(被上层 except 静默吞掉,
# 致 _metrics.json 不生成)。统一去掉 tz 保证对齐。
if getattr(bench_df["date"].dt, "tz", None) is not None:
bench_df["date"] = bench_df["date"].dt.tz_localize(None)
bench_df = bench_df.sort_values("date")
benchmark_returns = bench_df["close"].pct_change().dropna()
benchmark_returns.index = pd.to_datetime(bench_df["date"].iloc[1:])
@@ -285,8 +290,12 @@ def run_cta_backtest(
json.dump({"series": series_data}, f, indent=2)
except Exception as metrics_error:
# Log but don't fail backtest if metrics calculation fails
logging.warning("相对指标计算失败(回测结果不受影响): %s", metrics_error)
# Log but don't fail backtest if metrics calculation fails.
# 附 traceback 以便定位(_metrics.json 不生成时这里是根因)。
logging.warning(
"相对指标计算失败(回测结果不受影响): %s\n%s",
metrics_error, traceback.format_exc(),
)
# Build equity curve DataFrame (S1.2): use the daily_df returned by
# calculate_result (index=date, has a 'balance' column). get_all_daily_results
+7
View File
@@ -3,6 +3,13 @@ from dataclasses import dataclass, field
from typing import Dict, Literal
import math
import numpy as np
# empyrical 0.5.5 引用了 NumPy 2.0 已移除的别名(np.NINF / np.NaN / np.Inf / np.PINF),
# 不补回会在 sortino_ratio/downside_risk 等函数里抛 AttributeError,导致整块相对指标计算
# 失败、_metrics.json 不生成、结果页回退到 vnpy 原始字段(单位混乱)。
for _alias, _val in (("NINF", -np.inf), ("Inf", np.inf), ("PINF", np.inf),
("NaN", np.nan), ("NAN", np.nan), ("infty", np.inf)):
if not hasattr(np, _alias):
setattr(np, _alias, _val)
import pandas as pd
import empyrical
+24 -3
View File
@@ -1,11 +1,14 @@
"""Backtest result storage using SQLite + parquet files."""
import sqlite3
import json
import logging
from dataclasses import dataclass
from pathlib import Path
from typing import Optional
import pandas as pd
logger = logging.getLogger(__name__)
@dataclass
class BacktestResult:
@@ -109,6 +112,22 @@ def save_result(result: BacktestResult, db_path: str, file_dir: Optional[str] =
conn.close()
def _safe_read_json(path: Optional[str]) -> Optional[pd.DataFrame]:
"""Read a JSON equity/trades file; return None if missing or unreadable.
Historical records may reference paths from a previous host (e.g. NAS
absolute paths after migration to VPS). Swallow those failures so the
record stays listable instead of crashing list_results.
"""
if not path:
return None
try:
return pd.read_json(path, orient="records")
except (FileNotFoundError, ValueError, OSError) as e:
logger.warning("result_store: skipping unreadable JSON file %s: %s", path, e)
return None
def load_result(rid: int, db_path: str) -> BacktestResult:
"""
Load backtest result by ID from database.
@@ -133,9 +152,11 @@ def load_result(rid: int, db_path: str) -> BacktestResult:
cols = [d[0] for d in conn.execute("SELECT * FROM backtest_stats LIMIT 0").description]
d = dict(zip(cols, row))
# Load JSON files if paths exist (equity_curve/trades persisted as JSON)
equity = pd.read_json(d["equity_path"], orient="records") if d.get("equity_path") else None
trades = pd.read_json(d["trades_path"], orient="records") if d.get("trades_path") else None
# Load JSON files if paths exist (equity_curve/trades persisted as JSON).
# Tolerate stale paths (e.g. NAS absolute paths left after VPS migration):
# missing/unreadable file -> None, so record still appears in list_results.
equity = _safe_read_json(d.get("equity_path"))
trades = _safe_read_json(d.get("trades_path"))
return BacktestResult(
task_id=d["task_id"],
+46 -27
View File
@@ -105,41 +105,60 @@ def read_parquet_15min(symbol: str, start: str, end: str, cfg, dir_key: str = "m
return bars
def read_index_daily(code: str, start: date, end: date, cfg) -> pd.DataFrame:
def read_index_daily(code: str, start, end, cfg) -> pd.DataFrame:
"""
读指数日线数据(sh000300/sz000905),复用 read_parquet_daily 的年分片 parquet 路径
读指数日线数据(sh000300/sz399001 等),从 vnpy DB 读取(统一数据源)。
parquet 仅作原始备份,不再读取。返回类型保持 pd.DataFramecta_engine 消费不变)。
Args:
code: 指数代码,如 "sh000300"(沪深300)、"sz000905"中证500
start: 起始日期
end: 结束日期
code: 指数代码,带交易所前缀,"sh000300"(沪深300)、"sz399001"深证成指
start: 起始日期str "YYYY-MM-DD" / date / datetime
end: 结束日期str "YYYY-MM-DD" / date / datetime
cfg: 数据配置对象
Returns:
pd.DataFrame: 包含 date/open/high/low/close/volume 列的日线数据
pd.DataFrame: date/open/high/low/close/volume 列;无数据返回空 DataFrame。
"""
daily_dir = Path(cfg.data_paths["daily_dir"])
start_dt = start if isinstance(start, datetime) else datetime.combine(start, datetime.min.time())
end_dt = end if isinstance(end, datetime) else datetime.combine(end, datetime.max.time())
from vnpy.trader.database import get_database # lazy:避免模块 import 依赖数据库驱动
dfs: list[pd.DataFrame] = []
# 前缀解析交易所(指数不能用 guess_exchange000300 以 0 开头会被误判成 SZSE,
# 但 000300 实际属于 SSE)。sh → SSEsz → SZSE。
symbol = code[2:]
exchange = Exchange.SSE if code.startswith("sh") else Exchange.SZSE
# 按年分片读取(与 read_parquet_daily 相同路径逻辑)
for year in range(start_dt.year, end_dt.year + 1):
f = daily_dir / str(year) / f"{code}_daily.parquet"
if not f.exists():
continue
df = pd.read_parquet(f)
# 过滤日期范围
df["date"] = pd.to_datetime(df["date"])
mask = (df["date"] >= start_dt) & (df["date"] <= end_dt)
filtered_df = df[mask].copy()
if not filtered_df.empty:
dfs.append(filtered_df)
# 日期归一化:str → parse, date → combine, datetime → as-is
def _to_dt(s, is_start: bool) -> datetime:
if isinstance(s, datetime):
return s
if isinstance(s, date):
return datetime.combine(s, datetime.min.time() if is_start else datetime.max.time())
return datetime.strptime(s, "%Y-%m-%d")
if dfs:
result = pd.concat(dfs, ignore_index=True)
result = result.sort_values("date")
return result.reset_index(drop=True)
else:
start_dt = _to_dt(start, True)
end_dt = _to_dt(end, False)
# 配置 vnpy DB(与 read_db_daily 同模式)
SETTINGS["database.name"] = "sqlite"
SETTINGS["database.database"] = cfg.data_paths["vnpy_db"]
db = get_database()
bars = db.load_bar_data(
symbol=symbol,
exchange=exchange,
interval=Interval.DAILY,
start=start_dt,
end=end_dt,
)
if not bars:
return pd.DataFrame(columns=["date", "open", "high", "low", "close", "volume"])
df = pd.DataFrame([{
"date": b.datetime,
"open": b.open_price,
"high": b.high_price,
"low": b.low_price,
"close": b.close_price,
"volume": b.volume,
} for b in bars])
return df.sort_values("date").reset_index(drop=True)