fix(regime): aggregate fallback batches incrementally

This commit is contained in:
0112020179
2026-09-04 12:19:35 +08:00
parent 8c28132361
commit 885e693118
2 changed files with 140 additions and 26 deletions
+58 -26
View File
@@ -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)
+82
View File
@@ -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) ─────────────────────────