"""大窗口(>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