#!/usr/bin/env python3 """ kis_trader/backtest/optuna_search_space.py — Optuna 탐색 공간 (Grid 축 재사용) ============================================================================== 각 전략 Grid 와 동일한 이산 축을 trial.suggest_categorical 로 샘플링. """ from __future__ import annotations from typing import Any, Dict, List import optuna from kis_trader.backtest.param_search_breakout import _breakout_grids from kis_trader.backtest.param_search_momentum import ( _momentum_combo_grid_valid, _momentum_grids, ) from kis_trader.backtest.param_search_scalping import _scalp_grids from kis_trader.backtest.tail_param_search import _tail_grids def _dedupe_preserve_order(values: List[Any]) -> List[Any]: seen = set() out: List[Any] = [] for v in values: key = v if isinstance(v, (int, float, str, bool)) else repr(v) if key in seen: continue seen.add(key) out.append(v) return out def _suggest_from_grid(trial: optuna.Trial, grid: Dict[str, List[Any]]) -> Dict[str, Any]: from kis_trader.backtest.optuna_grid_narrow import load_narrow_grid_override narrow = load_narrow_grid_override() combo: Dict[str, Any] = {} for key, values in grid.items(): if not values: continue src = narrow.get(key) if narrow.get(key) else values choices = _dedupe_preserve_order(list(src)) if not choices: continue combo[key] = trial.suggest_categorical(key, choices) return combo def suggest_scalp_params(trial: optuna.Trial, mode: str) -> Dict[str, Any]: combo = _suggest_from_grid(trial, _scalp_grids()[mode]) combo["use_macd_cross"] = False return combo def suggest_tail_params(trial: optuna.Trial, mode: str) -> Dict[str, Any]: return _suggest_from_grid(trial, _tail_grids(mode)) def suggest_momentum_params( trial: optuna.Trial, mode: str, *, market: str = "KR", ) -> Dict[str, Any]: mk = (market or "KR").strip().upper() or "KR" combo = _suggest_from_grid(trial, _momentum_grids(market=mk)[mode]) if not _momentum_combo_grid_valid(combo): raise optuna.TrialPruned("momentum invalid combo") return combo def suggest_breakout_params(trial: optuna.Trial, mode: str) -> Dict[str, Any]: combo = _suggest_from_grid(trial, _breakout_grids()[mode]) if "prev_chg_min" in combo and "prev_chg_max" in combo: if float(combo["prev_chg_min"]) >= float(combo["prev_chg_max"]): raise optuna.TrialPruned("breakout prev_chg invalid") return combo def scalp_grid_axis_keys(mode: str) -> List[str]: return list(_scalp_grids()[mode].keys()) def tail_grid_axis_keys(mode: str) -> List[str]: return list(_tail_grids(mode).keys()) def momentum_grid_axis_keys(mode: str, *, market: str = "KR") -> List[str]: mk = (market or "KR").strip().upper() or "KR" return list(_momentum_grids(market=mk)[mode].keys()) def breakout_grid_axis_keys(mode: str) -> List[str]: return list(_breakout_grids()[mode].keys())