115 lines
4.5 KiB
Python
115 lines
4.5 KiB
Python
"""评估股票池:dbbardata 分块直读 → AlphaLab 格式 polars DataFrame.
|
|
|
|
口径(spec):主板+创业板(前缀 60/00/30,剔科创68/北交/ETF),退市股保留,
|
|
120 交易日预热(次新 + 因子窗口)。不做 ST 过滤(无 point-in-time 名称史,
|
|
当前名回溯过滤=前视偏差,比不过滤更糟)。
|
|
"""
|
|
import random
|
|
import sqlite3
|
|
from datetime import datetime, timedelta
|
|
|
|
import polars as pl
|
|
|
|
STOCK_PREFIXES = ("60", "00", "30")
|
|
WARMUP_BARS = 120
|
|
_CHUNK = 300 # 分块 IN 查询每块 symbol 数(控瞬时内存)
|
|
_LOOKBACK_DAYS = 300 # start 前缓冲(覆盖最大 60 日窗口 + 节假日)
|
|
_FORWARD_DAYS = 45 # end 后缓冲(覆盖 10 日前瞻收益)
|
|
|
|
_COLS = "symbol, exchange, datetime, volume, turnover, open_interest, open_price, high_price, low_price, close_price"
|
|
|
|
|
|
def load_universe_bars(
|
|
vnpy_db: str,
|
|
start: str,
|
|
end: str,
|
|
symbols: list[str] | None = None,
|
|
limit: int | None = None,
|
|
) -> pl.DataFrame:
|
|
"""读评估窗(含前后缓冲)全A日线 → AlphaLab 格式(含 vwap/bar_idx)."""
|
|
lookback_start = (
|
|
datetime.strptime(start, "%Y-%m-%d") - timedelta(days=_LOOKBACK_DAYS)
|
|
).strftime("%Y-%m-%d")
|
|
forward_end = (
|
|
datetime.strptime(end, "%Y-%m-%d") + timedelta(days=_FORWARD_DAYS)
|
|
).strftime("%Y-%m-%d")
|
|
|
|
if symbols is None:
|
|
conn = sqlite3.connect(vnpy_db, timeout=60)
|
|
try:
|
|
# 无过滤 DISTINCT symbol 走 (symbol,...) 前导索引顺序流式扫——
|
|
# 带 WHERE(interval/datetime/LIKE)的版本会退化为 26G 全表扫(NAS 实测>5min)。
|
|
# 前缀在 Python 侧滤;interval='d'/窗口过滤由下方分块数据查询天然承担
|
|
# (无日线数据的 symbol 返回 0 行,不进最终 df,语义不变)。
|
|
cur = conn.execute("SELECT DISTINCT symbol FROM dbbardata")
|
|
symbols = [r[0] for r in cur if str(r[0]).startswith(STOCK_PREFIXES)]
|
|
finally:
|
|
conn.close()
|
|
|
|
if not symbols:
|
|
return _empty_alpha_df()
|
|
if limit is not None and limit < len(symbols):
|
|
symbols = sorted(random.Random(42).sample(symbols, limit))
|
|
|
|
chunks: list[pl.DataFrame] = []
|
|
for i in range(0, len(symbols), _CHUNK):
|
|
part = symbols[i : i + _CHUNK]
|
|
ph = ",".join("?" * len(part))
|
|
conn = sqlite3.connect(vnpy_db, timeout=30)
|
|
conn.execute("PRAGMA busy_timeout=30000")
|
|
try:
|
|
cur = conn.execute(
|
|
f"SELECT {_COLS} FROM dbbardata "
|
|
f"WHERE interval='d' AND datetime>=? AND datetime<=? AND symbol IN ({ph})",
|
|
(lookback_start, forward_end, *part),
|
|
)
|
|
rows = cur.fetchall()
|
|
finally:
|
|
conn.close()
|
|
if rows:
|
|
chunks.append(pl.DataFrame(
|
|
rows,
|
|
schema={"symbol": pl.Utf8, "exchange": pl.Utf8, "dt": pl.Utf8,
|
|
"volume": pl.Float64, "turnover": pl.Float64, "open_interest": pl.Float64,
|
|
"open": pl.Float64, "high": pl.Float64, "low": pl.Float64, "close": pl.Float64},
|
|
orient="row",
|
|
))
|
|
if not chunks:
|
|
return _empty_alpha_df()
|
|
|
|
df = pl.concat(chunks)
|
|
df = (
|
|
df.with_columns(
|
|
pl.col("dt").str.slice(0, 10).str.to_datetime("%Y-%m-%d").alias("datetime"),
|
|
(pl.col("symbol") + "." + pl.col("exchange")).alias("vt_symbol"),
|
|
pl.when(pl.col("volume") > 0)
|
|
.then(pl.col("turnover") / pl.col("volume"))
|
|
.otherwise(None)
|
|
.alias("vwap"),
|
|
)
|
|
.sort(["vt_symbol", "datetime"])
|
|
.with_columns(pl.int_range(pl.len()).over("vt_symbol").alias("bar_idx"))
|
|
.select(["vt_symbol", "datetime", "open", "high", "low", "close",
|
|
"volume", "turnover", "open_interest", "vwap", "bar_idx"])
|
|
)
|
|
return df
|
|
|
|
|
|
def evaluation_filter(df: pl.DataFrame, start: str, end: str) -> pl.DataFrame:
|
|
"""IC 评估行集:窗口内 + 预热期已过."""
|
|
from datetime import datetime as _dt
|
|
return df.filter(
|
|
(pl.col("datetime") >= _dt.strptime(start, "%Y-%m-%d"))
|
|
& (pl.col("datetime") <= _dt.strptime(end, "%Y-%m-%d"))
|
|
& (pl.col("bar_idx") >= WARMUP_BARS)
|
|
)
|
|
|
|
|
|
def _empty_alpha_df() -> pl.DataFrame:
|
|
return pl.DataFrame(schema={
|
|
"vt_symbol": pl.Utf8, "datetime": pl.Datetime,
|
|
"open": pl.Float64, "high": pl.Float64, "low": pl.Float64, "close": pl.Float64,
|
|
"volume": pl.Float64, "turnover": pl.Float64, "open_interest": pl.Float64,
|
|
"vwap": pl.Float64, "bar_idx": pl.Int64,
|
|
})
|