Files
MoFin/scripts/research/step49_save_to_research.py
T

271 lines
12 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.
#!/usr/bin/env python3
"""step49_save_to_research.py — 复现 v5 定稿(年化18.57%)并保存到 strategy_research
在 step49_final_v5.py 基础上:记录逐笔 trades + 保存 results_json(研究Tab 展示)
输出:strategy_research 表 version='v_oversold' period_tag='10y'
"""
import numpy as np
import pandas as pd
import sqlite3
import sys, json
from datetime import datetime
sys.path.insert(0, "/tmp")
sys.path.insert(0, "/home/hmo/MoFin")
from sr_calculator import SRCalculator
import strategy_lab as lab
print("=== 加载面板 ===", flush=True)
panel = pd.read_pickle("/tmp/panel_12d.pkl")
panel = panel.sort_values(["code", "date"]).reset_index(drop=True)
panel["_key"] = panel["code"] + "_" + panel["date"]
pos_map = {k: i for i, k in enumerate(panel["_key"])}
dates = sorted(panel["date"].unique())
# 大盘指标(mkt_dd60 + mkt_down_days
conn = sqlite3.connect("file:/home/hmo/MoFin/data/mofin.db?mode=ro", uri=True)
idx_df = pd.read_sql("SELECT date, close, high FROM stock_daily WHERE code='sh000001' ORDER BY date", conn)
idx_df["date"] = idx_df["date"].astype(str)
idx_df["hi60"] = idx_df["high"].rolling(60).max()
idx_df["mkt_dd60"] = (idx_df["close"] / idx_df["hi60"] - 1) * 100
close_arr = idx_df["close"].values
down_days, streak, prev = [], 0, None
for c in close_arr:
streak = streak + 1 if (prev is not None and c < prev) else 0
down_days.append(streak)
prev = c
idx_df["mkt_down_days"] = down_days
dd_map = dict(zip(idx_df["date"], idx_df["mkt_dd60"]))
down_map = dict(zip(idx_df["date"], idx_df["mkt_down_days"]))
panel["mkt_dd60"] = panel["date"].map(dd_map)
panel["mkt_down_days"] = panel["date"].map(down_map)
# 信号(v5 定稿条件)
sig_cond = (
(panel["mkt_rsi"] < 50) & (panel["mcap_q"] < 0.2) & (panel["pe_q"] < 0.2) &
(panel["news3"] >= 1) & (panel["sec_ret20"] < 0) & (panel["bias60"] < -20) &
(panel["mkt_dd60"] <= -5)
)
yin_die = (panel["mkt_down_days"] >= 2) & (panel["mkt_adx"] <= 55) & (panel["mkt_rsi"] >= 33)
final_cond = sig_cond & ~yin_die
print("基础信号:", sig_cond.sum(), "| 阴跌跳过:", (sig_cond & yin_die).sum(), "| 最终:", final_cond.sum(), flush=True)
def build_cand(cond):
cand = panel[cond][["code", "date"]].copy()
cand = cand.sort_values(["code", "date"])
cand["prev"] = cand.groupby("code")["date"].shift(1)
cand["gap"] = (pd.to_datetime(cand["date"]) - pd.to_datetime(cand["prev"])).dt.days
cand = cand[(cand["prev"].isna()) | (cand["gap"] > 30)]
return cand
def prep(sigdf):
codes = sigdf["code"].unique().tolist()
ph = ",".join("?" * len(codes))
df = pd.read_sql("SELECT code, date, close, high, low FROM stock_daily WHERE code IN ({}) ORDER BY code, date".format(ph), conn, params=codes)
df["date"] = df["date"].astype(str)
df["code"] = df["code"].astype(str).str.zfill(6)
df = df.sort_values(["code", "date"]).reset_index(drop=True)
df["_key"] = df["code"] + "_" + df["date"]
dpos = {k: i for i, k in enumerate(df["_key"])}
sigdf["kidx"] = (sigdf["code"] + "_" + sigdf["date"]).map(dpos)
sigdf = sigdf.dropna(subset=["kidx"]).copy()
sigdf["kidx"] = sigdf["kidx"].astype(int)
sr = SRCalculator()
out = []
for r in sigdf.itertuples():
bars_code = sr.get_bars(r.code)
d_idx = bars_code.index[bars_code["date"] == r.date]
if len(d_idx) == 0:
continue
sr_full = sr.sr_full(r.code, d_idx[0])
pv = sr_full["pivot"]
chip = sr_full["chip"]
if not pv:
continue
sig_close = df["close"].values[r.kidx]
support = pv["s2"]
if chip and chip["chip_ss"] < sig_close:
support = max(pv["s2"], chip["chip_ss"])
resist = pv["r2"]
if chip and chip["chip_sr"] > sig_close:
resist = min(pv["r2"], chip["chip_sr"])
if support >= sig_close or resist <= sig_close:
continue
out.append({"code": r.code, "date": r.date, "kidx": r.kidx,
"sig_close": sig_close, "support": support, "resist": resist})
return pd.DataFrame(out), df
def run_sim(sigdf, df, max_daily=5, slots=10, hold=40, pos_frac=0.15, stop_buf=0.05):
"""同 step49 run_sim,但记录逐笔 trades"""
dpos = {k: i for i, k in enumerate(df["_key"])}
closes = df["close"].values
highs = df["high"].values
lows = df["low"].values
sig_by_date = {}
for r in sigdf.itertuples():
sig_by_date.setdefault(r.date, []).append(r)
INIT_CAP = 1_000_000
positions = {}
cash = INIT_CAP
navs = []
trades = []
for di, d in enumerate(dates):
for code in list(positions.keys()):
pos = positions[code]
k = dpos.get(code + "_" + d)
if k is None:
continue
hi, lo, cl = highs[k], lows[k], closes[k]
exited = False
if hi >= pos["tp"]:
sell_qty = pos["qty"] // 2
if sell_qty > 0:
cash += sell_qty * pos["tp"]
pos["qty"] -= sell_qty
pos["sold_parts"].append({"pct": sell_qty / pos["init_qty"], "price": pos["tp"]})
if pos["qty"] <= 0:
del positions[code]
exited = True
if not exited and lo <= pos["stop"]:
cash += pos["qty"] * lo
pos["exit_price"] = lo
pos["exit_reason"] = "stop"
pos["exit_date"] = d
pos["hold_days"] = di - pos["entry_di"]
trades.append(pos)
del positions[code]
exited = True
if not exited and di - pos["entry_di"] >= hold:
cash += pos["qty"] * cl
pos["exit_price"] = cl
pos["exit_reason"] = "time"
pos["exit_date"] = d
pos["hold_days"] = di - pos["entry_di"]
trades.append(pos)
del positions[code]
if d in sig_by_date:
bought = 0
for r in sig_by_date[d]:
if bought >= max_daily or len(positions) >= slots:
break
if r.code in positions:
continue
cur_nav = cash
for c, p in positions.items():
k = dpos.get(c + "_" + d)
px = closes[k] if k is not None else p["avg_cost"]
cur_nav += p["qty"] * px
pos_val = cur_nav * pos_frac
price = r.sig_close
if price <= 0:
continue
qty = int(pos_val / price)
if qty <= 0 or qty * price > cash:
continue
cash -= qty * price
positions[r.code] = {
"code": r.code, "entry_date": r.date, "entry_price": price,
"entry_di": di, "qty": qty, "init_qty": qty, "avg_cost": price,
"stop": r.support * (1 - stop_buf), "tp": r.resist,
"support": r.support, "resist": r.resist,
"sold_parts": [], "exit_price": None, "exit_reason": None,
"exit_date": None, "hold_days": None,
}
bought += 1
nav = cash
for c, p in positions.items():
k = dpos.get(c + "_" + d)
px = closes[k] if k is not None else p["avg_cost"]
nav += p["qty"] * px
navs.append(nav)
# 收盘时未平仓的强制平仓
for code in list(positions.keys()):
pos = positions[code]
k = dpos.get(code + "_" + dates[-1])
px = closes[k] if k is not None else pos["avg_cost"]
pos["exit_price"] = px
pos["exit_reason"] = "end"
pos["exit_date"] = dates[-1]
pos["hold_days"] = len(dates) - 1 - pos["entry_di"]
trades.append(pos)
nav_arr = np.array(navs)
final = nav_arr[-1]
years = len(navs) / 250
cagr = ((final / INIT_CAP) ** (1 / years) - 1) * 100 if final > 0 else -100
cummax = np.maximum.accumulate(nav_arr)
dd = ((nav_arr - cummax) / cummax).min() * 100
return cagr, dd, nav_arr, trades
# ── 支持 period_tag 参数(10y/5y/2y)──
PERIOD_TAG = sys.argv[1] if len(sys.argv) > 1 else '10y'
PERIODS = {
'10y': ('2016-07-01', '2026-07-24'),
'5y': ('2021-07-01', '2026-07-24'),
'2y': ('2024-07-01', '2026-07-24'),
}
BT_START, BT_END = PERIODS.get(PERIOD_TAG, PERIODS['10y'])
print(f"=== v5 定稿(阴跌判定+限流5)复现 + 保存 [{PERIOD_TAG}] {BT_START}~{BT_END} ===", flush=True)
# 按回测窗口过滤信号(同股30日去重保持)
bt_mask = (panel["date"] >= BT_START) & (panel["date"] <= BT_END)
s1, df1 = prep(build_cand(final_cond & bt_mask))
# run_sim 里 dates 也要限窗口
dates = sorted(d for d in dates if BT_START <= d <= BT_END)
c1, d1, nav1, trades = run_sim(s1, df1)
print("年化={:.2f}% 回撤={:.1f}% 信号={} trades={}".format(c1, d1, len(s1), len(trades)), flush=True)
# 构造 result dict(对齐 run_mr_backtest 的 result 格式)
trade_list = []
for t in trades:
profit_pct = (t["exit_price"] / t["entry_price"] - 1) * 100 if t["entry_price"] > 0 else 0
trade_list.append({
"code": t["code"], "name": t["code"],
"entry_date": t["entry_date"], "entry_price": round(t["entry_price"], 2),
"exit_price": round(t["exit_price"], 2), "profit_pct": round(profit_pct, 2),
"exit_reason": t["exit_reason"], "hold_days": t["hold_days"],
"score": 0, "score_comp": {}, "kelly": 0,
"stop_loss": round(t["stop"], 2), "target": round(t["tp"], 2),
"dna": False,
"factors": {"bias60": None, "mkt_rsi": None, "mcap_q": None, "pe_q": None,
"news3": None, "sec_ret20": None, "support": round(t["support"], 2),
"resist": round(t["resist"], 2)},
})
wins = [t for t in trade_list if t["profit_pct"] > 0]
losses = [t for t in trade_list if t["profit_pct"] <= 0]
summary = {
"total_trades": len(trade_list),
"win_rate": round(len(wins) / len(trade_list) * 100, 1) if trade_list else 0,
"avg_profit_pct": round(sum(t["profit_pct"] for t in trade_list) / len(trade_list), 2) if trade_list else 0,
"avg_win_pct": round(sum(t["profit_pct"] for t in wins) / len(wins), 2) if wins else 0,
"avg_loss_pct": round(sum(t["profit_pct"] for t in losses) / len(losses), 2) if losses else 0,
"avg_hold_days": round(sum(t["hold_days"] or 0 for t in trade_list) / len(trade_list), 1) if trade_list else 0,
"wins": len(wins), "losses": len(losses),
"portfolio": {"cagr_pct": round(c1, 1), "total_return_pct": round((nav1[-1] / 1000000 - 1) * 100, 1),
"portfolio_max_dd_pct": round(abs(d1), 1)},
"portfolio_full": {"cagr_pct": round(c1, 1), "total_return_pct": round((nav1[-1] / 1000000 - 1) * 100, 1),
"portfolio_max_dd_pct": round(abs(d1), 1)},
}
print("summary:", json.dumps(summary, ensure_ascii=False), flush=True)
# 保存到 strategy_research(先删旧的同版本同周期记录)
strat = lab.STRATEGIES["v_oversold"]
result = {
"strategy": "v_oversold",
"strategy_name": strat["name"],
"market": "a",
"period": f"{BT_START} ~ {BT_END}",
"period_tag": PERIOD_TAG,
"capital": 1000000,
"total_stocks_screened": len(df1["code"].unique()),
"scored_events": len(trade_list),
"trades": trade_list,
"summary": summary,
}
conn2 = sqlite3.connect("/home/hmo/MoFin/data/mofin.db")
conn2.execute("DELETE FROM strategy_research WHERE version='v_oversold' AND period_tag=?", (PERIOD_TAG,))
conn2.commit()
conn2.close()
lab.save_result(strat, result)
print(f"已保存 v_oversold {PERIOD_TAG} 到 strategy_research", flush=True)
print("=== 完成 ===", flush=True)