diff --git a/backend/app/backtest/engine.py b/backend/app/backtest/engine.py index 7c61b77..06d8244 100644 --- a/backend/app/backtest/engine.py +++ b/backend/app/backtest/engine.py @@ -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。""" diff --git a/backend/app/tickflow/repository.py b/backend/app/tickflow/repository.py index 899dfe2..a5cb285 100644 --- a/backend/app/tickflow/repository.py +++ b/backend/app/tickflow/repository.py @@ -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 查询内部方法 # ================================================================ diff --git a/backend/tests/backtest/test_minute_fill.py b/backend/tests/backtest/test_minute_fill.py index a71c369..bda4086 100644 --- a/backend/tests/backtest/test_minute_fill.py +++ b/backend/tests/backtest/test_minute_fill.py @@ -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 == {}