feat(factor): 评估股票池加载器——dbbardata分块直读polars,前缀60/00/30,300天lookback+45天forward缓冲,vwap派生,bar_idx预热 [vps]
This commit is contained in:
@@ -0,0 +1,116 @@
|
||||
"""评估股票池: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 = 60
|
||||
_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=30)
|
||||
try:
|
||||
likes = " OR ".join(f"symbol LIKE '{p}%'" for p in STOCK_PREFIXES)
|
||||
cur = conn.execute(
|
||||
f"SELECT DISTINCT symbol FROM dbbardata "
|
||||
f"WHERE interval='d' AND datetime>=? AND datetime<=? AND ({likes})",
|
||||
(lookback_start, forward_end),
|
||||
)
|
||||
symbols = [r[0] for r in cur.fetchall()]
|
||||
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 评估行集:窗口内 + 预热期已过."""
|
||||
# Convert datetime to string (YYYY-MM-DD) for string comparison
|
||||
# This avoids polars Datetime comparison issues
|
||||
return df.filter(
|
||||
(pl.col("datetime").dt.strftime("%Y-%m-%d") >= pl.lit(start))
|
||||
& (pl.col("datetime").dt.strftime("%Y-%m-%d") <= pl.lit(end))
|
||||
& (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,
|
||||
})
|
||||
@@ -0,0 +1,90 @@
|
||||
"""股票池加载:前缀过滤/时间缓冲/vwap/bar_idx/显式symbols/limit抽样."""
|
||||
import sqlite3
|
||||
import sys, os
|
||||
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")))
|
||||
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "vnpy_v4.4.0")))
|
||||
|
||||
import pytest
|
||||
from sanguo_factor.universe import load_universe_bars, evaluation_filter, WARMUP_BARS
|
||||
|
||||
_DDL = """
|
||||
CREATE TABLE dbbardata(
|
||||
symbol TEXT, exchange TEXT, datetime TEXT, interval TEXT,
|
||||
volume REAL, turnover REAL, open_interest REAL,
|
||||
open_price REAL, high_price REAL, low_price REAL, close_price REAL)
|
||||
"""
|
||||
|
||||
|
||||
def _mk_db(tmp_path, rows):
|
||||
db = str(tmp_path / "qt.db")
|
||||
conn = sqlite3.connect(db)
|
||||
conn.execute(_DDL)
|
||||
conn.executemany("INSERT INTO dbbardata VALUES(?,?,?,?,?,?,?,?,?,?,?)", rows)
|
||||
conn.commit(); conn.close()
|
||||
return db
|
||||
|
||||
|
||||
def _row(sym, ex, day, close, volume=100.0, turnover=None):
|
||||
return (sym, ex, f"{day} 00:00:00", "d", volume,
|
||||
turnover if turnover is not None else close * volume, 0,
|
||||
close, close, close, close)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def db(tmp_path):
|
||||
rows = []
|
||||
# 600000:60根(2017-10~2017-12) + 评估窗内 3 根
|
||||
for i in range(60):
|
||||
rows.append(_row("600000", "SSE", f"2017-10-{(i % 28) + 1:02d}", 10.0 + i * 0.01))
|
||||
for d in ("2018-01-02", "2018-01-03", "2018-01-04"):
|
||||
rows.append(_row("600000", "SSE", d, 11.0))
|
||||
# 000001:只有 2 根(次新,bar_idx<WARMUP 应被 evaluation_filter 排除)
|
||||
for d in ("2018-01-02", "2018-01-03"):
|
||||
rows.append(_row("000001", "SZSE", d, 5.0))
|
||||
# 应剔除:科创68/北交8开头/ETF 510300
|
||||
rows.append(_row("688001", "SSE", "2018-01-02", 20.0))
|
||||
rows.append(_row("830001", "BJSE", "2018-01-02", 3.0))
|
||||
rows.append(_row("510300", "SSE", "2018-01-02", 4.0))
|
||||
# 非日线 interval 应忽略
|
||||
rows.append(("600000", "SSE", "2018-01-02 09:35:00", "15m", 1, 1, 0, 1, 1, 1, 1))
|
||||
return _mk_db(tmp_path, rows)
|
||||
|
||||
|
||||
def test_load_filters_prefix_and_interval(db):
|
||||
df = load_universe_bars(db, "2018-01-01", "2018-01-31")
|
||||
syms = df["vt_symbol"].unique().to_list()
|
||||
assert set(syms) == {"600000.SSE", "000001.SZSE"}
|
||||
|
||||
|
||||
def test_load_lookback_buffer_and_vwap(db):
|
||||
df = load_universe_bars(db, "2018-01-01", "2018-01-31")
|
||||
# lookback 生效:2017-10 的 bar 也在(因子窗口预热需要)
|
||||
assert str(df["datetime"].min()).startswith("2017-10")
|
||||
# vwap = turnover/volume
|
||||
row = df.filter(df["vt_symbol"] == "600000.SSE").sort("datetime").row(0, named=True)
|
||||
assert row["vwap"] == pytest.approx(10.0)
|
||||
|
||||
|
||||
def test_bar_idx_per_symbol(db):
|
||||
df = load_universe_bars(db, "2018-01-01", "2018-01-31").sort(["vt_symbol", "datetime"])
|
||||
idx = df.filter(df["vt_symbol"] == "000001.SZSE")["bar_idx"].to_list()
|
||||
assert idx == [0, 1]
|
||||
|
||||
|
||||
def test_evaluation_filter_warmup(db):
|
||||
df = load_universe_bars(db, "2018-01-01", "2018-01-31")
|
||||
ev = evaluation_filter(df, "2018-01-01", "2018-01-31")
|
||||
# 000001 只有2根 < WARMUP_BARS → 全排除;600000 保留
|
||||
assert set(ev["vt_symbol"].unique().to_list()) == {"600000.SSE"}
|
||||
assert ev.height == 3
|
||||
|
||||
|
||||
def test_explicit_symbols(db):
|
||||
df = load_universe_bars(db, "2018-01-01", "2018-01-31", symbols=["000001"])
|
||||
assert set(df["vt_symbol"].unique().to_list()) == {"000001.SZSE"}
|
||||
|
||||
|
||||
def test_limit_deterministic(db):
|
||||
df1 = load_universe_bars(db, "2018-01-01", "2018-01-31", limit=1)
|
||||
df2 = load_universe_bars(db, "2018-01-01", "2018-01-31", limit=1)
|
||||
assert set(df1["vt_symbol"].unique()) == set(df2["vt_symbol"].unique())
|
||||
Reference in New Issue
Block a user