perf(factor): fast_ops v2 去over化——rolling.over(vt_symbol)在10M行分组滚动分钟级/算子(py-spy实锤ts_mean),改符号内shift展开统一模式(ts_mean/std/min/max/sum/corr/cov),等价性测试对齐原版+1M行速度护栏 [vps]

This commit is contained in:
2026-08-25 22:15:40 +08:00
parent a2fe51b4b0
commit 7bbba554f8
2 changed files with 412 additions and 5 deletions
+228 -2
View File
@@ -288,6 +288,220 @@ def fast_ts_quantile(feature: DataProxy, window: int, quantile: float) -> DataPr
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")
"""
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")
"""
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)
"""
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")
"""
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")
"""
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_merged = feature1.df.join(feature2.df, on=["datetime", "vt_symbol"])
x = pl.col("data")
y = pl.col("data_right")
# 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_merged = corr_result.df.join(
std_x_result.df.select(["datetime", "vt_symbol", pl.col("data").alias("std_x")]),
on=["datetime", "vt_symbol"]
).join(
std_y_result.df.select(["datetime", "vt_symbol", pl.col("data").alias("std_y")]),
on=["datetime", "vt_symbol"]
)
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.
@@ -299,9 +513,14 @@ def register_fast_ops() -> list[str]:
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_corr": fast_ts_corr,
"ts_cov": fast_ts_cov,
"ts_decay_linear": fast_ts_decay_linear,
"ts_slope": fast_ts_slope,
"ts_rsquare": fast_ts_rsquare,
@@ -328,4 +547,11 @@ __all__ = [
"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",
]