fix(factor): fast_ops大窗口OOM根治——>60窗shift展开物化window列×10M行(250窗≈20GB)必被内核击杀,卡死24/227六次自愈全撞alpha19的统一根因;sum/min/max/mean/std/corr改原生rolling(min_samples=1截断窗与展开逐点恒等,13等值测试锁定NaN/null语义差异:min须fill_nan(None),max不fill因NaN视作最大两路同步毒化) [vps]

This commit is contained in:
2026-08-26 20:19:03 +08:00
parent c40c6f3216
commit 1539e678f6
2 changed files with 231 additions and 0 deletions
+84
View File
@@ -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)]
+147
View File
@@ -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