#!/usr/bin/env python3
"""
Пакетная тренировка для всех MOEX/Crypto/Forex тикеров с per-market RR конфигами.

Запуск:
  python3 train_all_rr1x2.py                               # все 14 MOEX тикеров
  python3 train_all_rr1x2.py --ticker SBER                  # один MOEX тикер
  python3 train_all_rr1x2.py --ticker BITCOIN               # один Crypto тикер
  python3 train_all_rr1x2.py --ticker EURUSD                # один Forex тикер
  python3 train_all_rr1x2.py --ticker ALL                   # все 17 тикеров (MOEX + Crypto + Forex)

Per-market target config:
  MOEX:   SL=3×ATR, TP=6×ATR (RR 1:2)
  Crypto: SL=4×ATR, TP=8×ATR (RR 1:2)
  Forex:  SL=2×ATR, TP=4×ATR (RR 1:2)

Результат: models/saved/{ticker}_moe_v12_rr1x2.joblib
"""
import sys, os, time, json, warnings
import numpy as np
from datetime import datetime

warnings.filterwarnings('ignore')
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))

from config import MOEX_TICKERS, CRYPTO_TICKERS, FOREX_TICKERS, MONITOR_TICKERS, SAVE_DIR, \
    TICKER_SHORT_WEIGHTS, TICKER_LONG_WEIGHTS, DEFAULT_SHORT_WEIGHT, DEFAULT_LONG_WEIGHT, \
    get_target_config, get_market, MARKET_TARGET_CONFIGS
from models.moe import MultiTimeframeMoE, _prepare_df, _get_feature_cols, _build_dataset, _flat_mask
from models.experts import ExpertEnsemble
from models.crypto_moe import CryptoMoE, _crypto_prepare_df, _crypto_feature_cols, \
    _crypto_build_dataset, CryptoExpertEnsemble, CRYPTO_EXPERT_NAMES, CRYPTO_HORIZONS
import xgboost as xgb

LOG_FILE = f'{SAVE_DIR}/training_rr1x2_log.json'

ALL_TICKERS = MOEX_TICKERS + CRYPTO_TICKERS + FOREX_TICKERS  # 17 шт


def train_moe_for(ticker: str) -> dict | None:
    """Тренирует MoE с per-market RR конфигом для одного тикера.

    Для крипты (BITCOIN, BITCOINC) использует CryptoMoE с кастомными экспертами,
    горизонтами [2,4,8,16] и окном 40 свечей.
    Для MOEX/Forex — стандартный MultiTimeframeMoE.
    """
    t0 = time.time()
    suffix = '_rr1x2'
    market = get_market(ticker)
    target_cfg = get_target_config(ticker)
    is_crypto = market == 'crypto'
    market_label = 'CRYPTO' if is_crypto else market.upper()

    print(f'\n[{ticker}] {"=" * 55}')
    print(f'  Рынок: {market_label} | SL={target_cfg["atr_mult_sl"]}×ATR TP={target_cfg["atr_mult_tp"]}×ATR')

    # ── OOS фаза (train 85% + val 15%) ──
    if is_crypto:
        moe = CryptoMoE(ticker)
        prepare_df_fn = _crypto_prepare_df
        feature_cols_fn = _crypto_feature_cols
        build_dataset_fn = _crypto_build_dataset
        ensemble_class = CryptoExpertEnsemble
        expert_names = CRYPTO_EXPERT_NAMES
        print(f'  Архитектура: CryptoMoE ({len(expert_names)} экспертов, '
              f'горизонты {CRYPTO_HORIZONS})')
    else:
        moe = MultiTimeframeMoE(ticker)
        prepare_df_fn = _prepare_df
        feature_cols_fn = _get_feature_cols
        build_dataset_fn = _build_dataset
        ensemble_class = ExpertEnsemble

    ok = moe.train(limit=None, verbose=True, retrain_on_full=False)
    if not ok:
        return {'status': 'error', 'message': 'train() вернул False'}

    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'  OOS: val_acc={oos_val_acc:.1%}, L≥{oos_long_th:.2f}, S≥{oos_short_th:.2f}')

    # ── Retrain на 100% + walk-forward CV для threshold search ──
    print(f'  Retrain на 100%...')
    df_full = prepare_df_fn(ticker, limit=None)
    if df_full is None or len(df_full) < 200:
        return {'status': 'error', 'message': f'Недостаточно данных: {len(df_full) if df_full is not None else 0}'}

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

    ensemble_full = ensemble_class(n_features=len(available))
    print(f'  LSTM эксперты ({len(ensemble_full.expert_names)})...')
    ensemble_full.train_all(df_full, available, verbose=True, epochs=None)
    ensemble_full.set_norm_stats(*norm_stats)

    print(f'  Forward pass...')
    signals_full = ensemble_full.predict_all(df_full, available)
    X_full = build_dataset_fn(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:]

    # Сохраняем оригинальные outcomes ДО flat filter
    flat_mask = _flat_mask(df_full).values[-len(X_full):]
    n_flat = int(flat_mask.sum())
    y_orig = y_full.copy()  # оригинальные outcomes для threshold search

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

    # Flat filter для XGBoost training (на все 100%)
    y_full[flat_mask] = 0.0

    # XGBoost обучаем на 100%
    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,
    }

    # Recency-weighted sample_weight (последние строки важнее)
    n_full = len(X_full)
    sample_weight_full = np.exp(0.8 * np.linspace(0, 1, n_full))
    sample_weight_full = sample_weight_full / sample_weight_full.mean()

    # Long XGBoost (cap 15 + per-ticker multiplier)
    pl = params_xgb.copy()
    pos_l = y_full[:, 0].sum()
    neg_l = len(y_full) - pos_l
    scale_l = min(max(neg_l / max(pos_l, 1), 1.0), 15.0)
    scale_l *= TICKER_LONG_WEIGHTS.get(ticker, DEFAULT_LONG_WEIGHT)
    pl['scale_pos_weight'] = scale_l
    xgb_long = xgb.XGBClassifier(**pl)
    xgb_long.fit(X_full, y_full[:, 0], sample_weight=sample_weight_full)

    # Short XGBoost (cap 5 + per-ticker multiplier)
    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)
    xgb_short.fit(X_full, y_full[:, 1], sample_weight=sample_weight_full)

    # ── Walk-forward CV для поиска порогов ──
    N = len(X_full)
    n_folds = 4
    oof_probs_long = np.full(N, np.nan)
    oof_probs_short = np.full(N, np.nan)

    # Границы фолдов: [0, 25%, 50%, 75%, 100%]
    fold_bounds = [int(N * i / n_folds) for i in range(n_folds + 1)]

    print(f'  Walk-forward CV ({n_folds} folds):')
    for fold in range(n_folds):
        train_end = fold_bounds[fold + 1]
        test_start = train_end
        test_end = fold_bounds[fold + 2] if fold + 2 < len(fold_bounds) else N

        if test_end - test_start < 50:
            continue  # слишком маленький тестовый сет

        X_fold_train = X_full[:train_end]
        y_fold_train = y_full[:train_end].copy()
        X_fold_test = X_full[test_start:test_end]

        # Recency weights for this fold
        sw_fold = np.exp(0.8 * np.linspace(0, 1, len(X_fold_train)))
        sw_fold = sw_fold / sw_fold.mean()

        # Flat filter на train этой fold
        fm_fold = flat_mask[:train_end]
        y_fold_train[fm_fold] = 0.0

        pl_fold = params_xgb.copy()
        pos_f = y_fold_train[:, 0].sum()
        neg_f = len(y_fold_train) - pos_f
        scale_f = min(max(neg_f / max(pos_f, 1), 1.0), 15.0)
        scale_f *= TICKER_LONG_WEIGHTS.get(ticker, DEFAULT_LONG_WEIGHT)
        pl_fold['scale_pos_weight'] = scale_f
        xgb_l_fold = xgb.XGBClassifier(**pl_fold)
        xgb_l_fold.fit(X_fold_train, y_fold_train[:, 0], sample_weight=sw_fold)

        ps_fold = params_xgb.copy()
        pos_s = y_fold_train[:, 1].sum()
        neg_s = len(y_fold_train) - 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_fold['scale_pos_weight'] = scale_s
        xgb_s_fold = xgb.XGBClassifier(**ps_fold)
        xgb_s_fold.fit(X_fold_train, y_fold_train[:, 1], sample_weight=sw_fold)

        oof_probs_long[test_start:test_end] = xgb_l_fold.predict_proba(X_fold_test)[:, 1]
        oof_probs_short[test_start:test_end] = xgb_s_fold.predict_proba(X_fold_test)[:, 1]
        print(f'    Fold {fold+1}: train [0:{train_end}], test [{test_start}:{test_end}]')

    # Threshold search на OOF probabilities (пропуская flat)
    oof_valid = ~np.isnan(oof_probs_long)
    oof_probs_long = oof_probs_long[oof_valid]
    oof_probs_short = oof_probs_short[oof_valid]
    y_oof = y_orig[oof_valid]
    flat_oof = flat_mask[oof_valid]
    atr_oof = df_full['atr'].values[-len(X_full):][oof_valid]
    close_oof = df_full['Close'].values[-len(X_full):][oof_valid]

    tp_mult = target_cfg['atr_mult_tp']
    sl_mult = target_cfg['atr_mult_sl']
    breakeven_wr = sl_mult / (sl_mult + tp_mult) + 0.02

    def _find_best_thr_oof(probs, y_true, flat_m, atr_vals, close_vals):
        best_thr, best_net = 0.55, -float('inf')
        for thr in np.arange(0.55, 0.75, 0.02):
            net = 0.0
            wins, total = 0, 0
            for i in range(len(probs)):
                if flat_m[i]:
                    continue
                if probs[i] >= thr:
                    total += 1
                    success = (y_true[i] == 1)
                    ret = tp_mult * atr_vals[i] / close_vals[i] if success else -sl_mult * atr_vals[i] / close_vals[i]
                    net += ret
                    if success: wins += 1
            wr = wins / max(total, 1)
            print(f'    thr={thr:.2f}: trades={total}, WR={wr:.1%}, net={net:.4f}')
            if total >= 15 and net > best_net and wr > breakeven_wr:
                best_net = net
                best_thr = thr
        if best_net == -float('inf'):
            best_net = 0.0
        return np.clip(best_thr, 0.55, 0.75), best_net

    print(f'  Threshold search (OOF walk-forward, flat skipped):')
    best_thr_long, net_long = _find_best_thr_oof(oof_probs_long, y_oof[:, 0],
                                                  flat_oof, atr_oof, close_oof)
    best_thr_short, net_short = _find_best_thr_oof(oof_probs_short, y_oof[:, 1],
                                                    flat_oof, atr_oof, close_oof)
    print(f'  ✓ LONG≥{best_thr_long:.2f} (net={net_long:.4f}), '
          f'SHORT≥{best_thr_short:.2f} (net={net_short:.4f})')

    # Сборка финальной модели
    moe.ensemble = ensemble_full
    moe.rf_long = xgb_long
    moe.rf_short = xgb_short
    moe.rf_model = None
    moe.optimal_thresholds = {'long': best_thr_long, 'short': best_thr_short}
    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)

    elapsed = time.time() - t0
    print(f'  ✓ {save_path} ({elapsed/60:.1f}мин)')

    return {
        'status': 'ok',
        'ticker': ticker,
        'market': market,
        'rows': len(df_full),
        'val_acc': float(oos_val_acc),
        'long_th': float(best_thr_long),
        'short_th': float(best_thr_short),
        'long_sr': float(long_sr),
        'short_sr': float(short_sr),
        'flat_pct': round(n_flat / len(flat_mask) * 100, 1),
        'selected_experts': len(oos_selected) if oos_selected else 16,
        'elapsed_min': round(elapsed / 60, 1),
        'path': save_path,
    }


def print_market_header(tickers):
    """Печатает сводку по рынкам тренируемых тикеров."""
    markets = {}
    for t in tickers:
        m = get_market(t)
        markets.setdefault(m, []).append(t)
    lines = []
    for m, ts in sorted(markets.items()):
        cfg = MARKET_TARGET_CONFIGS[m]
        lines.append(f'    {m.upper():>8} ({len(ts):2d} шт): SL={cfg["atr_mult_sl"]}×ATR, TP={cfg["atr_mult_tp"]}×ATR')
    return '\n'.join(lines)


if __name__ == '__main__':
    import argparse
    parser = argparse.ArgumentParser(description='Train ALL tickers with per-market RR configs')
    parser.add_argument('--ticker', type=str, default=None,
                        help='Single ticker, or ALL for all 17 tickers')
    args = parser.parse_args()

    if args.ticker and args.ticker.upper() == 'ALL':
        tickers = ALL_TICKERS
    elif args.ticker:
        tickers = [args.ticker.upper()]
    else:
        tickers = MOEX_TICKERS  # по умолчанию MOEX

    # Определяем время для заголовка
    now_str = datetime.now().strftime("%d.%m.%Y %H:%M")

    print(f'\n{"█"*72}')
    print(f'  ПАКЕТНАЯ ТРЕНИРОВКА (per-market RR)')
    print(f'  {now_str}')
    print(f'  Тикеры: {", ".join(tickers)} ({len(tickers)} шт)')
    print(f'  Архитектура: 16 LSTM экспертов + XGBoost + MTF контекст')
    print(f'  Сохранение: *moe_v12_rr1x2.joblib')
    print(print_market_header(tickers))
    print(f'{"█"*72}\n')

    results_log = {}
    t_start = time.time()
    errors = 0

    for i, ticker in enumerate(tickers, 1):
        print(f'\n[{i}/{len(tickers)}] {ticker}')
        entry = {'ticker': ticker}

        try:
            result = train_moe_for(ticker)
            if result:
                entry = result
                print(f'  ✓ {ticker}: val_acc={result["val_acc"]:.1%}, '
                      f'L≥{result["long_th"]:.2f}, S≥{result["short_th"]:.2f}, '
                      f'SR={result["long_sr"]:.0%}/{result["short_sr"]:.0%}, '
                      f'флет={result["flat_pct"]:.1f}%')
            else:
                print(f'  ⏭ {ticker}: пропущен')
                entry['status'] = 'skipped'
        except Exception as e:
            import traceback
            errors += 1
            entry['status'] = 'error'
            entry['error'] = str(e)
            print(f'  ❌ {ticker}: {e}')
            traceback.print_exc()

        results_log[ticker] = entry

    total_time = time.time() - t_start
    ok = [v for v in results_log.values() if v.get('status') == 'ok']

    print(f'\n{"="*72}')
    print(f'  ТРЕНИРОВКА ЗАВЕРШЕНА')
    print(f'  Время: {total_time/60:.1f} мин')
    print(f'  Успешно: {len(ok)}/{len(tickers)}, Ошибок: {errors}')
    if ok:
        print(f'  Средняя val_acc: {np.mean([v["val_acc"] for v in ok]):.1%}')
        print(f'  Средняя flat%: {np.mean([v["flat_pct"] for v in ok]):.1f}%')
    print(f'{"="*72}')

    # Детальная таблица
    print(f'\n  {"="*64}')
    print(f'  {"Тикер":>7} {"Рынок":>7} {"Rows":>6} {"val_acc":>8} {"L≥th":>5} {"S≥th":>5} '
          f'{"L_SR":>5} {"S_SR":>5} {"Флет%":>6} {"Эксп":>4} {"мин":>5}')
    print(f'  {"-"*64}')
    for v in ok:
        print(f'  {v["ticker"]:>7} {v.get("market","-"):>7} {v["rows"]:>6} {v["val_acc"]:>7.1%} '
              f'{v["long_th"]:>4.2f} {v["short_th"]:>4.2f} '
              f'{v["long_sr"]:>4.0%} {v["short_sr"]:>4.0%} '
              f'{v["flat_pct"]:>5.1f}% {v["selected_experts"]:>4} {v["elapsed_min"]:>4.1f}')
    if ok:
        print(f'  {"-"*64}')
        print(f'  {"СРЕДНЕЕ":>7} {"":>7} {np.mean([v["rows"] for v in ok]):>6.0f} '
              f'{np.mean([v["val_acc"] for v in ok]):>7.1%} '
              f'{np.mean([v["long_th"] for v in ok]):>4.2f} {np.mean([v["short_th"] for v in ok]):>4.2f} '
              f'{np.mean([v["long_sr"] for v in ok]):>4.0%} {np.mean([v["short_sr"] for v in ok]):>4.0%} '
              f'{np.mean([v["flat_pct"] for v in ok]):>5.1f}% '
              f'{np.mean([v["selected_experts"] for v in ok]):>4.0f} {np.mean([v["elapsed_min"] for v in ok]):>4.1f}')

    with open(LOG_FILE, 'w') as f:
        json.dump(results_log, f, indent=2, default=str)
