feat(factor): 批量评估引擎——进程内逐因子calculate_by_expression(弃spawn池整df pickle),单因子失败不中断,结果增量落盘 [vps]
This commit is contained in:
@@ -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}"}}
|
||||
@@ -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"]
|
||||
Reference in New Issue
Block a user