Changes: - Added `apply_params_to_db` function to streamline parameter application to the database. - Introduced `evaluate_breakout_param_combo`, `evaluate_momentum_param_combo`, and `evaluate_tail_param_combo` functions to enhance the evaluation of parameter combinations for respective strategies. - Updated `requirements.txt` to include `optuna==4.2.1` for improved optimization capabilities. Impact: - These additions improve the modularity and efficiency of parameter evaluations across different trading strategies, facilitating better optimization and backtesting processes.
391 lines
14 KiB
Python
391 lines
14 KiB
Python
#!/usr/bin/env python3
|
|
"""kis_trader/backtest/optuna_breakout.py — 돌파 Optuna (Grid add-on)."""
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
import os
|
|
import time
|
|
from dataclasses import dataclass, field
|
|
from datetime import datetime
|
|
from typing import Any, Dict, List, Optional
|
|
|
|
import optuna
|
|
from optuna.samplers import RandomSampler, TPESampler
|
|
|
|
from database import TradeDB
|
|
from kis_trader.backtest import breakout_backtest_common as bbc
|
|
from kis_trader.backtest.optuna_search_space import breakout_grid_axis_keys, suggest_breakout_params
|
|
from kis_trader.backtest.param_search_breakout import (
|
|
_bo_fixed_defaults,
|
|
_load_candles_for_search,
|
|
_breakout_grids,
|
|
_ui_to_engine_params,
|
|
apply_params_to_db,
|
|
evaluate_breakout_param_combo,
|
|
)
|
|
from kis_trader.backtest.param_search_cli_common import (
|
|
apply_session_to_fixed,
|
|
combo_passes_search_filters,
|
|
format_session_hm,
|
|
)
|
|
from kis_trader.backtest.tail_param_search import _results_dir_for_write
|
|
from kis_trader.strategies.breakout import breakout_backtest_wants_tick_replay, breakout_entry_mode
|
|
from kis_trader.engine.indicator_cache import attach_indicator_caches_to_params
|
|
from kis_trader.backtest.breakout_tick_loader import load_breakout_ticks_by_code
|
|
|
|
logger = logging.getLogger("param_search_optuna")
|
|
|
|
_FAIL_OBJECTIVE = -1e18
|
|
|
|
|
|
@dataclass
|
|
class BreakoutSearchContext:
|
|
start: str
|
|
end: str
|
|
mode: str
|
|
base_fixed: Dict[str, Any]
|
|
codes_candles: Dict[str, List[Dict]]
|
|
universe_by_slot: Optional[Dict[str, List[str]]]
|
|
ticks_by_code: Any
|
|
orderbook_by_code: Dict[str, Any]
|
|
program_by_code: Dict[str, Any]
|
|
log_verdict_by_code: Dict[str, Any]
|
|
share_denom_by_code: Dict[str, float]
|
|
fee_rate: float
|
|
sell_tax: float
|
|
slot_money: float
|
|
max_stocks: int
|
|
total_budget_krw: float
|
|
period_days: int
|
|
portfolio: Dict[str, Any]
|
|
grid_keys: List[str]
|
|
start_key: str
|
|
end_key: str
|
|
cache_holder: Dict[str, Any] = field(default_factory=dict)
|
|
|
|
|
|
def prepare_breakout_search_context(
|
|
start: str,
|
|
end: str,
|
|
mode: str,
|
|
*,
|
|
use_fallback_universe: bool = False,
|
|
time_start_hm: Optional[int] = None,
|
|
time_end_hm: Optional[int] = None,
|
|
slot_money: Optional[float] = None,
|
|
max_stocks: Optional[int] = None,
|
|
total_budget_krw: Optional[float] = None,
|
|
orderbook_filter: str = "off",
|
|
) -> Optional[BreakoutSearchContext]:
|
|
grids = _breakout_grids()
|
|
if mode not in grids:
|
|
logger.error("❌ 돌파 mode: %s (fast/coarse/fine/full)", mode)
|
|
return None
|
|
|
|
base_fixed = _bo_fixed_defaults()
|
|
apply_session_to_fixed(base_fixed, time_start_hm=time_start_hm, time_end_hm=time_end_hm)
|
|
|
|
_ob_mode = (orderbook_filter or "off").strip().lower()
|
|
if _ob_mode == "off":
|
|
base_fixed["_orderbook_filter_enabled"] = False
|
|
elif _ob_mode == "on":
|
|
base_fixed["_orderbook_filter_enabled"] = True
|
|
ob_filter_on = bool(base_fixed.get("_orderbook_filter_enabled")) or _ob_mode == "auto"
|
|
logger.info(
|
|
"📌 호가필터: %s (%s)",
|
|
_ob_mode.upper(),
|
|
"적용" if ob_filter_on else "스킵 — 코어 파라미터 순수 탐색",
|
|
)
|
|
|
|
db = TradeDB()
|
|
try:
|
|
row = db.conn.execute("SELECT * FROM env_config ORDER BY id DESC LIMIT 1").fetchone()
|
|
env_row = dict(row) if row else {}
|
|
finally:
|
|
db.close()
|
|
|
|
fee_rate, sell_tax, slot_from_env = bbc.fee_and_slot_from_env(env_row)
|
|
portfolio = bbc.resolve_breakout_portfolio_params(
|
|
env_row, None,
|
|
slot_money=slot_money if slot_money is not None else slot_from_env,
|
|
max_stocks=max_stocks,
|
|
total_budget_krw=total_budget_krw,
|
|
)
|
|
slot_money_f = float(portfolio["slot_money"])
|
|
max_stocks_i = int(portfolio["max_stocks"])
|
|
total_budget_f = float(portfolio["total_budget_krw"])
|
|
period_days = max(
|
|
1,
|
|
(datetime.strptime(end, "%Y-%m-%d") - datetime.strptime(start, "%Y-%m-%d")).days + 1,
|
|
)
|
|
logger.info(
|
|
f"💼 포트폴리오: 1회 {slot_money_f:,.0f}원 | 동시 {max_stocks_i}종 | "
|
|
f"총한도 {total_budget_f:,.0f}원 | 매매 {format_session_hm(base_fixed)}"
|
|
)
|
|
logger.info("📌 진입 모드: %s", breakout_entry_mode())
|
|
|
|
codes_candles = _load_candles_for_search(
|
|
start, end, base_fixed.get("lookback_min", 1), base_fixed,
|
|
)
|
|
if not codes_candles:
|
|
logger.error("❌ 캔들 데이터 없음")
|
|
return None
|
|
logger.info("✅ 데이터 로드: %s종목", len(codes_candles))
|
|
|
|
share_denom_by_code: Dict[str, float] = {}
|
|
_share_db = TradeDB()
|
|
try:
|
|
from kis_trader.share.stock_share import load_share_denom_map
|
|
share_denom_by_code = load_share_denom_map(_share_db, codes_candles.keys())
|
|
finally:
|
|
_share_db.close()
|
|
|
|
start_key = (start.replace("-", "") + "0000") if start else "202601010000"
|
|
end_key = (end.replace("-", "") + "2359") if end else "999912312359"
|
|
start_ymd = start.replace("-", "") if start else ""
|
|
end_ymd = end.replace("-", "") if end else ""
|
|
|
|
ticks_by_code: Dict[str, Any] = {}
|
|
engine_probe = _ui_to_engine_params(base_fixed)
|
|
engine_probe["_orderbook_filter_enabled"] = base_fixed.get("_orderbook_filter_enabled")
|
|
if breakout_backtest_wants_tick_replay(engine_probe):
|
|
_tick_db = TradeDB()
|
|
try:
|
|
ticks_by_code, tick_rows = load_breakout_ticks_by_code(
|
|
_tick_db, start_key, end_key, set(codes_candles.keys()),
|
|
)
|
|
logger.info("✅ ws_ticks %s건", f"{tick_rows:,}")
|
|
finally:
|
|
_tick_db.close()
|
|
|
|
grid = grids[mode]
|
|
_ob_axes = ("max_spread_pct", "min_bid_ask_ratio", "ask_wall_max_qty")
|
|
_ob_sweeping = any(len(set(grid.get(k) or [])) > 1 for k in _ob_axes)
|
|
if ob_filter_on and _ob_sweeping:
|
|
base_fixed["backtest_use_kiwoom_body_snapshot"] = True
|
|
base_fixed["_backtest_use_kiwoom_body"] = True
|
|
|
|
orderbook_by_code: Dict[str, Any] = {}
|
|
program_by_code: Dict[str, Any] = {}
|
|
log_verdict_by_code: Dict[str, Any] = {}
|
|
_snap_db = TradeDB()
|
|
try:
|
|
from kis_trader.backtest.trigger_snapshot_loader import load_trigger_snapshots_by_code
|
|
orderbook_by_code, program_by_code, trigger_snap_meta = load_trigger_snapshots_by_code(
|
|
_snap_db, start_key, end_key, set(codes_candles.keys()),
|
|
engine_params=engine_probe, strategy="BREAKOUT",
|
|
)
|
|
log_verdict_by_code = trigger_snap_meta.get("log_verdict_by_code") or {}
|
|
finally:
|
|
_snap_db.close()
|
|
|
|
universe_by_slot = None
|
|
fallback_sim_interval = 5
|
|
if not use_fallback_universe and start_ymd and end_ymd:
|
|
try:
|
|
from kis_trader.database.db_manager import get_db as _get_ext_db
|
|
_ext = _get_ext_db()
|
|
history = _ext.get_universe_by_candle_time(
|
|
strategy_id="BREAKOUT", start_ymd=start_ymd, end_ymd=end_ymd,
|
|
)
|
|
if history:
|
|
universe_by_slot = history
|
|
n_bins = len(history)
|
|
avg = sum(len(v) for v in history.values()) / max(1, n_bins)
|
|
logger.info("✅ 유니버스: BREAKOUT 이력 | %s분봉 · 평균 %.1f종목", n_bins, avg)
|
|
except Exception as exc:
|
|
logger.debug("유니버스 이력 스킵: %s", exc)
|
|
|
|
if universe_by_slot is None:
|
|
from kis_trader.engine import scalping_engine as se
|
|
universe_top_n = int(os.environ.get("UPDATE_UNIVERSE_TOP_N", "20"))
|
|
universe_min_score = float(os.environ.get("UPDATE_UNIVERSE_MIN_SCORE", "4.0"))
|
|
universe_by_slot = se.build_universe_simulation(
|
|
codes_candles,
|
|
top_n=universe_top_n,
|
|
min_score=universe_min_score,
|
|
scan_interval_min=fallback_sim_interval,
|
|
)
|
|
base_fixed["scan_interval_min"] = fallback_sim_interval
|
|
logger.info("📌 유니버스: 시뮬 fallback (%d분)", fallback_sim_interval)
|
|
else:
|
|
base_fixed["scan_interval_min"] = 1
|
|
|
|
cache_holder: Dict[str, Any] = {}
|
|
attach_indicator_caches_to_params(cache_holder, codes_candles)
|
|
|
|
return BreakoutSearchContext(
|
|
start=start,
|
|
end=end,
|
|
mode=mode,
|
|
base_fixed=base_fixed,
|
|
codes_candles=codes_candles,
|
|
universe_by_slot=universe_by_slot,
|
|
ticks_by_code=ticks_by_code,
|
|
orderbook_by_code=orderbook_by_code,
|
|
program_by_code=program_by_code,
|
|
log_verdict_by_code=log_verdict_by_code,
|
|
share_denom_by_code=share_denom_by_code,
|
|
fee_rate=fee_rate,
|
|
sell_tax=sell_tax,
|
|
slot_money=slot_money_f,
|
|
max_stocks=max_stocks_i,
|
|
total_budget_krw=total_budget_f,
|
|
period_days=period_days,
|
|
portfolio=portfolio,
|
|
grid_keys=breakout_grid_axis_keys(mode),
|
|
start_key=start_key,
|
|
end_key=end_key,
|
|
cache_holder=cache_holder,
|
|
)
|
|
|
|
|
|
def _make_sampler(name: str, seed: Optional[int]):
|
|
n = (name or "tpe").strip().lower()
|
|
if n == "random":
|
|
return RandomSampler(seed=seed)
|
|
return TPESampler(seed=seed, multivariate=True)
|
|
|
|
|
|
def run_breakout_optuna(
|
|
ctx: BreakoutSearchContext,
|
|
*,
|
|
n_trials: int,
|
|
storage_url: str,
|
|
study_name: str,
|
|
min_trades: int,
|
|
min_win_rate: float,
|
|
min_pf: float,
|
|
sort_by: str = "pnl",
|
|
sampler_name: str = "tpe",
|
|
seed: Optional[int] = None,
|
|
n_jobs: int = 1,
|
|
show_progress: bool = True,
|
|
) -> optuna.Study:
|
|
study = optuna.create_study(
|
|
study_name=study_name,
|
|
storage=storage_url,
|
|
load_if_exists=True,
|
|
direction="maximize",
|
|
sampler=_make_sampler(sampler_name, seed),
|
|
)
|
|
|
|
def objective(trial: optuna.Trial) -> float:
|
|
combo = suggest_breakout_params(trial, ctx.mode)
|
|
result = evaluate_breakout_param_combo(
|
|
combo,
|
|
base_fixed=ctx.base_fixed,
|
|
grid_keys=ctx.grid_keys,
|
|
codes_candles=ctx.codes_candles,
|
|
min_trades=min_trades,
|
|
min_win_rate=min_win_rate,
|
|
min_pf=min_pf,
|
|
universe_by_slot=ctx.universe_by_slot,
|
|
slot_money=ctx.slot_money,
|
|
max_stocks=ctx.max_stocks,
|
|
total_budget_krw=ctx.total_budget_krw,
|
|
fee_rate=ctx.fee_rate,
|
|
sell_tax=ctx.sell_tax,
|
|
period_days=ctx.period_days,
|
|
cache_holder=ctx.cache_holder,
|
|
ticks_by_code=ctx.ticks_by_code,
|
|
orderbook_by_code=ctx.orderbook_by_code,
|
|
program_by_code=ctx.program_by_code,
|
|
log_verdict_by_code=ctx.log_verdict_by_code,
|
|
share_denom_by_code=ctx.share_denom_by_code,
|
|
)
|
|
if result is None:
|
|
trial.set_user_attr("gates_ok", False)
|
|
return _FAIL_OBJECTIVE
|
|
obj = float(result["win_rate"]) if sort_by == "win_rate" else float(result["total_pnl"])
|
|
trial.set_user_attr("gates_ok", True)
|
|
trial.set_user_attr("total_pnl", float(result["total_pnl"]))
|
|
trial.set_user_attr("win_rate", float(result["win_rate"]))
|
|
trial.set_user_attr("pf", float(result.get("pf") or 0))
|
|
trial.set_user_attr("total_trades", int(result["total_trades"]))
|
|
trial.set_user_attr("merged_json", json.dumps(result.get("merged_params") or {}, ensure_ascii=False))
|
|
return obj
|
|
|
|
logger.info("🔬 Optuna BREAKOUT | study=%s | trials=%d", study_name, n_trials)
|
|
t0 = time.time()
|
|
study.optimize(objective, n_trials=n_trials, n_jobs=n_jobs, show_progress_bar=show_progress)
|
|
elapsed = time.time() - t0
|
|
|
|
passing: List[Dict[str, Any]] = []
|
|
for trial in study.trials:
|
|
if trial.state != optuna.trial.TrialState.COMPLETE:
|
|
continue
|
|
if not trial.user_attrs.get("gates_ok"):
|
|
continue
|
|
merged_raw = trial.user_attrs.get("merged_json") or "{}"
|
|
try:
|
|
merged = json.loads(merged_raw)
|
|
except json.JSONDecodeError:
|
|
merged = dict(trial.params)
|
|
passing.append({
|
|
"params": dict(trial.params),
|
|
"merged_params": merged,
|
|
"total_trades": int(trial.user_attrs.get("total_trades") or 0),
|
|
"win_rate": float(trial.user_attrs.get("win_rate") or 0),
|
|
"total_pnl": float(trial.user_attrs.get("total_pnl") or 0),
|
|
"pf": float(trial.user_attrs.get("pf") or 0),
|
|
"optuna_trial_number": trial.number,
|
|
})
|
|
|
|
if sort_by == "win_rate":
|
|
passing.sort(key=lambda r: (-r["win_rate"], -r["total_pnl"]))
|
|
else:
|
|
passing.sort(key=lambda r: (-r["total_pnl"], -r["win_rate"]))
|
|
profitable = [r for r in passing if r["total_pnl"] > 0]
|
|
if profitable:
|
|
passing = profitable
|
|
|
|
hints: Dict[str, str] = {}
|
|
out_data = {
|
|
"engine": "optuna",
|
|
"strategy": "breakout",
|
|
"mode": ctx.mode,
|
|
"start": ctx.start,
|
|
"end": ctx.end,
|
|
"slot_money": int(ctx.slot_money),
|
|
"max_stocks": ctx.max_stocks,
|
|
"total_budget_krw": int(ctx.total_budget_krw),
|
|
"backtest_days": ctx.period_days,
|
|
"min_trades": min_trades,
|
|
"min_win_rate": min_win_rate,
|
|
"min_pf": min_pf,
|
|
"sort_by": sort_by,
|
|
"grid_keys": ctx.grid_keys,
|
|
"grid_axis_hints": {k: hints[k] for k in ctx.grid_keys if k in hints},
|
|
"optuna_study_name": study_name,
|
|
"optuna_storage": storage_url,
|
|
"optuna_n_trials_requested": n_trials,
|
|
"optuna_trials_completed": len(study.trials),
|
|
"optuna_best_value": study.best_value if study.best_trial else None,
|
|
"optuna_best_trial_number": study.best_trial.number if study.best_trial else None,
|
|
"elapsed_sec": round(elapsed, 1),
|
|
"results": passing[:5000],
|
|
}
|
|
ts = datetime.now().strftime("%Y%m%d_%H%M%S")
|
|
out_path = os.path.join(_results_dir_for_write(), f"optuna_breakout_{ctx.mode}_{ts}.json")
|
|
with open(out_path, "w", encoding="utf-8") as f:
|
|
json.dump(out_data, f, indent=2, ensure_ascii=False)
|
|
logger.info("💾 Optuna 결과 저장: %s", out_path)
|
|
study._kis_export_path = out_path # type: ignore[attr-defined]
|
|
return study
|
|
|
|
|
|
def apply_best_breakout_trial(study: optuna.Study) -> bool:
|
|
if not study.best_trial or study.best_value <= _FAIL_OBJECTIVE + 1:
|
|
logger.warning("⚠️ 적용할 best trial 없음")
|
|
return False
|
|
pnl = float(study.best_trial.user_attrs.get("total_pnl") or 0)
|
|
if pnl <= 0:
|
|
logger.warning("⚠️ Best trial 총손익 ≤ 0 — DB 미적용")
|
|
return False
|
|
merged = json.loads(study.best_trial.user_attrs.get("merged_json") or "{}")
|
|
apply_params_to_db(merged)
|
|
logger.info("🚀 [Optuna apply-best] breakout trial #%d → env_config", study.best_trial.number)
|
|
return True
|