perf: 外部因子改预加载模式——prepare_external_context一次载入基本面分位/行业动量/新闻到内存, 消除循环内逐日查库(403万行news全表扫卡死根因)

This commit is contained in:
hmo
2026-08-12 00:11:16 +08:00
parent dba24a2bde
commit 09b8eaa2cc
+79 -35
View File
@@ -664,7 +664,65 @@ def _load_index_ctx(index_code, start_date, end_date):
return ctx
# ── 2026-08-11 外部因子缓存(基本面/新闻/行业动量,回测超跌策略用)──
# 预加载模式:一次载入内存,避免循环内逐日查库(403万行news/62万行sector全表扫=卡死根因)
_EXTERNAL_CACHE = {}
_EXT_READY = False # prepare_external_context 是否已执行
_EXT_MCAP_Q = {} # code -> mcap 分位 (0~1)
_EXT_PE_Q = {} # code -> pe 分位 (0~1)
_EXT_SECTOR = {} # code -> sector_name
_EXT_SEC_RET20 = {} # code -> 20日行业动量(用最近收盘日)
_EXT_NEWS_DATES = {} # code -> sorted [日期串, ...]
def prepare_external_context():
"""预加载外部因子数据(回测前调用一次):
- stock_fundamentals(5546行) 全量载入算 mcap/pe 分位
- stock_sectors(581行) 载入 code→sector 映射
- sector_index_daily 载入行业20日动量
- stock_news 载入每code新闻日期列表(3日窗口用 bisect 查)
"""
global _EXT_READY, _EXT_MCAP_Q, _EXT_PE_Q, _EXT_SECTOR, _EXT_SEC_RET20, _EXT_NEWS_DATES, _EXTERNAL_CACHE
if _EXT_READY:
return
import sqlite3, bisect
conn = sqlite3.connect(DB_PATH)
try:
# 1. 基本面分位(当前快照全市场排序)
rows = conn.execute("SELECT code, mcap_total, pe FROM stock_fundamentals").fetchall()
mcaps = sorted(r[1] for r in rows if r[1] and r[1] > 0)
pes = sorted(r[2] for r in rows if r[2] and r[2] > 0)
n_mcap, n_pe = len(mcaps), len(pes)
for code, mcap, pe in rows:
if mcap and mcap > 0 and n_mcap:
_EXT_MCAP_Q[code] = round(bisect.bisect_left(mcaps, mcap) / n_mcap, 3)
if pe and pe > 0 and n_pe:
_EXT_PE_Q[code] = round(bisect.bisect_left(pes, pe) / n_pe, 3)
# 2. code→sector 映射
for code, sec in conn.execute("SELECT code, sector FROM stock_sectors_em").fetchall():
_EXT_SECTOR[code] = sec
for code, sec, src in conn.execute("SELECT code, sector_name, source FROM stock_sectors").fetchall():
if code not in _EXT_SECTOR and src in ('ths', 'hk_em', 'hk_manual'):
_EXT_SECTOR[code] = sec
# 3. 行业20日动量:每 sector 取最新收盘 + 20日前收盘
for sec in set(_EXT_SECTOR.values()):
closes = conn.execute(
"SELECT close FROM sector_index_daily WHERE sector=? ORDER BY date DESC LIMIT 21",
(sec,)).fetchall()
if len(closes) >= 20 and closes[-1][0] and closes[-1][0] > 0:
ret20 = round((closes[0][0] - closes[-1][0]) / closes[-1][0] * 100, 2)
for code, s in _EXT_SECTOR.items():
if s == sec:
_EXT_SEC_RET20[code] = ret20
# 4. 新闻日期列表(每 code 的所有新闻日期,3日窗口用 bisect 统计)
for code, d in conn.execute("SELECT code, SUBSTR(date,1,10) FROM stock_news").fetchall():
_EXT_NEWS_DATES.setdefault(code, []).append(d)
for code in _EXT_NEWS_DATES:
_EXT_NEWS_DATES[code] = sorted(_EXT_NEWS_DATES[code])
except sqlite3.OperationalError:
pass
finally:
conn.close()
_EXT_READY = True
_EXTERNAL_CACHE.clear()
def _days_ago(dt, n):
"""返回 dt 前 n 天的日期字符串"""
@@ -675,45 +733,29 @@ def _days_ago(dt, n):
return dt
def _get_external_factors(code, dt):
"""惰性查询 + 缓存:返回该 code 在 dt 日的外部因子(mcap_q/pe_q/news3/sec_ret20"""
"""纯内存查询:返回该 code 在 dt 日的外部因子(mcap_q/pe_q/news3/sec_ret20"""
key = (code, dt)
if key in _EXTERNAL_CACHE:
return _EXTERNAL_CACHE[key]
result = {}
try:
import sqlite3
_conn = sqlite3.connect(DB_PATH)
# 行业动量 sec_ret20sector_index_daily 20日涨跌)
_srow = _conn.execute("SELECT sector_name FROM stock_sectors WHERE code=? LIMIT 1", (code,)).fetchone()
if _srow and _srow[0]:
_sector = _srow[0]
_prev = _conn.execute(
"SELECT close FROM sector_index_daily WHERE sector=? AND date<=? ORDER BY date DESC LIMIT 21",
(_sector, dt)).fetchall()
if len(_prev) >= 20 and _prev[-1][0] and _prev[-1][0] > 0:
result['sec_ret20'] = round((_prev[0][0] - _prev[-1][0]) / _prev[-1][0] * 100, 2)
# 新闻 3 日计数 news3
_nrow = _conn.execute(
"SELECT COUNT(*) FROM stock_news WHERE code=? AND date>? AND date<=?",
(code, _days_ago(dt, 3), dt)).fetchone()
if _nrow:
result['news3'] = _nrow[0]
# 基本面分位 mcap_q/pe_q
_frow = _conn.execute(
"SELECT mcap_total, pe FROM stock_fundamentals WHERE code=?", (code,)).fetchone()
if _frow and _frow[0]:
_mcap = _frow[0]
_t = _conn.execute("SELECT COUNT(*) FROM stock_fundamentals WHERE mcap_total>0 AND mcap_total<?", (_mcap,)).fetchone()[0]
_tot = _conn.execute("SELECT COUNT(*) FROM stock_fundamentals WHERE mcap_total>0").fetchone()[0]
if _tot:
result['mcap_q'] = round(_t / _tot, 3)
if _frow and _frow[1]:
_pe = _frow[1]
_t = _conn.execute("SELECT COUNT(*) FROM stock_fundamentals WHERE pe>0 AND pe<?", (_pe,)).fetchone()[0]
_tot = _conn.execute("SELECT COUNT(*) FROM stock_fundamentals WHERE pe>0").fetchone()[0]
if _tot:
result['pe_q'] = round(_t / _tot, 3)
_conn.close()
if _EXT_READY:
# 基本面分位(预计算,O(1)
if code in _EXT_MCAP_Q:
result['mcap_q'] = _EXT_MCAP_Q[code]
if code in _EXT_PE_Q:
result['pe_q'] = _EXT_PE_Q[code]
# 行业20日动量(预计算)
if code in _EXT_SEC_RET20:
result['sec_ret20'] = _EXT_SEC_RET20[code]
# 新闻3日计数(bisect 统计 3日窗口)
if code in _EXT_NEWS_DATES:
import bisect
dates = _EXT_NEWS_DATES[code]
lo = bisect.bisect_right(dates, _days_ago(dt, 3))
hi = bisect.bisect_right(dates, dt)
result['news3'] = hi - lo
_EXTERNAL_CACHE[key] = result
return result
except Exception:
pass
_EXTERNAL_CACHE[key] = result
@@ -1169,6 +1211,7 @@ def run_backtest(strategy_version, start_date, end_date, capital=913000, save=Tr
prepare_flow_context(fetch_start, end_date)
prepare_weekly_context(fetch_start, end_date)
prepare_news_context(start_date, end_date)
prepare_external_context() # 2026-08-11: 预加载外部因子(基本面分位/行业动量/新闻),避免循环内查库
conn = sqlite3.connect(DB_PATH)
stocks = conn.execute("""
@@ -2119,6 +2162,7 @@ def run_mr_backtest(strategy_version, start_date, end_date, capital=913000,
prepare_flow_context(fetch_start, end_date)
prepare_weekly_context(fetch_start, end_date)
prepare_news_context(start_date, end_date)
prepare_external_context() # 2026-08-11: 预加载外部因子(基本面分位/行业动量/新闻),避免循环内查库
conn = sqlite3.connect(DB_PATH)
stocks = conn.execute("""