From cfb604588ccecf1e4e795000150928b0935933b7 Mon Sep 17 00:00:00 2001 From: claude_dev Date: Tue, 25 Aug 2026 00:07:37 +0800 Subject: [PATCH] =?UTF-8?q?feat(factor):=20=E8=AF=84=E4=BC=B0=E8=82=A1?= =?UTF-8?q?=E7=A5=A8=E6=B1=A0=E5=8A=A0=E8=BD=BD=E5=99=A8=E2=80=94=E2=80=94?= =?UTF-8?q?dbbardata=E5=88=86=E5=9D=97=E7=9B=B4=E8=AF=BBpolars,=E5=89=8D?= =?UTF-8?q?=E7=BC=8060/00/30,300=E5=A4=A9lookback+45=E5=A4=A9forward?= =?UTF-8?q?=E7=BC=93=E5=86=B2,vwap=E6=B4=BE=E7=94=9F,bar=5Fidx=E9=A2=84?= =?UTF-8?q?=E7=83=AD=20[vps]?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- sanguo_factor/universe.py | 116 ++++++++++++++++++++++++++++++++++ tests/factor/test_universe.py | 90 ++++++++++++++++++++++++++ 2 files changed, 206 insertions(+) create mode 100644 sanguo_factor/universe.py create mode 100644 tests/factor/test_universe.py diff --git a/sanguo_factor/universe.py b/sanguo_factor/universe.py new file mode 100644 index 0000000..7194c91 --- /dev/null +++ b/sanguo_factor/universe.py @@ -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, + }) diff --git a/tests/factor/test_universe.py b/tests/factor/test_universe.py new file mode 100644 index 0000000..7a52427 --- /dev/null +++ b/tests/factor/test_universe.py @@ -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