Files
qlib/backend/app/quant/service.py
T
Simon 195f5d41f4 perf(backend): 内存优化三项——全市场研究不再占满 8G
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 干净。
2026-09-06 22:12:44 +08:00

146 lines
5.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""研究服务:把 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)