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]
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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()(后者取真实今天)。
|
||||
|
||||
Reference in New Issue
Block a user