mirror of
https://ghfast.top/https://github.com/aeroxw/tick-stock-panel.git
synced 2026-09-12 14:24:15 +08:00
fix(backtest): 修复市场环境过滤未生效
This commit is contained in:
@@ -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([
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
@@ -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 && (
|
||||
|
||||
Reference in New Issue
Block a user