Files
stock-advisor/engine/strategies.py
T

426 lines
16 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.
# -*- coding: utf-8 -*-
"""
量化策略引擎:主流策略信号生成 + 回测框架
- 6 个主流策略:双均线 / MACD / RSI / 布林带 / 动量 / N日突破
- 回测规则:收盘产生信号,次日开盘成交(全仓多头,避免未来函数)
- 指标:总收益 / 年化 / 最大回撤 / 夏普 / 胜率 / 盈亏比 / 交易次数
- 结果持久化到 strategy_backtests 表,支持全市场批量回测
扩展新策略:在 STRATEGIES 注册 name/desc + 实现信号生成函数即可
"""
import json
import math
import statistics
from database import executemany, query_one
# ===================================================================== 指标序列
def ma_series(closes, n):
out = []
s = 0.0
for i, c in enumerate(closes):
s += c
if i >= n:
s -= closes[i - n]
out.append(s / n if i >= n - 1 else None)
return out
def ema_series(vals, n):
out = []
k = 2 / (n + 1)
e = vals[0] if vals else 0
for i, v in enumerate(vals):
e = v if i == 0 else v * k + e * (1 - k)
out.append(e)
return out
def macd_series(closes):
ema12 = ema_series(closes, 12)
ema26 = ema_series(closes, 26)
dif = [a - b for a, b in zip(ema12, ema26)]
dea = ema_series(dif, 9)
return dif, dea
def rsi_series(closes, n=14):
out = [None] * len(closes)
if len(closes) <= n:
return out
gains, losses = [], []
for i in range(1, len(closes)):
chg = closes[i] - closes[i - 1]
gains.append(max(chg, 0))
losses.append(max(-chg, 0))
avg_g = sum(gains[:n]) / n
avg_l = sum(losses[:n]) / n
for i in range(n, len(gains)):
avg_g = (avg_g * (n - 1) + gains[i]) / n
avg_l = (avg_l * (n - 1) + losses[i]) / n
rs = 100 if avg_l == 0 else avg_g / avg_l
out[i + 1] = 100 - 100 / (1 + rs)
out[n] = 100 if avg_l == 0 else 100 - 100 / (1 + avg_g / max(avg_l, 1e-9))
return out
def boll_series(closes, n=20, k=2.0):
mid, upper, lower = [], [], []
for i in range(len(closes)):
if i >= n - 1:
seg = closes[i - n + 1:i + 1]
m = sum(seg) / n
sd = statistics.pstdev(seg)
mid.append(m); upper.append(m + k * sd); lower.append(m - k * sd)
else:
mid.append(None); upper.append(None); lower.append(None)
return mid, upper, lower
# ===================================================================== 策略注册
STRATEGIES = {
"ma_cross": {
"name": "双均线金叉",
"icon": "📐",
"params": "MA5 / MA20",
"desc": "短均线MA5上穿长均线MA20买入(金叉),下穿卖出(死叉)。经典趋势跟踪策略。",
"tags": ["趋势"],
},
"macd": {
"name": "MACD 金叉",
"icon": "🟢",
"params": "12/26/9",
"desc": "DIF 上穿 DEA 买入,下穿卖出。捕捉中线趋势拐点,过滤震荡噪音。",
"tags": ["趋势", "动量"],
},
"rsi_rev": {
"name": "RSI 超买超卖",
"icon": "🔄",
"params": "RSI(14) 30/70",
"desc": "RSI 低于 30 超卖买入、高于 70 超买卖出。均值回归型反转策略。",
"tags": ["反转"],
},
"boll": {
"name": "布林带回归",
"icon": "📦",
"params": "20日 / 2σ",
"desc": "价格跌破下轨买入、突破上轨卖出,赌价格向中轨回归。震荡市表现佳。",
"tags": ["回归"],
},
"momentum": {
"name": "20日动量",
"icon": "🚀",
"params": "20日涨幅 / MA20",
"desc": "20日涨幅超阈值且站上MA20买入,跌破MA20卖出。顺势强者恒强。",
"tags": ["动量"],
},
"breakout": {
"name": "N日新高突破",
"icon": "🧗",
"params": "20日高低点",
"desc": "收盘突破20日新高买入,跌破20日新低卖出。海龟式突破策略。",
"tags": ["突破"],
},
}
def _signals(bars, key):
"""生成信号列表 [{date, action:'buy'/'sell', close, reason}]"""
closes = [b["close"] for b in bars]
n = len(bars)
sigs = []
if key == "ma_cross":
ma5, ma20 = ma_series(closes, 5), ma_series(closes, 20)
prev_state = None
for i in range(n):
if ma5[i] is None or ma20[i] is None:
continue
state = ma5[i] > ma20[i]
if prev_state is not None and state != prev_state:
if state:
sigs.append({"date": bars[i]["date"], "action": "buy", "close": closes[i],
"reason": f"MA5({ma5[i]:.2f})上穿MA20({ma20[i]:.2f})金叉"})
else:
sigs.append({"date": bars[i]["date"], "action": "sell", "close": closes[i],
"reason": f"MA5({ma5[i]:.2f})下穿MA20({ma20[i]:.2f})死叉"})
prev_state = state
elif key == "macd":
dif, dea = macd_series(closes)
prev = None
for i in range(n):
if dif[i] is None or dea[i] is None:
continue
state = dif[i] > dea[i]
if prev is not None and state != prev:
if state:
sigs.append({"date": bars[i]["date"], "action": "buy", "close": closes[i],
"reason": f"MACD金叉 DIF({dif[i]:.3f})上穿DEA({dea[i]:.3f})"})
else:
sigs.append({"date": bars[i]["date"], "action": "sell", "close": closes[i],
"reason": f"MACD死叉 DIF({dif[i]:.3f})下穿DEA({dea[i]:.3f})"})
prev = state
elif key == "rsi_rev":
rsi = rsi_series(closes)
for i in range(1, n):
if rsi[i] is None or rsi[i - 1] is None:
continue
if rsi[i - 1] >= 30 and rsi[i] < 30:
sigs.append({"date": bars[i]["date"], "action": "buy", "close": closes[i],
"reason": f"RSI({rsi[i]:.1f})下穿30超卖"})
elif rsi[i - 1] <= 70 and rsi[i] > 70:
sigs.append({"date": bars[i]["date"], "action": "sell", "close": closes[i],
"reason": f"RSI({rsi[i]:.1f})上穿70超买"})
elif key == "boll":
mid, upper, lower = boll_series(closes)
prev_state = None
for i in range(n):
if lower[i] is None:
continue
if closes[i] < lower[i]:
state = "buy"
elif closes[i] > upper[i]:
state = "sell"
else:
state = prev_state
if state != prev_state and state is not None:
if state == "buy":
sigs.append({"date": bars[i]["date"], "action": "buy", "close": closes[i],
"reason": f"收盘跌破下轨({lower[i]:.2f})"})
else:
sigs.append({"date": bars[i]["date"], "action": "sell", "close": closes[i],
"reason": f"收盘突破上轨({upper[i]:.2f})"})
prev_state = state
elif key == "momentum":
ma20 = ma_series(closes, 20)
prev_state = None
for i in range(n):
if i < 20 or ma20[i] is None:
continue
mom = closes[i] / closes[i - 20] - 1
state = "buy" if (mom > 0.03 and closes[i] > ma20[i]) else ("sell" if closes[i] < ma20[i] else prev_state)
if state != prev_state and state is not None:
if state == "buy":
sigs.append({"date": bars[i]["date"], "action": "buy", "close": closes[i],
"reason": f"20日动量{mom*100:+.1f}%且站上MA20"})
else:
sigs.append({"date": bars[i]["date"], "action": "sell", "close": closes[i],
"reason": "跌破MA20止盈/止损"})
prev_state = state
elif key == "breakout":
N = 20
for i in range(N, n):
window = closes[i - N:i]
if closes[i] > max(window) and closes[i] > closes[i - 1]:
sigs.append({"date": bars[i]["date"], "action": "buy", "close": closes[i],
"reason": f"突破{N}日新高({max(window):.2f})"})
elif closes[i] < min(window):
sigs.append({"date": bars[i]["date"], "action": "sell", "close": closes[i],
"reason": f"跌破{N}日新低({min(window):.2f})"})
return sigs
# ===================================================================== 回测
def backtest(bars, key):
"""全仓多头回测。收盘信号 → 次日开盘成交。返回 metrics/equity/trades"""
signals = _signals(bars, key)
n = len(bars)
opens = [b["open"] for b in bars]
closes = [b["close"] for b in bars]
dates = [b["date"] for b in bars]
# 信号按日期归类(同一天可能多个 buy/sell,取最后一个有效方向)
by_day = {}
for s in signals:
by_day[s["date"]] = s
position = 0.0 # 持股数量(按买入价格折算)
cash = 1.0 # 初始资金=1
trades = []
entry = None
equity = []
trade_idx = 0
for i in range(n):
d = dates[i]
# 当日开盘执行前一日收盘信号
sig = by_day.get(d)
exec_price = opens[i]
if sig and sig["action"] == "buy" and position == 0:
position = cash / exec_price
cash = 0.0
entry = {"date": d, "price": exec_price, "reason": sig["reason"]}
elif sig and sig["action"] == "sell" and position > 0:
cash = position * exec_price
position = 0.0
if entry:
ret = (exec_price / entry["price"] - 1)
trades.append({"entry_date": entry["date"], "entry_price": round(entry["price"], 3),
"exit_date": d, "exit_price": round(exec_price, 3),
"return": round(ret * 100, 2),
"days": _day_diff(dates, entry["date"], d),
"reason": entry["reason"]})
entry = None
value = cash + position * closes[i]
equity.append({"date": d, "value": round(value, 4)})
# 期末仍持仓则平仓(按最后收盘)
if position > 0 and entry:
last = closes[-1]
cash = position * last
ret = last / entry["price"] - 1
trades.append({"entry_date": entry["date"], "entry_price": round(entry["price"], 3),
"exit_date": dates[-1], "exit_price": round(last, 3),
"return": round(ret * 100, 2), "days": _day_diff(dates, entry["date"], dates[-1]),
"reason": entry["reason"] + "(期末平仓)"})
position = 0.0
equity[-1]["value"] = round(cash, 4)
metrics = _calc_metrics(equity, trades, dates, bars)
# 买入持有基准
bh = bars[0]["close"]
for e in equity:
e["bh"] = round(bars[n - 1]["close"] / bh, 4) if bh else 1.0
return {"metrics": metrics, "equity": equity, "trades": trades}
def _day_diff(dates, start, end):
try:
from datetime import date
ds = date.fromisoformat(start)
de = date.fromisoformat(end)
return (de - ds).days
except Exception:
return 0
def _calc_metrics(equity, trades, dates, bars):
n = len(equity)
final = equity[-1]["value"] if equity else 1.0
total = final - 1
years = n / 252.0
ann = (final ** (1 / years) - 1) if (years > 0 and final > 0) else 0
# 最大回撤
peak, mdd = equity[0]["value"], 0.0
for e in equity:
peak = max(peak, e["value"])
mdd = min(mdd, (e["value"] - peak) / peak)
# 日收益 → 夏普
rets = []
for i in range(1, n):
prev = equity[i - 1]["value"]
if prev > 0:
rets.append(equity[i]["value"] / prev - 1)
sharpe = 0.0
if rets and statistics.stdev(rets) > 0:
sharpe = statistics.mean(rets) / statistics.stdev(rets) * math.sqrt(252)
# 交易统计
n_tr = len(trades)
wins = [t for t in trades if t["return"] > 0]
losses = [t for t in trades if t["return"] <= 0]
win_rate = len(wins) / n_tr if n_tr else 0.0
gp = sum(t["return"] for t in wins)
gl = abs(sum(t["return"] for t in losses))
pf = (gp / gl) if gl > 0 else (gp if gp > 0 else 0)
avg_hold = sum(t["days"] for t in trades) / n_tr if n_tr else 0
# 基准(买入持有)
bh_ret = bars[-1]["close"] / bars[0]["close"] - 1
return {
"total_return": round(total * 100, 2),
"annualized": round(ann * 100, 2),
"max_drawdown": round(mdd * 100, 2),
"sharpe": round(sharpe, 2),
"win_rate": round(win_rate * 100, 1),
"profit_factor": round(pf, 2),
"trades": n_tr,
"avg_hold_days": round(avg_hold, 1),
"benchmark": round(bh_ret * 100, 2),
"excess": round((total - bh_ret) * 100, 2),
"days": n,
}
# ===================================================================== 批量回测
def run_one(code, name, bars, key):
"""单只股票单策略回测,返回入库行"""
res = backtest(bars, key)
return {
"strategy": key, "code": code, "stock_name": name,
"metrics": json.dumps(res["metrics"], ensure_ascii=False),
"equity": json.dumps(res["equity"], ensure_ascii=False),
"trades": json.dumps(res["trades"], ensure_ascii=False),
}
def build_all(progress=None):
"""全市场 × 全策略批量回测(覆盖写入 strategy_backtests"""
from database import query, executemany
stocks = query("SELECT code, name FROM stocks")
rows = []
for si, s in enumerate(stocks):
bars = query("SELECT date, open, high, low, close, volume FROM stock_daily "
"WHERE code=? ORDER BY date ASC", (s["code"],))
if len(bars) < 30:
continue
for key in STRATEGIES:
rows.append(run_one(s["code"], s["name"], bars, key))
if progress:
progress(si + 1, len(stocks))
executemany(
"INSERT OR REPLACE INTO strategy_backtests(strategy, code, stock_name, metrics, equity, trades, run_at) "
"VALUES(?,?,?,?,?,?,datetime('now','localtime'))",
[(r["strategy"], r["code"], r["stock_name"], r["metrics"], r["equity"], r["trades"]) for r in rows])
return len(rows)
def strategy_summary(key):
"""某策略的全市场统计(用于列表页头部)"""
from database import query
rows = query("SELECT metrics FROM strategy_backtests WHERE strategy=?", (key,))
if not rows:
return {}
best = None
sums = {"total": 0, "sharpe": 0, "win": 0, "n": 0}
for r in rows:
m = json.loads(r["metrics"])
sums["total"] += m["total_return"]
sums["sharpe"] += m["sharpe"]
sums["win"] += m["win_rate"]
sums["n"] += 1
if best is None or m["total_return"] > best["metrics"]["total_return"]:
best = {"code": None, "metrics": m}
# 顺便找最优个股(第二遍,轻量)
bcode, bname, bret = None, None, -1e9
rows2 = query("SELECT code, stock_name, metrics FROM strategy_backtests WHERE strategy=?", (key,))
for r in rows2:
m = json.loads(r["metrics"])
if m["total_return"] > bret:
bret, bcode, bname = m["total_return"], r["code"], r["stock_name"]
nn = sums["n"]
return {
"avg_return": round(sums["total"] / nn, 2) if nn else 0,
"avg_sharpe": round(sums["sharpe"] / nn, 2) if nn else 0,
"avg_win_rate": round(sums["win"] / nn, 1) if nn else 0,
"n": nn,
"best_code": bcode, "best_name": bname, "best_return": round(bret, 2) if bcode else None,
}
def market_rank(key, limit=100):
"""某策略全市场收益榜"""
from database import query
rows = query("SELECT code, stock_name, metrics FROM strategy_backtests WHERE strategy=? ORDER BY run_at DESC", (key,))
items = []
for r in rows:
m = json.loads(r["metrics"])
items.append({"code": r["code"], "name": r["stock_name"], **m})
items.sort(key=lambda x: x["total_return"], reverse=True)
return items[:limit]