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:
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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-naive,benchmark 若 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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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.DataFrame(cta_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_exchange:000300 以 0 开头会被误判成 SZSE,
|
||||
# 但 000300 实际属于 SSE)。sh → SSE,sz → 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)
|
||||
|
||||
Reference in New Issue
Block a user