Files
kis_trader/scripts/momentum_optuna_trade_diff.py

211 lines
6.6 KiB
Python

#!/usr/bin/env python3
"""#199 Optuna 기록 daily_pnl vs 재시뮬 체결 diff (DB 미변경)."""
from __future__ import annotations
import argparse
import json
import time
import traceback
from collections import defaultdict
from pathlib import Path
def _trade_key(t: dict) -> str:
code = str(t.get("code") or t.get("ticker") or "")
buy = str(t.get("buy_time") or t.get("entry_time") or t.get("entry_ts") or "")
sell = str(t.get("sell_time") or t.get("exit_time") or t.get("exit_ts") or "")
return f"{code}|{buy}|{sell}"
def _day_of(t: dict) -> str:
for k in ("sell_time", "exit_time", "buy_time", "entry_time"):
v = str(t.get(k) or "")
if len(v) >= 8 and v[:8].isdigit():
return f"{v[:4]}-{v[4:6]}-{v[6:8]}"
return "?"
def main() -> int:
ap = argparse.ArgumentParser()
ap.add_argument(
"--json",
default="kis_trader/backtest/results/optuna_momentum_tpe_20260821_220227.json",
)
ap.add_argument("--out", required=True)
ap.add_argument("--trial", type=int, default=199)
args = ap.parse_args()
t0 = time.time()
d = json.loads(Path(args.json).read_text())
trial = next(
x for x in d["results_all"] if x.get("optuna_trial_number") == args.trial
)
grid_keys = list(d["grid_keys"])
params = dict(trial["params"])
mp = dict(trial["merged_params"])
fixed = {k: v for k, v in mp.items() if k not in params}
fixed["_orderbook_filter_enabled"] = False
from kis_trader.backtest.optuna_momentum import prepare_momentum_search_context
from kis_trader.backtest.param_search_momentum import evaluate_momentum_param_combo
print(
f"prepare {d['start']}~{d['end']} trial=#{args.trial} OB=off include_trades",
flush=True,
)
ctx = prepare_momentum_search_context(
d["start"],
d["end"],
"tpe",
slot_money=float(d["slot_money"]),
max_stocks=int(d["max_stocks"]),
total_budget_krw=float(d["total_budget_krw"]),
orderbook_filter="off",
market="KR",
history_source="kiwoom",
)
if ctx is None:
print("prepare failed", flush=True)
return 1
base = dict(ctx.base_fixed)
base.update(fixed)
base["_orderbook_filter_enabled"] = False
r = evaluate_momentum_param_combo(
params,
base_fixed=base,
grid_keys=grid_keys,
codes_candles=ctx.codes_candles,
min_trades=1,
min_win_rate=0.0,
min_pf=0.0,
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,
include_trades=True,
)
if not r:
print("evaluate None", flush=True)
return 1
trades = list(r.get("_trades") or [])
daily_reeval: dict = defaultdict(float)
by_code: dict = defaultdict(lambda: {"n": 0, "pnl": 0.0})
slim = []
for t in trades:
pnl = float(t.get("pnl") or 0)
day = _day_of(t)
daily_reeval[day] += pnl
code = str(t.get("code") or "")
by_code[code]["n"] += 1
by_code[code]["pnl"] += pnl
slim.append(
{
"key": _trade_key(t),
"code": code,
"day": day,
"pnl": pnl,
"buy": t.get("buy_time") or t.get("entry_time"),
"sell": t.get("sell_time") or t.get("exit_time"),
"reason": t.get("sell_reason") or t.get("reason") or t.get("exit_reason"),
}
)
recorded_daily = dict(trial.get("daily_pnl") or {})
days = sorted(set(recorded_daily) | set(daily_reeval))
daily_diff = []
for day in days:
a = float(recorded_daily.get(day) or 0)
b = float(daily_reeval.get(day) or 0)
daily_diff.append(
{
"day": day,
"optuna_recorded": a,
"reeval": b,
"delta": b - a,
}
)
# 코드별 상위 |pnl|
code_rows = sorted(
(
{"code": c, "n": v["n"], "pnl": round(v["pnl"], 1)}
for c, v in by_code.items()
),
key=lambda x: abs(x["pnl"]),
reverse=True,
)[:25]
report = {
"db_touched": False,
"trial": args.trial,
"note": (
"Optuna JSON에 체결원본 없음 → 기록 daily_pnl vs 재시뮬 체결 집계 diff. "
"체결 키 목록은 재시뮬만."
),
"recorded": {
"total_pnl": trial["total_pnl"],
"total_trades": trial["total_trades"],
"mdd": trial.get("mdd"),
"win_rate": trial.get("win_rate"),
"daily_pnl": recorded_daily,
},
"reeval": {
"total_pnl": r.get("total_pnl"),
"total_trades": r.get("total_trades"),
"mdd": r.get("mdd"),
"win_rate": r.get("win_rate"),
"pf": r.get("pf"),
"daily_pnl": {k: round(v, 1) for k, v in sorted(daily_reeval.items())},
"n_trades_list": len(trades),
},
"delta_total_pnl": float(r.get("total_pnl") or 0) - float(trial["total_pnl"]),
"delta_trades": int(r.get("total_trades") or 0) - int(trial["total_trades"]),
"daily_diff": daily_diff,
"reeval_top_codes": code_rows,
"reeval_trades": slim,
"elapsed_sec": round(time.time() - t0, 1),
}
out = Path(args.out)
out.parent.mkdir(parents=True, exist_ok=True)
out.write_text(json.dumps(report, ensure_ascii=False, indent=2))
print("WROTE", out, flush=True)
print("DAILY_DIFF", json.dumps(daily_diff, ensure_ascii=False), flush=True)
print(
"SUMMARY",
json.dumps(
{
"recorded_pnl": trial["total_pnl"],
"reeval_pnl": r.get("total_pnl"),
"delta_pnl": report["delta_total_pnl"],
"recorded_tr": trial["total_trades"],
"reeval_tr": r.get("total_trades"),
"mdd_rec": trial.get("mdd"),
"mdd_reeval": r.get("mdd"),
},
ensure_ascii=False,
),
flush=True,
)
return 0
if __name__ == "__main__":
try:
raise SystemExit(main())
except Exception:
traceback.print_exc()
raise SystemExit(1)