mirror of
https://ghfast.top/https://github.com/aeroxw/easy-tdx.git
synced 2026-09-12 16:54:17 +08:00
Merge pull request #51 from handsomejustin/fix/issue-39-param-order-constraints
fix(backtest): 策略参数加跨语义约束,杜绝寻优选出快慢倒挂组合(issue #39)
This commit is contained in:
@@ -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"])
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -111,6 +111,47 @@ 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
|
||||
|
||||
Reference in New Issue
Block a user