From 3f9a4fc681e78d62fbb096e956a6684885d54468 Mon Sep 17 00:00:00 2001 From: claude_dev Date: Tue, 25 Aug 2026 00:21:54 +0800 Subject: [PATCH] =?UTF-8?q?feat(factor):=20=E6=89=B9=E9=87=8F=E8=AF=84?= =?UTF-8?q?=E4=BC=B0=E5=BC=95=E6=93=8E=E2=80=94=E2=80=94=E8=BF=9B=E7=A8=8B?= =?UTF-8?q?=E5=86=85=E9=80=90=E5=9B=A0=E5=AD=90calculate=5Fby=5Fexpression?= =?UTF-8?q?(=E5=BC=83spawn=E6=B1=A0=E6=95=B4df=20pickle),=E5=8D=95?= =?UTF-8?q?=E5=9B=A0=E5=AD=90=E5=A4=B1=E8=B4=A5=E4=B8=8D=E4=B8=AD=E6=96=AD?= =?UTF-8?q?,=E7=BB=93=E6=9E=9C=E5=A2=9E=E9=87=8F=E8=90=BD=E7=9B=98=20[vps]?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- sanguo_factor/batch_eval.py | 124 ++++++++++++++++++++++++++++++++ tests/factor/test_batch_eval.py | 74 +++++++++++++++++++ 2 files changed, 198 insertions(+) create mode 100644 sanguo_factor/batch_eval.py create mode 100644 tests/factor/test_batch_eval.py diff --git a/sanguo_factor/batch_eval.py b/sanguo_factor/batch_eval.py new file mode 100644 index 0000000..0fb3790 --- /dev/null +++ b/sanguo_factor/batch_eval.py @@ -0,0 +1,124 @@ +# sanguo_factor/batch_eval.py +"""批量评估引擎:全A bars → 进程内逐因子 calculate_by_expression → 指标 → 落盘. + +不走 AlphaDataset.prepare_data(其 spawn 池对每个表达式 pickle 整个 DataFrame, +9M 行 × 258 因子的传输开销不可接受);calculate_by_expression 纯进程内 polars, +内存随单因子天然有界。现有单因子分析链路(analyzer/alphalens tears)零改动. +""" +import sys +import os +import time + +_VNPY_SRC = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "vnpy_v4.4.0")) +if _VNPY_SRC not in sys.path: + sys.path.insert(0, _VNPY_SRC) + +import pandas as pd +import polars as pl + +from .universe import load_universe_bars, evaluation_filter +from .registry import get_factor +from . import eval_store +from .metrics import summarize_factor + + +def _forward_return_matrices(close_wide: pd.DataFrame, periods=(1, 5, 10)) -> dict[int, pd.DataFrame]: + out = {} + for p in periods: + out[p] = close_wide.shift(-p) / close_wide - 1.0 + return out + + +def run_batch_eval( + factor_names: list[str], + start: str, + end: str, + eval_db: str, + label: str, + universe: str = "all_a", + symbols: list[str] | None = None, + limit: int | None = None, + cfg=None, + progress_cb=None, + vnpy_db_override: str | None = None, +) -> dict: + """跑一轮批量评估,结果增量写入 eval_db,返回摘要.""" + from vnpy.alpha.dataset.utility import calculate_by_expression + + if cfg is None: + from sanguo_data.config import load_config, find_config_path + cfg = load_config(find_config_path()) + vnpy_db = vnpy_db_override or cfg.data_paths["vnpy_db"] + + t0 = time.time() + bars = load_universe_bars(vnpy_db, start, end, symbols=symbols, limit=limit) + if bars.height == 0: + raise ValueError(f"股票池为空: vnpy_db={vnpy_db} window={start}~{end}") + + alpha_df = bars.select(["vt_symbol", "datetime", "open", "high", "low", "close", + "volume", "turnover", "open_interest", "vwap"]) + eval_rows = evaluation_filter(bars, start, end).select(["vt_symbol", "datetime"]) + + close_wide = ( + bars.select(["datetime", "vt_symbol", "close"]) + .pivot(index="datetime", on="vt_symbol", values="close") + .sort("datetime").to_pandas().set_index("datetime") + ) + close_wide.index = pd.to_datetime(close_wide.index) + rets = _forward_return_matrices(close_wide) + + universe_label = universe if symbols is None else "custom" + eval_store.init_db(eval_db) + run_id = eval_store.create_run( + eval_db, label=label, universe=universe_label, + symbols_count=bars["vt_symbol"].n_unique(), factors_total=len(factor_names), + start=start, end=end, params={"limit": limit, "symbols": symbols[:20] if symbols else None}, + ) + + errors: list[str] = [] + done = 0 + buffer: list[dict] = [] + for i, name in enumerate(factor_names): + row = _eval_one(name, alpha_df, eval_rows, rets, calculate_by_expression) + if "error" in row["metrics"]: + errors.append(name) + buffer.append(row) + done += 1 + if len(buffer) >= 20 or done == len(factor_names): + eval_store.save_results(eval_db, run_id, buffer) + buffer = [] + if progress_cb: + progress_cb(done, len(factor_names), name) + + eval_store.finish_run(eval_db, run_id, "done", factors_done=done) + return { + "run_id": run_id, + "factors_total": len(factor_names), + "factors_done": done, + "errors": errors, + "elapsed_sec": round(time.time() - t0, 1), + "symbols_count": int(bars["vt_symbol"].n_unique()), + } + + +def _eval_one(name: str, alpha_df: pl.DataFrame, eval_rows: pl.DataFrame, + rets: dict[int, pd.DataFrame], calculate_by_expression) -> dict: + factor = get_factor(name) + if factor is None: + return {"factor": name, "category": "unknown", "expression": "", + "metrics": {"error": f"因子未注册: {name}"}} + try: + res = calculate_by_expression(alpha_df, factor["expression"]) + f_long = res.join(eval_rows, on=["vt_symbol", "datetime"], how="inner") + F = ( + f_long.pivot(index="datetime", on="vt_symbol", values="data") + .sort("datetime").to_pandas().set_index("datetime") + ) + F.index = pd.to_datetime(F.index) + metrics = summarize_factor(F, rets[1], rets[5], rets[10]) + return {"factor": name, "category": factor["category"], + "expression": factor["expression"], "metrics": metrics} + except Exception as e: # 单因子失败不拖垮整批 + return {"factor": name, "category": factor["category"], + "expression": factor["expression"], + "metrics": {"error": f"{type(e).__name__}: {e}"}} diff --git a/tests/factor/test_batch_eval.py b/tests/factor/test_batch_eval.py new file mode 100644 index 0000000..0e9d665 --- /dev/null +++ b/tests/factor/test_batch_eval.py @@ -0,0 +1,74 @@ +# tests/factor/test_batch_eval.py +"""批量引擎端到端:合成小库 → 2因子计算 → 指标落盘可查.""" +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 import alpha_datasets # 挂载 +from sanguo_factor.batch_eval import run_batch_eval +from sanguo_factor import eval_store + +_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) +""" + + +@pytest.fixture(scope="module") +def db(tmp_path_factory): + """3 只 × ~400 交易日(2017-06~2018-12)合成库:一只单调涨/一只震荡/一只反着走.""" + # Mount alpha datasets first to ensure factors are available + alpha_datasets.mount_all() + + p = tmp_path_factory.mktemp("b") + db = str(p / "qt.db") + conn = sqlite3.connect(db) + conn.execute(_DDL) + import pandas as pd + days = pd.bdate_range("2017-06-01", "2018-12-28") + for i, day in enumerate(days): + d = day.strftime("%Y-%m-%d") + rows = [ + ("600000", "SSE", f"{d} 00:00:00", "d", 100.0, 1_000_000.0, 0, + 10 + i * 0.01, 10 + i * 0.01, 10 + i * 0.01, 10 + i * 0.01), # 单调涨 + ("000001", "SZSE", f"{d} 00:00:00", "d", 100.0, 500_000.0, 0, + 5, 5, 5, 5), # 平盘(截面另一端) + ("300001", "SZSE", f"{d} 00:00:00", "d", 200.0, 900_000.0, 0, + 20 - i * 0.005, 20 - i * 0.005, 20 - i * 0.005, 20 - i * 0.005), # 缓跌 + ] + conn.executemany("INSERT INTO dbbardata VALUES(?,?,?,?,?,?,?,?,?,?,?)", rows) + conn.commit(); conn.close() + return db + + +def test_run_batch_eval_end_to_end(db, tmp_path, monkeypatch): + eval_db = str(tmp_path / "factor_eval.db") + out = run_batch_eval( + factor_names=["ma_20", "roc_5"], start="2018-01-01", end="2018-06-30", + eval_db=eval_db, label="t", symbols=None, cfg=None, vnpy_db_override=db, + ) + assert out["factors_done"] == 2 + assert out["errors"] == [] + rows = eval_store.get_rows(eval_db, out["run_id"]) + assert {r["factor"] for r in rows} == {"ma_20", "roc_5"} + # ma_20 = ts_mean(close,20)/close:涨股该值持续低(均价低于现价) 跌股高 → 与次日收益负相关为主 + m = eval_store.get_detail(eval_db, out["run_id"], "ma_20")["metrics"] + assert "1" in m and "5" in m and "10" in m and "turnover" in m + assert isinstance(m["1"]["ic_mean"], float) + + +def test_bad_factor_recorded_not_fatal(db, tmp_path): + eval_db = str(tmp_path / "factor_eval.db") + out = run_batch_eval( + factor_names=["不存在的因子", "kmid"], start="2018-01-01", end="2018-03-31", + eval_db=eval_db, label="t2", cfg=None, vnpy_db_override=db, + ) + assert out["factors_done"] == 2 + assert len(out["errors"]) == 1 + detail = eval_store.get_detail(eval_db, out["run_id"], "不存在的因子") + assert "error" in detail["metrics"]