Files
kis_bot/kis_trader/backtest/optuna_search_space.py
2026-07-30 18:05:07 +09:00

90 lines
2.8 KiB
Python

#!/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]:
combo: Dict[str, Any] = {}
for key, values in grid.items():
if not values:
continue
choices = _dedupe_preserve_order(list(values))
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())