fix(backtest): 策略参数加跨语义约束,杜绝寻优选出快慢倒挂组合(issue #39)

寻优器网格探索不区分语义:ma_cross 的 fast∈[5..60] × slow∈[10..250]
笛卡尔积包含 {"fast":30,"slow":20} 这类倒挂组合,倒挂的双均线交叉
本质是反向策略,回测成绩可能反而突出,从而被选为"最优参数"展示
(即 issue #39 截图中 {"fast":30,"slow":20} 的来源)。presets.py 的
注释"快线<慢线才有意义;去重无效组合"早已写下意图但从未实现。

根因是参数校验只有单参数 min/max,缺跨参数语义约束:

- ParametrizedStrategy 新增 param_constraints 类属性 [(a,b), ...]
  表示要求 a<b,在 __init__ 解析后统一校验,报错带中文标签;
  语义约束不受 skip_bounds 影响(寻优跳过的只是数值边界)
- 7 个策略声明约束:ma_cross/ema_cross(fast<slow)、macd
  (short<long)、triple_ma(short<mid<long)、rsi_reversal/cci/
  wr_reversal(oversold<overbought)
- 寻优器 build 阶段的 ValueError 降为 info 级跳过(无堆栈噪音),
  回测异常仍走 warning;预设网格 36 组合 → 有效 25 个
- 新增 10 个回归测试(旧代码全失败、新代码全通过)
This commit is contained in:
GitHub
2026-08-19 18:55:48 +08:00
parent cc17ab650d
commit 3180ab3002
5 changed files with 94 additions and 5 deletions
+7 -1
View File
@@ -181,8 +181,14 @@ class ParamGridOptimizer:
for combo in itertools.product(*value_lists):
params = dict(zip(param_names, combo, strict=True))
try:
# 寻优时跳过参数范围检查——探索超范围值是寻优的目的
# 寻优时跳过参数范围检查——探索超范围值是寻优的目的
# 但跨参数语义约束(如 fast<slow)仍生效,倒挂组合在此被跳过
strategy = entry.build(params, skip_bounds=True)
except ValueError:
# 预期内的无效组合(语义倒挂/相等),info 级即可,无需堆栈
logger.info("网格点 %s 语义无效,跳过", params)
continue
try:
engine = BacktestEngine(
strategy=strategy,
cash=self._cash,
@@ -54,6 +54,7 @@ class MaCrossStrategy(ParametrizedStrategy):
Param("fast", int, default=5, min_value=1, max_value=60, label="快线周期"),
Param("slow", int, default=20, min_value=5, max_value=250, label="慢线周期"),
]
param_constraints = [("fast", "slow")]
def init(self) -> None:
self.ma_fast = self.I(MA, self.data.close, self.p["fast"])
@@ -85,6 +86,7 @@ class MacdStrategy(ParametrizedStrategy):
Param("long", int, default=26, min_value=5, max_value=100, label="长期EMA"),
Param("signal", int, default=9, min_value=2, max_value=50, label="信号周期"),
]
param_constraints = [("short", "long")]
def init(self) -> None:
self.dif, self.dea, self._hist = self.I(
@@ -150,6 +152,7 @@ class RsiReversalStrategy(ParametrizedStrategy):
Param("oversold", int, default=30, min_value=5, max_value=45, label="超卖线"),
Param("overbought", int, default=70, min_value=55, max_value=95, label="超买线"),
]
param_constraints = [("oversold", "overbought")]
def init(self) -> None:
self.rsi = self.I(RSI, self.data.close, self.p["n"])
@@ -210,6 +213,7 @@ class EmaCrossStrategy(ParametrizedStrategy):
Param("fast", int, default=12, min_value=2, max_value=60, label="快线周期"),
Param("slow", int, default=26, min_value=5, max_value=120, label="慢线周期"),
]
param_constraints = [("fast", "slow")]
def init(self) -> None:
self.ema_fast = self.I(EMA, self.data.close, self.p["fast"])
@@ -239,6 +243,7 @@ class TripleMaStrategy(ParametrizedStrategy):
Param("mid", int, default=20, min_value=5, max_value=60, label="中期"),
Param("long", int, default=60, min_value=20, max_value=250, label="长期"),
]
param_constraints = [("short", "mid"), ("mid", "long")]
def init(self) -> None:
self.ma_s = self.I(MA, self.data.close, self.p["short"])
@@ -350,6 +355,7 @@ class CciStrategy(ParametrizedStrategy):
Param("oversold", int, default=-100, min_value=-200, max_value=0, label="超卖线"),
Param("overbought", int, default=100, min_value=0, max_value=200, label="超买线"),
]
param_constraints = [("oversold", "overbought")]
def init(self) -> None:
self.cci = self.I(CCI, self.data.close, self.data.high, self.data.low, self.p["n"])
@@ -377,6 +383,7 @@ class WrReversalStrategy(ParametrizedStrategy):
Param("oversold", int, default=-80, min_value=-100, max_value=-40, label="超卖线"),
Param("overbought", int, default=-20, min_value=-60, max_value=0, label="超买线"),
]
param_constraints = [("oversold", "overbought")]
def init(self) -> None:
self.wr, self._wr1 = self.I(WR, self.data.close, self.data.high, self.data.low, self.p["n"])
+21 -3
View File
@@ -151,7 +151,8 @@ class Param:
) from exc
# skip_bounds=True 时跳过范围/取值集合检查(供寻优器探索超范围值),
# 仍保留类型转换 + NaN/Inf 拦截。
# 仍保留类型转换 + NaN/Inf 拦截。跨参数语义约束在
# ParametrizedStrategy.__init__ 里另行校验,同样不受 skip_bounds 影响。
if not skip_bounds:
if self.choices is not None and self.type is str and converted not in self.choices:
raise ValueError(
@@ -188,6 +189,11 @@ class ParametrizedStrategy(Strategy):
# 类属性:参数 schema(子类覆盖)。ClassVar 表明这是类级配置而非实例字段。
params: ClassVar[list[Param]] = []
# 类属性:跨参数语义约束 (a, b) 列表,要求 a < b(如 ("fast", "slow"))。
# 与单参数边界不同,这是语义级校验:skip_bounds(寻优探索超范围值)也不
# 跳过——否则网格寻优会把"快线30慢线20"这类倒挂组合跑完还可能当选最优
# issue #39)。
param_constraints: ClassVar[list[tuple[str, str]]] = []
# 实例属性:已校验的参数值字典(init/next 中通过 self.p[name] 访问)。
p: dict[str, Any]
@@ -195,7 +201,8 @@ class ParametrizedStrategy(Strategy):
"""从 kwargs 构造策略参数。
多余的未知参数抛 ValueError;缺失参数取默认值。
skip_bounds=True 时跳过参数范围检查(供寻优器探索超范围值)
skip_bounds=True 时跳过参数范围检查(供寻优器探索超范围值)
但跨参数语义约束(``param_constraints``)仍然生效。
"""
super().__init__()
declared = {param.name: param for param in self.params}
@@ -207,6 +214,16 @@ class ParametrizedStrategy(Strategy):
for name, param in declared.items():
raw = kwargs.get(name, param.default)
resolved[name] = param.validate(raw, skip_bounds=skip_bounds)
labels = {param.name: param.label or param.name for param in self.params}
for smaller, larger in self.param_constraints:
a, b = resolved[smaller], resolved[larger]
if not a < b:
raise ValueError(
f"参数 '{labels[smaller]}'({smaller})={a} 必须小于 "
f"'{labels[larger]}'({larger})={b},周期/阈值倒挂在语义上无效"
)
self.p = resolved
@@ -244,7 +261,8 @@ class RegisteredStrategy:
) -> ParametrizedStrategy:
"""用给定参数构造策略实例,缺失参数取默认值。
skip_bounds=True 时跳过参数范围检查(供寻优器探索超范围值)
skip_bounds=True 时跳过参数范围检查(供寻优器探索超范围值)
跨参数语义约束(``param_constraints``)不受影响仍然生效。
"""
return self.strategy_cls(**(params or {}), skip_bounds=skip_bounds)
+16 -1
View File
@@ -79,7 +79,8 @@ class TestParamGridOptimizer:
assert result.heatmap["y_name"] == "slow"
assert len(result.heatmap["x"]) == 3
assert len(result.heatmap["y"]) == 3
assert len(result.heatmap["data"]) == 9 # 3×3
# 3×3=9 中 fast=20×slow=15(倒挂)与 20×20(相等)被语义约束跳过(issue #39)
assert len(result.heatmap["data"]) == 7
def test_no_heatmap_for_single_param(self) -> None:
"""1 参数时 heatmap 应为 None。"""
@@ -92,6 +93,20 @@ class TestParamGridOptimizer:
assert result.heatmap is None
assert len(result.results) == 3
def test_inverted_combos_skipped(self) -> None:
"""网格含倒挂组合(fast≥slow)时应全部跳过,best 不可能是倒挂(issue #39)。"""
opt = ParamGridOptimizer(
strategy_name="ma_cross",
param_grid={"fast": [5, 20, 30], "slow": [10, 20]},
df=_make_df(),
)
result = opt.run()
# 6 组合中仅 (5,10) 与 (5,20) 语义有效
assert len(result.results) == 2
assert all(r.params["fast"] < r.params["slow"] for r in result.results)
assert result.best is not None
assert result.best.params["fast"] < result.best.params["slow"]
def test_grid_size_limit_exceeded(self) -> None:
"""超过 MAX_GRID_POINTS 应抛 ValueError。"""
big_grid = {f"p{i}": list(range(5)) for i in range(6)} # 5^6 = 15625
+43
View File
@@ -111,6 +111,49 @@ def test_strategy_rejects_out_of_range_param():
get_registry().get("ma_cross").build({"fast": 999})
def test_strategy_rejects_inverted_period_params():
"""快/慢周期倒挂(fast≥slow)应抛 ValueErrorissue #39)。"""
from easy_tdx.backtest.strategies import get_registry
with pytest.raises(ValueError, match="必须小于"):
get_registry().get("ma_cross").build({"fast": 30, "slow": 20})
# 相等同样无效:两线重合,交叉永不触发
with pytest.raises(ValueError, match="必须小于"):
get_registry().get("ma_cross").build({"fast": 20, "slow": 20})
@pytest.mark.parametrize(
"name,params",
[
("ma_cross", {"fast": 30, "slow": 20}),
("ema_cross", {"fast": 30, "slow": 20}),
("macd", {"short": 30, "long": 20}),
("triple_ma", {"short": 30, "mid": 20}),
("triple_ma", {"short": 5, "mid": 80, "long": 60}),
("rsi_reversal", {"oversold": 60, "overbought": 50}),
("cci", {"oversold": 50, "overbought": -50}),
("wr_reversal", {"oversold": -50, "overbought": -60}),
],
)
def test_param_constraints_enforced_despite_skip_bounds(name, params):
"""跨参数语义约束不受 skip_bounds 影响(寻优防倒挂组合,issue #39)。"""
from easy_tdx.backtest.strategies import get_registry
with pytest.raises(ValueError, match="必须小于"):
get_registry().get(name).build(params, skip_bounds=True)
def test_skip_bounds_still_allows_out_of_range_values():
"""skip_bounds 仍应放行超范围取值(只是不放过语义倒挂)。"""
from easy_tdx.backtest.strategies import get_registry
inst = get_registry().get("ma_cross").build(
{"fast": 100, "slow": 200}, skip_bounds=True
)
assert inst.p["fast"] == 100
assert inst.p["slow"] == 200
def test_strategy_rejects_unknown_name():
"""未知策略名应抛 KeyError。"""
from easy_tdx.backtest.strategies import get_registry