mirror of
https://ghfast.top/https://github.com/aeroxw/tick-stock-panel.git
synced 2026-09-12 16:44:15 +08:00
fix(backtest): 分钟K回测爆内存 — 按触发日精确读取 + 紧凑数组缓存
上一提交 (365f1cc) 修 index bug 时把 cache 从紧凑 numpy 数组改成臃肿的
polars DataFrame, 分钟K回测内存翻几倍导致 OOM 死机。且 _load_minute_for_fills
用 get_minute_range 扫描整个触发区间 (start~end), 触发日稀疏时读了大量无关
日期, 全市场长周期回测直接爆内存。
修复:
1. repository 新增 get_minute_by_dates — 按日期列表精确读 date=YYYY-MM-DD
分区文件, 不扫描区间。内存与回测区间长度解耦, 只随触发日数量增长。
2. _load_minute_for_fills 改用 get_minute_by_dates, 每批 50 天分批读取。
3. cache 改回 float64 紧凑 2D 数组 (6 列: open/high/low/close/volume/amount),
.cast(Float64) 保证 to_numpy 不退化回 object, 原来的 index bug 不复发。
4. _resolve_minute_fill 改用整数列索引 (arr[:,0]) 替代列名索引。
内存对比 (全市场 1 年):
之前: get_minute_range 扫全年 ~365 文件 + DataFrame 缓存 → 5-10GB
现在: 只读触发日文件 + float64 数组 → 几十 MB
验证: 191 passed (含 7 个分钟K回归测试)
This commit is contained in:
@@ -525,9 +525,9 @@ class BacktestEngine:
|
||||
_loaded = self._load_minute_for_fills(
|
||||
self.repo, list(_trigger_symbols), _trigger_dates, "stock",
|
||||
)
|
||||
for _key, _mdf in _loaded.items():
|
||||
if not _mdf.is_empty():
|
||||
minute_cache[_key] = _mdf
|
||||
for _key, _marr in _loaded.items():
|
||||
if _marr is not None and len(_marr) > 0:
|
||||
minute_cache[_key] = _marr
|
||||
|
||||
def _refill_price(idx: int, side: str, daily_price: float) -> float:
|
||||
if not config.minute_fill or not minute_cache:
|
||||
@@ -839,33 +839,32 @@ class BacktestEngine:
|
||||
|
||||
@staticmethod
|
||||
def _resolve_minute_fill(
|
||||
minute_df: pl.DataFrame,
|
||||
minute_arr: np.ndarray,
|
||||
ref_price: float | None,
|
||||
side: str,
|
||||
) -> float | None:
|
||||
"""用当日分钟K确定精确成交价。
|
||||
|
||||
Args:
|
||||
minute_df: 当日分钟K polars DataFrame, 列含 open/high/low/close/volume/amount
|
||||
minute_arr: float64 2D 数组, 列顺序 = _MINUTE_NUMERIC_COLS
|
||||
[open(0), high(1), low(2), close(3), volume(4), amount(5)]
|
||||
(缺失的尾部列直接不存在, 用 shape 判断)
|
||||
ref_price: 信号参考线价格 (如 MA5 值); None 表示无参考线
|
||||
side: "buy" 或 "sell", 决定穿越方向
|
||||
|
||||
Returns:
|
||||
精确成交价, 或 None (降级到日K口径)
|
||||
|
||||
注: 原实现用 df.to_numpy() 转 structured array 再按字段名索引, 但当列类型不
|
||||
一致时 (如 datetime 列 + float 列) to_numpy() 退化为 object 二维数组, 字段名
|
||||
索引 arr["open"] 会抛 IndexError。改为直接按列取 Series, 稳定且更快。
|
||||
"""
|
||||
if minute_df is None or minute_df.is_empty():
|
||||
if minute_arr is None or len(minute_arr) == 0:
|
||||
return None
|
||||
|
||||
opens = minute_df["open"].to_numpy().astype(float)
|
||||
highs = minute_df["high"].to_numpy().astype(float)
|
||||
lows = minute_df["low"].to_numpy().astype(float)
|
||||
closes = minute_df["close"].to_numpy().astype(float)
|
||||
volumes = minute_df["volume"].to_numpy().astype(float) if "volume" in minute_df.columns else None
|
||||
amounts = minute_df["amount"].to_numpy().astype(float) if "amount" in minute_df.columns else None
|
||||
ncols = minute_arr.shape[1] if minute_arr.ndim == 2 else 1
|
||||
opens = minute_arr[:, 0]
|
||||
highs = minute_arr[:, 1] if ncols > 1 else opens
|
||||
lows = minute_arr[:, 2] if ncols > 2 else opens
|
||||
closes = minute_arr[:, 3] if ncols > 3 else opens
|
||||
volumes = minute_arr[:, 4] if ncols > 4 else None
|
||||
amounts = minute_arr[:, 5] if ncols > 5 else None
|
||||
|
||||
# 有参考线 → 穿越价成交 (逻辑同止损: 找价格穿越参考线的时刻)
|
||||
if ref_price is not None and ref_price > 0 and np.isfinite(ref_price):
|
||||
@@ -893,6 +892,9 @@ class BacktestEngine:
|
||||
|
||||
return float(closes[-1]) if np.isfinite(closes[-1]) else None
|
||||
|
||||
# 分钟K cache 存储的数值列及固定顺序 (_resolve_minute_fill 按此顺序整数索引)。
|
||||
_MINUTE_NUMERIC_COLS = ["open", "high", "low", "close", "volume", "amount"]
|
||||
|
||||
@staticmethod
|
||||
def _load_minute_for_fills(
|
||||
repo,
|
||||
@@ -900,36 +902,49 @@ class BacktestEngine:
|
||||
dates_needed: set,
|
||||
asset_type: str,
|
||||
) -> dict:
|
||||
"""批量加载回测区间内触发日的分钟K, 返回 {(symbol, date_str): minute_df}。
|
||||
"""按触发日加载分钟K, 返回 {(symbol, date_str): float64 2D ndarray}。
|
||||
|
||||
dates_needed: 需要分钟数据的日期集合 (set of date strings "YYYY-MM-DD")
|
||||
|
||||
按触发日分批读取对应分区文件 (get_minute_by_dates), 而非扫描整个区间
|
||||
(get_minute_range)。内存与回测区间长度解耦, 只随触发日数量增长 ——
|
||||
触发日稀疏时避免读取区间内大量无关日期导致爆内存。
|
||||
|
||||
cache 值为 float64 紧凑 2D 数组 (列顺序见 _MINUTE_NUMERIC_COLS), 而非
|
||||
完整 DataFrame, 避免每条记录携带 polars 元数据开销导致内存膨胀。
|
||||
"""
|
||||
if not symbols or not dates_needed:
|
||||
return {}
|
||||
from datetime import date as _date
|
||||
sorted_dates = sorted(dates_needed)
|
||||
start = _date.fromisoformat(sorted_dates[0])
|
||||
end = _date.fromisoformat(sorted_dates[-1])
|
||||
try:
|
||||
df = repo.get_minute_range(symbols, start, end, asset_type=asset_type)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("minute fill data load failed: %s", e)
|
||||
return {}
|
||||
if df.is_empty():
|
||||
return {}
|
||||
sorted_date_strs = sorted(dates_needed)
|
||||
date_objs = [_date.fromisoformat(s) for s in sorted_date_strs]
|
||||
|
||||
# 按 (symbol, 日期) 向量化分组, 替代原 iter_rows 逐行 Python 循环。
|
||||
# 原实现对每行做 dict 转换再重建 DataFrame, 回测区间内触发股数多时极慢。
|
||||
df = df.with_columns(
|
||||
pl.col("datetime").dt.strftime("%Y-%m-%d").alias("_d_str")
|
||||
)
|
||||
cache: dict = {}
|
||||
for sub in df.partition_by(["symbol", "_d_str"], as_dict=False):
|
||||
if sub.is_empty():
|
||||
# 分批读取: 每批 50 个交易日, 处理完拼进 cache, 避免单批过大。
|
||||
BATCH = 50
|
||||
numeric_cols = BacktestEngine._MINUTE_NUMERIC_COLS
|
||||
for i in range(0, len(date_objs), BATCH):
|
||||
batch = date_objs[i:i + BATCH]
|
||||
try:
|
||||
df = repo.get_minute_by_dates(symbols, batch, asset_type=asset_type)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("minute fill data load failed (batch %d-%d): %s", i, i + len(batch), e)
|
||||
continue
|
||||
sym = sub["symbol"][0]
|
||||
d_str = sub["_d_str"][0]
|
||||
cache[(sym, d_str)] = sub.drop("_d_str")
|
||||
if df.is_empty():
|
||||
continue
|
||||
# 按 (symbol, 日期) 分组, 每组转紧凑 float64 数组存入 cache
|
||||
df = df.with_columns(
|
||||
pl.col("datetime").dt.strftime("%Y-%m-%d").alias("_d_str")
|
||||
)
|
||||
for sub in df.partition_by(["symbol", "_d_str"], as_dict=False):
|
||||
if sub.is_empty():
|
||||
continue
|
||||
sym = sub["symbol"][0]
|
||||
d_str = sub["_d_str"][0]
|
||||
cols = [c for c in numeric_cols if c in sub.columns]
|
||||
cache[(sym, d_str)] = sub.select(
|
||||
[pl.col(c).cast(pl.Float64) for c in cols]
|
||||
).to_numpy()
|
||||
return cache
|
||||
|
||||
def simulate_portfolio(
|
||||
@@ -1051,9 +1066,9 @@ class BacktestEngine:
|
||||
loaded = self._load_minute_for_fills(
|
||||
self.repo, list(trigger_symbols), trigger_dates, asset_type,
|
||||
)
|
||||
for key, mdf in loaded.items():
|
||||
if not mdf.is_empty():
|
||||
minute_cache[key] = mdf
|
||||
for key, marr in loaded.items():
|
||||
if marr is not None and len(marr) > 0:
|
||||
minute_cache[key] = marr
|
||||
|
||||
def _refill_price(idx: int, side: str, daily_price: float) -> float:
|
||||
"""分钟K精确成交价; 无数据则降级为 daily_price。"""
|
||||
|
||||
@@ -1307,6 +1307,47 @@ class KlineRepository:
|
||||
logger.warning("分钟K范围查询失败: %s", e)
|
||||
return pl.DataFrame()
|
||||
|
||||
def get_minute_by_dates(
|
||||
self,
|
||||
symbols: list[str],
|
||||
dates: list[date],
|
||||
asset_type: str = "stock",
|
||||
) -> pl.DataFrame:
|
||||
"""按日期列表精确读取分钟K分区文件 (分钟K精确回测用)。
|
||||
|
||||
与 get_minute_range 的区别: 后者扫描 [start, end] 区间全部日期的 parquet
|
||||
(触发日稀疏时会读大量无关日期 → 爆内存); 本方法只读 dates 里列出的日期
|
||||
对应的分区文件 (date=YYYY-MM-DD/part.parquet), 内存与回测区间长度解耦,
|
||||
只随触发日数量增长。
|
||||
|
||||
缺失的日期分区直接跳过 (该日无分钟K数据)。
|
||||
返回列: symbol, datetime, open, high, low, close, volume, amount。
|
||||
"""
|
||||
if not symbols or not dates:
|
||||
return pl.DataFrame()
|
||||
base = self._etf_minute_glob.rsplit("/", 2)[0] if asset_type == "etf" else self._minute_glob.rsplit("/", 2)[0]
|
||||
# 收集存在的分区文件路径, 避免对不存在的文件 scan 报错
|
||||
parts: list[str] = []
|
||||
for d in dates:
|
||||
p = f"{base}/date={d.isoformat()}/part.parquet"
|
||||
if Path(p).exists():
|
||||
parts.append(p)
|
||||
if not parts:
|
||||
return pl.DataFrame()
|
||||
try:
|
||||
lf = pl.scan_parquet(parts)
|
||||
available = set(lf.collect_schema().names())
|
||||
select_cols = [c for c in ["symbol", "datetime", "open", "high", "low", "close", "volume", "amount"] if c in available]
|
||||
return (
|
||||
lf.select(select_cols)
|
||||
.filter(pl.col("symbol").is_in(symbols))
|
||||
.sort(["symbol", "datetime"])
|
||||
.collect(streaming=True)
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("分钟K按日期查询失败: %s", e)
|
||||
return pl.DataFrame()
|
||||
|
||||
# ================================================================
|
||||
# Polars 查询内部方法
|
||||
# ================================================================
|
||||
|
||||
@@ -3,19 +3,28 @@
|
||||
背景: 原实现用 df.to_numpy() 转 structured array 再按字段名索引 (arr["open"])。
|
||||
当 DataFrame 含 datetime 列 + float 列时, to_numpy() 退化为 dtype=object 的二维
|
||||
数组, 字段名索引抛 IndexError: "only integers, slices... are valid indices"。
|
||||
开启 minute_fill 的回测从未成功跑通过。此测试锁定该 bug 不再复发。
|
||||
开启 minute_fill 的回测从未成功跑通过。
|
||||
|
||||
当前设计:
|
||||
- _load_minute_for_fills 按触发日分批读取分区 (get_minute_by_dates), 返回
|
||||
{(symbol, date_str): float64 2D ndarray} (列顺序 = _MINUTE_NUMERIC_COLS)。
|
||||
- _resolve_minute_fill 接收 float64 2D 数组, 按整数列索引访问。
|
||||
.cast(Float64) 保证 to_numpy 不退化, 整数索引避免列名依赖。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
|
||||
import numpy as np
|
||||
import polars as pl
|
||||
|
||||
from app.backtest.engine import BacktestEngine
|
||||
|
||||
NUMERIC_COLS = BacktestEngine._MINUTE_NUMERIC_COLS # open/high/low/close/volume/amount
|
||||
|
||||
|
||||
def _sample_minute_df(symbol: str = "000001.SZ") -> pl.DataFrame:
|
||||
"""构造一份带 datetime 列 + float 列的分钟K (复现 to_numpy 退化的场景)。"""
|
||||
"""构造一份带 datetime 列 + float 列的分钟K (get_minute_by_dates 的返回形态)。"""
|
||||
base = datetime(2024, 1, 2, 9, 31)
|
||||
return pl.DataFrame({
|
||||
"symbol": [symbol] * 4,
|
||||
@@ -30,64 +39,69 @@ def _sample_minute_df(symbol: str = "000001.SZ") -> pl.DataFrame:
|
||||
})
|
||||
|
||||
|
||||
def test_resolve_minute_fill_with_mixed_columns_no_index_error():
|
||||
"""混合列类型 (datetime + float) 不再抛 IndexError。
|
||||
def _to_compact_arr(mdf: pl.DataFrame) -> np.ndarray:
|
||||
"""模拟 _load_minute_for_fills 的转换: 统一 float64 再 to_numpy。"""
|
||||
cols = [c for c in NUMERIC_COLS if c in mdf.columns]
|
||||
return mdf.select([pl.col(c).cast(pl.Float64) for c in cols]).to_numpy()
|
||||
|
||||
这是原 bug 的精确复现点: 旧实现 _resolve_minute_fill 接收 ndarray,
|
||||
arr["open"] 在 object 数组上会炸。现在接受 DataFrame, 按列取值。
|
||||
|
||||
def test_resolve_minute_fill_with_compact_array_no_index_error():
|
||||
"""紧凑 float64 数组 + 整数列索引不再抛 IndexError (锁定原 bug)。
|
||||
|
||||
原始 bug: arr["open"] 在 object 数组上炸。现在 .cast(Float64) 保证类型一致,
|
||||
to_numpy 返回规整 2D float64, 整数索引 arr[:,0] 稳定。
|
||||
"""
|
||||
mdf = _sample_minute_df()
|
||||
# 三种分支都应正常返回, 不抛 IndexError
|
||||
assert BacktestEngine._resolve_minute_fill(mdf, ref_price=10.5, side="buy") is not None
|
||||
assert BacktestEngine._resolve_minute_fill(mdf, ref_price=10.5, side="sell") is not None
|
||||
# 无参考线 → VWAP 分支
|
||||
vwap = BacktestEngine._resolve_minute_fill(mdf, ref_price=None, side="buy")
|
||||
arr = _to_compact_arr(_sample_minute_df())
|
||||
assert arr.dtype == np.float64 # 必须是 float64, 不能退化成 object
|
||||
# 三种分支都应正常返回
|
||||
assert BacktestEngine._resolve_minute_fill(arr, ref_price=10.5, side="buy") is not None
|
||||
assert BacktestEngine._resolve_minute_fill(arr, ref_price=10.5, side="sell") is not None
|
||||
vwap = BacktestEngine._resolve_minute_fill(arr, ref_price=None, side="buy")
|
||||
assert vwap is not None and vwap > 0
|
||||
|
||||
|
||||
def test_resolve_minute_fill_buy_cross_above_ref():
|
||||
"""买入: 价格涨破参考线 → 开盘已高于则按开盘。"""
|
||||
mdf = _sample_minute_df()
|
||||
# ref=9.5, 开盘 10.0 已高于 → 按开盘
|
||||
assert BacktestEngine._resolve_minute_fill(mdf, 9.5, "buy") == 10.0
|
||||
arr = _to_compact_arr(_sample_minute_df())
|
||||
assert BacktestEngine._resolve_minute_fill(arr, 9.5, "buy") == 10.0
|
||||
|
||||
|
||||
def test_resolve_minute_fill_sell_cross_below_ref():
|
||||
"""卖出: 价格跌破参考线 → 开盘已低于则按开盘。"""
|
||||
mdf = _sample_minute_df()
|
||||
# ref=10.5, 开盘 10.0 已低于 → 按开盘
|
||||
assert BacktestEngine._resolve_minute_fill(mdf, 10.5, "sell") == 10.0
|
||||
arr = _to_compact_arr(_sample_minute_df())
|
||||
assert BacktestEngine._resolve_minute_fill(arr, 10.5, "sell") == 10.0
|
||||
|
||||
|
||||
def test_resolve_minute_fill_vwap():
|
||||
"""无参考线 → VWAP = 总成交额 / 总成交量。"""
|
||||
mdf = _sample_minute_df()
|
||||
arr = _to_compact_arr(_sample_minute_df())
|
||||
total_amt = 1020.0 + 2120.0 + 1627.0 + 1278.0
|
||||
total_vol = 100 + 200 + 150 + 120
|
||||
expected = total_amt / total_vol
|
||||
assert BacktestEngine._resolve_minute_fill(mdf, None, "buy") == expected
|
||||
assert BacktestEngine._resolve_minute_fill(arr, None, "buy") == total_amt / total_vol
|
||||
|
||||
|
||||
def test_resolve_minute_fill_empty_returns_none():
|
||||
"""空 DataFrame → None (降级到日K口径)。"""
|
||||
assert BacktestEngine._resolve_minute_fill(pl.DataFrame(), None, "buy") is None
|
||||
"""空数组 → None (降级到日K口径)。"""
|
||||
assert BacktestEngine._resolve_minute_fill(np.array([]).reshape(0, 6), None, "buy") is None
|
||||
assert BacktestEngine._resolve_minute_fill(None, None, "buy") is None
|
||||
|
||||
|
||||
class _FakeRepo:
|
||||
"""最小 repo 桩: get_minute_range 直接返回预构造的混合列 DataFrame。"""
|
||||
"""最小 repo 桩: get_minute_by_dates 直接返回预构造的混合列 DataFrame。"""
|
||||
|
||||
def __init__(self, df: pl.DataFrame) -> None:
|
||||
self._df = df
|
||||
|
||||
def get_minute_range(self, symbols, start, end, asset_type="stock") -> pl.DataFrame: # noqa: ANN001
|
||||
def get_minute_by_dates(self, symbols, dates, asset_type="stock"): # noqa: ANN001
|
||||
return self._df
|
||||
|
||||
|
||||
def test_load_minute_for_fills_returns_dataframe_dict():
|
||||
"""_load_minute_for_fills 返回 {(symbol, date_str): DataFrame}, 而非 ndarray。
|
||||
def test_load_minute_for_fills_returns_compact_arrays():
|
||||
"""_load_minute_for_fills 返回 {(symbol, date_str): float64 ndarray}。
|
||||
|
||||
锁定: cache 值类型必须是 pl.DataFrame (旧实现返回的对象后续被 .to_numpy()
|
||||
退化成 object 数组触发 bug)。
|
||||
锁定两个关键性质:
|
||||
1) cache 值类型必须是 float64 ndarray (不能是 object, 也不能是臃肿的 DataFrame)。
|
||||
2) load 时已做 .cast(Float64), _resolve_minute_fill 直接整数索引可用。
|
||||
"""
|
||||
df = _sample_minute_df()
|
||||
repo = _FakeRepo(df)
|
||||
@@ -95,10 +109,21 @@ def test_load_minute_for_fills_returns_dataframe_dict():
|
||||
repo, ["000001.SZ"], {"2024-01-02"}, "stock",
|
||||
)
|
||||
assert ("000001.SZ", "2024-01-02") in result
|
||||
val = result[("000001.SZ", "2024-01-02")]
|
||||
# 关键断言: 返回的是 DataFrame, 可直接喂给 _resolve_minute_fill
|
||||
assert isinstance(val, pl.DataFrame)
|
||||
assert not val.is_empty()
|
||||
arr = result[("000001.SZ", "2024-01-02")]
|
||||
# 关键断言: 紧凑 float64 数组, 非 object, 非 DataFrame
|
||||
assert isinstance(arr, np.ndarray)
|
||||
assert arr.dtype == np.float64
|
||||
# 列顺序 = _MINUTE_NUMERIC_COLS (open=0, high=1, low=2, close=3, volume=4, amount=5)
|
||||
assert arr.shape == (4, 6)
|
||||
# 端到端: load → resolve 不抛异常
|
||||
price = BacktestEngine._resolve_minute_fill(val, None, "buy")
|
||||
price = BacktestEngine._resolve_minute_fill(arr, None, "buy")
|
||||
assert price is not None and price > 0
|
||||
|
||||
|
||||
def test_load_minute_for_fills_handles_missing_dates():
|
||||
"""缺失的日期分区不报错, 直接跳过。"""
|
||||
repo = _FakeRepo(pl.DataFrame()) # 空返回
|
||||
result = BacktestEngine._load_minute_for_fills(
|
||||
repo, ["000001.SZ"], {"2024-01-02", "2024-01-03"}, "stock",
|
||||
)
|
||||
assert result == {}
|
||||
|
||||
Reference in New Issue
Block a user