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