Files
sanguo_vnpy_v2/tests/factor/test_fast_ops_bigwindow.py
T

148 lines
5.4 KiB
Python

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