#!/usr/bin/env python3
"""
Эксперимент: SBER с SL=6×ATR, TP=12×ATR (RR 1:2, x2 шире).

Запуск:
  python3 experiment_rr_x2.py
"""
import sys, os, time, warnings
import numpy as np
warnings.filterwarnings('ignore')
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))

import config as cfg_module
cfg_module.set_device('train')  # CUDA for training

TARGET_CFG = {
    'atr_mult_sl': 6.0,
    'atr_mult_tp': 12.0,
    'max_bars': 300,
}
TARGET_KEYS = list(TARGET_CFG.keys())

old_cfg = cfg_module.TARGET_CONFIG.copy()
cfg_module.TARGET_CONFIG.clear()
cfg_module.TARGET_CONFIG.update(TARGET_CFG)

from config import SAVE_DIR, TICKER_SHORT_WEIGHTS, DEFAULT_SHORT_WEIGHT
from models.moe import MultiTimeframeMoE, _prepare_df, _get_feature_cols, _build_dataset, _flat_mask
from models.experts import ExpertEnsemble
import xgboost as xgb
import joblib

TICKER = 'SBER'
suffix = '_rr1x2_x2'
t0 = time.time()

print(f'\n{"█"*60}')
print(f'  ЭКСПЕРИМЕНТ: {TICKER} SL=6×ATR, TP=12×ATR (RR 1:2)')
print(f'{"█"*60}\n')

# ── OOS фаза ──
moe = MultiTimeframeMoE(TICKER)
ok = moe.train(limit=None, verbose=True, retrain_on_full=False)
if not ok:
    print('train() failed')
    sys.exit(1)

oos_val_acc = getattr(moe, 'oos_val_acc', 0)
oos_long_th = moe.optimal_thresholds.get('long', 0.55)
oos_short_th = moe.optimal_thresholds.get('short', 0.55)
oos_selected = moe.selected_experts
print(f'\n  OOS: val_acc={oos_val_acc:.1%}, L≥{oos_long_th:.2f}, S≥{oos_short_th:.2f}')

# ── Retrain 100% ──
print(f'\n  Retrain на 100%...')
df_full = _prepare_df(TICKER)
if df_full is None or len(df_full) < 200:
    print('no data')
    sys.exit(1)

all_features = _get_feature_cols()
available = [c for c in all_features if c in df_full.columns]
norm_stats = moe.ensemble.get_norm_stats()

ensemble_full = ExpertEnsemble(n_features=len(available))
ensemble_full.train_all(df_full, available, verbose=True, epochs=None)
ensemble_full.set_norm_stats(*norm_stats)

signals_full = ensemble_full.predict_all(df_full, available)
X_full = _build_dataset(df_full, signals_full, ensemble_full, available, keep_experts=oos_selected)
y_full = np.column_stack([df_full['outcome_long'].values, df_full['outcome_short'].values])
min_len = min(len(X_full), len(y_full))
X_full, y_full = X_full[-min_len:], y_full[-min_len:]

flat_mask = _flat_mask(df_full).values[-len(X_full):]
n_flat = int(flat_mask.sum())
y_full[flat_mask] = 0.0

long_sr = df_full['outcome_long'].values[-len(X_full):].mean()
short_sr = df_full['outcome_short'].values[-len(X_full):].mean()
print(f'  Флет: {n_flat}/{len(flat_mask)} ({n_flat/len(flat_mask):.1%}), '
      f'SR: long={long_sr:.1%}, short={short_sr:.1%}')

params_xgb = {
    'max_depth': 4, 'learning_rate': 0.03, 'n_estimators': 300,
    'subsample': 0.7, 'colsample_bytree': 0.7,
    'reg_alpha': 0.1, 'reg_lambda': 1.0, 'min_child_weight': 5,
    'random_state': 42, 'verbosity': 0, 'n_jobs': -1,
}

pl = params_xgb.copy()
pl['scale_pos_weight'] = min(max((len(y_full)-y_full[:,0].sum())/max(y_full[:,0].sum(),1), 1.0), 5.0)
xgb_long = xgb.XGBClassifier(**pl).fit(X_full, y_full[:, 0])

ps = params_xgb.copy()
pos_s = y_full[:, 1].sum()
neg_s = len(y_full) - pos_s
scale_s = min(max(neg_s / max(pos_s, 1), 1.0), 5.0)
scale_s *= TICKER_SHORT_WEIGHTS.get(TICKER, DEFAULT_SHORT_WEIGHT)
ps['scale_pos_weight'] = scale_s
xgb_short = xgb.XGBClassifier(**ps).fit(X_full, y_full[:, 1])

moe.ensemble = ensemble_full
moe.rf_long = xgb_long
moe.rf_short = xgb_short
moe.rf_model = None
moe.optimal_thresholds = {'long': oos_long_th, 'short': oos_short_th}
moe.selected_experts = oos_selected
moe.oos_val_acc = oos_val_acc
moe.val_acc = oos_val_acc
moe.feature_cols = available
moe.n_base_features = len(available)

save_path = os.path.join(SAVE_DIR, f'{TICKER.lower()}_moe_v12{suffix}.joblib')
moe.save(save_path)
print(f'\n  ✓ Сохранено: {save_path} ({time.time()-t0:.0f}с)')

# ── БЭКТЕСТ ──
print(f'\n  ── БЭКТЕСТ ──')
df = _prepare_df(TICKER)
result = moe.predict_proba_aligned(df)
n_res = len(result)
close_vals = df['Close'].values[-n_res:]
atr_vals = df['atr'].values[-n_res:]
high_vals = df['High'].values[-n_res:]
low_vals = df['Low'].values[-n_res:]

sl_mult, tp_mult, max_bars = 6.0, 12.0, 300
trades = []
in_trade = False
entry_price = 0.0
entry_idx = 0
direction = ''
sl_price = 0.0
tp_price = 0.0

for i in range(n_res):
    signal = result.iloc[i]['signal']
    close = close_vals[i]
    atr = atr_vals[i]

    if in_trade:
        bars_held = i - entry_idx
        high = high_vals[i]
        low = low_vals[i]
        hit = None
        exit_price = close
        if direction == 'LONG':
            if high >= tp_price: hit, exit_price = 'TP', tp_price
            elif low <= sl_price: hit, exit_price = 'SL', sl_price
        else:
            if low <= tp_price: hit, exit_price = 'TP', tp_price
            elif high >= sl_price: hit, exit_price = 'SL', sl_price

        if hit is None and bars_held >= max_bars:
            hit, exit_price = 'TIME_STOP', close

        if hit:
            pnl = ((exit_price - entry_price) / entry_price) if direction == 'LONG' else ((entry_price - exit_price) / entry_price)
            trades.append({'dir': direction, 'pnl_pct': pnl*100, 'reason': hit, 'bars': bars_held})
            in_trade = False

    if not in_trade and signal in ('BUY', 'SELL'):
        direction = 'LONG' if signal == 'BUY' else 'SHORT'
        entry_price = close
        entry_idx = i
        in_trade = True
        sl_dist = atr * sl_mult
        tp_dist = atr * tp_mult
        sl_price = close - sl_dist if direction == 'LONG' else close + sl_dist
        tp_price = close + tp_dist if direction == 'LONG' else close - tp_dist

total = len(trades)
wins = sum(1 for t in trades if t['pnl_pct'] > 0)
wr = wins/max(total,1)*100
total_ret = sum(t['pnl_pct'] for t in trades)
avg_w = np.mean([t['pnl_pct'] for t in trades if t['pnl_pct']>0]) if wins else 0
avg_l = np.mean([t['pnl_pct'] for t in trades if t['pnl_pct']<0]) if total-wins else 0
pf = abs(sum(t['pnl_pct'] for t in trades if t['pnl_pct']>0) / max(abs(sum(t['pnl_pct'] for t in trades if t['pnl_pct']<0)), 0.01))
returns = np.array([t['pnl_pct'] for t in trades])
sharpe = returns.mean()/max(returns.std(),0.01)*np.sqrt(365*24/np.mean([t['bars'] for t in trades])) if total>1 else 0
tp_c = sum(1 for t in trades if t['reason']=='TP')
sl_c = sum(1 for t in trades if t['reason']=='SL')
ts_c = sum(1 for t in trades if t['reason']=='TIME_STOP')

print(f'\n  {"="*50}')
print(f'  РЕЗУЛЬТАТЫ: {TICKER} SL=6×ATR TP=12×ATR')
print(f'  {"="*50}')
print(f'  OOS val_acc:     {oos_val_acc:.1%}')
print(f'  Success Rate L:  {long_sr:.1%}  (breakeven: 33.3%)')
print(f'  Success Rate S:  {short_sr:.1%}  (breakeven: 33.3%)')
print(f'  Сделок всего:    {total}')
print(f'  WinRate:         {wr:.1f}% ({wins}/{total})')
print(f'  TP / SL / TS:    {tp_c} / {sl_c} / {ts_c}')
print(f'  Сумм доходность: {total_ret:+.2f}%')
print(f'  Средняя прибыль: {avg_w:.2f}%')
print(f'  Средний убыток:  {avg_l:.2f}%')
print(f'  Profit Factor:   {pf:.2f}')
print(f'  Sharpe:          {sharpe:.2f}')
print(f'  Время:           {(time.time()-t0)/60:.1f} мин')
print(f'  {"="*50}')

# ── СРАВНЕНИЕ с RR 1:2 (3/6) ──
print(f'\n  СРАВНЕНИЕ С RR 1:2 (SL=3, TP=6):')
print(f'  {"Параметр":<25} {"3×/6×":<18} {"6×/12×":<18}')
print(f'  {"-"*25} {"-"*18} {"-"*18}')
for param, val_36, val_612 in [
    ('val_acc', '55.17%', f'{oos_val_acc:.1%}'),
    ('Success Rate L', '34.9%', f'{long_sr:.1%}'),
    ('Success Rate S', '31.8%', f'{short_sr:.1%}'),
    ('Сделок', '285', str(total)),
    ('WinRate', '56.1%', f'{wr:.1f}%'),
    ('TP', '157', str(tp_c)),
    ('SL', '123', str(sl_c)),
    ('Доходность', '+371.86%', f'{total_ret:+.2f}%'),
    ('Profit Factor', '2.69', f'{pf:.2f}'),
    ('Sharpe', '4.87', f'{sharpe:.2f}'),
]:
    print(f'  {param:<25} {val_36:<18} {val_612:<18}')

# Восстанавливаем конфиг
cfg_module.TARGET_CONFIG.clear()
cfg_module.TARGET_CONFIG.update(old_cfg)
print()
