1) 数据装配流式+列裁剪:Repository 新增 stream_range_many_columns(只 SELECT 所需列、SQL 侧转 REAL、yield_per 分批),引擎按 required_columns 取数 (LocalEngine 仅 close+因子字段),消除 ORM/Decimal 全量物化; 2) 研究 Job 独立子进程执行(job.mode=subprocess):python -m app.cli.run_job 在子进程内设 RLIMIT_AS 上限,OOM 归档 failed 而非拖垮 API worker; 子进程异常退出由父进程补记 failed;并发上限 2; 3) 服务启动清理:残留 queued/running Job 标记 failed(防永久 running)。 实测同款全市场回测:uvicorn worker RSS 稳定 ~220MB,任务峰值内存由 4.1GB+ 降至 ~470MB,24s 完成并归档(此前 43s 未完成即 OOM)。 新增/更新测试 96 passed,ruff 干净。
146 lines
5.5 KiB
Python
146 lines
5.5 KiB
Python
"""研究服务:把 Research Specification 编排为数据获取 + 引擎执行。
|
||
|
||
本层是业务入口:API / Agent 只能调用这里的用例(AGENT.md §16/§17),
|
||
禁止直接拼接引擎配置。数据一律经 Repository 获取(防未来函数由查询层保证)。
|
||
|
||
内存优化:大数据面板优先走 Repository 的流式列裁剪查询
|
||
(stream_range_many_columns,SQL 侧转 REAL、分批拉取),避免 ORM 对象 /
|
||
Decimal 全量物化;老实现回退到 get_range_many 逐实体路径。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from collections.abc import Iterable
|
||
from datetime import date, timedelta
|
||
|
||
import pandas as pd
|
||
|
||
from app.domain.entities.market import Stock
|
||
from app.domain.entities.research import (
|
||
BacktestResult,
|
||
FactorTestReport,
|
||
ResearchSpec,
|
||
UniverseSpec,
|
||
)
|
||
from app.domain.repositories.market import (
|
||
DailyBarRepository,
|
||
StockRepository,
|
||
)
|
||
from app.quant.engine import QuantEngine
|
||
|
||
# 流式路径每攒多少行落一个 DataFrame 分片(控制 concat 峰值)
|
||
_FRAME_CHUNK_ROWS = 50_000
|
||
|
||
|
||
def filter_stocks(stocks: list[Stock], universe: UniverseSpec, as_of: date) -> list[Stock]:
|
||
"""按股票池口径过滤(名称含 ST 判定 —— 名称快照为当日口径,属历史可追溯数据)。"""
|
||
out: list[Stock] = []
|
||
for s in stocks:
|
||
if s.delist_date is not None and s.delist_date < as_of:
|
||
continue
|
||
if universe.exclude_st and s.name and "ST" in s.name.upper():
|
||
continue
|
||
if (
|
||
universe.min_listing_days
|
||
and s.list_date
|
||
and (as_of - s.list_date).days < universe.min_listing_days
|
||
):
|
||
continue
|
||
out.append(s)
|
||
return out
|
||
|
||
|
||
def bars_to_daily_df(bars) -> pd.DataFrame:
|
||
"""DailyBar 列表 → 引擎长表 DataFrame(symbol/trade_date/ohlc/volume/amount)。
|
||
|
||
领域实体中的 Decimal 在此转 float,供 pandas 数值运算(保持 DataFrame 全数值列)。
|
||
"""
|
||
df = pd.DataFrame([b.model_dump() for b in bars])
|
||
if not df.empty:
|
||
for col in ("open", "high", "low", "close", "volume", "amount"):
|
||
if col in df.columns:
|
||
df[col] = df[col].astype(float)
|
||
return df
|
||
|
||
|
||
def _frame_from_stream(rows: Iterable[tuple], columns: list[str]) -> pd.DataFrame:
|
||
"""把流式 (symbol, trade_date_iso, *float_cols) 分批拼成 float 长表。
|
||
|
||
全程只保留分批 DataFrame + 最终一份结果,避免整批 tuple/Decimal 同时驻留。
|
||
返回列:symbol, trade_date(datetime64), *columns(float64)。
|
||
"""
|
||
cols = ["symbol", "trade_date", *columns]
|
||
buf: list[tuple] = []
|
||
pieces: list[pd.DataFrame] = []
|
||
for row in rows:
|
||
buf.append(row)
|
||
if len(buf) >= _FRAME_CHUNK_ROWS:
|
||
pieces.append(pd.DataFrame(buf, columns=cols))
|
||
buf = []
|
||
if buf:
|
||
pieces.append(pd.DataFrame(buf, columns=cols))
|
||
if not pieces:
|
||
return pd.DataFrame()
|
||
df = pd.concat(pieces, ignore_index=True)
|
||
df["trade_date"] = pd.to_datetime(df["trade_date"])
|
||
for col in columns: # NULL → NaN,统一 float64
|
||
df[col] = pd.to_numeric(df[col], errors="coerce")
|
||
return df
|
||
|
||
|
||
class ResearchService:
|
||
"""研究用例入口(因子测试 / 回测)。依赖注入 Repository 与引擎。"""
|
||
|
||
def __init__(
|
||
self,
|
||
stock_repo: StockRepository,
|
||
daily_repo: DailyBarRepository,
|
||
engine: QuantEngine,
|
||
) -> None:
|
||
self._stock_repo = stock_repo
|
||
self._daily_repo = daily_repo
|
||
self._engine = engine
|
||
|
||
def run_factor_test(self, spec: ResearchSpec, horizon_days: int = 21) -> FactorTestReport:
|
||
if spec.type != "factor_test":
|
||
raise ValueError("factor_test 用例需要 spec.type=factor_test")
|
||
daily = self._load_daily(spec)
|
||
return self._engine.run_factor_test(daily, spec, horizon_days=horizon_days)
|
||
|
||
def run_backtest(self, spec: ResearchSpec) -> BacktestResult:
|
||
if spec.type != "backtest":
|
||
raise ValueError("backtest 用例需要 spec.type=backtest")
|
||
daily = self._load_daily(spec)
|
||
return self._engine.run_backtest(daily, spec)
|
||
|
||
# ---- 数据装配 ----
|
||
|
||
def _load_daily(self, spec: ResearchSpec) -> pd.DataFrame:
|
||
start, end = spec.period
|
||
# 回测前预留因子 warmup(lookback≤120 交易日,取 300 自然日余量)
|
||
data_start = start - timedelta(days=300)
|
||
stocks = filter_stocks(self._stock_repo.list(), spec.universe, as_of=start)
|
||
if not stocks:
|
||
return pd.DataFrame()
|
||
symbols = [s.symbol for s in stocks]
|
||
|
||
# 引擎所需列裁剪(LocalEngine 只取 close + 因子字段;Qlib 回测取全 OHLCV)
|
||
required = self._engine.required_columns(spec)
|
||
streamer = getattr(self._daily_repo, "stream_range_many_columns", None)
|
||
if streamer is not None:
|
||
try:
|
||
return _frame_from_stream(
|
||
streamer(symbols, data_start, end, sorted(required)), sorted(required)
|
||
)
|
||
except Exception: # noqa: BLE001 —— 流式路径失败回退旧路径(兼容非 SQL 实现)
|
||
pass
|
||
# 旧路径:逐实体(供内存 / Fake 仓储等实现使用)
|
||
get_many = getattr(self._daily_repo, "get_range_many", None)
|
||
if get_many is not None:
|
||
bars = list(get_many(symbols, data_start, end))
|
||
else: # 兜底:逐只查询
|
||
bars = []
|
||
for s in stocks:
|
||
bars.extend(self._daily_repo.get_range(s.symbol, data_start, end))
|
||
return bars_to_daily_df(bars)
|