"""
Обучение MultiTimeframeFusion модели (H1 + D1).

Архитектура: MTFBasic (GRU encoders + fusion concat + classifier).
Loss: WeightedFocalLoss (γ=2.0) — proven effective на X5.
Этапы:
  1. Загрузка H1 и D1 данных → alignment по timestamp
  2. Создание PyTorch Dataset + DataLoader
  3. Pre-train на multi_ticker (опционально) + Fine-tune на X5
  4. Threshold calibration
  5. Сохранение в ModelRegistry

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

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

import numpy as np
import torch
import torch.nn as nn
from torch.optim import AdamW
from torch.optim.lr_scheduler import CosineAnnealingWarmRestarts
from torch.utils.data import DataLoader

# Add project root
PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent
sys.path.insert(0, str(PROJECT_ROOT))

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

from src.ml.data.mtf_dataset import MultiTimeframeFusionDataset
from src.ml.models.mtf import MTFBasic
from src.ml.models.registry import ModelRegistry
from src.ml.train.losses import WeightedFocalLoss
from src.ml.train.metrics import (
    calculate_accuracy,
    calculate_confusion_matrix,
    calculate_macro_f1,
)


def set_seed(seed: int = 42):
    import random
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)


def compute_class_weights(y: np.ndarray, sell_multiplier: float = 2.0) -> torch.Tensor:
    """Вычислить веса классов для WeightedFocalLoss."""
    classes, counts = np.unique(y, return_counts=True)
    total = len(y)
    weights = total / (len(classes) * counts.astype(float))
    # Усиление SELL
    weights = torch.FloatTensor(weights)
    if 2 in classes:
        idx = list(classes).index(2)
        weights[idx] *= sell_multiplier
    return weights


def train_epoch_mtf(
    model: nn.Module,
    dataloader: DataLoader,
    optimizer: torch.optim.Optimizer,
    criterion: nn.Module,
    device: torch.device,
) -> float:
    """Обучить MTF модель на одну эпоху."""
    model.train()
    total_loss = 0.0
    num_batches = 0

    for batch in dataloader:
        h1_x, d1_x, batch_y = batch
        h1_x = h1_x.to(device)
        d1_x = d1_x.to(device)
        batch_y = batch_y.to(device)

        optimizer.zero_grad()
        outputs = model(h1_x, d1_x)
        # Клиппинг логитов для предотвращения NaN в Focal Loss
        outputs = torch.clamp(outputs, min=-15.0, max=15.0)
        loss = criterion(outputs, batch_y)
        loss.backward()
        nn.utils.clip_grad_norm_(model.parameters(), 1.0)
        optimizer.step()

        total_loss += loss.item() * h1_x.size(0)
        num_batches += 1

    return total_loss / len(dataloader.dataset)


@torch.no_grad()
def validate_mtf(
    model: nn.Module,
    dataloader: DataLoader,
    criterion: nn.Module,
    device: torch.device,
) -> Dict:
    """Валидация MTF модели."""
    model.eval()
    total_loss = 0.0
    all_preds, all_labels = [], []

    for batch in dataloader:
        h1_x, d1_x, batch_y = batch
        h1_x = h1_x.to(device)
        d1_x = d1_x.to(device)
        batch_y = batch_y.to(device)

        outputs = model(h1_x, d1_x)
        loss = criterion(outputs, batch_y)
        total_loss += loss.item() * h1_x.size(0)

        preds = torch.argmax(outputs, dim=1)
        all_preds.extend(preds.cpu().numpy())
        all_labels.extend(batch_y.cpu().numpy())

    all_preds = np.array(all_preds)
    all_labels = np.array(all_labels)

    return {
        'loss': total_loss / len(dataloader.dataset),
        'accuracy': calculate_accuracy(all_labels, all_preds),
        'cm': calculate_confusion_matrix(all_labels, all_preds),
        'predictions': all_preds,
        'labels': all_labels,
    }


def calibrate_thresholds_mtf(
    model: nn.Module,
    dataloader: DataLoader,
    device: torch.device,
) -> Tuple[float, Dict[int, float]]:
    """
    Калибровка порогов для BUY/SELL на валидации.
    Возвращает лучший threshold (для macro F1) и словарь порогов.
    """
    model.eval()
    all_probs, all_labels = [], []

    for batch in dataloader:
        h1_x, d1_x, batch_y = batch
        h1_x = h1_x.to(device)
        d1_x = d1_x.to(device)

        with torch.no_grad():
            outputs = model(h1_x, d1_x)
            probs = torch.softmax(outputs, dim=1)
        all_probs.extend(probs.detach().cpu().numpy())
        all_labels.extend(batch_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
        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, {0: 0.0, 1: best_threshold, 2: best_threshold}


def evaluate_test_mtf(
    model: nn.Module,
    dataloader: DataLoader,
    device: torch.device,
    thresholds: Dict[int, float],
) -> Dict:
    """Оценка на тесте с порогами."""
    model.eval()
    all_probs, all_labels, all_preds_raw = [], [], []

    for batch in dataloader:
        h1_x, d1_x, batch_y = batch
        h1_x = h1_x.to(device)
        d1_x = d1_x.to(device)

        with torch.no_grad():
            outputs = model(h1_x, d1_x)
            probs = torch.softmax(outputs, dim=1)
            preds = torch.argmax(outputs, dim=1)

        all_probs.extend(probs.detach().cpu().numpy())
        all_labels.extend(batch_y.cpu().numpy())
        all_preds_raw.extend(preds.detach().cpu().numpy())

    all_probs = np.array(all_probs)
    all_labels = np.array(all_labels)

    # Применяем пороги
    preds_calibrated = np.argmax(all_probs, axis=1)
    for i in range(len(preds_calibrated)):
        cls = preds_calibrated[i]
        if cls in [1, 2] and all_probs[i, cls] < thresholds.get(cls, 0.0):
            preds_calibrated[i] = 0

    cm = calculate_confusion_matrix(all_labels, preds_calibrated)
    sell_recall = cm[2, 2] / max(1, cm[2].sum()) * 100
    sell_precision = cm[2, 2] / max(1, cm[:, 2].sum()) * 100

    return {
        'accuracy': float(calculate_accuracy(all_labels, preds_calibrated)),
        'macro_f1': float(calculate_macro_f1(all_labels, preds_calibrated)),
        'confusion_matrix': cm.tolist(),
        'sell_recall': float(sell_recall),
        'sell_precision': float(sell_precision),
        'thresholds': thresholds,
    }


def main():
    logger.info("=" * 70)
    logger.info("MULTI-TIMEFRAME FUSION (H1 + D1) — ОБУЧЕНИЕ")
    logger.info("=" * 70)

    device = torch.device('cpu')
    MODEL_NAME = 'x5_mtf_v1'
    TICKER = 'X5'

    # ── 1. Создание датасета ──────────────────────────────────────
    logger.info("\n[1/5] Создание MTF датасета...")
    dataset = MultiTimeframeFusionDataset(
        ticker=TICKER,
        h1_limit=5000,
        d1_limit=2000,
        h1_seq_len=120,   # увеличили 40→120 (15 торговых дней H1)
        d1_seq_len=30,
        forecast_horizon=1,
        threshold_pct=0.5,
        val_split=0.15,
        test_split=0.15,
        top_k_features=30,  # feature selection: 61→30
    )

    batch_size = 32
    train_loader = DataLoader(dataset.get_split('train'), batch_size=batch_size, shuffle=True)
    val_loader = DataLoader(dataset.get_split('val'), batch_size=batch_size)
    test_loader = DataLoader(dataset.get_split('test'), batch_size=batch_size)

    input_size = len(dataset.get_feature_cols())
    logger.info(f"  Input size: {input_size}")
    logger.info(f"  Train: {len(train_loader.dataset)} | Val: {len(val_loader.dataset)} | Test: {len(test_loader.dataset)}")

    # ── 2. Создание модели ───────────────────────────────────────
    logger.info("\n[2/5] Создание модели MTFBasic...")
    model = MTFBasic(
        input_size=input_size,
        h1_hidden=64,
        d1_hidden=64,
        h1_num_layers=2,
        d1_num_layers=1,
        dropout=0.4,
    )
    logger.info(f"  Параметров: {sum(p.numel() for p in model.parameters()):,}")

    # ── 3. Обучение ──────────────────────────────────────────────
    logger.info("\n[3/5] Обучение...")

    # Веса классов
    train_y = dataset.data['train']['y']
    class_weights = compute_class_weights(train_y, sell_multiplier=2.0)
    logger.info(f"  Class weights: {class_weights}")

    # Loss + Optimizer + Scheduler
    criterion = WeightedFocalLoss(gamma=2.0, class_weights=class_weights)
    optimizer = AdamW(model.parameters(), lr=1e-3, weight_decay=1e-3)
    scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=10, T_mult=2)

    model = model.to(device)
    best_val_loss = float('inf')
    best_model_state = None
    patience = 15
    patience_counter = 0
    max_epochs = 100

    history = {'train_loss': [], 'val_loss': [], 'val_acc': []}

    for epoch in range(1, max_epochs + 1):
        train_loss = train_epoch_mtf(model, train_loader, optimizer, criterion, device)
        val_result = validate_mtf(model, val_loader, criterion, device)

        scheduler.step()

        history['train_loss'].append(train_loss)
        history['val_loss'].append(val_result['loss'])
        history['val_acc'].append(val_result['accuracy'])

        if (epoch - 1) % 5 == 0 or epoch == 1:
            logger.info(f"  Эпоха {epoch:3d}/{max_epochs} | Train Loss: {train_loss:.4f} "
                        f"| Val Loss: {val_result['loss']:.4f} | Val Acc: {val_result['accuracy']:.4f}")

        # Early stopping
        if val_result['loss'] < best_val_loss:
            best_val_loss = val_result['loss']
            best_model_state = model.state_dict().copy()
            patience_counter = 0
        else:
            patience_counter += 1
            if patience_counter >= patience:
                logger.info(f"  Early stopping на эпохе {epoch}. Лучшая: {epoch - patience}")
                break

    # Восстанавливаем лучшую модель
    if best_model_state:
        model.load_state_dict(best_model_state)

    # ── 4. Калибровка порогов ────────────────────────────────────
    logger.info("\n[4/5] Калибровка порогов...")
    best_threshold, thresholds = calibrate_thresholds_mtf(model, val_loader, device)
    logger.info(f"  Best threshold: {best_threshold:.2f} (macro F1 on val)")

    # ── 5. Оценка на тесте + сохранение ──────────────────────────
    logger.info("\n[5/5] Оценка на тесте...")
    test_results = evaluate_test_mtf(model, test_loader, device, thresholds)
    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']:.1f}%")
    logger.info(f"  SELL precision: {test_results['sell_precision']:.1f}%")

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

    # Сохранение
    saved_dir = Path('src/ml/models/saved') / MODEL_NAME
    saved_dir.mkdir(parents=True, exist_ok=True)

    torch.save(model.state_dict(), saved_dir / 'model.pt')
    logger.info(f"  Модель сохранена: {saved_dir / 'model.pt'}")

    config = {
        'model_name': MODEL_NAME,
        'model_type': 'MTFBasic',
        'ticker': TICKER,
        'timeframe': 'H1+D1',
        'params': {
            'input_size': input_size,
            'h1_hidden': 64,
            'd1_hidden': 64,
            'h1_num_layers': 2,
            'd1_num_layers': 1,
            'dropout': 0.4,
            'h1_seq_len': 40,
            'd1_seq_len': 30,
            'forecast_horizon': 1,
            'threshold_pct': 0.5,
            'focal_gamma': 2.0,
            'weight_decay': 0.001,
            'sell_multiplier': 2.0,
            'loss_function': 'WeightedFocalLoss',
        },
        '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),
            'val_accuracy': float(history['val_acc'][-1]) if history['val_acc'] else 0,
            'training_time_sec': 0,
            'total_epochs': epoch,
        },
        'calibrated_thresholds': thresholds,
        'feature_cols': dataset.get_feature_cols(),
        'model_path': str(saved_dir),
    }

    (saved_dir / 'config.json').write_text(json.dumps(config, indent=2, ensure_ascii=False))
    logger.info(f"  Config saved: {saved_dir / 'config.json'}")

    # Регистрация
    registry = ModelRegistry()
    registry.register(
        model_name=MODEL_NAME,
        model_type='MTFBasic',
        ticker=TICKER,
        timeframe='H1+D1',
        params=config['params'],
        metrics=config['metrics'],
        model_path=str(saved_dir),
    )
    logger.info(f"  ✅ {MODEL_NAME} зарегистрирована в реестре")

    # Итог
    logger.info("\n" + "=" * 70)
    logger.info(f"ИТОГ: {MODEL_NAME}")
    logger.info(f"  Архитектура: GRU(H1→64) + GRU(D1→64) + Concat → FC(128→64→3)")
    logger.info(f"  Test accuracy: {test_results['accuracy']*100:.2f}%")
    logger.info(f"  SELL recall:   {test_results['sell_recall']:.1f}%")
    logger.info(f"  SELL precision: {test_results['sell_precision']:.1f}%")
    logger.info(f"  Параметры: {sum(p.numel() for p in model.parameters()):,}")
    logger.info("=" * 70)


if __name__ == '__main__':
    main()
