148 lines
5.4 KiB
Python
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
|