fix(backtest): NaN/Inf浮点→None消毒(vnpy统计+empyrical指标+时序)修JSON序列化500

This commit is contained in:
2026-07-11 14:33:41 +08:00
parent f7c2e2eea3
commit 3c6f25d1b0
2 changed files with 10 additions and 3 deletions
+7 -3
View File
@@ -1,6 +1,7 @@
"""CTA strategy backtesting engine wrapper using vnpy_ctastrategy.backtesting.""" """CTA strategy backtesting engine wrapper using vnpy_ctastrategy.backtesting."""
import sys import sys
import os import os
import math
import traceback import traceback
import uuid import uuid
from datetime import datetime from datetime import datetime
@@ -120,9 +121,11 @@ def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end:
# calculate_statistics(df) returns the stats dict (sharpe/drawdown/etc.) # calculate_statistics(df) returns the stats dict (sharpe/drawdown/etc.)
daily_df = engine.calculate_result() daily_df = engine.calculate_result()
raw_stats = engine.calculate_statistics(daily_df, output=False) or {} raw_stats = engine.calculate_statistics(daily_df, output=False) or {}
# Ensure JSON-serializable (vnpy may include Timestamp / non-numeric values) # Ensure JSON-serializable (vnpy may include Timestamp / non-numeric / NaN values)
statistics = { statistics = {
k: (v if isinstance(v, (int, float, str, bool)) or v is None else str(v)) k: (None if (isinstance(v, float) and not math.isfinite(v))
else v if isinstance(v, (int, float, str, bool)) or v is None
else str(v))
for k, v in raw_stats.items() for k, v in raw_stats.items()
} }
@@ -173,7 +176,8 @@ def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end:
if isinstance(series, pd.Series): if isinstance(series, pd.Series):
series_data[key] = { series_data[key] = {
"dates": series.index.astype(str).tolist(), "dates": series.index.astype(str).tolist(),
"values": series.tolist() "values": [None if (isinstance(x, float) and not math.isfinite(x)) else x
for x in series.tolist()]
} }
# Write metrics series to JSON file # Write metrics series to JSON file
+3
View File
@@ -1,6 +1,7 @@
"""回测相对/绝对指标计算(empyrical,聚宽同源口径)。纯函数。""" """回测相对/绝对指标计算(empyrical,聚宽同源口径)。纯函数。"""
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Dict, Literal from typing import Dict, Literal
import math
import numpy as np import numpy as np
import pandas as pd import pandas as pd
import empyrical import empyrical
@@ -44,6 +45,8 @@ def compute_metrics(
"benchmark_return": float(empyrical.cum_returns_final(b)), "benchmark_return": float(empyrical.cum_returns_final(b)),
"benchmark_volatility": float(empyrical.annual_volatility(b, period='daily')), "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) equity = empyrical.cum_returns(s)
bench_curve = empyrical.cum_returns(b) bench_curve = empyrical.cum_returns(b)