"""
Скрипт обучения моделей для ансамбля X5 (v4).

Обучает две новые модели на основе успеха x5_gru_v3:
    1. x5_lstm_v3_focal — LSTM + Focal Loss (γ=2.0) + dropout 0.4
    2. x5_gru_v4        — GRU + Focal Loss (γ=3.0) + seed=123 + hidden=128

Каждая модель проходит двухэтапное обучение:
    A — Pre-train на 9 тикерах MOEX (15,669 семплов)
    B — Fine-tune на X5 D1 augmented (900 семплов)

Использование:
    python src/ml/train_ensemble_v4.py [model_name]

    model_name: 'x5_lstm_v3_focal' | 'x5_gru_v4' | 'all' (по умолчанию)
"""

import json
import logging
import sys
import time
from pathlib import Path
from typing import Dict, Optional

import numpy as np
import torch
import torch.nn as nn

logging.basicConfig(
    level=logging.INFO,
    format='%(asctime)s [%(levelname)s] %(message)s',
    datefmt='%H:%M:%S',
)
logger = logging.getLogger(__name__)

project_root = Path(__file__).resolve().parent.parent.parent
sys.path.insert(0, str(project_root))

from src.ml.models.lstm import LSTMPredictor, GRUPredictor
from src.ml.models.registry import ModelRegistry
from src.ml.train.trainer import (
    train_model, compute_class_weights, set_seed,
    load_hdf5_dataset, create_dataloaders, print_metrics_report,
)
from src.ml.train.metrics import calculate_all_metrics, calculate_macro_f1
from src.ml.train.losses import WeightedFocalLoss

PRETRAIN_DATA = 'src/ml/models/saved/multi_ticker_d1_pretrain.h5'
FINETUNE_DATA = 'src/ml/models/saved/x5_d1_augmented.h5'
SAVE_DIR = 'src/ml/models/saved'
TICKER = 'X5'
TIMEFRAME = 'D1'

# ─── Конфигурации моделей ───────────────────────────────────────────────
MODEL_CONFIGS = {
    'x5_lstm_v3_focal': {
        'model_class': 'LSTMPredictor',
        'input_size': 61,
        'hidden_size': 128,
        'num_layers': 2,
        'dropout': 0.4,
        'bidirectional': False,
        'focal_gamma': 2.0,
        'weight_decay': 1e-3,
        'pretrain_lr': 1e-3,
        'finetune_lr': 5e-5,
        'pretrain_epochs': 100,
        'finetune_epochs': 50,
        'pretrain_patience': 15,
        'finetune_patience': 20,
        'batch_size': 32,
        'seed': 42,
        'sell_multiplier': 2.0,
        'desc': 'LSTM + Focal Loss γ=2.0 + dropout 0.4 (как v3, но LSTM)',
    },
    'x5_gru_v4': {
        'model_class': 'GRUPredictor',
        'input_size': 61,
        'hidden_size': 128,
        'num_layers': 2,
        'dropout': 0.4,
        'bidirectional': False,
        'focal_gamma': 3.0,
        'weight_decay': 1e-3,
        'pretrain_lr': 1e-3,
        'finetune_lr': 5e-5,
        'pretrain_epochs': 100,
        'finetune_epochs': 50,
        'pretrain_patience': 15,
        'finetune_patience': 20,
        'batch_size': 32,
        'seed': 123,
        'sell_multiplier': 2.5,
        'desc': 'GRU + Focal Loss γ=3.0 + seed=123 + hidden=128 (сильнее фокус на SELL)',
    },
}


def create_model(model_class_name: str, params: dict) -> nn.Module:
    """Создать модель по имени класса."""
    if model_class_name == 'LSTMPredictor':
        return LSTMPredictor(
            input_size=params['input_size'],
            hidden_size=params['hidden_size'],
            num_layers=params['num_layers'],
            dropout=params['dropout'],
            bidirectional=params.get('bidirectional', False),
        )
    elif model_class_name == 'GRUPredictor':
        return GRUPredictor(
            input_size=params['input_size'],
            hidden_size=params['hidden_size'],
            num_layers=params['num_layers'],
            dropout=params['dropout'],
        )
    else:
        raise ValueError(f"Неизвестный класс модели: {model_class_name}")


def train_ensemble_model(model_name: str) -> bool:
    """Обучить одну модель для ансамбля."""
    config = MODEL_CONFIGS[model_name]
    checkpoint_path = str(Path(SAVE_DIR) / f'{model_name}_pretrained.pt')

    logger.info("\n" + "=" * 70)
    logger.info(f"ОБУЧЕНИЕ МОДЕЛИ: {model_name}")
    logger.info(f"  {config['desc']}")
    logger.info(f"  Архитектура: {config['model_class']} "
                f"(hidden={config['hidden_size']}, dropout={config['dropout']})")
    logger.info(f"  Focal Loss: gamma={config['focal_gamma']}")
    logger.info(f"  SELL weight: ×{config['sell_multiplier']}")
    logger.info(f"  Seed: {config['seed']}")
    logger.info("=" * 70)

    # Seed
    set_seed(config['seed'])
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    logger.info(f"Устройство: {device}")

    # Load datasets
    if not Path(PRETRAIN_DATA).exists():
        logger.error(f"Pre-train датасет не найден: {PRETRAIN_DATA}")
        return False
    if not Path(FINETUNE_DATA).exists():
        logger.error(f"Fine-tune датасет не найден: {FINETUNE_DATA}")
        return False

    pretrain_data = load_hdf5_dataset(PRETRAIN_DATA)
    finetune_data = load_hdf5_dataset(FINETUNE_DATA)

    # Dataloaders
    pretrain_loader, pretrain_val_loader, _ = create_dataloaders(
        pretrain_data, batch_size=config['batch_size'],
        shuffle_train=True, val_split=0.1,
    )
    finetune_loader, val_loader, test_loader = create_dataloaders(
        finetune_data, batch_size=config['batch_size'], shuffle_train=True,
    )

    # Weights with SELL усиление
    raw_weights = compute_class_weights(finetune_data['y_train'], num_classes=3)
    sell_mult = config['sell_multiplier']
    finetune_weights = raw_weights.clone()
    finetune_weights[2] *= sell_mult
    finetune_weights = finetune_weights / finetune_weights.mean()
    logger.info(f"Веса классов (SELL×{sell_mult}): "
                f"HOLD={finetune_weights[0]:.4f}, "
                f"BUY={finetune_weights[1]:.4f}, "
                f"SELL={finetune_weights[2]:.4f}")

    # Create model
    model = create_model(config['model_class'], config)
    total_params = sum(p.numel() for p in model.parameters())
    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
    logger.info(f"Параметры: {total_params:,} всего, {trainable_params:,} обучаемых")

    # ─── Pre-train ──────────────────────────────────────────────────────
    logger.info("\n--- ЭТАП A: PRE-TRAIN ---")
    if Path(checkpoint_path).exists():
        logger.info(f"Загружаем чекпоинт: {checkpoint_path}")
        checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
        model.load_state_dict(checkpoint['model_state_dict'])
        pretrain_epochs_done = checkpoint.get('pretrain_epochs', 0)
    else:
        from src.ml.train.trainer import train_model as tm
        # Pre-train со стандартным CE (данные сбалансированы)
        pretrain_result = tm(
            model=model,
            train_loader=pretrain_loader,
            val_loader=pretrain_val_loader,
            n_epochs=config['pretrain_epochs'],
            learning_rate=config['pretrain_lr'],
            weight_decay=config['weight_decay'],
            patience=config['pretrain_patience'],
            gradient_clip=1.0,
            device='cpu',
            model_name=f'{model_name}_pretrain',
            save_dir=SAVE_DIR,
            class_weights=None,
            verbose=True,
        )
        logger.info(f"Pre-train: {pretrain_result.total_epochs} эпох, "
                     f"val_acc={pretrain_result.val_metrics['accuracy']:.4f}")
        pretrain_epochs_done = pretrain_result.total_epochs

        # Сохраняем чекпоинт
        torch.save({
            'model_state_dict': model.state_dict(),
            'val_loss': pretrain_result.best_val_loss,
            'val_accuracy': pretrain_result.val_metrics['accuracy'],
            'pretrain_epochs': pretrain_result.total_epochs,
        }, checkpoint_path)
        logger.info(f"Чекпоинт сохранён: {checkpoint_path}")

    # ─── Fine-tune с Focal Loss ─────────────────────────────────────────
    logger.info("\n--- ЭТАП B: FINE-TUNE С FOCAL LOSS ---")

    focal_criterion = WeightedFocalLoss(
        gamma=config['focal_gamma'],
        class_weights=finetune_weights,
        reduction='mean',
    )

    # Патчим train_model для поддержки custom_criterion
    from torch.optim import AdamW
    from torch.optim.lr_scheduler import CosineAnnealingWarmRestarts
    from src.ml.train.trainer import train_epoch, validate, TrainingResult

    optimizer = AdamW(model.parameters(), lr=config['finetune_lr'],
                      weight_decay=config['weight_decay'])
    scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=10, T_mult=2)
    criterion = focal_criterion.to(device)
    model = model.to(device)

    history = {'train_loss': [], 'val_loss': [], 'val_accuracy': []}
    best_val_loss = float('inf')
    early_stop_counter = 0
    best_state_dict = None
    best_epoch = 0
    start_time = time.time()

    save_path = Path(SAVE_DIR) / model_name
    save_path.mkdir(parents=True, exist_ok=True)

    for epoch in range(config['finetune_epochs']):
        train_loss = train_epoch(model, finetune_loader, optimizer, criterion, device, 1.0)
        val_metrics = validate(model, val_loader, criterion, device)
        val_loss = val_metrics['loss']
        val_accuracy = val_metrics['accuracy']
        scheduler.step()

        history['train_loss'].append(train_loss)
        history['val_loss'].append(val_loss)
        history['val_accuracy'].append(val_accuracy)

        if val_loss < best_val_loss:
            best_val_loss = val_loss
            best_epoch = epoch
            early_stop_counter = 0
            best_state_dict = model.state_dict()
            torch.save({
                'epoch': epoch,
                'model_state_dict': best_state_dict,
                'val_loss': best_val_loss,
                'val_accuracy': val_accuracy,
                'history': history,
            }, str(save_path / 'model_checkpoint.pt'))
        else:
            early_stop_counter += 1

        if (epoch + 1) % 5 == 0:
            logger.info(
                f"Эпоха {epoch+1:3d}/{config['finetune_epochs']} | "
                f"Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | "
                f"Val Acc: {val_accuracy:.4f} | ES: {early_stop_counter}/{config['finetune_patience']}"
            )

        if early_stop_counter >= config['finetune_patience']:
            logger.info(f"Early stopping на эпохе {epoch+1}, лучшая: {best_epoch+1}")
            break

    training_time = time.time() - start_time
    model.load_state_dict(best_state_dict)

    # Финальные метрики
    final_val = validate(model, val_loader, criterion, device, return_probs=True)
    test_result = validate(model, test_loader, criterion, device, return_probs=True) if test_loader else None

    logger.info(f"\nFine-tune: {epoch+1} эпох, best_epoch={best_epoch+1}, "
                f"val_acc={final_val['accuracy']:.4f}")
    logger.info(f"Test acc: {test_result['accuracy']:.4f}" if test_result else "")

    # ─── Threshold calibration ──────────────────────────────────────────
    logger.info("\n--- КАЛИБРОВКА ПОРОГОВ ---")
    model.eval()
    all_probs, all_labels = [], []
    with torch.no_grad():
        for bx, by in val_loader:
            out = model(bx.to(device))
            all_probs.append(torch.softmax(out, dim=1).cpu().numpy())
            all_labels.append(by.numpy())
    all_probs = np.concatenate(all_probs)
    all_labels = np.concatenate(all_labels)

    best_f1, best_t = 0.0, 0.4
    for t in np.arange(0.1, 0.95, 0.05):
        preds = []
        for probs in all_probs:
            cls = int(probs.argmax())
            if cls in [1, 2] and probs[cls] < t:
                cls = 0
            preds.append(cls)
        preds = np.array(preds)
        metrics = calculate_all_metrics(all_labels, preds, num_classes=3)
        if metrics['macro_f1'] > best_f1:
            best_f1 = metrics['macro_f1']
            best_t = t

    thresholds = {0: 0.0, 1: best_t, 2: best_t}
    logger.info(f"Порог: {best_t:.2f}, Macro F1 (val): {best_f1:.4f}")

    # Применяем на тесте
    if test_loader:
        model.eval()
        all_probs_t, all_labels_t = [], []
        with torch.no_grad():
            for bx, by in test_loader:
                out = model(bx.to(device))
                all_probs_t.append(torch.softmax(out, dim=1).cpu().numpy())
                all_labels_t.append(by.numpy())
        all_probs_t = np.concatenate(all_probs_t)
        all_labels_t = np.concatenate(all_labels_t)

        cal_preds = []
        for probs in all_probs_t:
            cls = int(probs.argmax())
            if cls in [1, 2] and probs[cls] < best_t:
                cls = 0
            cal_preds.append(cls)
        cal_preds = np.array(cal_preds)
        cal_metrics = calculate_all_metrics(all_labels_t, cal_preds, num_classes=3)

        logger.info(f"\nТест (калиброванный): acc={cal_metrics['accuracy']:.4f}, "
                    f"macro_f1={cal_metrics['macro_f1']:.4f}")
        cm = cal_metrics['confusion_matrix']
        logger.info(f"CM: HOLD→{cm[0]} BUY→{cm[1]} SELL→{cm[2]}")

    # ─── Сохраняем ──────────────────────────────────────────────────────
    logger.info("\n--- СОХРАНЕНИЕ ---")
    torch.save(model.state_dict(), str(save_path / 'model.pt'))

    # Собираем метрики
    val_preds = final_val['predictions']
    val_labels_arr = final_val['labels']
    val_metrics = {
        'val_loss': float(final_val['loss']),
        'val_accuracy': float(final_val['accuracy']),
        'val_macro_f1': float(calculate_macro_f1(val_labels_arr, val_preds)),
        'val_confusion_matrix': final_val['confusion_matrix'].tolist(),
        'best_epoch': best_epoch + 1,
        'total_epochs': epoch + 1,
        'pretrain_epochs_done': pretrain_epochs_done,
        'training_time_sec': training_time,
    }

    if test_result:
        test_labels_arr = test_result['labels']
        test_preds_arr = test_result['predictions']
        val_metrics['test_loss'] = float(test_result['loss'])
        val_metrics['test_accuracy'] = float(test_result['accuracy'])
        val_metrics['test_macro_f1'] = float(calculate_macro_f1(test_labels_arr, test_preds_arr))
        val_metrics['test_confusion_matrix'] = test_result['confusion_matrix'].tolist()
        val_metrics['calibrated_test_accuracy'] = float(cal_metrics['accuracy'])
        val_metrics['calibrated_test_macro_f1'] = float(cal_metrics['macro_f1'])
        val_metrics['calibrated_test_confusion_matrix'] = cal_metrics['confusion_matrix']

    # Сохраняем конфиг
    config_data = {
        'model_type': config['model_class'],
        'params': {
            'input_size': config['input_size'],
            'hidden_size': config['hidden_size'],
            'num_layers': config['num_layers'],
            'dropout': config['dropout'],
            'bidirectional': config.get('bidirectional', False),
            'focal_gamma': config['focal_gamma'],
            'weight_decay': config['weight_decay'],
            'pretrain_lr': config['pretrain_lr'],
            'finetune_lr': config['finetune_lr'],
            'pretrain_epochs': config['pretrain_epochs'],
            'finetune_epochs': config['finetune_epochs'],
            'batch_size': config['batch_size'],
            'seed': config['seed'],
            'sell_multiplier': config['sell_multiplier'],
            'loss_function': 'WeightedFocalLoss',
        },
        'feature_cols': list(finetune_data.get('feature_names', [])),
        'n_classes': 3,
        'class_names': {0: 'HOLD', 1: 'BUY', 2: 'SELL'},
        'metrics': val_metrics,
        'ticker': TICKER,
        'timeframe': TIMEFRAME,
        'training_type': 'pretrain_then_finetune_v4',
        'pretrain_data': PRETRAIN_DATA,
        'finetune_data': FINETUNE_DATA,
        'calibrated_thresholds': thresholds,
        'loss_function': 'WeightedFocalLoss',
        'focal_gamma': config['focal_gamma'],
    }

    config_path = save_path / 'config.json'
    config_path.write_text(json.dumps(config_data, indent=2, ensure_ascii=False))

    # Сохраняем нормализатор
    if 'scaler_mean' in finetune_data and 'scaler_scale' in finetune_data:
        norm = {
            'mean_': finetune_data['scaler_mean'].tolist(),
            'std_': finetune_data['scaler_scale'].tolist(),
        }
        (save_path / 'normalizer.json').write_text(json.dumps(norm, indent=2))

    # Сохраняем фичи
    if 'feature_names' in finetune_data:
        (save_path / 'features.json').write_text(
            json.dumps(list(finetune_data['feature_names']), indent=2))

    # Регистрируем
    registry = ModelRegistry()
    registry.register(
        model_name=model_name,
        model_type=config['model_class'],
        ticker=TICKER,
        timeframe=TIMEFRAME,
        params=config_data['params'],
        metrics=val_metrics,
        model_path=str(save_path),
    )

    logger.info(f"\n✅ Модель {model_name} сохранена и зарегистрирована")
    logger.info(f"   Test acc: {val_metrics.get('test_accuracy', 0)*100:.2f}%")
    logger.info(f"   Calibrated test acc: {val_metrics.get('calibrated_test_accuracy', 0)*100:.2f}%")
    logger.info(f"   SELL recall: см. confusion matrix")
    logger.info(f"   Время: {training_time:.1f} сек")

    return True


def print_confusion_matrix(cm, title="Confusion Matrix"):
    """Вывести confusion matrix."""
    logger.info(f"\n{title}:")
    logger.info(f"{'':>10} {'HOLD':>8} {'BUY':>8} {'SELL':>8}")
    for i, name in enumerate(['HOLD', 'BUY', 'SELL']):
        row = f"{name:<10}"
        for j in range(3):
            row += f"{cm[i][j]:>8}"
        logger.info(row)

    # SELL recall/precision
    sell_true = sum(cm[i][2] for i in range(3))
    sell_correct = cm[2][2]
    sell_pred = sum(cm[2][j] for j in range(3))
    sell_recall = sell_correct / sell_true * 100 if sell_true > 0 else 0
    sell_prec = sell_correct / sell_pred * 100 if sell_pred > 0 else 0
    logger.info(f"   SELL recall: {sell_recall:.1f}% | SELL precision: {sell_prec:.1f}%")

    return sell_recall, sell_prec


def main():
    models_to_train = sys.argv[1] if len(sys.argv) > 1 else 'all'

    if models_to_train == 'all':
        model_list = list(MODEL_CONFIGS.keys())
    elif models_to_train in MODEL_CONFIGS:
        model_list = [models_to_train]
    else:
        logger.error(f"Неизвестная модель: {models_to_train}")
        logger.info(f"Доступны: {list(MODEL_CONFIGS.keys())}, или 'all'")
        sys.exit(1)

    logger.info("=" * 70)
    logger.info(f"ЗАПУСК ОБУЧЕНИЯ {len(model_list)} МОДЕЛЕЙ ДЛЯ АНСАМБЛЯ X5")
    logger.info(f"Модели: {model_list}")
    logger.info("=" * 70)

    for model_name in model_list:
        success = train_ensemble_model(model_name)
        if not success:
            logger.error(f"Ошибка обучения {model_name}")

    # После всех тренировок — сводка
    logger.info("\n" + "=" * 70)
    logger.info("СВОДКА ОБУЧЕННЫХ МОДЕЛЕЙ")
    logger.info("=" * 70)

    for model_name in model_list:
        try:
            config = MODEL_CONFIGS[model_name]
            cfg_path = Path(SAVE_DIR) / model_name / 'config.json'
            if cfg_path.exists():
                cfg = json.loads(cfg_path.read_text())
                met = cfg['metrics']
                cm = met.get('calibrated_test_confusion_matrix', met.get('test_confusion_matrix', []))
                logger.info(f"\n{model_name} ({config['desc']}):")
                logger.info(f"  Test acc: {met.get('test_accuracy', 0)*100:.2f}% | "
                           f"Calibrated: {met.get('calibrated_test_accuracy', 0)*100:.2f}%")
                if cm and len(cm) == 3:
                    print_confusion_matrix(cm, "Test CM (calibrated)")
        except Exception as e:
            logger.warning(f"Не удалось прочитать метрики {model_name}: {e}")

    logger.info("\n" + "=" * 70)
    logger.info("ГОТОВО! Ансамбль: x5_gru_v3 + x5_lstm_v3_focal + x5_gru_v4")
    logger.info("=" * 70)


if __name__ == '__main__':
    main()
