diff --git a/backend/app/services/regime_builder.py b/backend/app/services/regime_builder.py index 17605ed..d43af23 100644 --- a/backend/app/services/regime_builder.py +++ b/backend/app/services/regime_builder.py @@ -356,17 +356,30 @@ def _compute_batch(repo, enriched_dir, instruments, historical_shares, return df.filter((pl.col("date") >= batch_start) & (pl.col("date") <= batch_end)) -def _scan_enriched_fallback(repo, start: date, end: date) -> pl.DataFrame | None: - """缓存不覆盖时的慢路径: scan enriched parquet + 重算所需指标列。 +def _filter_excluded_symbols(df: pl.DataFrame, excluded_symbols: list[str]) -> pl.DataFrame: + if excluded_symbols and "symbol" in df.columns: + return df.filter(~pl.col("symbol").str.to_uppercase().is_in(excluded_symbols)) + return df - 仅在 regime 首次全量回填或缓存未预热时触发。返回含信号列的多日 DataFrame。 + +def _scan_enriched_fallback( + repo, + start: date, + end: date, + *, + index_pct_map: dict | None = None, + excluded_symbols: list[str] | None = None, +) -> pl.DataFrame | None: + """缓存不覆盖时的慢路径: 分批扫描并直接返回日级环境聚合。 + + 仅在 regime 首次全量回填或缓存未预热时触发。 内存控制(关键, 两层优化): 1. needed 白名单: regime 只需 change_pct/ma20/涨跌停信号等少数列, 不用 compute_all 算 72 列全套指标(那会让全量峰值达 6.8GB)。 - 2. 分批: 范围超过 batch_days 个交易日时按批切片, 每批带 warmup 前缀算完后 concat。 + 2. 分批: 每批带 warmup 前缀算完后立即聚合为日级行, 不保留跨批个股明细。 batch_days / warmup_days 由用户偏好控制(数据页「市场环境」卡片设置), - 实测默认值(60/40)全量(515万行)峰值约 1.9GB, 4GB 内存机器可稳跑。 + 峰值随单批大小受控, 不再随完整历史长度线性增长。 必须传入 instruments(涨跌停价表), 否则 compute_limit_signals 会跳过涨跌停信号。 """ try: @@ -381,6 +394,7 @@ def _scan_enriched_fallback(repo, start: date, end: date) -> pl.DataFrame | None enriched_dir = repo.store.data_dir / "kline_daily_enriched" if not enriched_dir.exists(): return None + excluded_symbols = excluded_symbols or [] instruments = repo.get_instruments() historical_shares = repo.get_historical_shares() @@ -394,23 +408,31 @@ def _scan_enriched_fallback(repo, start: date, end: date) -> pl.DataFrame | None if len(target_dates) <= batch_days: df = _compute_batch(repo, enriched_dir, instruments, historical_shares, target_dates[0], target_dates[-1], warmup_days) - return df if not df.is_empty() else None + if df.is_empty(): + return None + df = _filter_excluded_symbols(df, excluded_symbols) + result = _aggregate_daily(df, index_pct_map) + return result if not result.is_empty() else None - # 大范围: 按交易日分批, 逐批算 + concat + # 大范围: 每批个股明细立即压缩为日级行, 只保留小型聚合结果。 batches = [ (target_dates[i], target_dates[min(i + batch_days - 1, len(target_dates) - 1)]) for i in range(0, len(target_dates), batch_days) ] - logger.info("regime fallback: %d 天分 %d 批 (每批≤%d天 + %d天warmup)", + logger.info("regime fallback: %d 天分 %d 批逐批聚合 (每批≤%d天 + %d天warmup)", len(target_dates), len(batches), batch_days, warmup_days) - parts: list[pl.DataFrame] = [] + daily_parts: list[pl.DataFrame] = [] for bs, be in batches: df = _compute_batch(repo, enriched_dir, instruments, historical_shares, bs, be, warmup_days) - if not df.is_empty(): - parts.append(df) - if not parts: + if df.is_empty(): + continue + df = _filter_excluded_symbols(df, excluded_symbols) + daily = _aggregate_daily(df, index_pct_map) + if not daily.is_empty(): + daily_parts.append(daily) + if not daily_parts: return None - return pl.concat(parts, how="vertical_relaxed") + return pl.concat(daily_parts, how="vertical_relaxed") except Exception as e: # noqa: BLE001 logger.warning("regime scan_enriched_fallback failed: %s", e) return None @@ -440,15 +462,6 @@ def run_regime_batch(repo, start: date, end: date) -> pl.DataFrame: # 指数涨幅(主力指数) index_pct_map = _load_index_pct(repo, start, end) - # enriched 多日数据(优先缓存) - df = repo.get_enriched_range(start, end) - if df is None or df.is_empty(): - logger.info("regime batch: enriched cache miss [%s~%s], fallback to scan", start, end) - df = _scan_enriched_fallback(repo, start, end) - if df is None or df.is_empty(): - logger.info("regime batch: no enriched data for [%s~%s]", start, end) - return pl.DataFrame() - # 口径: 默认剔除风险警示(ST)股(与主线统计同一开关) — 主板 ST 在 2026-07 前 # 享 5% 涨跌幅且是跨行业状态桶, 混入会系统性抬高涨停宽度/高度(弱市炒 ST 尤甚)。 # 涨跌家数/MA20 占比等宽度指标几乎不受影响。切换口径需全量重算 regime。 @@ -457,14 +470,33 @@ def run_regime_batch(repo, start: date, end: date) -> pl.DataFrame: exclude_st = _prefs_st.get_sentiment_exclude_st() except Exception: exclude_st = True + excluded_symbols: list[str] = [] if exclude_st: from app.services.market_mainline import load_risk_warning_symbols st_syms = load_risk_warning_symbols(repo.store.data_dir) - if st_syms and "symbol" in df.columns: - df = df.filter( - ~pl.col("symbol").str.to_uppercase().is_in(sorted(st_syms)) - ) + excluded_symbols = sorted(st_syms) + + # enriched 多日数据(优先缓存)。慢路径在每批内部完成过滤和日级聚合, + # 避免把所有批次的个股明细同时保留到最终 group_by。 + df = repo.get_enriched_range(start, end) + if df is None or df.is_empty(): + logger.info("regime batch: enriched cache miss [%s~%s], fallback to scan", start, end) + result = _scan_enriched_fallback( + repo, + start, + end, + index_pct_map=index_pct_map, + excluded_symbols=excluded_symbols, + ) + if result is None or result.is_empty(): + logger.info("regime batch: no enriched data for [%s~%s]", start, end) + return pl.DataFrame() + return result + + df = _filter_excluded_symbols(df, excluded_symbols) + if df.is_empty(): + return pl.DataFrame() return _aggregate_daily(df, index_pct_map) diff --git a/backend/tests/test_regime_builder.py b/backend/tests/test_regime_builder.py index 45641f2..bc8c50f 100644 --- a/backend/tests/test_regime_builder.py +++ b/backend/tests/test_regime_builder.py @@ -14,6 +14,7 @@ from datetime import date import polars as pl import pytest +from polars.testing import assert_frame_equal from app.services import regime_builder @@ -183,6 +184,87 @@ def test_run_regime_batch_excludes_st(tmp_path, monkeypatch): assert r1_all["up_count"] == 2 # A,C 上涨 +def test_run_regime_batch_aggregates_fallback_per_batch(tmp_path, monkeypatch): + """回归 #241: fallback 不得保留所有批次的个股明细后再统一聚合。""" + from app.services import market_mainline, preferences + + target_dates = [date(2026, 1, day) for day in range(2, 6)] + enriched_dir = tmp_path / "kline_daily_enriched" + for trade_date in target_dates: + partition = enriched_dir / f"date={trade_date.isoformat()}" + partition.mkdir(parents=True) + (partition / "part.parquet").write_bytes(b"") + + class _Store: + data_dir = tmp_path + + class _FakeRepo: + store = _Store + + def get_enriched_range(self, start, end): + return None + + def get_instruments(self): + return pl.DataFrame() + + def get_historical_shares(self): + return pl.DataFrame() + + def _batch_frame(batch_start, batch_end): + dates = [d for d in target_dates if batch_start <= d <= batch_end] + rows = [] + for trade_date in dates: + for symbol, change_pct in (("A", 0.05), ("B", -0.03)): + rows.append({ + "date": trade_date, + "symbol": symbol, + "close": 10.0, + "change_pct": change_pct, + "amount": 1e8, + "ma20": 9.0, + "signal_limit_up": symbol == "A", + "signal_limit_down": False, + "signal_broken_limit_up": False, + "consecutive_limit_ups": 1 if symbol == "A" else 0, + "_prev_consec": 0, + }) + return pl.DataFrame(rows) + + monkeypatch.setattr(preferences, "get_regime_batch_days", lambda: 2) + monkeypatch.setattr(preferences, "get_regime_warmup_days", lambda: 40) + monkeypatch.setattr(preferences, "get_sentiment_exclude_st", lambda: True) + monkeypatch.setattr(market_mainline, "load_risk_warning_symbols", lambda *a, **k: {"A"}) + monkeypatch.setattr(regime_builder, "_load_index_pct", lambda *a, **k: {}) + monkeypatch.setattr( + regime_builder, + "_compute_batch", + lambda repo, enriched, instruments, shares, start, end, warmup: _batch_frame(start, end), + ) + + aggregate_input_heights = [] + original_aggregate = regime_builder._aggregate_daily + + def _recording_aggregate(df, index_pct_map=None): + aggregate_input_heights.append(df.height) + return original_aggregate(df, index_pct_map) + + monkeypatch.setattr(regime_builder, "_aggregate_daily", _recording_aggregate) + + result = regime_builder.run_regime_batch(_FakeRepo(), target_dates[0], target_dates[-1]) + expected_detail = pl.concat([ + _batch_frame(target_dates[0], target_dates[1]), + _batch_frame(target_dates[2], target_dates[3]), + ]).filter(pl.col("symbol") == "B") + expected = original_aggregate(expected_detail, {}) + + assert result.height == len(target_dates) + assert aggregate_input_heights == [2, 2] + assert result["up_count"].to_list() == [0, 0, 0, 0] + assert result["down_count"].to_list() == [1, 1, 1, 1] + assert result["limit_up"].sum() == 0 + assert_frame_equal(result.sort("date"), expected.sort("date")) + + # ───────────────────────── 持久化(upsert) ─────────────────────────