#!/usr/bin/env python3
"""
Walk-Forward Cross-Validation для MultiTimeframeMoE (v12.5).

Оценивает устойчивость модели на последовательных отрезках времени
(rolling-origin CV). Использует v12.5 архитектуру: раздельные rf_long/rf_short,
per-ticker class_weight и пороги, единый predict_all на всём df.

Запуск:
    python3 walkforward_cv.py --ticker SBER --folds 7
    python3 walkforward_cv.py --ticker SBER --folds 5 --epochs 40
    python3 walkforward_cv.py --all                   # все 14 тикеров
"""
import argparse
import os
import sys
import warnings
import numpy as np
import pandas as pd

warnings.filterwarnings('ignore')
os.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))

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

from models.moe import _prepare_df, _get_feature_cols, _build_dataset, _split_signals
from models.experts import ExpertEnsemble
from config import TARGET_CONFIG, get_market, should_disable_short_training
from sklearn.multioutput import MultiOutputClassifier
from sklearn.ensemble import RandomForestClassifier
from config import TICKER_SHORT_WEIGHTS, TICKER_LONG_WEIGHTS, TICKER_THRESHOLDS, DEFAULT_SHORT_WEIGHT, DEFAULT_LONG_WEIGHT


CACHE_DIR = '/tmp/opencode/moe_cache'


def _cache_path(ticker: str) -> str:
    return os.path.join(CACHE_DIR, f'{ticker}_experts_latest.joblib')


def _save_expert_cache(ticker: str, ensemble: ExpertEnsemble):
    """Сохраняет веса экспертов в кэш для transfer learning между folds."""
    os.makedirs(CACHE_DIR, exist_ok=True)
    cache = {
        'state': ensemble.state_dict(),
        'expert_names': ensemble.expert_names,
        'n_features': ensemble.n_features,
    }
    import joblib
    joblib.dump(cache, _cache_path(ticker))
    return cache


def _load_expert_cache(ticker: str) -> ExpertEnsemble | None:
    """Загружает кэш экспертов для transfer learning."""
    path = _cache_path(ticker)
    if not os.path.exists(path):
        return None
    try:
        import joblib
        cache = joblib.load(path)
        ensemble = ExpertEnsemble(
            n_features=cache['n_features'],
            experts=cache['expert_names'],
        )
        ensemble.load_state_dict(cache['state'])
        return ensemble
    except Exception:
        return None


def walk_forward_cv(
    ticker: str,
    n_folds: int = 5,
    epochs: int = 60,
    val_pct: float = 0.10,
    min_train_pct: float = 0.40,
    limit: int | None = 5000,
    verbose: bool = True,
    fast: bool = False,
) -> dict:
    """
    Walk-Forward CV: rolling-origin expanding window (v12.5).

    Разбивает хронологические данные на n_folds блоков. На каждой итерации
    обучает модель на данных от начала до конца блока k, валидирует на блоке k+1.
    Окно train расширяется.

    Архитектура v12.5: раздельные rf_long/rf_short, per-ticker thresholds,
    единый forward pass экспертов на всём df.

    Returns:
        dict: val_accs, long_recalls, short_recalls, per_fold, summary stats
    """
    df = _prepare_df(ticker, limit=limit)
    if df is None or len(df) < 1000:
        return {'error': f'Недостаточно данных для {ticker}: {len(df) if df is not None else 0}'}

    n = len(df)
    all_features = _get_feature_cols()
    available = [c for c in all_features if c in df.columns]
    n_val = max(int(n * val_pct), 50)
    n_train_min = max(int(n * min_train_pct), 500)

    # Per-ticker классовые веса
    short_weight = TICKER_SHORT_WEIGHTS.get(ticker, DEFAULT_SHORT_WEIGHT)
    long_weight = TICKER_LONG_WEIGHTS.get(ticker, DEFAULT_LONG_WEIGHT)
    th = TICKER_THRESHOLDS.get(ticker, {'long': 0.48, 'short': 0.42})

    # Генерируем фолды: val блоки расположены в конце, train растёт
    MAX_BARS = TARGET_CONFIG.get('max_bars', 100)
    folds = []
    train_end = n_train_min
    val_start = train_end + MAX_BARS  # purged gap: train не смотрит в val
    val_end = val_start + n_val
    while val_end <= n and len(folds) < n_folds:
        folds.append((0, train_end, val_start, val_end))
        train_end = val_end
        val_start = val_end + MAX_BARS
        val_end = val_start + n_val

    if len(folds) == 0:
        return {'error': 'Не удалось сгенерировать валидные фолды'}

    if verbose:
        print(f'\n  Walk-Forward CV v12.5: {ticker}')
        print(f'  Data: {n} rows, {len(folds)} folds')
        print(f'  Val size: {n_val}, min train: {n_train_min}')
        print(f'  Thresholds: long={th["long"]}, short={th["short"]}')
        print()

    val_accs = []
    long_recalls = []
    short_recalls = []
    long_f1s = []
    short_f1s = []
    per_fold = []

    for i, (t_start, t_end, v_start, v_end) in enumerate(folds, 1):
        if verbose:
            print(f'  Fold {i}/{len(folds)}: train=[{t_start}:{t_end}] ({t_end-t_start} rows), '
                  f'val=[{v_start}:{v_end}] ({v_end-v_start} rows)')

        df_train = df.iloc[t_start:t_end].copy()
        df_val = df.iloc[v_start:v_end].copy()

        if len(df_train) < 200 or len(df_val) < 30:
            if verbose:
                print(f'    ⏭ Skip (train={len(df_train)}, val={len(df_val)})')
            continue

        try:
            # Эксперты: warm-start из кэша (если есть) — transfer learning
            # между последовательными folds, так как train расширяется
            cached = _load_expert_cache(ticker)
            if cached is not None and cached.n_features == len(available):
                if fast:
                    # Fast mode: fine-tune с уменьшенным числом эпох
                    ft_epochs = max(epochs // 3, 10)
                    if verbose:
                        print(f'    [fast] cache → fine-tune {ft_epochs} ep')
                    cached.train_all(df_train, available, verbose=False, epochs=ft_epochs)
                    ensemble = cached
                else:
                    # Normal: warm-start из кэша, продолжаем обучение
                    if verbose:
                        print(f'    [warm] cache → continue training {epochs} ep')
                    cached.train_all(df_train, available, verbose=False, epochs=epochs)
                    ensemble = cached
            else:
                # Нет кэша: полное обучение с нуля
                if verbose:
                    print(f'    [cold] full training from scratch ({epochs} ep)')
                ensemble = ExpertEnsemble(n_features=len(available))
                ensemble.train_all(df_train, available, verbose=False, epochs=epochs)
                # Сохраняем первый fold в кэш для последующих
                _save_expert_cache(ticker, ensemble)

            # Единый forward pass на всём df, разделение train/val
            signals_full = ensemble.predict_all(df, available)
            signals_train = _split_signals(signals_full, t_start, t_end)
            signals_val = _split_signals(signals_full, v_start, v_end)

            X_train = _build_dataset(df_train, signals_train, ensemble, available)
            _wfc_disable_short = should_disable_short_training(ticker)
            if _wfc_disable_short:
                y_train = df_train['outcome_long'].values.reshape(-1, 1)
                y_val = df_val['outcome_long'].values.reshape(-1, 1)
            else:
                y_train = np.column_stack([
                    df_train['outcome_long'].values, df_train['outcome_short'].values
                ])
                y_val = np.column_stack([
                    df_val['outcome_long'].values, df_val['outcome_short'].values
                ])
            min_len_t = min(len(X_train), len(y_train))
            X_train, y_train = X_train[-min_len_t:], y_train[-min_len_t:]

            min_len_v = min(len(X_val), len(y_val))
            X_val, y_val = X_val[-min_len_v:], y_val[-min_len_v:]

            if len(X_train) < 50 or len(X_val) < 10:
                if verbose:
                    print(f'    ⏭ Skip (X_train={len(X_train)}, X_val={len(X_val)})')
                continue

            # Раздельные RF — v12.5
            # long_weight может быть числом (напр. 2.0) или 'balanced'
            # MOEX-only-long: rf_short не создаётся
            rf_long_cw = {0: 1.0, 1: long_weight} if isinstance(long_weight, (int, float)) else long_weight
            rf_long = RandomForestClassifier(
                n_estimators=200, max_depth=10,
                random_state=42, n_jobs=-1,
                class_weight=rf_long_cw,
            )
            rf_long.fit(X_train, y_train[:, 0])
            rf_short = None

            if not _wfc_disable_short:
                rf_short = RandomForestClassifier(
                    n_estimators=200, max_depth=10,
                    random_state=42, n_jobs=-1,
                    class_weight={0: 1.0, 1: short_weight},
                )
                rf_short.fit(X_train, y_train[:, 1])

            # Predict на val с per-ticker порогами
            p_long = rf_long.predict_proba(X_val)[:, 1]
            if rf_short is not None:
                p_short = rf_short.predict_proba(X_val)[:, 1]
                y_long = y_val[:, 0]
                y_short = y_val[:, 1]
            else:
                p_short = np.zeros(len(X_val))
                y_long = y_val[:, 0]
                y_short = None

            pred_long = (p_long >= th['long']).astype(int)
            pred_short = (p_short >= th['short']).astype(int)

            # Long metrics
            lp = (pred_long == 1).sum()
            ltp = ((pred_long == 1) & (y_long == 1)).sum()
            lfn = (y_long == 1).sum() - ltp
            long_prec = ltp / lp if lp > 0 else 0
            long_rec = ltp / (ltp + lfn) if (ltp + lfn) > 0 else 0
            long_f1 = 2 * long_prec * long_rec / (long_prec + long_rec) if (long_prec + long_rec) > 0 else 0

            # Short metrics
            sp = (pred_short == 1).sum()
            stp = ((pred_short == 1) & (y_short == 1)).sum()
            sfn = (y_short == 1).sum() - stp
            short_prec = stp / sp if sp > 0 else 0
            short_rec = stp / (stp + sfn) if (stp + sfn) > 0 else 0
            short_f1 = 2 * short_prec * short_rec / (short_prec + short_rec) if (short_prec + short_rec) > 0 else 0

            # val_acc = средняя F1
            val_acc = (long_f1 + short_f1) / 2.0

            val_accs.append(val_acc)
            long_recalls.append(long_rec)
            short_recalls.append(short_rec)
            long_f1s.append(long_f1)
            short_f1s.append(short_f1)

            fold_info = {
                'fold': i,
                'n_train': len(X_train),
                'n_val': len(X_val),
                'val_acc': float(val_acc),
                'long_f1': float(long_f1), 'short_f1': float(short_f1),
                'long_recall': float(long_rec), 'short_recall': float(short_rec),
                'long_prec': float(long_prec), 'short_prec': float(short_prec),
                'long_pred_rate': float((y_long == 1).mean()),
                'short_pred_rate': float((y_short == 1).mean()),
            }
            per_fold.append(fold_info)

            if verbose:
                print(f'    val_acc={val_acc:.2%}  long: F1={long_f1:.1%} R={long_rec:.1%} P={long_prec:.1%} | '
                      f'short: F1={short_f1:.1%} R={short_rec:.1%} P={short_prec:.1%}')

        except Exception as e:
            if verbose:
                print(f'    ❌ Fold {i} error: {e}')
            import traceback; traceback.print_exc()
            per_fold.append({'fold': i, 'error': str(e)})

    if not val_accs:
        return {'error': 'Все фолды упали', 'per_fold': per_fold}

    result = {
        'ticker': ticker,
        'n_folds': len(folds),
        'n_folds_ok': len(val_accs),
        'val_accs': val_accs,
        'val_acc_mean': float(np.mean(val_accs)),
        'val_acc_std': float(np.std(val_accs)),
        'val_acc_min': float(np.min(val_accs)),
        'val_acc_max': float(np.max(val_accs)),
        'long_recall_mean': float(np.mean(long_recalls)) if long_recalls else 0,
        'short_recall_mean': float(np.mean(short_recalls)) if short_recalls else 0,
        'long_f1_mean': float(np.mean(long_f1s)) if long_f1s else 0,
        'short_f1_mean': float(np.mean(short_f1s)) if short_f1s else 0,
        'per_fold': per_fold,
    }
    return result


def print_cv_report(result: dict):
    """Форматированный отчёт по результатам CV."""
    if 'error' in result and 'val_accs' not in result:
        print(f'\n  ❌ Ошибка: {result["error"]}')
        return

    print(f'\n{"=" * 70}')
    print(f'  WALK-FORWARD CV v12.5: {result["ticker"]}')
    print(f'{"=" * 70}')
    print(f'  Фолдов: {result["n_folds_ok"]}/{result["n_folds"]}')
    print(f'  Val accuracy (avg F1): {result["val_acc_mean"]:.2%} ± {result["val_acc_std"]:.2%}')
    print(f'  Min/Max:               {result["val_acc_min"]:.2%} / {result["val_acc_max"]:.2%}')
    print(f'  Long F1 / Recall:      {result["long_f1_mean"]:.1%} / {result["long_recall_mean"]:.1%}')
    print(f'  Short F1 / Recall:     {result["short_f1_mean"]:.1%} / {result["short_recall_mean"]:.1%}')

    print(f'\n  Per-fold:')
    print(f'  {"Fold":>4} {"Train":>6} {"Val":>5} {"ValAcc":>7} {"L.F1":>5} {"S.F1":>5} {"L.Rec":>5} {"S.Rec":>5}')
    for f in result['per_fold']:
        if 'error' in f:
            print(f'  {f["fold"]:>4} ERROR: {f["error"]}')
        else:
            print(f'  {f["fold"]:>4} {f["n_train"]:>6} {f["n_val"]:>5} '
                  f'{f["val_acc"]:>6.2%} {f["long_f1"]:>4.1%} {f["short_f1"]:>4.1%} '
                  f'{f["long_recall"]:>4.1%} {f["short_recall"]:>4.1%}')

    print(f'{"=" * 70}')


if __name__ == '__main__':
    parser = argparse.ArgumentParser(description='Walk-Forward CV v12.5 for MoE')
    parser.add_argument('--ticker', type=str, default='SBER')
    parser.add_argument('--folds', type=int, default=5)
    parser.add_argument('--epochs', type=int, default=60)
    parser.add_argument('--limit', type=int, default=5000, help='Max rows to load')
    parser.add_argument('--all', action='store_true', help='Run all 14 tickers')
    parser.add_argument('--fast', action='store_true',
                        help='Fast mode: transfer learning между folds (fine-tune, меньше эпох)')
    parser.add_argument('--quiet', action='store_true')
    args = parser.parse_args()

    # Очистка кэша при --all для честного CV (каждый тикер с нуля)
    if args.all:
        import shutil
        if os.path.exists(CACHE_DIR):
            shutil.rmtree(CACHE_DIR)

    if args.all:
        tickers = ['ASTR', 'GAZP', 'LKOH', 'MOEX', 'MTSS', 'NSVZ', 'NVTK',
                   'PHOR', 'PLZL', 'ROSN', 'SBER', 'SNGSP', 'VTBR', 'X5']
        results = []
        for t in tickers:
            print(f'\n{"─" * 60}')
            result = walk_forward_cv(
                ticker=t, n_folds=args.folds,
                epochs=args.epochs, limit=args.limit,
                verbose=not args.quiet, fast=args.fast,
            )
            results.append(result)
            if 'val_acc_mean' in result:
                print(f'  {t}: mean={result["val_acc_mean"]:.2%} ± {result["val_acc_std"]:.2%} '
                      f'(long F1={result["long_f1_mean"]:.1%}, short F1={result["short_f1_mean"]:.1%})')

        print(f'\n{"=" * 70}')
        print(f'  ALL 14 TICKERS SUMMARY (v12.5)')
        print(f'{"=" * 70}')
        for r in results:
            if 'val_acc_mean' in r:
                print(f'  {r["ticker"]:>6}: {r["val_acc_mean"]:.2%} ± {r["val_acc_std"]:.2%}  '
                      f'(L={r["long_f1_mean"]:.1%}/S={r["short_f1_mean"]:.1%})')
            else:
                print(f'  {r.get("ticker", "?"):>6}: ERROR — {r.get("error", "?")}')
    else:
        result = walk_forward_cv(
            ticker=args.ticker,
            n_folds=args.folds,
            epochs=args.epochs,
            limit=args.limit,
            verbose=not args.quiet,
            fast=args.fast,
        )
        print_cv_report(result)