fix(backtest): 修复市场环境过滤未生效

This commit is contained in:
shy3130
2026-08-11 12:43:46 +08:00
parent 8c00f7a3fd
commit 57eb641275
5 changed files with 195 additions and 23 deletions
+70 -12
View File
@@ -600,6 +600,7 @@ class StrategyBacktestService:
config.holding_days,
config.minute_fill,
json.dumps(config.overrides or {}, sort_keys=True, ensure_ascii=False, default=str),
json.dumps(config.regime_filter or {}, sort_keys=True, ensure_ascii=False, default=str),
)
def _resolve_composite_feature_plan(
@@ -854,6 +855,8 @@ class StrategyBacktestService:
_rm = self._build_regime_mask(
market_data.timestamp_labels, first.regime_filter,
getattr(getattr(self.engine.repo, "store", None), "data_dir", None),
required_start=first.start,
required_end=first.end,
)
if _rm is not None:
entry_time_mask = entry_time_mask & _rm
@@ -1139,10 +1142,15 @@ class StrategyBacktestService:
config.end,
)
# 市场环境过滤(强制 T-1): 只叠加 entry, 不影响 exit
_rm = self._build_regime_mask(
market_data.timestamp_labels, config.regime_filter,
getattr(getattr(self.engine.repo, "store", None), "data_dir", None),
)
try:
_rm = self._build_regime_mask(
market_data.timestamp_labels, config.regime_filter,
getattr(getattr(self.engine.repo, "store", None), "data_dir", None),
required_start=config.start,
required_end=config.end,
)
except ValueError as e:
return _err(str(e))
if _rm is not None:
entry_time_mask = entry_time_mask & _rm
exit_time_mask = self._matrix_date_range_mask(
@@ -1238,10 +1246,15 @@ class StrategyBacktestService:
config.start,
config.end,
)
_rm = self._build_regime_mask(
market_data.timestamp_labels, config.regime_filter,
getattr(getattr(self.engine.repo, "store", None), "data_dir", None),
)
try:
_rm = self._build_regime_mask(
market_data.timestamp_labels, config.regime_filter,
getattr(getattr(self.engine.repo, "store", None), "data_dir", None),
required_start=config.start,
required_end=config.end,
)
except ValueError as e:
return _err(str(e))
if _rm is not None:
entry_time_mask = entry_time_mask & _rm
exit_time_mask = self._matrix_date_range_mask(
@@ -1348,6 +1361,27 @@ class StrategyBacktestService:
formal_candidate_mask = candidate_mask & formal_range
entry_mask = self._build_entry_mask_from_candidate(panel, candidate_mask, s, entry_signals)
entry_mask = entry_mask & formal_range
if config.regime_filter:
date_values = panel.get_column("date").unique().sort().to_list()
date_labels = tuple(str(value)[:10] for value in date_values)
try:
regime_time_mask = self._build_regime_mask(
date_labels,
config.regime_filter,
getattr(getattr(self.engine.repo, "store", None), "data_dir", None),
required_start=config.start,
required_end=config.end,
)
except ValueError as e:
return _err(str(e))
if regime_time_mask is not None:
allowed_dates = [
value for value, allowed in zip(date_values, regime_time_mask, strict=True)
if allowed
]
regime_row_mask = panel.get_column("date").is_in(allowed_dates).fill_null(False)
formal_candidate_mask = formal_candidate_mask & regime_row_mask
entry_mask = entry_mask & regime_row_mask
raw_exit_mask = self._build_signal_mask(panel, exit_signals, "_exit")
exit_range = self._date_range_mask(panel, config.start, load_end) if config.mode == "full" else formal_range
exit_mask = raw_exit_mask & exit_range
@@ -1649,25 +1683,31 @@ class StrategyBacktestService:
timestamp_labels: tuple[str, ...],
regime_filter: dict | None,
data_dir: Path | None,
*,
required_start: date | None = None,
required_end: date | None = None,
) -> np.ndarray | None:
"""构造逐日 regime mask。强制 T-1 防未来函数: regime[T-1] 决定 entry[T]。
timestamp_labels[i] 的入场资格 = 它的"前一交易日"的 regime 是否满足条件。
"前一交易日"用 timestamp_labels 自身的顺序确定(回测时间轴上的前一天)。
边界: 首日无前一日环境 → 默认允许(不阻断)。
regime_filter 为 None 或无 regime 数据时返回 None(不过滤)。
regime_filter 为 None 时返回 None(不过滤)。启用过滤后缺少正式区间所需的
环境数据则 fail-closed, 避免界面显示已过滤但实际静默放行。
"""
if not regime_filter or data_dir is None:
if not regime_filter:
return None
allowed_states = set(regime_filter.get("states") or [])
min_score = regime_filter.get("min_score")
if not allowed_states and min_score is None:
return None
if data_dir is None:
raise ValueError("市场环境过滤不可用: 未找到环境数据目录")
from app.services import regime_builder
regime_df = regime_builder.load_regime_history(data_dir)
if regime_df.is_empty():
return None
raise ValueError("市场环境数据为空, 请先在数据页完成市场环境计算后再回测")
# 构建 date(ISO) → (state, score) 映射
regime_map: dict[str, tuple[str, int]] = {}
@@ -1680,11 +1720,21 @@ class StrategyBacktestService:
# 对每个 label, 找它的前一交易日的 regime(timestamp_labels 顺序里的前一天)
n = len(timestamp_labels)
mask = np.ones(n, dtype=bool) # 默认允许
required_start_text = str(required_start) if required_start is not None else None
required_end_text = str(required_end) if required_end is not None else None
missing_dates: list[str] = []
for i in range(1, n):
current_label = timestamp_labels[i][:10]
prev_label = timestamp_labels[i - 1][:10]
entry = regime_map.get(prev_label)
if entry is None:
continue # 无前一日环境数据 → 允许(不阻断)
required = (
(required_start_text is None or current_label >= required_start_text)
and (required_end_text is None or current_label <= required_end_text)
)
if required:
missing_dates.append(prev_label)
continue
state, score = entry
ok = True
if allowed_states and state not in allowed_states:
@@ -1692,6 +1742,13 @@ class StrategyBacktestService:
if min_score is not None and score < min_score:
ok = False
mask[i] = ok
if missing_dates:
first_missing = missing_dates[0]
suffix = f"{len(missing_dates)}" if len(missing_dates) > 1 else ""
raise ValueError(
f"市场环境数据覆盖不完整: 缺少前一交易日环境 {first_missing}{suffix}, "
"请先补算对应区间"
)
return mask
def _build_candidate_filter_mask(
@@ -1975,6 +2032,7 @@ class StrategyBacktestService:
"mode": c.mode,
"holding_days": c.holding_days,
"minute_fill": c.minute_fill,
"regime_filter": c.regime_filter,
}
@staticmethod
@@ -9,6 +9,7 @@ import polars as pl
from app.backtest.engine import BacktestEngine, SimResult
from app.backtest.matrix import build_market_data_matrix, make_signal_matrix, rolling_mean
from app.backtest.strategy import StrategyBacktestConfig, StrategyBacktestService
from app.services import regime_builder
from app.strategy.engine import StrategyDef
@@ -43,14 +44,17 @@ class _StrategyEngineStub:
class _RepoStub:
def __init__(self, data_dir=None) -> None:
self.store = SimpleNamespace(data_dir=data_dir)
def get_index_daily(self, *args, **kwargs) -> pl.DataFrame:
return pl.DataFrame()
class _EngineStub:
def __init__(self, panel: pl.DataFrame) -> None:
def __init__(self, panel: pl.DataFrame, data_dir=None) -> None:
self.panel = panel
self.repo = _RepoStub()
self.repo = _RepoStub(data_dir)
self.load_args = None
self.load_count = 0
self.sim_panel: pl.DataFrame | None = None
@@ -162,6 +166,55 @@ def test_basic_filter_only_limits_entries_not_panel_rows():
}
def test_non_matrix_strategy_applies_regime_filter_and_reports_config(tmp_path):
start = date(2024, 1, 1)
panel = pl.DataFrame([
{
"symbol": "A",
"name": "A",
"date": start + timedelta(days=offset),
"open": 10.0,
"high": 10.0,
"low": 10.0,
"close": 10.0,
"volume": 1000.0,
"amount": 1000.0,
"signal_limit_up": False,
"signal_limit_down": False,
}
for offset in range(3)
]).sort(["symbol", "date"])
regime_builder.upsert_regime_history(tmp_path, pl.DataFrame({
"date": [start, start + timedelta(days=1)],
"state": ["weak", "strong"],
"score": [10, 85],
}))
engine = _EngineStub(panel, data_dir=tmp_path)
service = StrategyBacktestService(engine=engine, strategy_engine=_StrategyEngineStub(_strategy()))
regime_filter = {"states": ["strong"]}
result = service.run(StrategyBacktestConfig(
strategy_id="test",
symbols=None,
start=start,
end=start + timedelta(days=2),
matching="close_t",
mode="position",
regime_filter=regime_filter,
))
assert result.error is None
assert engine.sim_matrix is not None
assert engine.sim_matrix.entry[:, 0].tolist() == [1, 0, 1]
assert result.config["regime_filter"] == regime_filter
assert result.stats["selection"] == {
"strategy_matches": 2,
"entry_candidates": 2,
"entry_trigger_filtered": 0,
"entry_trigger_enabled": False,
}
def test_selection_stats_explain_entry_trigger_filtering():
start = date(2024, 1, 1)
panel = pl.DataFrame([
@@ -423,6 +476,21 @@ def test_matrix_optimizer_preparation_loads_and_builds_base_data_once():
assert all(result.stats["shared_market_data_bytes"] == prepared.market_data.nbytes for result in results)
def test_matrix_prepare_signature_includes_regime_filter():
base = dict(
strategy_id="native",
symbols=None,
start=date(2024, 1, 1),
end=date(2024, 1, 2),
)
without_filter = StrategyBacktestConfig(**base)
with_filter = StrategyBacktestConfig(**base, regime_filter={"states": ["strong"]})
assert StrategyBacktestService._matrix_prepare_signature(without_filter) != (
StrategyBacktestService._matrix_prepare_signature(with_filter)
)
def test_matrix_cache_preserves_trades_daily_equity_and_core_stats():
start = date(2024, 1, 1)
panel = pl.DataFrame([
+25 -7
View File
@@ -13,6 +13,7 @@ import time
from datetime import date
import polars as pl
import pytest
from app.services import regime_builder
@@ -312,14 +313,32 @@ def test_build_regime_mask_none_when_no_filter():
assert StrategyBacktestService._build_regime_mask(("2026-01-01",), None, None) is None
def test_build_regime_mask_none_when_no_data(tmp_path):
"""无 regime 历史数据 → 返回 None(不阻断回测)"""
def test_build_regime_mask_fails_when_no_data(tmp_path):
"""启用过滤但无 regime 历史数据时必须阻止回测。"""
from app.backtest.strategy import StrategyBacktestService
mask = StrategyBacktestService._build_regime_mask(
("2026-01-01", "2026-01-02"), {"states": ["strong"]}, tmp_path,
)
assert mask is None
with pytest.raises(ValueError, match="市场环境数据为空"):
StrategyBacktestService._build_regime_mask(
("2026-01-01", "2026-01-02"), {"states": ["strong"]}, tmp_path,
)
def test_build_regime_mask_fails_when_required_t1_date_is_missing(tmp_path):
"""正式区间内任一入场日缺少 T-1 环境时必须阻止回测。"""
from app.backtest.strategy import StrategyBacktestService
regime_builder.upsert_regime_history(tmp_path, pl.DataFrame({
"date": [date(2026, 1, 1)],
"state": ["strong"],
"score": [85],
}))
with pytest.raises(ValueError, match="缺少前一交易日环境"):
StrategyBacktestService._build_regime_mask(
("2026-01-01", "2026-01-02", "2026-01-03"),
{"states": ["strong"]},
tmp_path,
)
def test_build_regime_mask_first_day_allowed(tmp_path):
@@ -336,4 +355,3 @@ def test_build_regime_mask_first_day_allowed(tmp_path):
)
# 1/1 首日 → True; 1/2 由 1/1(weak) → False
assert mask.tolist() == [True, False]
+2
View File
@@ -113,6 +113,8 @@ export const storage = {
mode: 'position' | 'full'
holdingDays: string
minuteFill?: boolean
regimeStates?: string[]
regimeMinScore?: number | ''
params?: Record<string, any>
overrides?: Record<string, any>
strategyConfigSignature?: string
@@ -947,8 +947,8 @@ export function StrategyBacktest() {
const [holdingDays, setHoldingDays] = useState(saved?.holdingDays ?? '5')
const [highGranularity, setHighGranularity] = useState(saved?.minuteFill ?? false)
// 市场环境过滤(空=不过滤)
const [regimeStates, setRegimeStates] = useState<string[]>([])
const [regimeMinScore, setRegimeMinScore] = useState<number | ''>('')
const [regimeStates, setRegimeStates] = useState<string[]>(saved?.regimeStates ?? [])
const [regimeMinScore, setRegimeMinScore] = useState<number | ''>(saved?.regimeMinScore ?? '')
const [settingsOpen, setSettingsOpen] = useState(false)
// 分钟K成交价细化: 不改变信号日或成交日, 需 Pro+ 分钟K能力
const { data: caps } = useCapabilities()
@@ -1097,6 +1097,8 @@ export function StrategyBacktest() {
mode: simMode,
holdingDays,
minuteFill: highGranularity,
regimeStates,
regimeMinScore,
params: strategyParams,
overrides,
strategyConfigSignature: strategyDetail.data
@@ -1373,6 +1375,18 @@ export function StrategyBacktest() {
const resultStartDate = result?.config?.start ?? result?.equity_curve?.[0]?.date ?? start
const resultEndDate = result?.config?.end ?? result?.equity_curve?.[result.equity_curve.length - 1]?.date ?? end
const resultTradeDays = result?.equity_curve?.length ?? 0
const resultRegimeFilter = result?.config?.regime_filter as {
states?: string[]
min_score?: number
} | null | undefined
const resultRegimeSummary = resultRegimeFilter
? [
resultRegimeFilter.states?.length
? resultRegimeFilter.states.map(state => REGIME_STATE_LABELS[state as keyof typeof REGIME_STATE_LABELS] ?? state).join('/')
: null,
resultRegimeFilter.min_score != null ? `最低 ${resultRegimeFilter.min_score}` : null,
].filter(Boolean).join(' · ')
: ''
const selectionStats = result?.stats?.selection as Record<string, number | boolean> | undefined
const selectionStages = selectionStats
? [
@@ -1925,6 +1939,12 @@ export function StrategyBacktest() {
<div className="flex items-center gap-3">
<span className="text-sm font-medium text-foreground">{result.strategy_info?.name ?? '策略'}</span>
<span className="text-[10px] px-1 py-px rounded border border-accent/30 bg-accent/10 text-accent"></span>
{resultRegimeSummary && (
<span className="inline-flex items-center gap-1 rounded border border-accent/25 bg-accent/10 px-1.5 py-px text-[10px] text-accent">
<Gauge className="h-2.5 w-2.5" />
{resultRegimeSummary}
</span>
)}
<span className="text-[10px] text-secondary"> {result.config?.holding_days ?? 5} </span>
<span className="ml-auto text-[11px] text-muted font-mono">
{String(result.config?.start).slice(0,10)} ~ {String(result.config?.end).slice(0,10)}
@@ -1995,6 +2015,12 @@ export function StrategyBacktest() {
{SRC_MAP[result.strategy_info.source] ?? ''}
</span>
)}
{resultRegimeSummary && (
<span className="inline-flex items-center gap-1 rounded border border-accent/25 bg-accent/10 px-1.5 py-px text-[9px] text-accent">
<Gauge className="h-2.5 w-2.5" />
{resultRegimeSummary}
</span>
)}
</div>
{/* 叠加策略: 子策略构成归因 */}
{result.strategy_info.composite_children && result.strategy_info.composite_children.length > 0 && (