#!/usr/bin/env python3
"""
Полная тренировка MoE v12 для всех 14 MOEX тикеров.
Запуск:  python3 train_experts_all.py
         python3 train_experts_all.py --ticker SBER   # один тикер

Архитектура: MultiTimeframeMoE (16 LSTM экспертов + XGBoost, H1 + MTF контекст D1/W1)
Flow: 85% train → OOS thresholds → retrain 100% → save models/saved/{ticker}_moe_v12.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__)))

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

from config import MOEX_TICKERS, SAVE_DIR, TARGET_CONFIG, DEVICE
from models.moe import MultiTimeframeMoE, _prepare_df, _get_feature_cols, _build_dataset, _flat_mask
from models.experts import ExpertEnsemble

import xgboost as xgb
from config import MODEL_CONFIG, TICKER_SHORT_WEIGHTS, DEFAULT_SHORT_WEIGHT

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

def _get_scale_weights(y_full, ticker):
    """Вычисляет scale_pos_weight для XGBoost."""
    n = len(y_full)
    pos = float(y_full.sum())
    neg = n - pos
    scale = min(max(neg / max(pos, 1), 1.0), 5.0)
    if ticker in TICKER_SHORT_WEIGHTS:
        scale *= TICKER_SHORT_WEIGHTS[ticker]
    return scale

def train_moe_for(ticker):
    """Тренирует MoE v12 для одного тикера. Возвращает dict с метриками."""
    t0 = time.time()
    print(f'  [{datetime.now().strftime("%H:%M")}] Загрузка данных...')

    # ── Step 1: Train on 85% → OOS ──
    moe = MultiTimeframeMoE(ticker)
    ok = moe.train(limit=None, verbose=True, retrain_on_full=False)
    if not ok:
        return {'status': 'error', 'message': 'train() вернул False'}

    oos_long_th = moe.optimal_thresholds.get('long', 0.55)
    oos_short_th = moe.optimal_thresholds.get('short', 0.55)
    oos_val_acc = getattr(moe, 'oos_val_acc', 0)
    oos_selected = moe.selected_experts

    print(f'  OOS: val_acc={oos_val_acc:.1%}, L≥{oos_long_th:.2f}, S≥{oos_short_th:.2f}')
    if oos_selected:
        print(f'  Отобрано экспертов: {len(oos_selected)}/16')

    # ── Step 2: Retrain LSTM + XGBoost на 100% ──
    print(f'  Retrain на 100% данных...')
    df_full = _prepare_df(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 = _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))
    print(f'  LSTM эксперты ({len(ensemble_full.expert_names)}) на {len(df_full)} строк...')
    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(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
    print(f'  Флет-фильтр: {n_flat}/{len(flat_mask)} ({n_flat/len(flat_mask):.1%})')

    # Статистика по успешности исходных таргетов (до фильтрации)
    long_sr = df_full['outcome_long'].mean()
    short_sr = df_full['outcome_short'].mean()
    print(f'  Исходные success rate: long={long_sr:.1%}, short={short_sr:.1%}')
    print(f'  Размер датасета: {len(X_full)} samples')

    print(f'  XGBoost на {len(X_full)} samples...')
    params_xgb = MODEL_CONFIG['xgb'].copy()
    params_xgb.update({
        '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,
    })

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

    pos_short = y_full[:, 1].sum()
    neg_short = len(y_full) - pos_short
    scale_short = min(max(neg_short / max(pos_short, 1), 1.0), 5.0)
    scale_short *= TICKER_SHORT_WEIGHTS.get(ticker, DEFAULT_SHORT_WEIGHT)

    p = params_xgb.copy()
    p['scale_pos_weight'] = scale_short
    xgb_short = xgb.XGBClassifier(**p)
    xgb_short.fit(X_full, y_full[:, 1])

    # ── Оценка in-sample ──
    train_preds_long = xgb_long.predict(X_full)
    train_preds_short = xgb_short.predict(X_full)
    train_acc_long = (train_preds_long == y_full[:, 0]).mean()
    train_acc_short = (train_preds_short == y_full[:, 1]).mean()
    train_pos_rate_long = train_preds_long.mean()
    train_pos_rate_short = train_preds_short.mean()

    # Процент flat-баров, которые модель классифицирует как 1 (должны быть 0)
    flat_preds_long = train_preds_long[flat_mask]
    flat_preds_short = train_preds_short[flat_mask]
    flat_false_long = flat_preds_long.mean() if len(flat_preds_long) > 0 else 0
    flat_false_short = flat_preds_short.mean() if len(flat_preds_short) > 0 else 0

    # ── Step 3: Сборка production-модели ──
    moe.ensemble = ensemble_full
    moe.rf_long = xgb_long
    moe.rf_short = xgb_short
    moe.rf_model = None
    moe.calibrators = 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 ──
    save_path = os.path.join(SAVE_DIR, f'{ticker.lower()}_moe_v12.joblib')
    moe.save(save_path)

    elapsed = time.time() - t0
    return {
        'status': 'ok',
        'ticker': ticker,
        'rows': len(df_full),
        'val_acc': oos_val_acc,
        'long_th': oos_long_th,
        'short_th': oos_short_th,
        'selected_experts': len(oos_selected) if oos_selected else 16,
        'elapsed_min': round(elapsed / 60, 1),
        'path': save_path,
        # Детальная статистика
        'flat_pct': round(n_flat / len(flat_mask) * 100, 1),
        'long_sr': round(long_sr * 100, 1),
        'short_sr': round(short_sr * 100, 1),
        'train_acc_long': round(train_acc_long * 100, 1),
        'train_acc_short': round(train_acc_short * 100, 1),
        'train_pos_rate_long': round(train_pos_rate_long * 100, 1),
        'train_pos_rate_short': round(train_pos_rate_short * 100, 1),
        'flat_false_long': round(flat_false_long * 100, 1),
        'flat_false_short': round(flat_false_short * 100, 1),
    }


if __name__ == '__main__':
    import argparse
    parser = argparse.ArgumentParser(description='Train MoE v12 for all MOEX tickers')
    parser.add_argument('--ticker', type=str, default=None, help='Single ticker only')
    args = parser.parse_args()

    tickers = [args.ticker.upper()] if args.ticker else MOEX_TICKERS

    print(f'\n{"="*70}')
    print(f'  MoE v12 TRAINING — {datetime.now().strftime("%d.%m.%Y %H:%M")}')
    print(f'  Трекеры: {", ".join(tickers)}')
    print(f'  Архитектура: 16 LSTM экспертов + XGBoost + MTF контекст')
    print(f'  TARGET_CONFIG: SL={TARGET_CONFIG["atr_mult_sl"]}×ATR TP={TARGET_CONFIG["atr_mult_tp"]}×ATR')
    print(f'  Device: {DEVICE}')
    print(f'{"="*70}\n')

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

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

        try:
            result = train_moe_for(ticker)
            if result is None:
                print(f'  ⏭ Пропущен')
                entry['status'] = 'skipped'
            else:
                entry = result
                print(f'  ✓ {result["elapsed_min"]}мин — '
                      f'val_acc={result["val_acc"]:.1%}, '
                      f'L≥{result["long_th"]:.2f}, S≥{result["short_th"]:.2f}, '
                      f'экспертов={result["selected_experts"]}/16, '
                      f'флет={result.get("flat_pct", "?")}%')
        except Exception as e:
            import traceback
            errors += 1
            entry['status'] = 'error'
            entry['error'] = str(e)
            print(f'  ❌ Ошибка: {e}')
            traceback.print_exc()

        results_log[ticker] = entry

    total_time = time.time() - t_start
    print(f'\n{"="*70}')
    print(f'  ТРЕНИРОВКА ЗАВЕРШЕНА')
    print(f'  Время: {total_time/60:.1f} мин')
    print(f'  Ошибок: {errors}')
    ok = [v for v in results_log.values() if v.get('status') == 'ok']
    if ok:
        accs = [v['val_acc'] for v in ok]
        print(f'  Успешно: {len(ok)}/{len(tickers)}')
        print(f'  Средняя val_acc: {np.mean(accs):.1%}')
    print(f'{"="*70}')
    print()
    print(f'  ДЕТАЛЬНАЯ СТАТИСТИКА ПО КАЖДОМУ ТИКЕРУ:')
    print(f'  {"="*60}')
    print(f'  {"Тикер":>6} {"Rows":>6} {"val_acc":>8} {"L≥thr":>6} {"S≥thr":>6} '
          f'{"Эксп":>4} {"Флет%":>6} {"L_SR%":>6} {"S_SR%":>6} '
          f'{"Acc_L":>6} {"Acc_S":>6} {"ТоргL":>6} {"ТоргS":>6} '
          f'{"flong":>6} {"fshort":>6} {"мин":>5}')
    print(f'  {"-"*100}')
    for v in ok:
        print(f'  {v["ticker"]:>6} {v["rows"]:>6} {v["val_acc"]:>7.1%} '
              f'{v["long_th"]:>5.2f} {v["short_th"]:>5.2f} '
              f'{v["selected_experts"]:>4} '
              f'{v.get("flat_pct",0):>5.1f}% '
              f'{v.get("long_sr",0):>5.1f}% {v.get("short_sr",0):>5.1f}% '
              f'{v.get("train_acc_long",0):>5.1f}% {v.get("train_acc_short",0):>5.1f}% '
              f'{v.get("train_pos_rate_long",0):>5.1f}% {v.get("train_pos_rate_short",0):>5.1f}% '
              f'{v.get("flat_false_long",0):>5.1f}% {v.get("flat_false_short",0):>5.1f}% '
              f'{v["elapsed_min"]:>4.1f}')
    # Итоговые средние
    if ok:
        print(f'  {"-"*100}')
        print(f'  {"СРЕДНЕЕ":>6} {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]):>5.2f} {np.mean([v["short_th"] for v in ok]):>5.2f} '
              f'{np.mean([v["selected_experts"] for v in ok]):>4.0f} '
              f'{np.mean([v.get("flat_pct",0) for v in ok]):>5.1f}% '
              f'{np.mean([v.get("long_sr",0) for v in ok]):>5.1f}% {np.mean([v.get("short_sr",0) for v in ok]):>5.1f}% '
              f'{np.mean([v.get("train_acc_long",0) for v in ok]):>5.1f}% {np.mean([v.get("train_acc_short",0) for v in ok]):>5.1f}% '
              f'{np.mean([v.get("train_pos_rate_long",0) for v in ok]):>5.1f}% {np.mean([v.get("train_pos_rate_short",0) for v in ok]):>5.1f}% '
              f'{np.mean([v.get("flat_false_long",0) for v in ok]):>5.1f}% {np.mean([v.get("flat_false_short",0) for v in ok]):>5.1f}% '
              f'{np.mean([v["elapsed_min"] for v in ok]):>4.1f}')
    print(f'{"="*70}')
    print(f'  Лог: {LOG_FILE}')

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