diff --git a/sanguo_portfolio/strategies/momentum_timing.py b/sanguo_portfolio/strategies/momentum_timing.py index 0e7046f..6f4a9eb 100644 --- a/sanguo_portfolio/strategies/momentum_timing.py +++ b/sanguo_portfolio/strategies/momentum_timing.py @@ -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) diff --git a/tests/portfolio/test_momentum_timing.py b/tests/portfolio/test_momentum_timing.py index 8205363..3d9f96a 100644 --- a/tests/portfolio/test_momentum_timing.py +++ b/tests/portfolio/test_momentum_timing.py @@ -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()(后者取真实今天)。