fix(backtest): CTA metrics benchmark 缺失降级(不阻塞整组指标图)
benchmark 数据缺失时原实现 if 跳过整段 compute_metrics → {task_id}_metrics.json 不写 → 策略净值/回撤/波动图也空(不只基准图)。改为:benchmark 空时传空 Series + log WARN,compute_metrics 内部 reindex→fillna(0) 容空(基准类指标 NaN→None,策略指标正常),metrics.json 照写。效果:策略图照常出,只基准/Alpha/Beta 图空(前端 benchmark-curve/risk-series 已容错)。加测试 test_run_cta_backtest_benchmark_missing_degrades(mock read_index_daily 返空 → 验 compute_metrics 仍 called + metrics.json 仍写 + 策略指标有值)。6 测试全绿(Mac pytest)。
This commit is contained in:
@@ -317,32 +317,42 @@ def run_cta_backtest(
|
|||||||
bench_df = bench_df.sort_values("date")
|
bench_df = bench_df.sort_values("date")
|
||||||
benchmark_returns = bench_df["close"].pct_change().dropna()
|
benchmark_returns = bench_df["close"].pct_change().dropna()
|
||||||
benchmark_returns.index = pd.to_datetime(bench_df["date"].iloc[1:])
|
benchmark_returns.index = pd.to_datetime(bench_df["date"].iloc[1:])
|
||||||
|
else:
|
||||||
|
# 基准缺失:降级。compute_metrics 内部 reindex→ffill→fillna(0) 能容空
|
||||||
|
# benchmark(基准类指标返 NaN→None,策略指标正常算),故仍执行 + metrics.json
|
||||||
|
# 照写。避免原实现整段跳过 → 策略净值/回撤/波动图也空(不只基准图)。
|
||||||
|
benchmark_returns = pd.Series(dtype=float)
|
||||||
|
logging.warning(
|
||||||
|
"基准数据缺失,降级计算(基准/Alpha/Beta 图将空,策略指标不受影响): "
|
||||||
|
"benchmark=%s %s~%s",
|
||||||
|
benchmark_code, start_date.date(), end_date.date(),
|
||||||
|
)
|
||||||
|
|
||||||
# H4: compute_metrics 内部从 daily_df["balance"] 自算 simple return,
|
# H4: compute_metrics 内部从 daily_df["balance"] 自算 simple return,
|
||||||
# 不再依赖 vnpy 的 log return 列(删原三路 fallback)
|
# 不再依赖 vnpy 的 log return 列(删原三路 fallback)
|
||||||
|
|
||||||
# Compute relative metrics
|
# Compute relative metrics(benchmark 有无都执行;空时基准类指标 NaN→None)
|
||||||
metrics_result = compute_metrics(daily_df, benchmark_returns)
|
metrics_result = compute_metrics(daily_df, benchmark_returns)
|
||||||
|
|
||||||
# Merge scalars into statistics (for API response)
|
# Merge scalars into statistics (for API response)
|
||||||
statistics.update(metrics_result.scalars)
|
statistics.update(metrics_result.scalars)
|
||||||
|
|
||||||
# Serialize series to JSON (separate file, same as equity_curve/trades)
|
# Serialize series to JSON (separate file, same as equity_curve/trades)
|
||||||
import json
|
import json
|
||||||
series_data = {}
|
series_data = {}
|
||||||
for key, series in metrics_result.series.items():
|
for key, series in metrics_result.series.items():
|
||||||
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": [None if (isinstance(x, float) and not math.isfinite(x)) else x
|
"values": [None if (isinstance(x, float) and not math.isfinite(x)) else x
|
||||||
for x in series.tolist()]
|
for x in series.tolist()]
|
||||||
}
|
}
|
||||||
|
|
||||||
# Write metrics series to JSON file
|
# Write metrics series to JSON file
|
||||||
file_dir = os.path.dirname(os.path.abspath(db_path))
|
file_dir = os.path.dirname(os.path.abspath(db_path))
|
||||||
metrics_file = os.path.join(file_dir, f"{task_id}_metrics.json")
|
metrics_file = os.path.join(file_dir, f"{task_id}_metrics.json")
|
||||||
with open(metrics_file, "w") as f:
|
with open(metrics_file, "w") as f:
|
||||||
json.dump({"series": series_data}, f, indent=2)
|
json.dump({"series": series_data}, f, indent=2)
|
||||||
|
|
||||||
except Exception as metrics_error:
|
except Exception as metrics_error:
|
||||||
# Log but don't fail backtest if metrics calculation fails.
|
# Log but don't fail backtest if metrics calculation fails.
|
||||||
|
|||||||
@@ -360,4 +360,87 @@ class TestRunCtaBacktest:
|
|||||||
# Verify read_index_daily was called with hs300 code (sh000300)
|
# Verify read_index_daily was called with hs300 code (sh000300)
|
||||||
mock_read.assert_called_once()
|
mock_read.assert_called_once()
|
||||||
call_args = mock_read.call_args
|
call_args = mock_read.call_args
|
||||||
assert call_args[0][0] == "sh000300", "Default benchmark should be hs300 (sh000300)"
|
assert call_args[0][0] == "sh000300", "Default benchmark should be hs300 (sh000300)"
|
||||||
|
|
||||||
|
def test_run_cta_backtest_benchmark_missing_degrades(self, temp_db_path):
|
||||||
|
"""benchmark 缺失时 compute_metrics 仍执行、metrics.json 照写(降级)。
|
||||||
|
|
||||||
|
根因:read_index_daily 返空时原实现 if 跳过整段 compute_metrics →
|
||||||
|
{task_id}_metrics.json 不写 → 策略净值/回撤/波动图也空(不只基准图)。
|
||||||
|
降级:benchmark 空时传空 Series 给 compute_metrics(内部 reindex→fillna(0),
|
||||||
|
基准类指标返 NaN→None,策略指标正常),metrics.json 照写。
|
||||||
|
"""
|
||||||
|
mock_strategy_class = Mock()
|
||||||
|
mock_strategy_class.__name__ = "DegBmStrategy"
|
||||||
|
|
||||||
|
dates = pd.date_range("2024-01-01", "2024-03-31", freq="D")
|
||||||
|
daily_df = pd.DataFrame(
|
||||||
|
{"balance": [1_000_000.0 * (1.001 ** i) for i in range(len(dates))]},
|
||||||
|
index=dates,
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_engine = MagicMock()
|
||||||
|
mock_engine.calculate_result.return_value = daily_df
|
||||||
|
mock_engine.calculate_statistics.return_value = {"total_return": 0.15}
|
||||||
|
mock_engine.trades = {"t1": MagicMock()} # 非空 trades 避免 degenerate
|
||||||
|
mock_engine.history_data = [] # 空 history → 定寸跳过
|
||||||
|
|
||||||
|
mock_cfg = Mock()
|
||||||
|
mock_cfg.data_paths = {"daily_dir": "/mock/daily_dir"}
|
||||||
|
|
||||||
|
mock_ashare = MagicMock()
|
||||||
|
mock_ashare.AShareBacktestingEngine = Mock(return_value=mock_engine)
|
||||||
|
|
||||||
|
mock_metrics_result = Mock()
|
||||||
|
mock_metrics_result.scalars = {
|
||||||
|
"total_return": 0.15, # 策略指标(不依赖基准)有值
|
||||||
|
"max_drawdown": -0.05,
|
||||||
|
"alpha": None, # 基准类指标 NaN→None(基准缺失)
|
||||||
|
"benchmark_return": None,
|
||||||
|
}
|
||||||
|
mock_metrics_result.series = {
|
||||||
|
"equity_curve": pd.Series([1.0, 1.05, 1.1]),
|
||||||
|
"drawdown": pd.Series([0.0, -0.01, -0.02]),
|
||||||
|
}
|
||||||
|
|
||||||
|
mock_tzlocal = MagicMock()
|
||||||
|
mock_tzlocal.get_localzone_name = Mock(return_value="UTC")
|
||||||
|
|
||||||
|
with patch.dict("sys.modules", {
|
||||||
|
"sanguo_backtest.ashare_engine": mock_ashare,
|
||||||
|
"tzlocal": mock_tzlocal,
|
||||||
|
"vnpy.trader.setting": MagicMock(),
|
||||||
|
}):
|
||||||
|
# read_index_daily 返空 DataFrame = benchmark 缺失
|
||||||
|
with patch("sanguo_backtest.cta_engine.read_index_daily", return_value=pd.DataFrame()):
|
||||||
|
with patch("sanguo_backtest.cta_engine.compute_metrics", return_value=mock_metrics_result) as mock_cm:
|
||||||
|
result = run_cta_backtest(
|
||||||
|
strategy_class=mock_strategy_class,
|
||||||
|
symbol="600000",
|
||||||
|
params={},
|
||||||
|
start="2024-01-01",
|
||||||
|
end="2024-03-31",
|
||||||
|
cfg=mock_cfg,
|
||||||
|
db_path=temp_db_path,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 1. compute_metrics 仍被调用(降级,没因 benchmark 空跳过)
|
||||||
|
mock_cm.assert_called_once()
|
||||||
|
# 2. 传入的 benchmark_returns 是空 Series(降级标记)
|
||||||
|
benchmark_arg = mock_cm.call_args[0][1]
|
||||||
|
assert isinstance(benchmark_arg, pd.Series)
|
||||||
|
assert benchmark_arg.empty, "benchmark 缺失时应传空 Series 给 compute_metrics"
|
||||||
|
|
||||||
|
# 3. metrics.json 仍写(策略净值/回撤/波动图数据源)
|
||||||
|
file_dir = Path(temp_db_path).parent
|
||||||
|
metrics_file = file_dir / f"{result.task_id}_metrics.json"
|
||||||
|
assert metrics_file.exists(), f"降级时 metrics.json 应照写: {metrics_file}"
|
||||||
|
with open(metrics_file) as f:
|
||||||
|
metrics_data = json.load(f)
|
||||||
|
series_keys = list(metrics_data.get("series", {}).keys())
|
||||||
|
assert "equity_curve" in series_keys
|
||||||
|
assert "drawdown" in series_keys
|
||||||
|
|
||||||
|
# 4. 策略指标有值(基准类 None)
|
||||||
|
assert result.statistics.get("total_return") == 0.15
|
||||||
|
assert result.statistics.get("alpha") is None
|
||||||
Reference in New Issue
Block a user