Files
sanguo_vnpy_v2/sanguo_portfolio/strategies/momentum_timing.py
T
claude_dev 32dbcb8958 feat(portfolio): P2/P3 策略层向量化 + fundamentals批量提速解锁长回测
P2 行情向量化(get_price→get_closes_panel,口径实证 max_abs_diff=0.0 零偏差):
- momentum_timing: _cal_rps/_select_stocks/_cal_buy_sign 三处向量化
- small_cap: _cal_momentum_score 用 close.min/max 代理 low/high(方案A)

P3 fundamentals 批量(small_cap _pick_stocks 加 fields=[market_cap,eps],
对接数据session f416a17 get_fundamentals_df 按需短路):
- 02 000985 全市场 5128只 ~19min卡死 → 195s 跑通解锁

验收: 72单测全过; VPS 03短回测+138%(口径与get_price一致diff=0)/
02全市场-11%(2024Q1小盘股灾期合理)/01沪深300可跑
2026-07-29 07:49:35 +08:00

418 lines
17 KiB
Python

"""聚宽"牛熊分界+取强舍弱+均线动量"策略(post905)翻译到 BulletTrade 框架。
聚宽源码完整保留在 ``docs/research/joinquant_strategies/03_momentum_timing/source.py``,
这里做**结构等价 + bug 修复**翻译:
- ``initialize`` → ``MomentumTimingStrategy.initialize``
- ``calRPS`` → ``MomentumTimingStrategy._cal_rps`` (**修复取数区间**)
- ``findStockPool`` → ``MomentumTimingStrategy._find_stock_pool``
- ``selectStocks`` → ``MomentumTimingStrategy._select_stocks``
- ``calBuySign`` → ``MomentumTimingStrategy._cal_buy_sign``
- ``handle_data`` → ``MomentumTimingStrategy.handle_data`` (**修复 date.today()**)
策略层不直接 import bullet-trade 顶层 API(避免 Mac dev 环境装不全崩),
通过两个注入点接入(照 all_weather 模式):
1. ``self.provider`` → LocalUnifiedProvider / 任意满足接口的 provider
2. ``self.broker`` → ``BrokerFacade``(注入聚宽风格全局函数)
⚠️ 已修复原始策略的两个致命 bug(详见 notes.md「移植记录」):
1. **calRPS 取数区间错** — 原代码 ``get_price(start=curDate, end=curDate)`` 只取 1 天,
``iloc[0]==iloc[-1]``,涨跌幅恒 0,RPS 排名完全失效 → 改为 ``start=preDate, end=curDate``
取真实区间算百分比涨跌幅。
2. **date.today() 用错** — 回测里取真实今天而非回测当前日 → 改用 ``context.current_dt``。
"""
from __future__ import annotations
import datetime
import logging
from dataclasses import dataclass, field
from typing import Any, List, Optional
import numpy as np
import pandas as pd
from .. import filters
from .all_weather import (
BrokerFacade,
_available_cash,
_current_dt,
_dedup,
_get_positions,
)
logger = logging.getLogger(__name__)
# ------------------------ Config ------------------------
# ✅ 板块选择说明(2026-07-28 G1 数据补全后切回原版):
# 原策略用 11 个中证行业指数(000928-000938)'index' 模式。此前因 constituent_unified 表
# 无行业指数成份股,降级用 9 个宽基指数替代;现 G1 已补全 000928-000937 共 10 个
# (000938 仍缺,记为遗留),恢复行业轮动原版。
# 逻辑机制(择时+取强舍弱+均线动量)不动,仅切回行业指数列表。
_DEFAULT_INDEX_LIST: List[str] = [
"000928.XSHG", # 中证能源
"000929.XSHG", # 中证材料
"000930.XSHG", # 中证工业
"000931.XSHG", # 中证可选消费
"000932.XSHG", # 中证主要消费
"000933.XSHG", # 中证医药卫生
"000934.XSHG", # 中证金融地产
"000935.XSHG", # 中证信息技术
"000936.XSHG", # 中证电信业务
"000937.XSHG", # 中证公用事业
]
@dataclass
class MomentumTimingConfig:
"""牛熊分界+取强舍弱+均线动量 策略参数(聚宽 g.* 全局变量抽出便于调参)。"""
# 板块列表(默认 10 个中证行业指数 000928-000937,G1 补全后切回原版,见模块顶部说明)
index_list: List[str] = field(default_factory=lambda: list(_DEFAULT_INDEX_LIST))
index_thre: float = 0.2 # g.indexThre:站上 past_day 日均线的行业比重阈值
past_day: int = 30 # g.pastDay:RPS + 牛熊分界回看窗口
top_k: int = 6 # g.topK:每行业 RPS top K + 最终持仓上限
benchmark: str = "000300.XSHG"
new_stock_days: int = 375 # 次新股过滤阈值
max_pool: int = 0 # 0=不限;MVP 验证用,限制 _stock_pool 返回前 N 只
ma_short: int = 5 # selectStocks 短均线窗口(原 mavg(5,'close'))
ma_long: int = 15 # selectStocks 长均线窗口(原 mavg(15,'close'))
# ------------------------ 策略 ------------------------
class MomentumTimingStrategy:
"""牛熊分界+取强舍弱+均线动量策略(纯量价,无基本面)。
实例化时不连数据/不下单,所有 IO 走注入的 ``provider`` 和 ``broker``。
runner 负责注入,测试用 mock。
"""
def __init__(
self,
provider: Any,
broker: Optional[BrokerFacade] = None,
config: Optional[MomentumTimingConfig] = None,
) -> None:
self.provider = provider
self.broker = broker or BrokerFacade()
self.config = config or MomentumTimingConfig()
# =================== initialize ===================
def initialize(self, context: Any) -> None:
"""聚宽 initialize 等价物:set_benchmark / 成本滑点 / 定时任务。"""
b = self.broker
b.set_benchmark(self.config.benchmark)
b.set_option("use_real_price", True)
b.set_option("avoid_future_data", True)
try:
from bullet_trade.core import FixedSlippage # type: ignore
b.set_slippage(FixedSlippage(0))
except Exception:
pass
try:
from bullet_trade.core import OrderCost # type: ignore
b.set_order_cost(
OrderCost(
open_tax=0, close_tax=0.001,
open_commission=0.0003, close_commission=0.0003,
close_today_commission=0, min_commission=5,
),
type="stock",
)
except Exception:
pass
# 定时任务:每日 9:30 触发 handle_data(原策略 handle_data 单位时间触发)
b.run_daily(self.handle_data, "9:30")
# =================== handle_data (主流程) ===================
def handle_data(self, context: Any) -> None:
"""每日调仓:牛熊分界 → 取强舍弱 → 均线动量 → 调仓下单。
⚠️ **修复原始 bug** — 用 ``context.current_dt`` 而非 ``datetime.date.today()``。
"""
cfg = self.config
cur_dt = _current_dt(context)
if cur_dt is None:
logger.warning("handle_data: context.current_dt 为 None,跳过")
return
cur_date = _to_date_str(cur_dt)
pre_date = _to_date_str(cur_dt - datetime.timedelta(days=cfg.past_day))
# 1) 牛熊分界
buy_sign = self._cal_buy_sign(cfg.index_list, cfg.past_day, cur_date)
logger.info("[%s] buy_sign=%s", cur_date, buy_sign)
positions = _get_positions(context)
if not buy_sign:
# 熊市:全部清仓(原策略语义)
logger.info("[%s] 熊市信号,清仓 %d", cur_date, len(positions))
for stock in list(positions.keys()):
self._close_position(stock)
return
# 2) 牛市:取强舍弱(每行业 RPS top_k 并集) → 候选池
candidates = self._find_stock_pool(cfg.index_list, cur_date, pre_date)
# 3) 均线动量过滤(close > MA_short > MA_long)
stocks = self._select_stocks(candidates, cur_date)
# 4) 候选过多时再按 RPS 取前 top_k (原策略 handle_data 第 171-175 行)
if len(stocks) > cfg.top_k:
rps_df = self._cal_rps(stocks, cur_date, pre_date)
stocks = list(rps_df["code"])[: cfg.top_k]
# 5) 过滤涨停/跌停/停牌(复用 sanguo_portfolio.filters)
stocks = filters.filter_limitup_stock(
stocks, self.provider, positions=list(positions.keys())
)
stocks = filters.filter_limitdown_stock(
stocks, self.provider, positions=list(positions.keys())
)
stocks = filters.filter_paused_stock(stocks, self.provider)
stocks = _dedup(stocks)
# 6) 调仓:先清掉不在 stocks 的
for stock in list(positions.keys()):
if stock in stocks:
continue
self._close_position(stock)
# 7) 等额买入 stocks 里的新股(原策略 cash/countStocks 语义)
positions = _get_positions(context) # 卖出后刷新
target_num = len(stocks)
if target_num == 0:
return
cash = _available_cash(context)
if cash <= 0:
return
per_value = cash / target_num
for stock in stocks:
if stock in positions:
continue
if self._open_position(stock, per_value):
positions = _get_positions(context) # 刷新
if len(positions) >= target_num:
break
logger.info("[%s] 牛市调仓结束: target=%s", cur_date, stocks)
# =================== calRPS (修复:取 preDate~curDate 区间) ===================
def _cal_rps(
self,
stocks: List[str],
cur_date: str,
pre_date: str,
) -> pd.DataFrame:
"""计算 RPS(相对强弱)排名。
⚠️ **修复原始 bug** — 原策略 ``get_price(start=curDate, end_date=curDate)``
只取 1 天,``iloc[0]==iloc[-1]``,涨跌幅恒 0,RPS 排名完全失效 →
改为取 ``preDate ~ curDate`` 区间算**百分比涨跌幅**(更符合 RPS 语义,
原代码用绝对差值排序会偏向高价股,见 notes.md「移植记录」)。
性能改造:用 ``get_closes_panel`` 一次批量取宽表(替代 get_price+pivot,
5128 只从逐只循环秒级降到 UNION ALL 批量)。数值口径不变(fq='raw')。
Returns:
DataFrame[code, rps_value],按 rps_value 降序;``rps_value = 99 - 100*i/n``。
"""
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": []})
# 每只股票涨跌幅(末值/首值 - 1) — 向量化
first = panel.iloc[0]
last = panel.iloc[-1]
with np.errstate(divide="ignore", invalid="ignore"):
returns = (last / first) - 1.0
# 过滤 NaN/Inf(数据不全或首值为 0)
valid = returns.replace([np.inf, -np.inf], np.nan).dropna()
if valid.empty:
return pd.DataFrame({"code": [], "rps_value": []})
# 降序:涨幅大的排前
sorted_codes = valid.sort_values(ascending=False).index.tolist()
m = len(sorted_codes)
rps_value = [99 - (100 * i / m) for i in range(m)]
return pd.DataFrame({"code": sorted_codes, "rps_value": rps_value})
# =================== findStockPool (取强舍弱) ===================
def _find_stock_pool(
self,
index_list: List[str],
cur_date: str,
pre_date: str,
) -> List[str]:
"""每个行业取 RPS top_k → 候选池并集。
原策略 ``findStockPool`` 第 67-82 行:逐行业 get_index_stocks → calRPS → 前 topK。
"""
cfg = self.config
out: List[str] = []
for each_index in index_list:
stocks = self._stock_pool(each_index, cur_date)
if not stocks:
continue
rps_df = self._cal_rps(stocks, cur_date, pre_date)
top = list(rps_df["code"])[: cfg.top_k]
out.extend(top)
return _dedup(out)
# =================== selectStocks (均线动量) ===================
def _select_stocks(self, stocks: List[str], cur_date: str) -> List[str]:
"""均线动量过滤:``close > MA_short`` 且 ``MA_short > MA_long``。
原策略 ``data[security].mavg(5,'close')`` (聚宽 Security.mavg),
翻译为 ``get_closes_panel`` 批量取 close 宽表,向量化算双均线。
性能改造:用 ``get_closes_panel`` 替代 ``get_price(count=N)``。
count=N → ``start = cur - N*2 自然日``、``end = cur``、宽表 ``.tail(N)`` 切片
(避免自然日 vs 交易日的换算歧义,``.tail(N)`` 最稳)。数值口径不变(fq='raw')。
"""
cfg = self.config
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 []
if panel is None or panel.empty:
return []
panel = panel.tail(cfg.ma_long)
if panel.empty:
return []
# 向量化:对每列算 close/ma_short/ma_long,过滤条件 close>ma_short>ma_long
valid_count = panel.notna().sum() # 每列非 NaN 计数,等价于原 dropna 后长度
close = panel.iloc[-1]
ma_short = panel.tail(cfg.ma_short).mean()
ma_long = panel.mean() # panel 已 tail(ma_long),整体均值即 MA_long
mask = (
(valid_count >= cfg.ma_long)
& close.notna()
& ma_short.notna()
& ma_long.notna()
& (close > ma_short)
& (ma_short > ma_long)
)
return list(mask.index[mask])
# =================== calBuySign (牛熊分界) ===================
def _cal_buy_sign(
self,
index_list: List[str],
past_day: int,
cur_date: str,
) -> bool:
"""统计 past_day 均线上方的指数占比 > index_thre → 牛市(True)。
原策略 'index' 模式(第 110-115 行):对每个指数算 ``mavg(past_day,'close')``
与 ``mavg(1,'close')`` 比较。翻译为 ``get_closes_panel`` 批量取 close 宽表,
向量化算 past_day 均线。
性能改造:用 ``get_closes_panel`` 替代 ``get_price(count=N)``。
count=N → ``start = cur - N*2 自然日``、``end = cur``、宽表 ``.tail(N)`` 切片。
数值口径不变(fq='raw')。
⚠️ 原代码 ``float(count)/len(indexList)`` 在 py2 是浮点除法(因 float()强转),
与 py3 一致。这里保留浮点除法语义。
"""
cfg = self.config
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
if panel is None or panel.empty:
return False
panel = panel.tail(past_day)
if panel.empty:
return False
# 向量化:对每列算 cur_close 与 past_day 均值,统计 close > ma_past 的占比
valid_count = panel.notna().sum()
cur_close = panel.iloc[-1]
ma_past = panel.mean()
mask = (
(valid_count >= 2)
& cur_close.notna()
& ma_past.notna()
& (cur_close > ma_past)
)
count = int(mask.sum())
return (count / len(index_list)) > cfg.index_thre
# =================== 调仓辅助 ===================
def _close_position(self, code: str) -> bool:
order = self.broker.order_target_value(code, 0)
return order is not None
def _open_position(self, code: str, value: float) -> bool:
order = self.broker.order_target_value(code, value)
return order is not None
# =================== 数据辅助 ===================
def _stock_pool(self, index_symbol: str, cur_date: str) -> List[str]:
"""成分股 + 过滤 ST/科创北交/次新。"""
try:
stocks = self.provider.get_index_stocks(index_symbol, cur_date)
except Exception as exc:
logger.warning("get_index_stocks(%s) 失败: %s", index_symbol, exc)
return []
stocks = filters.filter_kcbj_stock(stocks)
if self.config.max_pool > 0:
stocks = stocks[: self.config.max_pool]
stocks = filters.filter_st_stock(stocks, self.provider)
stocks = filters.filter_new_stock(
stocks, self.provider, cur_date, self.config.new_stock_days
)
return stocks
# ======================== 日期辅助 ========================
def _to_date_str(value: Any) -> str:
"""datetime/date/str → YYYY-MM-DD str。
聚宽风格 get_price 的 start/end_date 接受 'YYYY-MM-DD' 字符串。
"""
if isinstance(value, str):
return value[:10]
try:
return value.strftime("%Y-%m-%d")
except AttributeError:
return str(value)[:10]
def _shift_date(date_str: str, days: int) -> str:
"""字符串日期加减天数,返回 YYYY-MM-DD。
用于 ``count=N`` → ``start = end - N*2 自然日`` 的换算(配合宽表 ``.tail(N)`` 切片,
避免自然日 vs 交易日的歧义)。
"""
try:
dt = datetime.datetime.strptime(date_str[:10], "%Y-%m-%d")
except (ValueError, TypeError):
return date_str
return (dt + datetime.timedelta(days=days)).strftime("%Y-%m-%d")
__all__ = ["MomentumTimingStrategy", "MomentumTimingConfig"]