"""
Скрипт обучения D1 Transformer для X5.

Архитектурно отличается от GRU/LSTM (self-attention вместо рекуррентности),
что добавляет диверсификацию в ансамбль.

Использует существующие HDF5 датасеты (multi_ticker_d1_pretrain.h5 +
x5_d1_augmented.h5) для честного сравнения с LSTM/GRU.

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

import json
import logging
import sys
import time
from pathlib import Path
from typing import 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.transformer import D1Transformer
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,
    calculate_all_metrics,
)

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_d1_transformer_v1'


def calibrate_thresholds(
    model: nn.Module,
    val_loader: torch.utils.data.DataLoader,
    device: torch.device,
) -> float:
    """Подобрать порог уверенности для BUY/SELL на валидации."""
    model.eval()
    all_probs, all_labels = [], []
    for batch in val_loader:
        x, y = batch
        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_macro_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  # → HOLD
        macro_f1 = calculate_macro_f1(all_labels, preds)
        if macro_f1 > best_macro_f1:
            best_macro_f1 = macro_f1
            best_threshold = threshold
    return best_threshold


def evaluate_test(
    model: nn.Module,
    test_loader: torch.utils.data.DataLoader,
    device: torch.device,
    threshold: float = 0.0,
) -> dict:
    """Оценка модели на тестовом датасете."""
    model.eval()
    all_probs, all_labels = [], []
    for batch in test_loader:
        x, y = batch
        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)

    # С threshold-фильтрацией
    preds_raw = np.argmax(all_probs, axis=1)
    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] < threshold:
            preds_cal[i] = 0

    # Метрики
    cm = calculate_confusion_matrix(all_labels, preds_cal)
    sell_recall = cm[2, 2] / max(1, cm[2].sum())
    sell_precision = cm[2, 2] / max(1, cm[:, 2].sum())

    return {
        'accuracy': float(calculate_accuracy(all_labels, preds_cal)),
        'macro_f1': float(calculate_macro_f1(all_labels, preds_cal)),
        'confusion_matrix': cm.tolist(),
        'sell_recall': float(sell_recall),
        'sell_precision': float(sell_precision),
        'threshold': threshold,
        'distribution_raw': {
            'HOLD': int((preds_raw == 0).sum()),
            'BUY': int((preds_raw == 1).sum()),
            'SELL': int((preds_raw == 2).sum()),
        },
        'distribution_calibrated': {
            'HOLD': int((preds_cal == 0).sum()),
            'BUY': int((preds_cal == 1).sum()),
            'SELL': int((preds_cal == 2).sum()),
        },
    }


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

    set_seed(42)

    # ─── Конфигурация ────────────────────────────────────────────────────
    config = {
        'input_size': 61,
        'd_model': 64,
        'nhead': 4,
        'num_layers': 2,
        'dim_feedforward': 128,
        'dropout': 0.4,
        'activation': 'gelu',
        'focal_gamma': 2.0,
        'weight_decay': 1e-3,
        'sell_multiplier': 2.0,
        '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,
    }

    logger.info(f"Конфигурация: {json.dumps(config, indent=2)}")
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    logger.info(f"Устройство: {device}")

    # ─── Загрузка датасетов ─────────────────────────────────────────────
    if not Path(PRETRAIN_DATA).exists():
        logger.error(f"Pre-train датасет не найден: {PRETRAIN_DATA}")
        return
    if not Path(FINETUNE_DATA).exists():
        logger.error(f"Fine-tune датасет не найден: {FINETUNE_DATA}")
        return

    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 = D1Transformer(
        input_size=config['input_size'],
        d_model=config['d_model'],
        nhead=config['nhead'],
        num_layers=config['num_layers'],
        dim_feedforward=config['dim_feedforward'],
        dropout=config['dropout'],
        activation=config['activation'],
    )
    total_params = sum(p.numel() for p in model.parameters())
    logger.info(f"Параметров: {total_params:,}")
    logger.info(f"Архитектура: Transformer(d_model={config['d_model']}, "
                f"nhead={config['nhead']}, layers={config['num_layers']}, "
                f"dropout={config['dropout']})")

    # ─── Pre-train со стандартным CE ────────────────────────────────────
    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 на X5 с Focal Loss ────────────────────────────────────
    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,
    )

    # Веса классов с усилением 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}")

    # Focal Loss на fine-tune
    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
    history = {'train_loss': [], 'val_loss': [], 'val_acc': []}

    for epoch in range(1, config['finetune_epochs'] + 1):
        # Train
        model.train()
        total_loss = 0.0
        for batch in finetune_loader:
            x, y = batch
            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 batch in val_loader:
            x, y = batch
            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()
        history['train_loss'].append(train_loss)
        history['val_loss'].append(val_loss)
        history['val_acc'].append(val_acc)

        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)
    best_threshold = calibrate_thresholds(model, val_loader, device)
    logger.info(f"  Best threshold: {best_threshold:.2f}")

    # ─── Оценка на тесте ─────────────────────────────────────────────────
    logger.info(f"\n{'='*70}")
    logger.info("ЭТАП 4: Оценка на тесте")
    logger.info("=" * 70)
    test_results = evaluate_test(model, test_loader, device, threshold=best_threshold)

    logger.info(f"  Test accuracy: {test_results['accuracy']*100:.2f}%")
    logger.info(f"  Test macro F1: {test_results['macro_f1']:.4f}")
    logger.info(f"  SELL recall: {test_results['sell_recall']*100:.1f}%")
    logger.info(f"  SELL precision: {test_results['sell_precision']*100:.1f}%")
    logger.info(f"  Распред (raw):       {test_results['distribution_raw']}")
    logger.info(f"  Распред (calibrated): {test_results['distribution_calibrated']}")

    cm = test_results['confusion_matrix']
    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.update({
        'model_name': MODEL_NAME,
        'model_type': 'D1Transformer',
        'ticker': TICKER,
        'timeframe': TIMEFRAME,
        'total_params': total_params,
    })
    config['metrics'] = {
        'test_accuracy': test_results['accuracy'],
        'test_macro_f1': test_results['macro_f1'],
        'test_confusion_matrix': cm,
        'sell_recall': test_results['sell_recall'],
        'sell_precision': test_results['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='D1Transformer',
        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_results['accuracy']*100:.2f}%")
    logger.info(f"  SELL recall:   {test_results['sell_recall']*100:.1f}%")
    logger.info(f"  SELL precision: {test_results['sell_precision']*100:.1f}%")
    if test_results['sell_recall'] > 0:
        logger.info(f"  ✅ SELL recall > 0% — можно добавить в ансамбль")
    else:
        logger.warning(f"  ❌ SELL recall = 0% — модель не пригодна для ансамбля")
    logger.info("=" * 70)


if __name__ == '__main__':
    main()
