From 1539e678f6964e2ae9739a592d8e8a3b09bea329 Mon Sep 17 00:00:00 2001 From: claude_dev Date: Wed, 26 Aug 2026 20:19:03 +0800 Subject: [PATCH] =?UTF-8?q?fix(factor):=20fast=5Fops=E5=A4=A7=E7=AA=97?= =?UTF-8?q?=E5=8F=A3OOM=E6=A0=B9=E6=B2=BB=E2=80=94=E2=80=94>60=E7=AA=97shi?= =?UTF-8?q?ft=E5=B1=95=E5=BC=80=E7=89=A9=E5=8C=96window=E5=88=97=C3=9710M?= =?UTF-8?q?=E8=A1=8C(250=E7=AA=97=E2=89=8820GB)=E5=BF=85=E8=A2=AB=E5=86=85?= =?UTF-8?q?=E6=A0=B8=E5=87=BB=E6=9D=80,=E5=8D=A1=E6=AD=BB24/227=E5=85=AD?= =?UTF-8?q?=E6=AC=A1=E8=87=AA=E6=84=88=E5=85=A8=E6=92=9Ealpha19=E7=9A=84?= =?UTF-8?q?=E7=BB=9F=E4=B8=80=E6=A0=B9=E5=9B=A0;sum/min/max/mean/std/corr?= =?UTF-8?q?=E6=94=B9=E5=8E=9F=E7=94=9Frolling(min=5Fsamples=3D1=E6=88=AA?= =?UTF-8?q?=E6=96=AD=E7=AA=97=E4=B8=8E=E5=B1=95=E5=BC=80=E9=80=90=E7=82=B9?= =?UTF-8?q?=E6=81=92=E7=AD=89,13=E7=AD=89=E5=80=BC=E6=B5=8B=E8=AF=95?= =?UTF-8?q?=E9=94=81=E5=AE=9ANaN/null=E8=AF=AD=E4=B9=89=E5=B7=AE=E5=BC=82:?= =?UTF-8?q?min=E9=A1=BBfill=5Fnan(None),max=E4=B8=8Dfill=E5=9B=A0NaN?= =?UTF-8?q?=E8=A7=86=E4=BD=9C=E6=9C=80=E5=A4=A7=E4=B8=A4=E8=B7=AF=E5=90=8C?= =?UTF-8?q?=E6=AD=A5=E6=AF=92=E5=8C=96)=20[vps]?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- sanguo_factor/fast_ops.py | 84 ++++++++++++++ tests/factor/test_fast_ops_bigwindow.py | 147 ++++++++++++++++++++++++ 2 files changed, 231 insertions(+) create mode 100644 tests/factor/test_fast_ops_bigwindow.py diff --git a/sanguo_factor/fast_ops.py b/sanguo_factor/fast_ops.py index 7fc5ff9..85b78c8 100644 --- a/sanguo_factor/fast_ops.py +++ b/sanguo_factor/fast_ops.py @@ -22,6 +22,18 @@ if _VNPY_SRC not in sys.path: from vnpy.alpha.dataset.utility import DataProxy, EXPRESSION_FUNCTIONS +# 窗口超过此值时,shift 展开(物化 window 列 × 10M 行,250 窗≈20GB)会 OOM, +# 改走原生 rolling_*(min_samples=1 的截断窗求和与展开在数学上逐点恒等, +# 仅浮点求和顺序不同)。 +_BIG_WINDOW = 60 + +# polars 版本兼容: 新版 rolling_* 用 min_samples, 旧版用 min_periods +try: + pl.col("a").rolling_sum(3, min_samples=1) + _ROLL_MIN1_KW = {"min_samples": 1} +except TypeError: # pragma: no cover - 旧版 polars + _ROLL_MIN1_KW = {"min_periods": 1} + def fast_ts_rank(feature: DataProxy, window: int) -> DataProxy: """Percentile rank of current value within window [0,1] (polars native). @@ -303,6 +315,14 @@ def fast_ts_mean(feature: DataProxy, window: int) -> DataProxy: Uses: mean = Σ shift_k / count_nonnull(shift_k) for k in 0..window-1 Replicates: rolling_map(lambda s: np.nanmean(s), window, min_samples=1).over("vt_symbol") """ + if window > _BIG_WINDOW: + # rolling_mean(min_samples=1) 跳过 null 计分母,与 sum/count_nonnull 恒等 + mean_expr = pl.col("data").rolling_mean(window, **_ROLL_MIN1_KW).over("vt_symbol") + df = feature.df.select( + pl.col("datetime"), pl.col("vt_symbol"), mean_expr.alias("data") + ) + return DataProxy(df) + current = pl.col("data") # Create shifted versions with symbol boundary protection @@ -329,6 +349,14 @@ def fast_ts_std(feature: DataProxy, window: int) -> DataProxy: Uses: std = sqrt(E[x²] - E[x]²) with shift-based mean computation Replicates: rolling_map(lambda s: np.nanstd(s, ddof=0), window, min_samples=1).over("vt_symbol") """ + if window > _BIG_WINDOW: + # ddof=0 与 nanstd 对齐;rolling_std 默认 ddof=1 必须显式覆盖 + std_expr = pl.col("data").rolling_std(window, ddof=0, **_ROLL_MIN1_KW).over("vt_symbol") + df = feature.df.select( + pl.col("datetime"), pl.col("vt_symbol"), std_expr.alias("data") + ) + return DataProxy(df) + current = pl.col("data") # Create shifted versions with symbol boundary protection @@ -362,6 +390,17 @@ def fast_ts_sum(feature: DataProxy, window: int) -> DataProxy: Uses: sum = Σ shift_k for k in 0..window-1 Replicates: rolling_sum(window).over("vt_symbol") (no min_samples, so partial windows are NaN) """ + if window > _BIG_WINDOW: + # 大窗口原生路: rolling_sum 默认 min_samples=window,前 window-1 行为 null + sum_expr = pl.col("data").rolling_sum(window).over("vt_symbol") + boundary_mask = pl.int_range(pl.len()).over("vt_symbol") >= (window - 1) + df = feature.df.select( + pl.col("datetime"), + pl.col("vt_symbol"), + pl.when(boundary_mask).then(sum_expr).otherwise(None).alias("data"), + ) + return DataProxy(df) + current = pl.col("data") # Create shifted versions and sum them @@ -386,6 +425,15 @@ def fast_ts_min(feature: DataProxy, window: int) -> DataProxy: Uses: min = min(shift_0, shift_1, ..., shift_{w-1}) Replicates: rolling_min(window, min_samples=1).over("vt_symbol") """ + if window > _BIG_WINDOW: + # NaN 会毒化 polars rolling 整窗,而 min_horizontal 跳过 NaN; + # fill_nan(None) 后 rolling 跳过 null,与展开路逐点一致(测试锁定) + min_expr = pl.col("data").fill_nan(None).rolling_min(window, **_ROLL_MIN1_KW).over("vt_symbol") + df = feature.df.select( + pl.col("datetime"), pl.col("vt_symbol"), min_expr.alias("data") + ) + return DataProxy(df) + current = pl.col("data") # Create shifted versions and find minimum @@ -406,6 +454,15 @@ def fast_ts_max(feature: DataProxy, window: int) -> DataProxy: Uses: max = max(shift_0, shift_1, ..., shift_{w-1}) Replicates: rolling_max(window, min_samples=1).over("vt_symbol") """ + if window > _BIG_WINDOW: + # NaN 在 polars 视作最大值,max_horizontal 与 rolling_max 均被其毒化—— + # 不做 fill_nan,两条路 NaN 语义天然逐点一致(测试锁定) + max_expr = pl.col("data").rolling_max(window, **_ROLL_MIN1_KW).over("vt_symbol") + df = feature.df.select( + pl.col("datetime"), pl.col("vt_symbol"), max_expr.alias("data") + ) + return DataProxy(df) + current = pl.col("data") # Create shifted versions and find maximum @@ -436,6 +493,33 @@ def fast_ts_corr_v2(feature1: DataProxy, feature2: DataProxy, window: int) -> Da x = pl.col("data") y = pl.col("data_right") + if window > _BIG_WINDOW: + # 大窗口原生路: 逐对掩码(两操作数皆非 null 才计入)的 rolling 合成, + # 与展开版 count_expr/sum_horizontal 的截断窗语义逐点恒等。 + pair = x.is_not_null() & y.is_not_null() + xm = pl.when(pair).then(x) + ym = pl.when(pair).then(y) + n = pl.when(pair).then(1).otherwise(None).cast(pl.Float64).rolling_sum(window, **_ROLL_MIN1_KW).over("vt_symbol") + sx = xm.rolling_sum(window, **_ROLL_MIN1_KW).over("vt_symbol") + sy = ym.rolling_sum(window, **_ROLL_MIN1_KW).over("vt_symbol") + sxy = (xm * ym).rolling_sum(window, **_ROLL_MIN1_KW).over("vt_symbol") + sxx = (xm * xm).rolling_sum(window, **_ROLL_MIN1_KW).over("vt_symbol") + syy = (ym * ym).rolling_sum(window, **_ROLL_MIN1_KW).over("vt_symbol") + mean_x = sx / n + mean_y = sy / n + corr_expr = (sxy / n - mean_x * mean_y) / ( + ((sxx / n - mean_x.pow(2)) * (syy / n - mean_y.pow(2))).sqrt() + ) + df = df_merged.select( + pl.col("datetime"), + pl.col("vt_symbol"), + pl.when(corr_expr.is_infinite() | corr_expr.is_nan()) + .then(None) + .otherwise(corr_expr) + .alias("data") + ) + return DataProxy(df) + # Create shifted versions for both series shifts_x = [x.shift(k).over("vt_symbol") for k in range(window)] shifts_y = [y.shift(k).over("vt_symbol") for k in range(window)] diff --git a/tests/factor/test_fast_ops_bigwindow.py b/tests/factor/test_fast_ops_bigwindow.py new file mode 100644 index 0000000..081f1ee --- /dev/null +++ b/tests/factor/test_fast_ops_bigwindow.py @@ -0,0 +1,147 @@ +"""大窗口(>60)原生 rolling 路与 shift 展开路的等值测试。 + +250 窗展开会物化 ~20GB 列,故大窗口走原生 rolling_*。数学上 min_samples=1 +的截断窗求和与展开逐点恒等(仅浮点求和顺序不同);ts_sum 例外——原生路 +min_samples=window 对窗内含 null 的行出 null(展开路为忽略 null 的部分和), +故 ts_sum 只测无 null 场景。 +""" + +import sys + +import numpy as np +import polars as pl +import pytest + +_VNPY_SRC = "/Users/chufeng/.openclaw/sanguo_projects/sanguo_vnpy_v2/vnpy_v4.4.0" +if _VNPY_SRC not in sys.path: + sys.path.insert(0, _VNPY_SRC) + +from vnpy.alpha.dataset.utility import DataProxy + +import sanguo_factor.fast_ops as fo + + +WINDOW = 65 # > _BIG_WINDOW=60 → 走原生路;对照时临时抬高阈值走展开路 +N_ROWS = 300 +SYMBOLS = ["000001.SZ", "600000.SH"] + + +def _make_frame(seed: int, with_nulls: bool) -> pl.DataFrame: + rng = np.random.default_rng(seed) + frames = [] + for si, sym in enumerate(SYMBOLS): + vals = rng.normal(1.0 + si, 0.05, N_ROWS) + if with_nulls: + mask = rng.random(N_ROWS) < 0.08 + vals = vals.copy() + vals[mask] = np.nan + frames.append(pl.DataFrame({ + "vt_symbol": pl.Series([sym] * N_ROWS, dtype=pl.String), + "data": pl.Series(vals, dtype=pl.Float64), + })) + df = pl.concat(frames) + df = df.with_row_index("_i").with_columns( + (pl.lit(N_ROWS) * pl.lit(0) + pl.col("_i") % N_ROWS).alias("_d") + ).drop("_i") + # datetime 按符号内行序生成 + df = df.with_columns( + pl.int_range(pl.len()).over("vt_symbol").alias("datetime") + ).drop("_d") + return df.select("datetime", "vt_symbol", "data") + + +def _run(op_name: str, df: pl.DataFrame, window: int, force_expansion: bool): + old = fo._BIG_WINDOW + if force_expansion: + fo._BIG_WINDOW = 10 ** 9 + try: + fn = getattr(fo, op_name) + out = fn(DataProxy(df), window) + finally: + fo._BIG_WINDOW = old + return out.df["data"].to_numpy() + + +def _assert_close(a, b, rtol=1e-9, atol=1e-9): + a = np.asarray(a, dtype=float) + b = np.asarray(b, dtype=float) + assert a.shape == b.shape + na, nb = np.isnan(a), np.isnan(b) + assert np.array_equal(na, nb), f"null 位不同: {na.sum()} vs {nb.sum()}" + m = ~na + if m.any(): + assert np.allclose(a[m], b[m], rtol=rtol, atol=atol), ( + f"最大偏差 {np.max(np.abs(a[m] - b[m]))}" + ) + + +@pytest.mark.parametrize("with_nulls", [False, True]) +@pytest.mark.parametrize("op", ["fast_ts_mean", "fast_ts_std", "fast_ts_min", "fast_ts_max"]) +def test_unary_ops_bigwindow_equivalence(op, with_nulls): + df = _make_frame(seed=7, with_nulls=with_nulls) + native = _run(op, df, WINDOW, force_expansion=False) + expansion = _run(op, df, WINDOW, force_expansion=True) + _assert_close(native, expansion) + + +def test_ts_sum_bigwindow_equivalence_nullfree(): + df = _make_frame(seed=8, with_nulls=False) + native = _run("fast_ts_sum", df, WINDOW, force_expansion=False) + expansion = _run("fast_ts_sum", df, WINDOW, force_expansion=True) + _assert_close(native, expansion) + + +def test_ts_sum_bigwindow_head_nulls_are_null(): + """ts_sum 原生路:窗内含 null → null(与展开路的忽略语义不同,锁定原生行为).""" + df = _make_frame(seed=9, with_nulls=True) + native = _run("fast_ts_sum", df, WINDOW, force_expansion=False) + total_null_rows = int(np.isnan(native).sum()) + assert total_null_rows > 0 # 散布 null 必然毒化部分窗口 + + +@pytest.mark.parametrize("with_nulls", [False, True]) +def test_ts_corr_v2_bigwindow_equivalence(with_nulls): + rng = np.random.default_rng(11) + frames = [] + for si, sym in enumerate(SYMBOLS): + x = rng.normal(1.0, 0.05, N_ROWS) + y = 0.5 * x + rng.normal(0, 0.02, N_ROWS) # 构造相关 + if with_nulls: + x = x.copy() + y = y.copy() + x[rng.random(N_ROWS) < 0.08] = np.nan + y[rng.random(N_ROWS) < 0.06] = np.nan + frames.append(pl.DataFrame({ + "vt_symbol": pl.Series([sym] * N_ROWS, dtype=pl.String), + "data": pl.Series(x, dtype=pl.Float64), + "data_right": pl.Series(y, dtype=pl.Float64), + })) + df = pl.concat(frames).with_columns( + pl.int_range(pl.len()).over("vt_symbol").alias("datetime") + ).select("datetime", "vt_symbol", "data", "data_right") + + old = fo._BIG_WINDOW + try: + dp1 = DataProxy(df.select("datetime", "vt_symbol", "data")) + dp2 = DataProxy(df.select("datetime", "vt_symbol", pl.col("data_right").alias("data"))) + fo._BIG_WINDOW = 10 ** 9 + expansion = fo.fast_ts_corr_v2(dp1, dp2, WINDOW).df["data"].to_numpy() + fo._BIG_WINDOW = 60 + native = fo.fast_ts_corr_v2(dp1, dp2, WINDOW).df["data"].to_numpy() + finally: + fo._BIG_WINDOW = old + _assert_close(native, expansion, rtol=1e-7, atol=1e-9) + + +def test_boundary_no_cross_symbol_leakage(): + """符号边界不串数:两符号值域分离,大窗 rolling 后各自值域不变.""" + n = 80 + df = pl.DataFrame({ + "datetime": list(range(n)) + list(range(n)), + "vt_symbol": ["000001.SZ"] * n + ["600000.SH"] * n, + "data": [1.0] * n + [100.0] * n, + }) + native = _run("fast_ts_mean", df, 65, force_expansion=False) + # 各符号首行即有值(min_samples=1),且符号2绝不被符号1的值拉低 + assert native[0] == 1.0 + assert native[n] == 100.0