659 lines
22 KiB
Python
659 lines
22 KiB
Python
"""
|
||
Fast polars-native drop-in replacements for vnpy's slow rolling operators.
|
||
|
||
This module provides native polars implementations that override vnpy's
|
||
rolling_map-based operators through the EXPRESSION_FUNCTIONS extension point.
|
||
Optimized for 10M+ row datasets on 2-core NAS infrastructure.
|
||
|
||
Registration: Call register_fast_ops() once before batch evaluation.
|
||
Idempotent and thread-safe.
|
||
"""
|
||
|
||
import sys
|
||
from typing import cast
|
||
|
||
import polars as pl
|
||
import numpy as np
|
||
|
||
# Inject vnpy source path (follow sanguo_factor/universe.py pattern)
|
||
_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, 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).
|
||
|
||
Replicates: scipy.stats.percentileofscore(s, s[-1]) / 100
|
||
Strategy: Count values <= current, divide by window size.
|
||
Handles NaN by treating NaN > everything (consistent with scipy).
|
||
Returns None for partial windows (like original).
|
||
"""
|
||
current = pl.col("data")
|
||
|
||
# Count values <= current for each shift k in 0..window-1
|
||
# Fill null with False to treat null comparisons as not <=
|
||
le_count = pl.sum_horizontal([
|
||
pl.when(current.shift(k).over("vt_symbol").le(current).fill_null(False))
|
||
.then(pl.lit(1))
|
||
.otherwise(pl.lit(0))
|
||
for k in range(window)
|
||
])
|
||
|
||
# Rank = count / window
|
||
result = pl.when(pl.int_range(pl.len()).over("vt_symbol") >= window - 1) \
|
||
.then(le_count / pl.lit(window)) \
|
||
.otherwise(None)
|
||
|
||
df = feature.df.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
result.alias("data")
|
||
)
|
||
|
||
return DataProxy(df)
|
||
|
||
|
||
def fast_ts_corr(feature1: DataProxy, feature2: DataProxy, window: int) -> DataProxy:
|
||
"""Correlation between two features over rolling window (optimized native).
|
||
|
||
Uses correlation identity: corr(x,y) = cov(x,y) / (std_x * std_y)
|
||
"""
|
||
# 两操作数源自同一父df且全算子保长保序(rolling/shift/cs_rank 均保长)——
|
||
# 按位拼接等价于按键join,免 10M 行 hash join(py-spy 实锤分钟级)
|
||
h1, h2 = feature1.df.height, feature2.df.height
|
||
if h1 != h2:
|
||
raise ValueError(f"ts_corr 操作数长度不一致: {h1} vs {h2}(行序对齐前提被破坏)")
|
||
df_merged = feature1.df.with_columns(feature2.df["data"].alias("data_right"))
|
||
|
||
# Use native polars rolling correlation for exact match
|
||
corr_result = pl.rolling_corr(
|
||
pl.col("data"),
|
||
pl.col("data_right"),
|
||
window_size=window,
|
||
min_samples=1
|
||
).over("vt_symbol")
|
||
|
||
df = df_merged.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
pl.when(corr_result.is_infinite() | corr_result.is_nan())
|
||
.then(None)
|
||
.otherwise(corr_result)
|
||
.alias("data")
|
||
)
|
||
|
||
return DataProxy(df)
|
||
|
||
|
||
def fast_ts_cov(feature1: DataProxy, feature2: DataProxy, window: int) -> DataProxy:
|
||
"""Covariance between two features over rolling window (native).
|
||
|
||
Uses identity: cov(x,y) = corr(x,y) * std_x * std_y
|
||
"""
|
||
# Get correlation and standard deviations, then compute covariance
|
||
corr_result = fast_ts_corr(feature1, feature2, window)
|
||
|
||
# Extract std deviations from the merged dataframe
|
||
# 两操作数源自同一父df且全算子保长保序——按位拼接等价于按键join
|
||
h1, h2 = feature1.df.height, feature2.df.height
|
||
if h1 != h2:
|
||
raise ValueError(f"ts_cov 操作数长度不一致: {h1} vs {h2}(行序对齐前提被破坏)")
|
||
df_merged = feature1.df.with_columns(feature2.df["data"].alias("data_right"))
|
||
|
||
std_x = pl.col("data").rolling_std(window, min_samples=1, ddof=0).over("vt_symbol")
|
||
std_y = pl.col("data_right").rolling_std(window, min_samples=1, ddof=0).over("vt_symbol")
|
||
|
||
# Covariance = corr * std_x * std_y
|
||
df = df_merged.with_columns([
|
||
std_x.alias("std_x"),
|
||
std_y.alias("std_y"),
|
||
corr_result.df["data"].alias("corr")
|
||
])
|
||
|
||
df = df.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
(pl.col("corr") * pl.col("std_x") * pl.col("std_y")).alias("data")
|
||
)
|
||
|
||
# Handle infinite/NaN values
|
||
df = df.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
pl.when(pl.col("data").is_infinite() | pl.col("data").is_nan())
|
||
.then(None)
|
||
.otherwise(pl.col("data"))
|
||
.alias("data")
|
||
)
|
||
|
||
return DataProxy(df)
|
||
|
||
|
||
def fast_ts_decay_linear(feature: DataProxy, window: int) -> DataProxy:
|
||
"""Linear decay weighted average: weights (w, w-1, ..., 1) / sum (polars native).
|
||
|
||
Optimized by expanding the weighted sum as:
|
||
Σ_k (k+1) * shift(k) for k in 0..window-1
|
||
where shift(0) gets weight 1 (last element), shift(window-1) gets weight w (first element)
|
||
Returns None for partial windows (like original).
|
||
"""
|
||
current = pl.col("data")
|
||
|
||
# Calculate weighted sum: Σ_k (k+1) * shift(k) for k in 0..window-1
|
||
# shift(k) goes backward in time, so higher k = earlier element = higher weight
|
||
weighted_sum = pl.sum_horizontal([
|
||
pl.lit(k + 1) * current.shift(k).over("vt_symbol")
|
||
for k in range(window)
|
||
])
|
||
|
||
denominator = window * (window + 1) // 2
|
||
|
||
result = pl.when(pl.int_range(pl.len()).over("vt_symbol") >= window - 1) \
|
||
.then(weighted_sum / pl.lit(denominator)) \
|
||
.otherwise(None)
|
||
|
||
df = feature.df.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
result.alias("data")
|
||
)
|
||
|
||
return DataProxy(df)
|
||
|
||
|
||
def fast_ts_slope(feature: DataProxy, window: int) -> DataProxy:
|
||
"""OLS slope over rolling window: cov(x,t) / var(t) where t = 0..w-1 (native).
|
||
|
||
Replicates the optimized original formula but with cleaner native polars syntax.
|
||
x = time index (0, 1, 2, ..., w-1)
|
||
slope = cov(y, x) / var(x)
|
||
"""
|
||
n = window
|
||
mean_x = (n - 1) / 2.0
|
||
var_x = (n**2 - 1) / 12.0 # Variance of [0, 1, ..., n-1]
|
||
|
||
# E[y*x] using weighted sum
|
||
mean_yx_expr = pl.sum_horizontal([
|
||
i * pl.col("data").shift(window - 1 - i).over("vt_symbol")
|
||
for i in range(n)
|
||
]) / n
|
||
|
||
mean_y = pl.col("data").rolling_mean(window, min_samples=window).over("vt_symbol")
|
||
|
||
# cov(y, x) = E[yx] - E[y]E[x]
|
||
cov_yx = mean_yx_expr - mean_y * mean_x
|
||
|
||
slope = cov_yx / var_x
|
||
|
||
df = feature.df.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
slope.alias("data")
|
||
)
|
||
|
||
return DataProxy(df)
|
||
|
||
|
||
def fast_ts_rsquare(feature: DataProxy, window: int) -> DataProxy:
|
||
"""R-squared of linear regression over rolling window (native).
|
||
|
||
r² = slope² * var(x) / var(y)
|
||
where var(x) is constant for window size.
|
||
"""
|
||
n = window
|
||
mean_x = (n - 1) / 2.0
|
||
var_x = (n**2 - 1) / 12.0 # Variance of [0, 1, ..., n-1]
|
||
|
||
# E[y*x]
|
||
mean_yx_expr = pl.sum_horizontal([
|
||
i * pl.col("data").shift(window - 1 - i).over("vt_symbol")
|
||
for i in range(n)
|
||
]) / n
|
||
|
||
mean_y = pl.col("data").rolling_mean(window, min_samples=window).over("vt_symbol")
|
||
|
||
cov_yx = mean_yx_expr - mean_y * mean_x
|
||
var_y = pl.col("data").rolling_var(window, min_samples=window, ddof=0).over("vt_symbol")
|
||
|
||
# r² = cov²(x,y) / (var(x) * var(y))
|
||
rsquare = (cov_yx.pow(2)) / (var_x * var_y)
|
||
|
||
df = feature.df.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
pl.when(rsquare.is_infinite() | rsquare.is_nan())
|
||
.then(None)
|
||
.otherwise(rsquare)
|
||
.alias("data")
|
||
)
|
||
|
||
return DataProxy(df)
|
||
|
||
|
||
def fast_ts_resi(feature: DataProxy, window: int) -> DataProxy:
|
||
"""Residuals from linear regression over rolling window (native).
|
||
|
||
residual = y - (intercept + slope * x_last)
|
||
where x_last = window - 1 (time index of last point)
|
||
"""
|
||
n = window
|
||
mean_x = (n - 1) / 2.0
|
||
var_x = (n**2 - 1) / 12.0
|
||
x_last = n - 1
|
||
|
||
# E[y*x]
|
||
mean_yx_expr = pl.sum_horizontal([
|
||
i * pl.col("data").shift(window - 1 - i).over("vt_symbol")
|
||
for i in range(n)
|
||
]) / n
|
||
|
||
mean_y = pl.col("data").rolling_mean(window, min_samples=window).over("vt_symbol")
|
||
|
||
cov_yx = mean_yx_expr - mean_y * mean_x
|
||
slope = cov_yx / var_x
|
||
intercept = mean_y - slope * mean_x
|
||
|
||
# residual = y - (intercept + slope * x_last)
|
||
residual = pl.col("data") - (intercept + slope * x_last)
|
||
|
||
df = feature.df.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
residual.alias("data")
|
||
)
|
||
|
||
return DataProxy(df)
|
||
|
||
|
||
def fast_ts_quantile(feature: DataProxy, window: int, quantile: float) -> DataProxy:
|
||
"""Quantile value over rolling window (native polars).
|
||
|
||
Uses rolling_quantile with linear interpolation to match original behavior.
|
||
Returns None for partial windows (like original).
|
||
"""
|
||
current = pl.col("data")
|
||
|
||
# Use native rolling_quantile with linear interpolation
|
||
quantile_result = current.rolling_quantile(
|
||
quantile=quantile,
|
||
interpolation="linear",
|
||
window_size=window
|
||
).over("vt_symbol")
|
||
|
||
# Apply window gating - null for partial windows
|
||
result = pl.when(pl.int_range(pl.len()).over("vt_symbol") >= window - 1) \
|
||
.then(quantile_result) \
|
||
.otherwise(None)
|
||
|
||
df = feature.df.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
result.alias("data")
|
||
)
|
||
|
||
return DataProxy(df)
|
||
|
||
|
||
def fast_ts_mean(feature: DataProxy, window: int) -> DataProxy:
|
||
"""Mean over rolling window with shift expansion (over-free).
|
||
|
||
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
|
||
shifts = [current.shift(k).over("vt_symbol") for k in range(window)]
|
||
|
||
# Sum all shifts, then count non-null values
|
||
sum_expr = pl.sum_horizontal(shifts)
|
||
count_expr = pl.sum_horizontal([pl.when(s.is_not_null()).then(1).otherwise(0) for s in shifts])
|
||
|
||
# Mean = sum / count (handles partial windows automatically)
|
||
mean_expr = sum_expr / count_expr
|
||
|
||
df = feature.df.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
mean_expr.alias("data")
|
||
)
|
||
return DataProxy(df)
|
||
|
||
|
||
def fast_ts_std(feature: DataProxy, window: int) -> DataProxy:
|
||
"""Standard deviation over rolling window with shift expansion (over-free).
|
||
|
||
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
|
||
shifts = [current.shift(k).over("vt_symbol") for k in range(window)]
|
||
|
||
# Count non-null values
|
||
count_expr = pl.sum_horizontal([pl.when(s.is_not_null()).then(1).otherwise(0) for s in shifts])
|
||
|
||
# Mean
|
||
sum_expr = pl.sum_horizontal(shifts)
|
||
mean_expr = sum_expr / count_expr
|
||
|
||
# E[x²] = sum(shift²) / count
|
||
sum_sq_expr = pl.sum_horizontal([s.pow(2) for s in shifts])
|
||
mean_sq_expr = sum_sq_expr / count_expr
|
||
|
||
# std = sqrt(E[x²] - E[x]²)
|
||
std_expr = (mean_sq_expr - mean_expr.pow(2)).sqrt()
|
||
|
||
df = feature.df.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
std_expr.alias("data")
|
||
)
|
||
return DataProxy(df)
|
||
|
||
|
||
def fast_ts_sum(feature: DataProxy, window: int) -> DataProxy:
|
||
"""Sum over rolling window with shift expansion (over-free).
|
||
|
||
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
|
||
shifts = [current.shift(k).over("vt_symbol") for k in range(window)]
|
||
sum_expr = pl.sum_horizontal(shifts)
|
||
|
||
# Apply boundary masking - null for partial windows (first window-1 rows per symbol)
|
||
boundary_mask = pl.int_range(pl.len()).over("vt_symbol") >= (window - 1)
|
||
result = pl.when(boundary_mask).then(sum_expr).otherwise(None)
|
||
|
||
df = feature.df.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
result.alias("data")
|
||
)
|
||
return DataProxy(df)
|
||
|
||
|
||
def fast_ts_min(feature: DataProxy, window: int) -> DataProxy:
|
||
"""Minimum over rolling window with shift expansion (over-free).
|
||
|
||
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
|
||
shifts = [current.shift(k).over("vt_symbol") for k in range(window)]
|
||
min_expr = pl.min_horizontal(shifts)
|
||
|
||
df = feature.df.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
min_expr.alias("data")
|
||
)
|
||
return DataProxy(df)
|
||
|
||
|
||
def fast_ts_max(feature: DataProxy, window: int) -> DataProxy:
|
||
"""Maximum over rolling window with shift expansion (over-free).
|
||
|
||
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
|
||
shifts = [current.shift(k).over("vt_symbol") for k in range(window)]
|
||
max_expr = pl.max_horizontal(shifts)
|
||
|
||
df = feature.df.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
max_expr.alias("data")
|
||
)
|
||
return DataProxy(df)
|
||
|
||
|
||
def fast_ts_corr_v2(feature1: DataProxy, feature2: DataProxy, window: int) -> DataProxy:
|
||
"""Correlation with shift expansion (over-free).
|
||
|
||
Uses: corr = (E[xy] - E[x]E[y]) / (std_x * std_y)
|
||
All statistics computed via shift-based mean/std
|
||
"""
|
||
# 两操作数源自同一父df且全算子保长保序(rolling/shift/cs_rank 均保长)——
|
||
# 按位拼接等价于按键join,免 10M 行 hash join(py-spy 实锤分钟级)
|
||
h1, h2 = feature1.df.height, feature2.df.height
|
||
if h1 != h2:
|
||
raise ValueError(f"ts_corr_v2 操作数长度不一致: {h1} vs {h2}(行序对齐前提被破坏)")
|
||
df_merged = feature1.df.with_columns(feature2.df["data"].alias("data_right"))
|
||
|
||
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)]
|
||
|
||
# Count non-null pairs
|
||
count_expr = pl.sum_horizontal([
|
||
pl.when(sx.is_not_null() & sy.is_not_null()).then(1).otherwise(0)
|
||
for sx, sy in zip(shifts_x, shifts_y)
|
||
])
|
||
|
||
# Means
|
||
mean_x = pl.sum_horizontal(shifts_x) / count_expr
|
||
mean_y = pl.sum_horizontal(shifts_y) / count_expr
|
||
|
||
# E[xy]
|
||
sum_xy = pl.sum_horizontal([sx * sy for sx, sy in zip(shifts_x, shifts_y)])
|
||
mean_xy = sum_xy / count_expr
|
||
|
||
# Standard deviations
|
||
sum_x_sq = pl.sum_horizontal([sx.pow(2) for sx in shifts_x])
|
||
sum_y_sq = pl.sum_horizontal([sy.pow(2) for sy in shifts_y])
|
||
var_x = (sum_x_sq / count_expr) - mean_x.pow(2)
|
||
var_y = (sum_y_sq / count_expr) - mean_y.pow(2)
|
||
std_x = var_x.sqrt()
|
||
std_y = var_y.sqrt()
|
||
|
||
# Correlation
|
||
corr_expr = (mean_xy - mean_x * mean_y) / (std_x * std_y)
|
||
|
||
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)
|
||
|
||
|
||
def fast_ts_cov_v2(feature1: DataProxy, feature2: DataProxy, window: int) -> DataProxy:
|
||
"""Covariance with shift expansion (over-free).
|
||
|
||
Uses: cov = corr * std_x * std_y
|
||
Delegates to fast_ts_corr_v2 for correlation
|
||
"""
|
||
# Get correlation
|
||
corr_result = fast_ts_corr_v2(feature1, feature2, window)
|
||
|
||
# Get individual std deviations
|
||
std_x_result = fast_ts_std(feature1, window)
|
||
std_y_result = fast_ts_std(feature2, window)
|
||
|
||
# Merge and compute covariance
|
||
# 三操作数源自同一父df且全算子保长保序——按位拼接等价于按键join
|
||
h_corr, h_x = corr_result.df.height, std_x_result.df.height
|
||
if h_corr != h_x:
|
||
raise ValueError(f"ts_cov_v2 corr与std_x长度不一致: {h_corr} vs {h_x}")
|
||
h_y = std_y_result.df.height
|
||
if h_corr != h_y:
|
||
raise ValueError(f"ts_cov_v2 corr与std_y长度不一致: {h_corr} vs {h_y}")
|
||
df_merged = corr_result.df.with_columns([
|
||
std_x_result.df["data"].alias("std_x"),
|
||
std_y_result.df["data"].alias("std_y")
|
||
])
|
||
|
||
df = df_merged.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
(pl.col("data") * pl.col("std_x") * pl.col("std_y")).alias("data")
|
||
)
|
||
|
||
# Handle infinite/NaN values
|
||
df = df.select(
|
||
pl.col("datetime"),
|
||
pl.col("vt_symbol"),
|
||
pl.when(pl.col("data").is_infinite() | pl.col("data").is_nan())
|
||
.then(None)
|
||
.otherwise(pl.col("data"))
|
||
.alias("data")
|
||
)
|
||
return DataProxy(df)
|
||
|
||
|
||
def register_fast_ops() -> list[str]:
|
||
"""Register fast polars operators into vnpy's EXPRESSION_FUNCTIONS.
|
||
|
||
Returns list of overridden operator names for verification.
|
||
Idempotent: safe to call multiple times.
|
||
|
||
Usage:
|
||
overrides = register_fast_ops()
|
||
print(f"Registered {len(overrides)} fast operators")
|
||
"""
|
||
operators = {
|
||
"ts_mean": fast_ts_mean,
|
||
"ts_std": fast_ts_std,
|
||
"ts_sum": fast_ts_sum,
|
||
"ts_min": fast_ts_min,
|
||
"ts_max": fast_ts_max,
|
||
"ts_corr": fast_ts_corr_v2,
|
||
"ts_cov": fast_ts_cov_v2,
|
||
"ts_rank": fast_ts_rank,
|
||
"ts_decay_linear": fast_ts_decay_linear,
|
||
"ts_slope": fast_ts_slope,
|
||
"ts_rsquare": fast_ts_rsquare,
|
||
"ts_resi": fast_ts_resi,
|
||
"ts_quantile": fast_ts_quantile,
|
||
}
|
||
|
||
# Register all operators
|
||
for name, func in operators.items():
|
||
EXPRESSION_FUNCTIONS[name] = func
|
||
|
||
return list(operators.keys())
|
||
|
||
|
||
# Auto-registration on import is intentional for batch_eval usage
|
||
# But also expose explicit registration for testing
|
||
__all__ = [
|
||
"register_fast_ops",
|
||
"fast_ts_rank",
|
||
"fast_ts_corr",
|
||
"fast_ts_cov",
|
||
"fast_ts_decay_linear",
|
||
"fast_ts_slope",
|
||
"fast_ts_rsquare",
|
||
"fast_ts_resi",
|
||
"fast_ts_quantile",
|
||
"fast_ts_mean",
|
||
"fast_ts_std",
|
||
"fast_ts_sum",
|
||
"fast_ts_min",
|
||
"fast_ts_max",
|
||
"fast_ts_corr_v2",
|
||
"fast_ts_cov_v2",
|
||
] |