Files
kis_bot/kis_trader/backtest/optuna_orderbook_recommend.py

428 lines
16 KiB
Python

"""
kis_trader/backtest/optuna_orderbook_recommend.py
=================================================
Optuna 차트 캔들 최적화(Stage 1)가 완료된 후, 후처리(Stage 2)로 1,000회 고속 호가 탐색을 수행하여
전략(모멘텀/돌파 등)별 최적의 진입 호가필터 & 수익구간 호가매도 합의 수치(Consensus)를 도출하고
Optuna out_data 및 Apply 패치에 자동으로 결합하는 핵심 모듈입니다.
"""
from __future__ import annotations
import logging
import math
from dataclasses import dataclass
from datetime import datetime, timedelta
from typing import Any, Dict, List, Optional, Tuple
import optuna
from database import TradeDB
logger = logging.getLogger("OptunaOBRecommend")
optuna.logging.set_verbosity(optuna.logging.WARNING)
@dataclass
class Snap:
t: datetime
total_bid: int
total_ask: int
best_bid: int
best_ask: int
@dataclass
class TradeInfo:
code: str
name: str
buy_dt: datetime
buy_price: float
sell_price: float
qty: int
actual_pnl: float
actual_profit_rate: float
entry_snaps: List[Snap]
holding_snaps: List[Snap]
orig_spread_pct: float
orig_bid_ask_ratio: float
passed_current: bool
def _parse_dt(v: Any) -> datetime:
if isinstance(v, datetime):
return v
return datetime.strptime(str(v).strip()[:19], "%Y-%m-%d %H:%M:%S")
def _snap_to_dt(snap_time: str) -> Optional[datetime]:
s = str(snap_time or "").strip()
if not s or len(s) < 14 or s == "None":
return None
try:
return datetime(int(s[:4]), int(s[4:6]), int(s[6:8]), int(s[8:10]), int(s[10:12]), int(s[12:14]))
except ValueError:
return None
def _parse_ratchet_tiers(val: str) -> List[Tuple[int, float]]:
t = []
for p in str(val or "").strip().split(","):
if ":" in p:
parts = p.split(":")
try:
t.append((int(parts[0].strip()), float(parts[1].strip())))
except ValueError:
pass
t.sort(key=lambda x: x[0])
if not t:
t = [(10, 2.6), (13, 2.2)]
return t
def get_orderbook_table_for_strategy(strategy: str) -> str:
s = strategy.upper()
if any(p in s for p in ("BREAKOUT", "SCALP", "LS")):
return "ls_ws_orderbook"
return "ws_orderbook"
def recommend_orderbook_parameters(
strategy: str = "MOMENTUM",
n_trials: int = 1000,
ob_table: Optional[str] = None,
log: Optional[logging.Logger] = None,
) -> Dict[str, Any]:
lg = log or logger
strat_upper = strategy.upper()
table = ob_table or get_orderbook_table_for_strategy(strat_upper)
db = TradeDB()
# 1. 테이블 존재 여부 및 컬럼 검사
try:
cols = [r["Field"] for r in db.conn.execute(f"SHOW COLUMNS FROM {table}").fetchall()]
need = {"code", "snap_time", "total_bid_qty", "total_ask_qty", "best_bid", "best_ask"}
if need - set(cols):
lg.warning("⚠️ [%s] 호가 테이블 필수 컬럼 부족. 추천 생략.", table)
return {"ok": False, "reason": "insufficient_columns"}
except Exception as exc:
lg.warning("⚠️ [%s] 테이블 조회 실패: %s. 추천 생략.", table, exc)
return {"ok": False, "reason": "table_not_found"}
date_rows = db.conn.execute(f"SELECT DISTINCT SUBSTR(snap_time, 1, 8) as dt FROM {table} ORDER BY dt").fetchall()
avail_dates = [str(r["dt"]) for r in date_rows if r["dt"] and str(r["dt"]) != "None"]
if not avail_dates:
return {"ok": False, "reason": "no_orderbook_snapshots"}
# 2. 전략별 config (래칫/어깨) 로딩
cfg_table = "config_breakout" if "BREAKOUT" in strat_upper else "config_momentum"
pfx = "BREAKOUT_" if "BREAKOUT" in strat_upper else "MOMENTUM_"
try:
row_cfg = db.conn.execute(
f"SELECT {pfx}RATCHET_TIERS, {pfx}SHOULDER_MIN_HIGH_PCT, {pfx}SHOULDER_CUT_PCT FROM {cfg_table} ORDER BY id DESC LIMIT 1"
).fetchone()
cfg_dict = dict(row_cfg) if row_cfg else {}
except Exception:
cfg_dict = {}
tiers = _parse_ratchet_tiers(str(cfg_dict.get(f"{pfx}RATCHET_TIERS") or "10:2.6,13:2.2"))
smh = float(cfg_dict.get(f"{pfx}SHOULDER_MIN_HIGH_PCT") or 0.05)
sc = float(cfg_dict.get(f"{pfx}SHOULDER_CUT_PCT") or 0.0055)
if smh > 1.0:
smh /= 100.0
if sc > 1.0:
sc /= 100.0
trades: List[TradeInfo] = []
for dt_str in avail_dates:
day_hyphen = f"{dt_str[:4]}-{dt_str[4:6]}-{dt_str[6:]}"
buys = db.conn.execute(
"""
SELECT id, code, name, buy_date, buy_price, sell_price, qty, profit_rate, realized_pnl
FROM trade_history
WHERE strategy=%s AND DATE(buy_date)=%s
ORDER BY buy_date
""",
(strat_upper, day_hyphen),
).fetchall()
for b in buys:
try:
buy_dt = _parse_dt(b["buy_date"])
except Exception:
continue
buy_price = float(b["buy_price"] or 0)
code = str(b["code"])
qty = int(b["qty"] or 0)
actual_pnl = float(b["realized_pnl"] or 0)
actual_profit_rate = float(b["profit_rate"] or 0)
lookback = timedelta(minutes=10)
horizon = timedelta(minutes=6)
s_rows = db.conn.execute(
f"SELECT snap_time, total_bid_qty, total_ask_qty, best_bid, best_ask FROM {table} "
"WHERE code=%s AND snap_time >= %s AND snap_time < %s ORDER BY snap_time ASC",
(code, (buy_dt - lookback).strftime("%Y%m%d%H%M%S"), (buy_dt + horizon).strftime("%Y%m%d%H%M%S")),
).fetchall()
snaps: List[Snap] = []
for sr in s_rows:
s_dt = _snap_to_dt(sr["snap_time"])
if not s_dt:
continue
snaps.append(
Snap(
t=s_dt,
total_bid=int(sr["total_bid_qty"] or 0),
total_ask=int(sr["total_ask_qty"] or 0),
best_bid=int(sr["best_bid"] or 0),
best_ask=int(sr["best_ask"] or 0),
)
)
if not snaps or buy_price <= 0 or qty <= 0:
continue
before = [s for s in snaps if s.t <= buy_dt]
entry_snaps = before if before else snaps
holding_snaps = [s for s in snaps if s.t > buy_dt]
# 진입 호가 지표 산정
best_entry_snap = entry_snaps[-1]
bid, ask = best_entry_snap.best_bid, best_entry_snap.best_ask
spread_pct = ((ask - bid) / ((ask + bid) / 2.0)) * 100.0 if (ask > 0 and bid > 0) else 0.0
tot_b, tot_a = best_entry_snap.total_bid, best_entry_snap.total_ask
ratio = (tot_b / tot_a) if tot_a > 0 else 999.0
trades.append(
TradeInfo(
code=code,
name=str(b.get("name") or code),
buy_dt=buy_dt,
buy_price=buy_price,
sell_price=float(b["sell_price"] or 0),
qty=qty,
actual_pnl=actual_pnl,
actual_profit_rate=actual_profit_rate,
entry_snaps=entry_snaps,
holding_snaps=holding_snaps,
orig_spread_pct=spread_pct,
orig_bid_ask_ratio=ratio,
passed_current=(spread_pct <= 0.45 and ratio >= 0.85),
)
)
if len(trades) < 3:
lg.warning("⚠️ [%s] 호가 연제 가능한 실제 매수 건수(%s건)가 부족하여 최적화 생략.", strat_upper, len(trades))
return {"ok": False, "reason": "not_enough_trades", "trade_count": len(trades)}
# 원본 실측값
orig_cnt = len(trades)
orig_win = sum(1 for t in trades if t.actual_pnl > 0) / orig_cnt * 100.0
orig_pnl = sum(t.actual_pnl for t in trades)
orig_rate = sum(t.actual_profit_rate for t in trades) / orig_cnt
def _sim_trade(tr: TradeInfo, p: Dict[str, Any]) -> Tuple[float, float, str]:
if tr.orig_spread_pct > p["max_spread_pct"] or tr.orig_bid_ask_ratio < p["min_bid_ask_ratio"]:
return (0.0, 0.0, "ENTRY_REJECTED")
if not p["use_ob_exit"]:
return (tr.actual_pnl, tr.actual_profit_rate, "ORIGINAL")
ma_win = p["ma_window"]
ratio_min = p["ob_ratio_min"]
min_prof = p["min_profit_pct"]
min_hold = p["min_hold_bars"]
history_ratio: List[float] = []
high_price = tr.buy_price
for s in tr.holding_snaps:
cur_p = float(s.best_bid if s.best_bid > 0 else s.best_ask)
if cur_p <= 0:
continue
if cur_p > high_price:
high_price = cur_p
r = (s.total_bid / s.total_ask) if s.total_ask > 0 else 1.0
history_ratio.append(r)
held_sec = (s.t - tr.buy_dt).total_seconds()
if held_sec < min_hold * 60:
continue
if len(history_ratio) < ma_win:
continue
ma_r = sum(history_ratio[-ma_win:]) / float(ma_win)
prof_rate = (cur_p - tr.buy_price) / tr.buy_price
if prof_rate < min_prof:
continue
# 래칫 & 어깨 체크
high_rate = (high_price - tr.buy_price) / tr.buy_price
if high_rate >= smh and ((high_price - cur_p) / high_price) >= sc:
continue
is_ratchet_blocked = False
for bar_m, target_r in tiers:
if (held_sec >= bar_m * 60) and (prof_rate >= (target_r / 100.0)):
is_ratchet_blocked = True
break
if is_ratchet_blocked:
continue
if ma_r < ratio_min:
realized = (cur_p - tr.buy_price) * tr.qty
return (realized, prof_rate * 100.0, "OB_EXIT")
return (tr.actual_pnl, tr.actual_profit_rate, "HOLD_TO_ORIG")
def _calc_suite(p: Dict[str, Any]) -> Tuple[int, float, float, float]:
t_cnt = 0
w_cnt = 0
tot_pnl = 0.0
tot_rate = 0.0
for t in trades:
pnl, rate, rtype = _sim_trade(t, p)
if rtype != "ENTRY_REJECTED":
t_cnt += 1
tot_pnl += pnl
tot_rate += rate
if pnl > 0:
w_cnt += 1
w_rate = (w_cnt / t_cnt * 100.0) if t_cnt > 0 else 0.0
avg_r = (tot_rate / t_cnt) if t_cnt > 0 else 0.0
return t_cnt, w_rate, tot_pnl, avg_r
valid_records: List[Dict[str, Any]] = []
can_use_ob_exit = strat_upper in ("MOMENTUM", "BREAKOUT")
def obj_func(trial: optuna.Trial) -> float:
params = {
"max_spread_pct": trial.suggest_float("max_spread_pct", 0.3, 3.5, step=0.1),
"min_bid_ask_ratio": trial.suggest_float("min_bid_ask_ratio", 0.1, 1.0, step=0.05),
"use_ob_exit": trial.suggest_categorical("use_ob_exit", [True, False]) if can_use_ob_exit else False,
"min_hold_bars": trial.suggest_int("min_hold_bars", 1, 5) if can_use_ob_exit else 3,
"ob_ratio_min": trial.suggest_float("ob_ratio_min", 0.2, 0.8, step=0.05) if can_use_ob_exit else 0.4,
"min_profit_pct": trial.suggest_float("min_profit_pct", 0.003, 0.02, step=0.001) if can_use_ob_exit else 0.005,
"ma_window": trial.suggest_int("ma_window", 3, 10) if can_use_ob_exit else 5,
}
cnt, win_r, pnl, rate = _calc_suite(params)
if cnt < max(3, int(orig_cnt * 0.3)):
return -999999999.0
w_p = (pnl / 100000.0)
w_w = win_r * 2.0
score = w_p + w_w
if win_r >= 60.0:
score += (win_r - 60.0) * 1.5
valid_records.append({"score": score, "pnl": pnl, "win_rate": win_r, "count": cnt, "rate": rate, "params": params})
return score
study = optuna.create_study(direction="maximize")
study.optimize(obj_func, n_trials=n_trials)
valid_records.sort(key=lambda x: x["score"], reverse=True)
top5 = valid_records[: min(5, len(valid_records))]
if not top5:
return {"ok": False, "reason": "no_valid_trials"}
avg_spread = round(sum(r["params"]["max_spread_pct"] for r in top5) / len(top5), 2)
avg_bid_ask = round(sum(r["params"]["min_bid_ask_ratio"] for r in top5) / len(top5), 2)
use_ob_votes = sum(1 for r in top5 if r["params"]["use_ob_exit"])
cons_ob_exit = use_ob_votes >= (len(top5) / 2.0)
avg_ob_ratio = round(sum(r["params"]["ob_ratio_min"] for r in top5) / len(top5), 2)
avg_hold = int(round(sum(r["params"]["min_hold_bars"] for r in top5) / len(top5)))
avg_ma = int(round(sum(r["params"]["ma_window"] for r in top5) / len(top5)))
avg_prof = round(sum(r["params"]["min_profit_pct"] for r in top5) / len(top5), 4)
cons_params = {
"max_spread_pct": avg_spread,
"min_bid_ask_ratio": avg_bid_ask,
"use_ob_exit": cons_ob_exit,
"min_hold_bars": avg_hold,
"ob_ratio_min": avg_ob_ratio,
"min_profit_pct": avg_prof,
"ma_window": avg_ma,
}
c_cnt, c_win, c_pnl, c_rate = _calc_suite(cons_params)
lg.info(
"⚡ [호가 수급 합의 추천] 전략=%s (모수=%d건, %d회 탐색) | 스프레드≤%.2f%% 잔량비≥%.2f | 호가익절=%s | 승률: %.1f%% 손익: %.0f",
strat_upper,
len(trades),
n_trials,
avg_spread,
avg_bid_ask,
"ON" if cons_ob_exit else "OFF",
c_win,
c_pnl,
)
return {
"ok": True,
"strategy": strat_upper,
"ob_table": table,
"n_trials": n_trials,
"trade_count": len(trades),
"orig_stats": {"count": orig_cnt, "win_rate": round(orig_win, 1), "pnl": orig_pnl, "avg_rate": round(orig_rate, 2)},
"recommended_stats": {"count": c_cnt, "win_rate": round(c_win, 1), "pnl": c_pnl, "avg_rate": round(c_rate, 2), "pnl_diff": c_pnl - orig_pnl},
"params": {
"orderbook_filter_enabled": True,
"orderbook_max_spread_pct": avg_spread,
"orderbook_min_bid_ask_ratio": avg_bid_ask,
"exit_ob_enabled": cons_ob_exit,
"exit_ob_min_hold_bars": avg_hold,
"exit_ob_min_profit_pct": avg_prof,
"exit_ob_ratio_min": avg_ob_ratio,
"exit_ob_ma_window": avg_ma,
},
}
def attach_orderbook_recommend(
out_data: Dict[str, Any],
*,
log: Optional[logging.Logger] = None,
) -> Dict[str, Any]:
"""out_data에 호가 진입/청산 합의 수치 추천(orderbook_recommend)을 첨부."""
lg = log or logger
strat = str(out_data.get("strategy") or "MOMENTUM").strip().upper()
rec = recommend_orderbook_parameters(strategy=strat, n_trials=1000, log=lg)
out_data["orderbook_recommend"] = rec
mc = out_data.get("mode_combo")
if isinstance(mc, dict):
mc["orderbook_recommend"] = rec
if not rec.get("ok"):
lg.info("⚡ [호가 수급 합의 추천] 생략 — %s", rec.get("reason") or "n/a")
return out_data
def build_orderbook_env_patch(rec: Dict[str, Any]) -> Dict[str, str]:
"""호가 수급 합의 추천 결과를 DB env 패치 dict로 변환."""
if not rec or not rec.get("ok"):
return {}
strat = str(rec.get("strategy") or "").strip().upper()
pfx = "TAIL" if strat in ("SHORT", "TAIL") else strat
p = rec.get("params", {})
if not pfx or not p:
return {}
patch = {
f"{pfx}_ORDERBOOK_FILTER_ENABLED": "true" if p.get("orderbook_filter_enabled", True) else "false",
f"{pfx}_ORDERBOOK_MAX_SPREAD_PCT": str(p["orderbook_max_spread_pct"]),
f"{pfx}_ORDERBOOK_MIN_BID_ASK_RATIO": str(p["orderbook_min_bid_ask_ratio"]),
}
if pfx in ("MOMENTUM", "BREAKOUT") and "exit_ob_enabled" in p:
patch[f"{pfx}_EXIT_OB_ENABLED"] = "true" if p.get("exit_ob_enabled") else "false"
if p.get("exit_ob_enabled"):
patch[f"{pfx}_EXIT_OB_RATIO_MIN"] = str(p["exit_ob_ratio_min"])
patch[f"{pfx}_EXIT_OB_MA_WINDOW"] = str(p["exit_ob_ma_window"])
patch[f"{pfx}_EXIT_OB_MIN_PROFIT_PCT"] = str(p["exit_ob_min_profit_pct"])
patch[f"{pfx}_EXIT_OB_MIN_HOLD_BARS"] = str(p["exit_ob_min_hold_bars"])
return patch