de04a8904b
三策略(聚宽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待做,解锁长回测)
422 lines
16 KiB
Python
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"]
|