看股功能更新
This commit is contained in:
265
backend/app/backtest/events.py
Normal file
265
backend/app/backtest/events.py
Normal file
@@ -0,0 +1,265 @@
|
||||
"""事件回测引擎:入场条件命中 -> 次日买入 -> 持有 N 日 -> 全市场汇总统计。
|
||||
|
||||
数据口径:
|
||||
- 行情底座是 candles(TDX 全量导入,不复权),全历史可用;
|
||||
- 指标计算用不复权价(与选股/看盘口径一致: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 lookback)AND。"""
|
||||
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
|
||||
Reference in New Issue
Block a user