From 09b8eaa2cc7518619c80c495a46c3e36834a05c9 Mon Sep 17 00:00:00 2001 From: hmo Date: Wed, 12 Aug 2026 00:11:16 +0800 Subject: [PATCH] =?UTF-8?q?perf:=20=E5=A4=96=E9=83=A8=E5=9B=A0=E5=AD=90?= =?UTF-8?q?=E6=94=B9=E9=A2=84=E5=8A=A0=E8=BD=BD=E6=A8=A1=E5=BC=8F=E2=80=94?= =?UTF-8?q?=E2=80=94prepare=5Fexternal=5Fcontext=E4=B8=80=E6=AC=A1?= =?UTF-8?q?=E8=BD=BD=E5=85=A5=E5=9F=BA=E6=9C=AC=E9=9D=A2=E5=88=86=E4=BD=8D?= =?UTF-8?q?/=E8=A1=8C=E4=B8=9A=E5=8A=A8=E9=87=8F/=E6=96=B0=E9=97=BB?= =?UTF-8?q?=E5=88=B0=E5=86=85=E5=AD=98,=20=E6=B6=88=E9=99=A4=E5=BE=AA?= =?UTF-8?q?=E7=8E=AF=E5=86=85=E9=80=90=E6=97=A5=E6=9F=A5=E5=BA=93(403?= =?UTF-8?q?=E4=B8=87=E8=A1=8Cnews=E5=85=A8=E8=A1=A8=E6=89=AB=E5=8D=A1?= =?UTF-8?q?=E6=AD=BB=E6=A0=B9=E5=9B=A0)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- strategy_lab.py | 114 +++++++++++++++++++++++++++++++++--------------- 1 file changed, 79 insertions(+), 35 deletions(-) diff --git a/strategy_lab.py b/strategy_lab.py index 0cf35511..23b7154d 100644 --- a/strategy_lab.py +++ b/strategy_lab.py @@ -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_total0").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 pe0").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("""