"""
Скрипт обучения BiGRU + Attention для X5.

Архитектурные отличия от существующих моделей:
  - Bidirectional GRU (обрабатывает в обе стороны)
  - Attention вместо последнего timestep
  - Добавляет диверсификацию в ансамбль

Использует pre-train (9 MOEX tickers) + fine-tune (X5 augmented).

Usage:
    python src/ml/train_x5_bigru_attn.py
"""

import json
import logging
import sys
import time
from pathlib import Path

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.gru_attention import BiGRUAttention
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,
)
from src.ml.train.losses import WeightedFocalLoss
from src.ml.train.metrics import (
    calculate_accuracy, calculate_confusion_matrix, calculate_macro_f1,
)

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_NAME = 'x5_bigru_attn_v1'


def main():
    logger.info("=" * 70)
    logger.info(f"ОБУЧЕНИЕ BiGRU + Attention: {MODEL_NAME}")
    logger.info("=" * 70)

    set_seed(42)
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    logger.info(f"Устройство: {device}")

    config = {
        'input_size': 61,
        'hidden_size': 64,
        'num_layers': 2,
        'dropout': 0.4,
        'focal_gamma': 2.0,
        'weight_decay': 1e-3,
        'sell_multiplier': 2.0,
        'pretrain_lr': 1e-3,
        'finetune_lr': 5e-5,
        'pretrain_epochs': 30,
        'finetune_epochs': 50,
        'pretrain_patience': 10,
        'finetune_patience': 20,
        'batch_size': 32,
        'seed': 42,
    }

    # ─── Pre-train ──────────────────────────────────────────────────────
    logger.info(f"\n{'='*70}")
    logger.info("ЭТАП 1: Pre-train на мульти-тикерах")
    logger.info("=" * 70)

    pretrain_data = load_hdf5_dataset(PRETRAIN_DATA)
    pretrain_loader, pretrain_val_loader, _ = create_dataloaders(
        pretrain_data, batch_size=config['batch_size'],
        shuffle_train=True, val_split=0.1,
    )

    model = BiGRUAttention(
        input_size=config['input_size'],
        hidden_size=config['hidden_size'],
        num_layers=config['num_layers'],
        dropout=config['dropout'],
    )
    total_params = sum(p.numel() for p in model.parameters())
    logger.info(f"Параметров: {total_params:,}")
    logger.info(f"Архитектура: BiGRU-Attention(hidden={config['hidden_size']}, "
                f"layers={config['num_layers']}, dropout={config['dropout']})")

    checkpoint_path = Path(SAVE_DIR) / f'{MODEL_NAME}_pretrained.pt'
    if 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'])
    else:
        pretrain_result = train_model(
            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,
        )
        torch.save({
            'model_state_dict': model.state_dict(),
            'pretrain_epochs': pretrain_result.total_epochs,
        }, checkpoint_path)
        logger.info(f"Pre-train завершён: {pretrain_result.total_epochs} эпох, "
                    f"best val_loss={pretrain_result.best_val_loss:.4f}")

    # ─── Fine-tune ──────────────────────────────────────────────────────
    logger.info(f"\n{'='*70}")
    logger.info("ЭТАП 2: Fine-tune на X5 (Focal Loss)")
    logger.info("=" * 70)

    finetune_data = load_hdf5_dataset(FINETUNE_DATA)
    finetune_loader, val_loader, test_loader = create_dataloaders(
        finetune_data, batch_size=config['batch_size'], shuffle_train=True,
    )

    # Weights
    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}")

    criterion = WeightedFocalLoss(gamma=config['focal_gamma'], class_weights=finetune_weights)
    optimizer = torch.optim.AdamW(
        model.parameters(), lr=config['finetune_lr'],
        weight_decay=config['weight_decay'],
    )
    scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(
        optimizer, T_0=10, T_mult=2,
    )

    model = model.to(device)
    best_val_loss = float('inf')
    best_state = None
    patience_counter = 0

    for epoch in range(1, config['finetune_epochs'] + 1):
        # Train
        model.train()
        total_loss = 0.0
        for x, y in finetune_loader:
            x, y = x.to(device), y.to(device)
            optimizer.zero_grad()
            outputs = model(x)
            loss = criterion(outputs, y)
            loss.backward()
            nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            optimizer.step()
            total_loss += loss.item() * x.size(0)
        train_loss = total_loss / len(finetune_loader.dataset)

        # Val
        model.eval()
        val_loss = 0.0
        all_preds, all_labels = [], []
        for x, y in val_loader:
            x, y = x.to(device), y.to(device)
            with torch.no_grad():
                outputs = model(x)
                loss = criterion(outputs, y)
            val_loss += loss.item() * x.size(0)
            preds = torch.argmax(outputs, dim=1)
            all_preds.extend(preds.cpu().numpy())
            all_labels.extend(y.cpu().numpy())
        val_loss /= max(1, len(val_loader.dataset))
        val_acc = calculate_accuracy(np.array(all_labels), np.array(all_preds))

        scheduler.step()

        if val_loss < best_val_loss:
            best_val_loss = val_loss
            best_state = model.state_dict().copy()
            patience_counter = 0
        else:
            patience_counter += 1

        if epoch % 5 == 0 or epoch == 1:
            logger.info(f"  Эпоха {epoch:3d}/{config['finetune_epochs']} | "
                        f"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | "
                        f"Val Acc: {val_acc:.4f} | patience: {patience_counter}")

        if patience_counter >= config['finetune_patience']:
            logger.info(f"  Early stopping на эпохе {epoch}")
            break

    if best_state:
        model.load_state_dict(best_state)
    logger.info(f"  ✅ Fine-tune завершён. Лучшая val_loss: {best_val_loss:.4f}")

    # ─── Калибровка ─────────────────────────────────────────────────────
    logger.info(f"\n{'='*70}")
    logger.info("ЭТАП 3: Калибровка порога")
    logger.info("=" * 70)

    model.eval()
    all_probs, all_labels = [], []
    for x, y in val_loader:
        x = x.to(device)
        with torch.no_grad():
            outputs = model(x)
            probs = torch.softmax(outputs, dim=1)
        all_probs.extend(probs.cpu().numpy())
        all_labels.extend(y.cpu().numpy())
    all_probs = np.array(all_probs)
    all_labels = np.array(all_labels)

    best_threshold = 0.0
    best_f1 = 0.0
    for threshold in np.arange(0.0, 0.7, 0.05):
        preds = np.argmax(all_probs, axis=1)
        for i in range(len(preds)):
            cls = preds[i]
            if cls in [1, 2] and all_probs[i, cls] < threshold:
                preds[i] = 0
        macro_f1 = calculate_macro_f1(all_labels, preds)
        if macro_f1 > best_f1:
            best_f1 = macro_f1
            best_threshold = threshold
    logger.info(f"  Best threshold: {best_threshold:.2f}")

    # ─── Test ────────────────────────────────────────────────────────────
    logger.info(f"\n{'='*70}")
    logger.info("ЭТАП 4: Оценка на тесте")
    logger.info("=" * 70)

    model.eval()
    all_probs, all_labels = [], []
    for x, y in test_loader:
        x = x.to(device)
        with torch.no_grad():
            outputs = model(x)
            probs = torch.softmax(outputs, dim=1)
        all_probs.extend(probs.cpu().numpy())
        all_labels.extend(y.cpu().numpy())
    all_probs = np.array(all_probs)
    all_labels = np.array(all_labels)

    # Raw
    preds_raw = np.argmax(all_probs, axis=1)
    logger.info(f"  Raw predictions: HOLD={int((preds_raw==0).sum())} "
                f"BUY={int((preds_raw==1).sum())} SELL={int((preds_raw==2).sum())}")

    # Calibrated
    preds_cal = preds_raw.copy()
    for i in range(len(preds_cal)):
        cls = preds_cal[i]
        if cls in [1, 2] and all_probs[i, cls] < best_threshold:
            preds_cal[i] = 0
    logger.info(f"  Calibrated (t={best_threshold:.2f}): HOLD={int((preds_cal==0).sum())} "
                f"BUY={int((preds_cal==1).sum())} SELL={int((preds_cal==2).sum())}")

    cm = calculate_confusion_matrix(all_labels, preds_cal)
    test_acc = calculate_accuracy(all_labels, preds_cal)
    sell_recall = cm[2, 2] / max(1, cm[2].sum())
    sell_precision = cm[2, 2] / max(1, cm[:, 2].sum())

    logger.info(f"  Test accuracy: {test_acc*100:.2f}%")
    logger.info(f"  Test macro F1: {calculate_macro_f1(all_labels, preds_cal):.4f}")
    logger.info(f"  SELL recall: {sell_recall*100:.1f}%")
    logger.info(f"  SELL precision: {sell_precision*100:.1f}%")
    logger.info("  Confusion Matrix:")
    logger.info(f"                HOLD    BUY    SELL")
    for i, label in enumerate(['HOLD', 'BUY', 'SELL']):
        logger.info(f"  {label:>8}  {cm[i][0]:6d} {cm[i][1]:6d} {cm[i][2]:6d}")

    # ─── Сохранение ─────────────────────────────────────────────────────
    logger.info(f"\n{'='*70}")
    logger.info("СОХРАНЕНИЕ")
    logger.info("=" * 70)

    saved_dir = Path(SAVE_DIR) / MODEL_NAME
    saved_dir.mkdir(parents=True, exist_ok=True)
    torch.save(model.state_dict(), saved_dir / 'model.pt')

    config['model_name'] = MODEL_NAME
    config['model_type'] = 'BiGRUAttention'
    config['ticker'] = TICKER
    config['timeframe'] = TIMEFRAME
    config['total_params'] = total_params
    config['metrics'] = {
        'test_accuracy': float(test_acc),
        'test_macro_f1': float(calculate_macro_f1(all_labels, preds_cal)),
        'test_confusion_matrix': cm.tolist(),
        'sell_recall': float(sell_recall),
        'sell_precision': float(sell_precision),
        'val_loss': float(best_val_loss),
    }
    config['calibrated_thresholds'] = {'1': best_threshold, '2': best_threshold}

    (saved_dir / 'config.json').write_text(
        json.dumps(config, indent=2, ensure_ascii=False)
    )

    registry = ModelRegistry()
    registry.register(
        model_name=MODEL_NAME,
        model_type='BiGRUAttention',
        ticker=TICKER,
        timeframe=TIMEFRAME,
        params=config,
        metrics=config['metrics'],
        model_path=str(saved_dir),
    )

    # ─── Итог ───────────────────────────────────────────────────────────
    logger.info("\n" + "=" * 70)
    logger.info(f"ИТОГ: {MODEL_NAME}")
    logger.info(f"  Параметров: {total_params:,}")
    logger.info(f"  Test accuracy: {test_acc*100:.2f}%")
    logger.info(f"  SELL recall:   {sell_recall*100:.1f}%")
    logger.info(f"  SELL precision: {sell_precision*100:.1f}%")
    if sell_recall > 0:
        logger.info(f"  ✅ SELL recall > 0% — можно добавить в ансамбль")
    else:
        logger.warning(f"  ❌ SELL recall = 0% — модель не пригодна для ансамбля")
    logger.info("=" * 70)


if __name__ == '__main__':
    main()
