Merge pull request #51 from handsomejustin/fix/issue-39-param-order-constraints

fix(backtest): 策略参数加跨语义约束,杜绝寻优选出快慢倒挂组合(issue #39)
This commit is contained in:
毛利哥
2026-08-19 19:05:42 +08:00
committed by GitHub
5 changed files with 92 additions and 5 deletions
+7 -1
View File
@@ -181,8 +181,14 @@ class ParamGridOptimizer:
for combo in itertools.product(*value_lists): for combo in itertools.product(*value_lists):
params = dict(zip(param_names, combo, strict=True)) params = dict(zip(param_names, combo, strict=True))
try: try:
# 寻优时跳过参数范围检查——探索超范围值是寻优的目的 # 寻优时跳过参数范围检查——探索超范围值是寻优的目的
# 但跨参数语义约束(如 fast<slow)仍生效,倒挂组合在此被跳过
strategy = entry.build(params, skip_bounds=True) strategy = entry.build(params, skip_bounds=True)
except ValueError:
# 预期内的无效组合(语义倒挂/相等),info 级即可,无需堆栈
logger.info("网格点 %s 语义无效,跳过", params)
continue
try:
engine = BacktestEngine( engine = BacktestEngine(
strategy=strategy, strategy=strategy,
cash=self._cash, cash=self._cash,
@@ -54,6 +54,7 @@ class MaCrossStrategy(ParametrizedStrategy):
Param("fast", int, default=5, min_value=1, max_value=60, label="快线周期"), 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("slow", int, default=20, min_value=5, max_value=250, label="慢线周期"),
] ]
param_constraints = [("fast", "slow")]
def init(self) -> None: def init(self) -> None:
self.ma_fast = self.I(MA, self.data.close, self.p["fast"]) 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("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("signal", int, default=9, min_value=2, max_value=50, label="信号周期"),
] ]
param_constraints = [("short", "long")]
def init(self) -> None: def init(self) -> None:
self.dif, self.dea, self._hist = self.I( 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("oversold", int, default=30, min_value=5, max_value=45, label="超卖线"),
Param("overbought", int, default=70, min_value=55, max_value=95, label="超买线"), Param("overbought", int, default=70, min_value=55, max_value=95, label="超买线"),
] ]
param_constraints = [("oversold", "overbought")]
def init(self) -> None: def init(self) -> None:
self.rsi = self.I(RSI, self.data.close, self.p["n"]) 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("fast", int, default=12, min_value=2, max_value=60, label="快线周期"),
Param("slow", int, default=26, min_value=5, max_value=120, label="慢线周期"), Param("slow", int, default=26, min_value=5, max_value=120, label="慢线周期"),
] ]
param_constraints = [("fast", "slow")]
def init(self) -> None: def init(self) -> None:
self.ema_fast = self.I(EMA, self.data.close, self.p["fast"]) 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("mid", int, default=20, min_value=5, max_value=60, label="中期"),
Param("long", int, default=60, min_value=20, max_value=250, label="长期"), Param("long", int, default=60, min_value=20, max_value=250, label="长期"),
] ]
param_constraints = [("short", "mid"), ("mid", "long")]
def init(self) -> None: def init(self) -> None:
self.ma_s = self.I(MA, self.data.close, self.p["short"]) 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("oversold", int, default=-100, min_value=-200, max_value=0, label="超卖线"),
Param("overbought", int, default=100, min_value=0, max_value=200, label="超买线"), Param("overbought", int, default=100, min_value=0, max_value=200, label="超买线"),
] ]
param_constraints = [("oversold", "overbought")]
def init(self) -> None: def init(self) -> None:
self.cci = self.I(CCI, self.data.close, self.data.high, self.data.low, self.p["n"]) 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("oversold", int, default=-80, min_value=-100, max_value=-40, label="超卖线"),
Param("overbought", int, default=-20, min_value=-60, max_value=0, label="超买线"), Param("overbought", int, default=-20, min_value=-60, max_value=0, label="超买线"),
] ]
param_constraints = [("oversold", "overbought")]
def init(self) -> None: def init(self) -> None:
self.wr, self._wr1 = self.I(WR, self.data.close, self.data.high, self.data.low, self.p["n"]) 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 ) from exc
# skip_bounds=True 时跳过范围/取值集合检查(供寻优器探索超范围值), # skip_bounds=True 时跳过范围/取值集合检查(供寻优器探索超范围值),
# 仍保留类型转换 + NaN/Inf 拦截。 # 仍保留类型转换 + NaN/Inf 拦截。跨参数语义约束在
# ParametrizedStrategy.__init__ 里另行校验,同样不受 skip_bounds 影响。
if not skip_bounds: if not skip_bounds:
if self.choices is not None and self.type is str and converted not in self.choices: if self.choices is not None and self.type is str and converted not in self.choices:
raise ValueError( raise ValueError(
@@ -188,6 +189,11 @@ class ParametrizedStrategy(Strategy):
# 类属性:参数 schema(子类覆盖)。ClassVar 表明这是类级配置而非实例字段。 # 类属性:参数 schema(子类覆盖)。ClassVar 表明这是类级配置而非实例字段。
params: ClassVar[list[Param]] = [] 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] 访问)。 # 实例属性:已校验的参数值字典(init/next 中通过 self.p[name] 访问)。
p: dict[str, Any] p: dict[str, Any]
@@ -195,7 +201,8 @@ class ParametrizedStrategy(Strategy):
"""从 kwargs 构造策略参数。 """从 kwargs 构造策略参数。
多余的未知参数抛 ValueError;缺失参数取默认值。 多余的未知参数抛 ValueError;缺失参数取默认值。
skip_bounds=True 时跳过参数范围检查(供寻优器探索超范围值) skip_bounds=True 时跳过参数范围检查(供寻优器探索超范围值)
但跨参数语义约束(``param_constraints``)仍然生效。
""" """
super().__init__() super().__init__()
declared = {param.name: param for param in self.params} declared = {param.name: param for param in self.params}
@@ -207,6 +214,16 @@ class ParametrizedStrategy(Strategy):
for name, param in declared.items(): for name, param in declared.items():
raw = kwargs.get(name, param.default) raw = kwargs.get(name, param.default)
resolved[name] = param.validate(raw, skip_bounds=skip_bounds) 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 self.p = resolved
@@ -244,7 +261,8 @@ class RegisteredStrategy:
) -> ParametrizedStrategy: ) -> ParametrizedStrategy:
"""用给定参数构造策略实例,缺失参数取默认值。 """用给定参数构造策略实例,缺失参数取默认值。
skip_bounds=True 时跳过参数范围检查(供寻优器探索超范围值) skip_bounds=True 时跳过参数范围检查(供寻优器探索超范围值)
跨参数语义约束(``param_constraints``)不受影响仍然生效。
""" """
return self.strategy_cls(**(params or {}), skip_bounds=skip_bounds) 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 result.heatmap["y_name"] == "slow"
assert len(result.heatmap["x"]) == 3 assert len(result.heatmap["x"]) == 3
assert len(result.heatmap["y"]) == 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: def test_no_heatmap_for_single_param(self) -> None:
"""1 参数时 heatmap 应为 None。""" """1 参数时 heatmap 应为 None。"""
@@ -92,6 +93,20 @@ class TestParamGridOptimizer:
assert result.heatmap is None assert result.heatmap is None
assert len(result.results) == 3 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: def test_grid_size_limit_exceeded(self) -> None:
"""超过 MAX_GRID_POINTS 应抛 ValueError。""" """超过 MAX_GRID_POINTS 应抛 ValueError。"""
big_grid = {f"p{i}": list(range(5)) for i in range(6)} # 5^6 = 15625 big_grid = {f"p{i}": list(range(5)) for i in range(6)} # 5^6 = 15625
+41
View File
@@ -111,6 +111,47 @@ def test_strategy_rejects_out_of_range_param():
get_registry().get("ma_cross").build({"fast": 999}) 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(): def test_strategy_rejects_unknown_name():
"""未知策略名应抛 KeyError。""" """未知策略名应抛 KeyError。"""
from easy_tdx.backtest.strategies import get_registry from easy_tdx.backtest.strategies import get_registry