#!/usr/bin/env python3 """kis_trader/backtest/optuna_momentum.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 momentum_backtest_common as mbc from kis_trader.backtest import scalping_backtest_common as sbc from kis_trader.backtest.optuna_search_space import momentum_grid_axis_keys, suggest_momentum_params from kis_trader.backtest.param_search_cli_common import ( apply_session_to_fixed, combo_passes_search_filters, format_session_hm, ) from kis_trader.backtest.param_search_momentum import ( MOMENTUM_GRID_AXIS_HINTS_KO, _load_candles_for_search, _mom_fixed_defaults, _momentum_grids, apply_params_to_db, evaluate_momentum_param_combo, ) from kis_trader.backtest.tail_param_search import _results_dir_for_write from kis_trader.engine import momentum_engine as me from kis_trader.engine.indicator_cache import attach_indicator_caches_to_params from kis_trader.utils.env import get_env_bool, get_env_float logger = logging.getLogger("param_search_optuna") _FAIL_OBJECTIVE = -1e18 @dataclass class MomentumSearchContext: 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] 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) shared_tick_store: Any = None # ws_ticks 공유메모리 핸들 (종료 시 unlink) def prepare_momentum_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[MomentumSearchContext]: grids = _momentum_grids() if mode not in grids: logger.error("❌ 모멘텀 mode: %s (fast/rr/coarse/fine/full)", mode) return None base_fixed = _mom_fixed_defaults() apply_session_to_fixed(base_fixed, time_start_hm=time_start_hm, time_end_hm=time_end_hm) from kis_trader.engine.momentum_tick_replay import ( momentum_backtest_use_tick_entry as _te, momentum_backtest_use_tick_exit as _tx, ) base_fixed["backtest_use_tick_entry"] = _te(None) base_fixed["backtest_use_tick_exit"] = _tx(None) _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 = sbc.fee_and_slot_from_env(env_row, strategy="MOMENTUM") portfolio = sbc.resolve_scalp_portfolio_params( env_row, None, strategy="MOMENTUM", 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)}" ) codes_candles = _load_candles_for_search(start, end, base_fixed.get("rsi_period", 3)) if not codes_candles: logger.error("❌ 캔들 데이터 없음") return None logger.info("✅ 데이터 로드: %s종목", len(codes_candles)) 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 "" universe_by_slot = None fallback_sim_interval = 5 if not use_fallback_universe and start_ymd and end_ymd: try: from kis_trader.backtest.momentum_backtest_common import resolve_momentum_universe history, src, n_bins, _scan_iv, timing = resolve_momentum_universe( start_ymd, end_ymd, use_saved_history=True, strategy_id="MOMENTUM", ) if history: universe_by_slot = history avg = sum(len(v) for v in history.values()) / max(1, n_bins) logger.info( "✅ 유니버스: MOMENTUM 이력 | %s분봉 · 평균 %.1f종목", n_bins, avg, ) except Exception as exc: logger.debug("유니버스 이력 스킵: %s", exc) if universe_by_slot is None: 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 = me.build_universe_simulation_momentum( 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 grid = grids[mode] _ob_axes = ("max_spread_pct", "min_bid_ask_ratio", "ask_max_mult") _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] = {} ticks_by_code: Dict[str, Any] = {} _snap_db = TradeDB() try: from kis_trader.backtest.trigger_snapshot_loader import ( backtest_needs_trigger_snapshot_load, load_trigger_snapshots_by_code, ) from kis_trader.engine.momentum_tick_replay import ( momentum_backtest_use_tick_entry, momentum_backtest_use_tick_exit, ) if backtest_needs_trigger_snapshot_load(base_fixed, strategy="MOMENTUM"): 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=base_fixed, strategy="MOMENTUM", ) log_verdict_by_code = trigger_snap_meta.get("log_verdict_by_code") or {} if momentum_backtest_use_tick_exit(base_fixed) or momentum_backtest_use_tick_entry(base_fixed): from kis_trader.backtest.momentum_tick_loader import load_momentum_ticks_by_code ticks_by_code, tick_rows = load_momentum_ticks_by_code( _snap_db, start_key, end_key, set(codes_candles.keys()), ) logger.info("✅ ws_ticks %s건", f"{tick_rows:,}") finally: _snap_db.close() # ── ws_ticks 공유메모리 (Optuna, opt-in) — dict→numpy 컬럼 shared_memory 로 RAM 절감 ── # 끄려면 OPTUNA_PARAM_SEARCH_SHARED_TICKS=0. numpy/shm 미지원·빌드 실패 시 자동 폴백. shared_tick_store = None if get_env_bool("OPTUNA_PARAM_SEARCH_SHARED_TICKS", True) and ticks_by_code: from kis_trader.backtest.shared_ticks import build_shared_ticks_view _view, shared_tick_store = build_shared_ticks_view(ticks_by_code, enabled=True) if shared_tick_store is not None: import atexit as _atexit _atexit.register(shared_tick_store.unlink) # 크래시 시 /dev/shm 누수 방지 logger.info("📦 ws_ticks 공유메모리 ON (Optuna) — dict 사본 제거, RAM 절감") ticks_by_code = _view import gc as _gc _gc.collect() try: import ctypes as _ctypes _ctypes.CDLL("libc.so.6").malloc_trim(0) except Exception: pass cache_holder: Dict[str, Any] = {} attach_indicator_caches_to_params(cache_holder, codes_candles) return MomentumSearchContext( 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, 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=momentum_grid_axis_keys(mode), start_key=start_key, end_key=end_key, cache_holder=cache_holder, shared_tick_store=shared_tick_store, ) 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 _momentum_objective_value(result: Dict[str, Any], sort_by: str) -> float: pnl = float(result["total_pnl"]) if sort_by == "score": mdd_floor = get_env_float("MOMENTUM_SCORE_MDD_FLOOR", 10000.0) mdd = float(result.get("mdd") or 0) return pnl / max(mdd, mdd_floor) if sort_by == "win_rate": return float(result["win_rate"]) return pnl def run_momentum_optuna( ctx: MomentumSearchContext, *, n_trials: int, storage_url: str, study_name: str, min_trades: int, min_win_rate: float, min_pf: float, sort_by: str = "score", 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_momentum_params(trial, ctx.mode) result = evaluate_momentum_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, start_key=ctx.start_key, end_key=ctx.end_key, ) if result is None: trial.set_user_attr("gates_ok", False) return _FAIL_OBJECTIVE obj = _momentum_objective_value(result, sort_by) 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("mdd", float(result.get("mdd") or 0)) trial.set_user_attr("score", float(obj if sort_by == "score" else _momentum_objective_value(result, "score"))) 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 float(obj) logger.info( "🔬 Optuna MOMENTUM | study=%s | trials=%d | sort=%s", study_name, n_trials, sort_by, ) t0 = time.time() try: study.optimize(objective, n_trials=n_trials, n_jobs=n_jobs, show_progress_bar=show_progress) finally: # 탐색 종료(또는 예외) 시 공유메모리 즉시 해제 (atexit 는 크래시 대비 이중 안전장치). _store = getattr(ctx, "shared_tick_store", None) if _store is not None: try: _store.unlink() except Exception: pass ctx.shared_tick_store = None 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) row = { "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), "mdd": float(trial.user_attrs.get("mdd") or 0), "score": float(trial.user_attrs.get("score") or 0), "optuna_trial_number": trial.number, } passing.append(row) if sort_by == "score": passing.sort(key=lambda r: (-r["score"], -r["total_pnl"], -r["win_rate"])) elif 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 out_data = { "engine": "optuna", "strategy": "momentum", "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: MOMENTUM_GRID_AXIS_HINTS_KO[k] for k in ctx.grid_keys if k in MOMENTUM_GRID_AXIS_HINTS_KO}, "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_momentum_{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_momentum_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_raw = study.best_trial.user_attrs.get("merged_json") or "{}" merged = json.loads(merged_raw) apply_params_to_db(merged) logger.info("🚀 [Optuna apply-best] momentum trial #%d → env_config", study.best_trial.number) return True