Files
stock/backend/app/backtest/events.py
2026-08-16 00:05:26 +08:00

266 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""事件回测引擎:入场条件命中 -> 次日买入 -> 持有 N 日 -> 全市场汇总统计。
数据口径:
- 行情底座是 candlesTDX 全量导入,不复权),全历史可用;
- 指标计算用不复权价(与选股/看盘口径一致J<10、RSI<30 等阈值均为归一化或惯例值);
- 收益率用 adj_factor 校正ret = 出场价×f出 / 入场价×f入 - 1消除除权除息失真
因子缺失的股退化为不复权收益(新股/缺因子,样本中占少数)。
信号语义:与选股引擎一致——每条条件在信号日 d 为终点、lookback 窗口内
match=all(连续满足)/any(曾经满足),多条件之间取 AND。
"""
from __future__ import annotations
from datetime import date, datetime, timedelta
import numpy as np
import pandas as pd
from sqlalchemy import and_, func, not_, or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from ..models import AdjFactor, Candle, StockBasic
from ..schemas import EventBacktestSpec
from ..screener.engine import (
FAMILIES,
_family_of,
_op_mask,
_params_for,
_resolve_params,
_series_for,
)
# 指标配热缓冲 bar 数MACD 等 EMA 类指标需要较长窗口才收敛)
BUFFER_BARS = 80
# 每批查询的股票数(全市场分块拉取,避免单条 SQL 过大)
BATCH_SIZE = 800
# 单次回测允许的最大样本数(超过则仅按日期取最近的,防内存失控)
MAX_TRADES = 200_000
class EventEngineError(RuntimeError):
"""事件回测可预期的业务错误(信息透传前端)。"""
def _signal_mask(g: pd.DataFrame, spec: EventBacktestSpec, cache: dict) -> pd.Series:
"""单股全序列信号掩码各条件rolling lookbackAND。"""
total = pd.Series(True, index=g.index)
for cond in spec.entry.indicator:
fam = _family_of(cond.indicator)
if len(g) < FAMILIES[fam].min_bars:
return pd.Series(False, index=g.index)
s = _series_for(g, cond.indicator, _params_for(cond.indicator, cond.params), cache)
if s is None:
return pd.Series(False, index=g.index)
if cond.value_indicator:
target = _series_for(g, cond.value_indicator,
_resolve_params(cond, cond.value_indicator), cache)
if target is None:
return pd.Series(False, index=g.index)
else:
target = pd.Series(cond.value, index=s.index)
m = _op_mask(s, target, cond).astype(int)
n = max(1, cond.lookback)
if n > 1:
rolled = m.rolling(n, min_periods=n).sum()
m = (rolled == n) if cond.match == "all" else (rolled > 0)
else:
m = m.astype(bool)
total = total & m.fillna(False).astype(bool)
return total
def _entry_exit_indices(sig_idx: int, spec: EventBacktestSpec, n: int) -> tuple[int, int] | None:
"""信号日索引 -> (入场索引, 出场索引)。前视/越界返回 None。"""
entry_i = sig_idx + 1 # 信号收盘后才动手:一律次日
exit_i = entry_i + spec.holding_days
if exit_i >= n:
return None
return entry_i, exit_i
def _price_at(row: pd.Series, timing: str) -> float:
return float(row["open"] if timing == "open" else row["close"])
def _stats_block(trades: list[dict]) -> dict:
"""样本集合 -> 汇总统计(空样本给零值)。"""
if not trades:
return {
"samples": 0, "stocks": 0,
"mean_pct": 0.0, "median_pct": 0.0, "win_rate": 0.0, "std_pct": 0.0,
"p10_pct": 0.0, "p25_pct": 0.0, "p75_pct": 0.0, "p90_pct": 0.0,
"max_pct": 0.0, "min_pct": 0.0, "by_year": [],
}
rets = np.array([t["ret_pct"] for t in trades], dtype=float)
by_year: list[dict] = []
df = pd.DataFrame(trades)
for year, grp in df.groupby(df["entry_date"].dt.year):
r = grp["ret_pct"].to_numpy()
by_year.append({
"year": int(year), "samples": int(len(r)),
"mean_pct": round(float(r.mean()), 3),
"median_pct": round(float(np.median(r)), 3),
"win_rate": round(float((r > 0).mean() * 100), 2),
})
by_year.sort(key=lambda x: x["year"])
return {
"samples": int(len(rets)),
"stocks": int(df["ts_code"].nunique()),
"mean_pct": round(float(rets.mean()), 3),
"median_pct": round(float(np.median(rets)), 3),
"win_rate": round(float((rets > 0).mean() * 100), 2),
"std_pct": round(float(rets.std(ddof=1)) if len(rets) > 1 else 0.0, 3),
"p10_pct": round(float(np.percentile(rets, 10)), 3),
"p25_pct": round(float(np.percentile(rets, 25)), 3),
"p75_pct": round(float(np.percentile(rets, 75)), 3),
"p90_pct": round(float(np.percentile(rets, 90)), 3),
"max_pct": round(float(rets.max()), 3),
"min_pct": round(float(rets.min()), 3),
"by_year": by_year,
}
async def run_event_backtest(
session: AsyncSession,
spec: EventBacktestSpec,
ts_code: str | None = None,
start: date | None = None,
end: date | None = None,
) -> dict:
"""主入口:返回 {spec, universe, start, end, stats, trades(sample), total}。"""
entry = spec.entry
if not entry.indicator:
raise EventEngineError("入场条件必须包含技术指标条件(如 J<10、RSI<30")
# 时间窗默认最近一年end 以 candles 最大日期为准
end_dt = end
if end_dt is None:
end_dt = (await session.scalar(select(func.max(Candle.ts)))) or date.today()
if isinstance(end_dt, datetime):
end_dt = end_dt.date()
start_dt = start or (end_dt - timedelta(days=365))
if start_dt >= end_dt:
raise EventEngineError("回测起始日期必须早于结束日期")
needed = _max_needed_bars_safe(entry) + BUFFER_BARS
buffer_start = start_dt - timedelta(days=int(needed * 1.7)) # 交易日->日历日近似
# 股票池ts_code+symbol 映射candles 按 symbol 存)
name_map: dict[str, str] = {}
if ts_code:
rows = (await session.execute(
select(StockBasic.ts_code, StockBasic.symbol, StockBasic.name)
.where(StockBasic.ts_code == ts_code)
)).all()
if not rows:
raise EventEngineError(f"未知股票代码: {ts_code}")
universe = [(r[0], r[1]) for r in rows]
name_map = {r[0]: r[2] for r in rows}
else:
stmt = select(StockBasic.ts_code, StockBasic.symbol, StockBasic.name).where(
StockBasic.list_status == "L"
)
if entry.exclude_st:
stmt = stmt.where(not_(or_(StockBasic.name.like("%ST%"), StockBasic.name.like("%退%"))))
if entry.exclude_bj:
stmt = stmt.where(not_(StockBasic.ts_code.like("%.BJ")))
rows = (await session.execute(stmt)).all()
universe = [(r[0], r[1]) for r in rows]
name_map = {r[0]: r[2] for r in rows}
start_ts = datetime(start_dt.year, start_dt.month, start_dt.day)
end_ts = datetime(end_dt.year, end_dt.month, end_dt.day, 23, 59, 59)
buffer_ts = datetime(buffer_start.year, buffer_start.month, buffer_start.day)
trades: list[dict] = []
for i in range(0, len(universe), BATCH_SIZE):
batch = universe[i : i + BATCH_SIZE]
symbols = [sym for _, sym in batch]
code_by_symbol = {sym: code for code, sym in batch}
candle_rows = (await session.execute(
select(Candle.symbol, Candle.ts, Candle.open, Candle.high,
Candle.low, Candle.close)
.where(and_(Candle.timeframe == "1d",
Candle.symbol.in_(symbols),
Candle.ts >= buffer_ts, Candle.ts <= end_ts))
.order_by(Candle.symbol, Candle.ts)
)).all()
if not candle_rows:
continue
codes = {code_by_symbol[s] for s in symbols}
adj_rows = (await session.execute(
select(AdjFactor.ts_code, AdjFactor.trade_date, AdjFactor.adj_factor)
.where(and_(AdjFactor.ts_code.in_(codes),
AdjFactor.trade_date >= buffer_ts, AdjFactor.trade_date <= end_ts))
)).all()
f_map = {(r[0], r[1].date()): float(r[2]) for r in adj_rows if r[2]}
bars = pd.DataFrame(
candle_rows, columns=["symbol", "ts", "open", "high", "low", "close"]
)
for symbol, g in bars.groupby("symbol", sort=False):
if len(g) < 30:
continue
g = g.reset_index(drop=True)
ts_code_l = code_by_symbol[symbol]
cache: dict = {"_families": set()}
mask = _signal_mask(g, spec, cache)
if not mask.any():
continue
for sig_i in np.flatnonzero(mask.to_numpy()):
ts_sig = g.at[sig_i, "ts"]
# 信号必须落在回测窗口内buffer 区只用于指标配热)
if ts_sig < start_ts:
continue
ie = _entry_exit_indices(int(sig_i), spec, len(g))
if ie is None:
continue
entry_i, exit_i = ie
e_row, x_row = g.iloc[entry_i], g.iloc[exit_i]
e_price = _price_at(e_row, "open" if spec.entry_timing == "next_open" else "close")
x_price = _price_at(x_row, "open" if spec.exit_timing == "open" else "close")
if not e_price or not x_price:
continue
f_in = f_map.get((ts_code_l, e_row["ts"].date()), 1.0)
f_out = f_map.get((ts_code_l, x_row["ts"].date()), 1.0)
ret_pct = (x_price * f_out) / (e_price * f_in) * 100 - 100
trades.append({
"ts_code": ts_code_l,
"name": name_map.get(ts_code_l),
"entry_date": e_row["ts"], "entry_price": round(e_price, 3),
"exit_date": x_row["ts"], "exit_price": round(x_price, 3),
"ret_pct": round(float(ret_pct), 3),
})
if len(trades) >= MAX_TRADES:
break
if len(trades) >= MAX_TRADES:
break
if len(trades) >= MAX_TRADES:
break
stats = _stats_block(trades)
# 明细样本:最好 100 + 最差 100其余统计已覆盖
trades_sorted = sorted(trades, key=lambda t: t["ret_pct"], reverse=True)
sample = trades_sorted[:100] + (trades_sorted[-100:] if len(trades_sorted) > 100 else [])
return {
"spec": spec,
"universe": ts_code or "all",
"start": start_ts,
"end": end_ts,
"stats": stats,
"trades": sample,
"total": stats["samples"],
}
# ---------- 小工具 ----------
def _max_needed_bars_safe(conds) -> int:
"""指标配热所需最大 bar 数(同 screener.engine._max_needed_bars"""
need = 1
for c in conds.indicator:
need = max(need, FAMILIES[_family_of(c.indicator)].min_bars + c.lookback)
if c.value_indicator:
need = max(need, FAMILIES[_family_of(c.value_indicator)].min_bars + c.lookback)
return need