From 3c6f25d1b097dc739d4656d8975877ee0d549a89 Mon Sep 17 00:00:00 2001 From: claude_dev Date: Sat, 11 Jul 2026 14:33:41 +0800 Subject: [PATCH] =?UTF-8?q?fix(backtest):=20NaN/Inf=E6=B5=AE=E7=82=B9?= =?UTF-8?q?=E2=86=92None=E6=B6=88=E6=AF=92(vnpy=E7=BB=9F=E8=AE=A1+empyrica?= =?UTF-8?q?l=E6=8C=87=E6=A0=87+=E6=97=B6=E5=BA=8F)=E4=BF=AEJSON=E5=BA=8F?= =?UTF-8?q?=E5=88=97=E5=8C=96500?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- sanguo_backtest/cta_engine.py | 10 +++++++--- sanguo_backtest/metrics.py | 3 +++ 2 files changed, 10 insertions(+), 3 deletions(-) diff --git a/sanguo_backtest/cta_engine.py b/sanguo_backtest/cta_engine.py index 54384eb..dd2ba87 100644 --- a/sanguo_backtest/cta_engine.py +++ b/sanguo_backtest/cta_engine.py @@ -1,6 +1,7 @@ """CTA strategy backtesting engine wrapper using vnpy_ctastrategy.backtesting.""" import sys import os +import math import traceback import uuid 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.) daily_df = engine.calculate_result() 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 = { - 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() } @@ -173,7 +176,8 @@ def run_cta_backtest(strategy_class, symbol: str, params: dict, start: str, end: if isinstance(series, pd.Series): series_data[key] = { "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 diff --git a/sanguo_backtest/metrics.py b/sanguo_backtest/metrics.py index 46a7b35..2afd7b9 100644 --- a/sanguo_backtest/metrics.py +++ b/sanguo_backtest/metrics.py @@ -1,6 +1,7 @@ """回测相对/绝对指标计算(empyrical,聚宽同源口径)。纯函数。""" from dataclasses import dataclass, field from typing import Dict, Literal +import math import numpy as np import pandas as pd import empyrical @@ -44,6 +45,8 @@ def compute_metrics( "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)