Files
sanguo_vnpy_v2/sanguo_portfolio/strategies/momentum_timing.py
T
claude_dev de04a8904b feat(portfolio): 移植3聚宽策略到BulletTrade + 8bug修正 + 数据缺口文档
三策略(聚宽py2→BulletTrade 0.9.2,BrokerFacade注入跨版本兼容):
- momentum_timing 动量择时(牛熊分界+行业RPS+均线,切回10中证行业指数)
- value_selection 价值精选(6条基本面过滤,切回沪深300)
- small_cap 小市值(去IC对冲,切回000985中证全指)

框架:
- runner_backtest 加 --strategy 分发(原硬编码all_weather)
- provider 加 get_value_metrics(价值精选6条基本面,NOTICE_DATE治前视偏差)
- 72单测全过(21+27+24)

修8个回测实测发现的真bug:
- 01第⑥条EPS绝对值0.08~0.5与①大盘矛盾→6条交集恒空致全程空仓,按注释本意改净利润同比8~50%
- 03原帖calRPS取数区间错(get_price start=end只取1天)→涨跌幅恒0 RPS失效;date.today()取真实今天非回测日
- 02 universe 000985不在constituent_unified→候选池空

VPS实测(短区间验证逻辑,非长期表现): 01价值+23%/03行业轮动+48%/02选出20只小盘

数据缺口(详见docs/research/joinquant_strategies/SUMMARY.md + data_gaps_fix_plan.md):
- 三表"1/3损坏"误报已撤回(全扫5530文件/表0损坏,沪深95%+健康,仅北交所920xxx空,不做北交所)
- 真实缺口: 行业成份股(G1已补)/000985(G2已补)/IC期货(02对冲去掉)/provider批量接口(G5待做,解锁长回测)
2026-07-28 22:20:49 +08:00

422 lines
16 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「移植记录」)。
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:
df = self.provider.get_price(
stocks,
start_date=pre_date,
end_date=cur_date,
frequency="daily",
fields=["close"],
panel=False,
fill_paused=False,
)
except Exception as exc:
logger.warning("_cal_rps get_price 失败: %s", exc)
return pd.DataFrame({"code": [], "rps_value": []})
if df is None or df.empty:
return pd.DataFrame({"code": [], "rps_value": []})
try:
pivot = df.pivot(index="time", columns="code", values="close")
except Exception as exc:
logger.warning("_cal_rps pivot 失败: %s", exc)
return pd.DataFrame({"code": [], "rps_value": []})
if pivot.empty or len(pivot) < 2:
return pd.DataFrame({"code": [], "rps_value": []})
# 每只股票涨跌幅(末值/首值 - 1)
first = pivot.iloc[0]
last = pivot.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),
翻译为 provider.get_price(count=ma_long) 后段求均值。
"""
cfg = self.config
if not stocks:
return []
try:
df = self.provider.get_price(
stocks,
end_date=cur_date,
frequency="daily",
fields=["close"],
count=cfg.ma_long,
panel=False,
fill_paused=False,
)
except Exception as exc:
logger.warning("_select_stocks get_price 失败: %s", exc)
return []
if df is None or df.empty:
return []
try:
pivot = df.pivot(index="time", columns="code", values="close")
except Exception:
return []
if pivot.empty:
return []
out: List[str] = []
for col in pivot.columns:
series = pivot[col].dropna()
if len(series) < cfg.ma_long:
continue
close = float(series.iloc[-1])
ma_short = float(series.tail(cfg.ma_short).mean())
ma_long = float(series.tail(cfg.ma_long).mean())
if np.isnan(close) or np.isnan(ma_short) or np.isnan(ma_long):
continue
if close > ma_short and ma_short > ma_long:
out.append(col)
return out
# =================== 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')`` 比较。翻译为取 past_day 日 close(含当日),
算均值与最后一根 close 比较。
⚠️ 原代码 ``float(count)/len(indexList)`` 在 py2 是浮点除法(因 float()强转),
与 py3 一致。这里保留浮点除法语义。
"""
cfg = self.config
if not index_list:
return False
try:
df = self.provider.get_price(
index_list,
end_date=cur_date,
frequency="daily",
fields=["close"],
count=past_day,
panel=False,
fill_paused=False,
)
except Exception as exc:
logger.warning("_cal_buy_sign get_price 失败: %s", exc)
return False
if df is None or df.empty:
return False
try:
pivot = df.pivot(index="time", columns="code", values="close")
except Exception:
return False
if pivot.empty:
return False
count = 0
for col in pivot.columns:
series = pivot[col].dropna()
if len(series) < 2:
continue
ma_past = float(series.tail(past_day).mean())
cur_close = float(series.iloc[-1])
if np.isnan(ma_past) or np.isnan(cur_close):
continue
if cur_close > ma_past:
count += 1
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]
__all__ = ["MomentumTimingStrategy", "MomentumTimingConfig"]