perf: 外部因子改预加载模式——prepare_external_context一次载入基本面分位/行业动量/新闻到内存, 消除循环内逐日查库(403万行news全表扫卡死根因)
This commit is contained in:
+79
-35
@@ -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_ret20(sector_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("""
|
||||
|
||||
Reference in New Issue
Block a user