From 3180ab30025f83757e5eb6eda33029b0b9933935 Mon Sep 17 00:00:00 2001 From: GitHub Date: Wed, 19 Aug 2026 18:55:48 +0800 Subject: [PATCH] =?UTF-8?q?fix(backtest):=20=E7=AD=96=E7=95=A5=E5=8F=82?= =?UTF-8?q?=E6=95=B0=E5=8A=A0=E8=B7=A8=E8=AF=AD=E4=B9=89=E7=BA=A6=E6=9D=9F?= =?UTF-8?q?=EF=BC=8C=E6=9D=9C=E7=BB=9D=E5=AF=BB=E4=BC=98=E9=80=89=E5=87=BA?= =?UTF-8?q?=E5=BF=AB=E6=85=A2=E5=80=92=E6=8C=82=E7=BB=84=E5=90=88=EF=BC=88?= =?UTF-8?q?issue=20#39=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 寻优器网格探索不区分语义: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 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"]) diff --git a/src/easy_tdx/backtest/strategies/registry.py b/src/easy_tdx/backtest/strategies/registry.py index 040c634..d13283f 100644 --- a/src/easy_tdx/backtest/strategies/registry.py +++ b/src/easy_tdx/backtest/strategies/registry.py @@ -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) diff --git a/tests/unit/test_optimizer.py b/tests/unit/test_optimizer.py index f6cff24..f008fa3 100644 --- a/tests/unit/test_optimizer.py +++ b/tests/unit/test_optimizer.py @@ -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 diff --git a/tests/unit/test_web_backtest.py b/tests/unit/test_web_backtest.py index 2f8f75c..4274184 100644 --- a/tests/unit/test_web_backtest.py +++ b/tests/unit/test_web_backtest.py @@ -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)应抛 ValueError(issue #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