看股功能更新

This commit is contained in:
2026-08-16 00:05:26 +08:00
parent 9cce670b74
commit fc86fe0674
28 changed files with 3823 additions and 96 deletions

View File

@@ -0,0 +1,265 @@
"""事件回测引擎:入场条件命中 -> 次日买入 -> 持有 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