perf(portfolio): #12 03 momentum_timing层2向量化—当日预取宽表+内存切片:handle_data牛市日1-2次批量SQL(熊市1次,原每日~13次:1 buy_sign+10行业rps+1-2 select,726天×13≈9000+次IO无跨日缓存,全周期>3h被_TIMEOUT kill);三处_cal_rps/_select_stocks/_cal_buy_sign切片优先+回退直查(预取失败/缺列/空切片→None回退,行为等价旧版);_panel_slice dropna(how=all)精确复刻provider直查行集语义(超集切片组外日期全NaN行会改iloc[0]/tail(N)口径,测试单stock场景抓出);_stock_pool当日缓存(预取+_find_stock_pool共享get_index_stocks+filters);+1回归测试断言牛市日get_closes_panel≤2次+每行业get_index_stocks 1次,22绿 [vps]
CI/CD / test (push) Successful in 12s
CI/CD / nas-deploy (push) Successful in 23s
CI/CD / nas-verify (push) Successful in 9s

This commit is contained in:
2026-08-15 12:03:29 +08:00
parent fa04475693
commit e34a82f6bc
2 changed files with 157 additions and 26 deletions
+114 -26
View File
@@ -95,6 +95,12 @@ class MomentumTimingStrategy:
self.provider = provider
self.broker = broker or BrokerFacade()
self.config = config or MomentumTimingConfig()
# 当日预取缓存(层2向量化 #12):每日 1-2 次批量 SQL,替代原 ~13 次
# (10 行业 rps + buy_sign + select);失败/未启用 → None,各调用回退直查
self._day_panel: Optional[pd.DataFrame] = None
self._day_panel_end: Optional[str] = None
self._pool_cache: Optional[dict] = None
self._pool_cache_date: Optional[str] = None
# =================== initialize ===================
def initialize(self, context: Any) -> None:
@@ -138,7 +144,8 @@ class MomentumTimingStrategy:
cur_date = _to_date_str(cur_dt)
pre_date = _to_date_str(cur_dt - datetime.timedelta(days=cfg.past_day))
# 1) 牛熊分界
# 1) 牛熊分界(当日预取:先只拉指数点位;熊市日不付股票池全量 IO,同原行为)
self._ensure_day_panel(cfg.index_list, cur_date)
buy_sign = self._cal_buy_sign(cfg.index_list, cfg.past_day, cur_date)
logger.info("[%s] buy_sign=%s", cur_date, buy_sign)
@@ -151,7 +158,13 @@ class MomentumTimingStrategy:
self._close_position(stock)
return
# 2) 牛市:取强舍弱(每行业 RPS top_k 并集) → 候选池
# 2) 牛市:当日预取十行业成份股池 close 宽表(层2向量化,每日 1 次批量 SQL)
union_stocks: List[str] = []
for each_index in cfg.index_list:
union_stocks.extend(self._stock_pool_cached(each_index, cur_date))
self._ensure_day_panel(union_stocks, cur_date)
# 3) 取强舍弱(每行业 RPS top_k 并集) → 候选池
candidates = self._find_stock_pool(cfg.index_list, cur_date, pre_date)
# 3) 均线动量过滤(close > MA_short > MA_long)
@@ -202,6 +215,72 @@ class MomentumTimingStrategy:
break
logger.info("[%s] 牛市调仓结束: target=%s", cur_date, stocks)
# =================== 当日预取缓存(层2向量化 #12) ===================
def _ensure_day_panel(self, symbols: List[str], cur_date: str) -> None:
"""确保当日预取宽表覆盖 ``symbols``(增量补列;跨日自动重置)。
预取窗口 = cur - max(past_day, ma_long)*2 自然日(覆盖 _cal_buy_sign/
_cal_rps/_select_stocks 三者最大窗口)。失败 → 保持原状,各调用回退直查
(行为等价旧版,只是慢)。
"""
if self._day_panel_end != cur_date or self._day_panel is None:
self._day_panel = pd.DataFrame()
self._day_panel_end = cur_date
p = self._day_panel
missing = [s for s in symbols if s not in p.columns]
if not missing:
return
span = max(self.config.past_day, self.config.ma_long) * 2
start = _shift_date(cur_date, -span)
try:
chunk = self.provider.get_closes_panel(
_dedup(missing), start, cur_date, fq="raw",
)
except Exception as exc:
logger.warning("预取 get_closes_panel 失败(回退逐调用直查): %s", exc)
return
if chunk is None or chunk.empty:
return
self._day_panel = chunk if p.empty else p.join(chunk, how="outer")
def _panel_slice(
self,
symbols: List[str],
start_date: str,
end_date: str,
) -> Optional[pd.DataFrame]:
"""从当日预取宽表切 ``[start_date, end_date] × symbols``。
不可用(无预取/请求符号全不在列/切片空/异常)→ ``None``,调用方回退
provider 直查(等价旧行为;生产场景预取窗口必覆盖,回退仅异常时触发)。
缺失符号列直接剔除(=provider 版该列全 NaN → 下游 dropna/notna 剔除,
语义一致)。
"""
p = self._day_panel
if self._day_panel_end is None or p is None or p.empty:
return None
cols = [s for s in symbols if s in p.columns]
if not cols:
return None
try:
sub = p.loc[start_date:end_date, cols]
except Exception:
return None
# dropna(how="all") 复刻 provider 直查行集语义:其 panel 行 = 该查询组内
# 任一 symbol 有 bar 的日期(超集切片会带进组外日期的全 NaN 行,改变
# iloc[0]/tail(N) 口径)。剔除后与直查逐值一致。
sub = sub.dropna(how="all")
return sub if not sub.empty else None
def _stock_pool_cached(self, index_symbol: str, cur_date: str) -> List[str]:
"""_stock_pool 当日缓存(同日 10 行业并集预取 + _find_stock_pool 复用)。"""
if self._pool_cache_date != cur_date or self._pool_cache is None:
self._pool_cache = {}
self._pool_cache_date = cur_date
if index_symbol not in self._pool_cache:
self._pool_cache[index_symbol] = self._stock_pool(index_symbol, cur_date)
return self._pool_cache[index_symbol]
# =================== calRPS (修复:取 preDate~curDate 区间) ===================
def _cal_rps(
self,
@@ -225,15 +304,20 @@ class MomentumTimingStrategy:
n = len(stocks)
if n == 0:
return pd.DataFrame({"code": [], "rps_value": []})
try:
panel = self.provider.get_closes_panel(
stocks, pre_date, cur_date, fq="raw",
)
except Exception as exc:
logger.warning("_cal_rps get_closes_panel 失败: %s", exc)
return pd.DataFrame({"code": [], "rps_value": []})
if panel is None or panel.empty or len(panel) < 2:
return pd.DataFrame({"code": [], "rps_value": []})
panel = self._panel_slice(stocks, pre_date, cur_date)
if panel is not None:
if panel.empty or len(panel) < 2:
return pd.DataFrame({"code": [], "rps_value": []})
else:
try:
panel = self.provider.get_closes_panel(
stocks, pre_date, cur_date, fq="raw",
)
except Exception as exc:
logger.warning("_cal_rps get_closes_panel 失败: %s", exc)
return pd.DataFrame({"code": [], "rps_value": []})
if panel is None or panel.empty or len(panel) < 2:
return pd.DataFrame({"code": [], "rps_value": []})
# 每只股票涨跌幅(末值/首值 - 1) — 向量化
first = panel.iloc[0]
@@ -265,7 +349,7 @@ class MomentumTimingStrategy:
cfg = self.config
out: List[str] = []
for each_index in index_list:
stocks = self._stock_pool(each_index, cur_date)
stocks = self._stock_pool_cached(each_index, cur_date)
if not stocks:
continue
rps_df = self._cal_rps(stocks, cur_date, pre_date)
@@ -288,13 +372,15 @@ class MomentumTimingStrategy:
if not stocks:
return []
start_date = _shift_date(cur_date, -cfg.ma_long * 2)
try:
panel = self.provider.get_closes_panel(
stocks, start_date, cur_date, fq="raw",
)
except Exception as exc:
logger.warning("_select_stocks get_closes_panel 失败: %s", exc)
return []
panel = self._panel_slice(stocks, start_date, cur_date)
if panel is None:
try:
panel = self.provider.get_closes_panel(
stocks, start_date, cur_date, fq="raw",
)
except Exception as exc:
logger.warning("_select_stocks get_closes_panel 失败: %s", exc)
return []
if panel is None or panel.empty:
return []
panel = panel.tail(cfg.ma_long)
@@ -340,13 +426,15 @@ class MomentumTimingStrategy:
if not index_list:
return False
start_date = _shift_date(cur_date, -past_day * 2)
try:
panel = self.provider.get_closes_panel(
index_list, start_date, cur_date, fq="raw",
)
except Exception as exc:
logger.warning("_cal_buy_sign get_closes_panel 失败: %s", exc)
return False
panel = self._panel_slice(index_list, start_date, cur_date)
if panel is None:
try:
panel = self.provider.get_closes_panel(
index_list, start_date, cur_date, fq="raw",
)
except Exception as exc:
logger.warning("_cal_buy_sign get_closes_panel 失败: %s", exc)
return False
if panel is None or panel.empty:
return False
panel = panel.tail(past_day)
+43
View File
@@ -395,6 +395,49 @@ class TestHandleData:
assert len(buy_calls) >= 1
assert any(c.args[0] == "CAND.XSHG" for c in buy_calls)
def test_bull_day_single_prefetch_queries(self):
"""#12 层2向量化:牛市全流程 provider.get_closes_panel ≤2 次
(指数预取 + 股票池预取;原每日 ~13 次),get_index_stocks 每行业仅 1 次
(_stock_pool 当日缓存,_find_stock_pool 复用)。
"""
cfg = MomentumTimingConfig(
index_list=["IDX1.XSHG", "IDX2.XSHG"], top_k=6,
ma_short=5, ma_long=15,
)
s = make_strategy(
index_stocks_map={
"IDX1.XSHG": ["A.XSHG"], "IDX2.XSHG": ["B.XSHG"],
},
config=cfg,
)
panel_calls: List[Any] = []
def _gcp(symbols, start=None, end=None, interval="d", fq="raw"):
panel_calls.append(list(symbols) if isinstance(symbols, list) else symbols)
syms = list(symbols)
if all(s.startswith("IDX") for s in syms):
return _make_close_wide(syms, [[20.0 + i for i in range(30)]] * len(syms), days=30)
return _make_close_wide(syms, [[10.0 + i for i in range(15)]] * len(syms), days=15)
s.provider.get_closes_panel.side_effect = _gcp
ctx = FakeContext(
current_dt=datetime(2024, 10, 8, 9, 30),
previous_date="2024-09-30", positions={}, cash=1_000_000,
)
s.handle_data(ctx)
# 预取收敛:牛市日 ≤2 次批量 SQL(原 1 buy_sign + 10 rps + 1-2 select ≈13)
assert len(panel_calls) <= 2, panel_calls
# _stock_pool 缓存:每行业 get_index_stocks 只查 1 次(handle_data 预取
# + _find_stock_pool 复用)
idx_calls = [
c for c in s.provider.get_index_stocks.call_args_list
if c.args and str(c.args[0]).startswith("IDX")
]
assert len(idx_calls) == 2, idx_calls
# 全流程仍有买入(链路活)
buys = [c for c in s.broker.order_target_value.call_args_list if c.args[1] != 0]
assert buys
def test_handle_data_uses_current_dt_not_today(self):
"""⚠️ 修复原始 bug 验证:handle_data 必须用 context.current_dt 计算 cur_date,
不能用 datetime.date.today()(后者取真实今天)。