"""
Скрипт обучения D1 Transformer для X5 (v2: без pre-train).

Упрощённая версия: только fine-tune на X5 augmented датасете.
Transformer — тяжёлая архитектура для CPU, поэтому pre-train пропущен.

Usage:
    python src/ml/train_x5_transformer_v2.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.transformer import D1Transformer
from src.ml.models.registry import ModelRegistry
from src.ml.train.trainer import set_seed
from src.ml.train.metrics import (
    calculate_accuracy, calculate_confusion_matrix, calculate_macro_f1,
)
from src.ml.train.losses import WeightedFocalLoss

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 load_hdf5(path: str):
    """Быстрая загрузка HDF5 без использования trainer."""
    import h5py
    logger.info(f"Загрузка {path}")
    with h5py.File(path, 'r') as f:
        data = {}
        for key in ['X_train', 'y_train', 'X_val', 'y_val', 'X_test', 'y_test']:
            if key in f:
                data[key] = f[key][:]
        if 'feature_names' in f:
            names = f['feature_names'][:]
            if len(names) > 0 and isinstance(names[0], bytes):
                data['feature_names'] = [n.decode() for n in names]
            else:
                data['feature_names'] = names
    for split in ['y_train', 'y_val', 'y_test']:
        if split in data:
            classes, counts = np.unique(data[split], return_counts=True)
            dist = {int(c): int(cnt) for c, cnt in zip(classes, counts)}
            logger.info(f"  {split}: {dist}")
    return data


def train_model_simple(
    model, train_loader, val_loader, criterion, optimizer, scheduler,
    n_epochs=50, patience=20, device='cpu',
):
    """Простой цикл обучения с early stopping."""
    best_val_loss = float('inf')
    best_state = None
    patience_counter = 0

    for epoch in range(1, n_epochs + 1):
        # Train
        model.train()
        total_loss = 0.0
        for x, y in train_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(train_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}/{n_epochs} | "
                        f"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | "
                        f"Val Acc: {val_acc:.4f} | patience: {patience_counter}")

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

    if best_state is not None:
        model.load_state_dict(best_state)
    return best_val_loss


def calibrate_threshold(model, val_loader, device):
    """Подбор порога уверенности на валидации."""
    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_t = 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
        f1 = calculate_macro_f1(all_labels, preds)
        if f1 > best_f1:
            best_f1 = f1
            best_t = threshold
    return best_t


def main():
    logger.info("=" * 70)
    logger.info(f"ОБУЧЕНИЕ D1 TRANSFORMER (v2, без pre-train): {MODEL_NAME}")
    logger.info("=" * 70)

    set_seed(42)
    device = torch.device('cpu')
    batch_size = 16  # меньше батч для Transformer (self-attention)

    # ─── Загрузка данных ────────────────────────────────────────────────
    data = load_hdf5(FINETUNE_DATA)

    seq_len = data['X_train'].shape[1]
    input_size = data['X_train'].shape[2]
    logger.info(f"Seq len: {seq_len}, Input size: {input_size}")

    # DataLoaders
    from torch.utils.data import DataLoader, TensorDataset
    train_ds = TensorDataset(
        torch.FloatTensor(data['X_train']),
        torch.LongTensor(data['y_train']),
    )
    val_ds = TensorDataset(
        torch.FloatTensor(data['X_val']),
        torch.LongTensor(data['y_val']),
    )
    test_ds = TensorDataset(
        torch.FloatTensor(data['X_test']),
        torch.LongTensor(data['y_test']),
    )
    train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True)
    val_loader = DataLoader(val_ds, batch_size=batch_size)
    test_loader = DataLoader(test_ds, batch_size=batch_size)

    # ─── Создание модели ─────────────────────────────────────────────────
    model = D1Transformer(
        input_size=input_size,
        d_model=64,
        nhead=4,
        num_layers=2,
        dim_feedforward=128,
        dropout=0.4,
        activation='gelu',
    )
    total_params = sum(p.numel() for p in model.parameters())
    logger.info(f"Параметров: {total_params:,}")
    logger.info(f"Архитектура: Transformer(d_model=64, nhead=4, layers=2, dropout=0.4)")

    # ─── Веса классов ───────────────────────────────────────────────────
    classes, counts = np.unique(data['y_train'], return_counts=True)
    total = len(data['y_train'])
    weights = total / (len(classes) * counts.astype(float))
    weights[2] *= 2.0  # SELL ×2
    weights = weights / weights.mean()
    class_weights = torch.FloatTensor(weights)
    logger.info(f"Веса классов: HOLD={weights[0]:.4f} BUY={weights[1]:.4f} SELL={weights[2]:.4f}")

    # ─── Оптимизация ─────────────────────────────────────────────────────
    criterion = WeightedFocalLoss(gamma=2.0, class_weights=class_weights)
    optimizer = torch.optim.AdamW(
        model.parameters(), lr=5e-5, weight_decay=1e-3,
    )
    scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(
        optimizer, T_0=10, T_mult=2,
    )

    model = model.to(device)

    # ─── Обучение ────────────────────────────────────────────────────────
    logger.info(f"\n{'='*70}")
    logger.info("ОБУЧЕНИЕ (Fine-tune на X5 Augmented)")
    logger.info("=" * 70)

    start = time.time()
    best_val_loss = train_model_simple(
        model, train_loader, val_loader, criterion, optimizer, scheduler,
        n_epochs=80, patience=25, device=device,
    )
    elapsed = time.time() - start
    logger.info(f"  ✅ Обучение завершено за {elapsed:.0f}с. "
                f"Лучшая val_loss: {best_val_loss:.4f}")

    # ─── Калибровка ──────────────────────────────────────────────────────
    logger.info(f"\n{'='*70}")
    logger.info("КАЛИБРОВКА ПОРОГА")
    logger.info("=" * 70)
    best_threshold = calibrate_threshold(model, val_loader, device)
    logger.info(f"  Best threshold: {best_threshold:.2f}")

    # ─── Оценка ──────────────────────────────────────────────────────────
    logger.info(f"\n{'='*70}")
    logger.info("ОЦЕНКА НА ТЕСТЕ")
    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): "
                f"HOLD={int((preds_raw==0).sum())} "
                f"BUY={int((preds_raw==1).sum())} "
                f"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"  Предсказания (threshold={best_threshold:.2f}): "
                f"HOLD={int((preds_cal==0).sum())} "
                f"BUY={int((preds_cal==1).sum())} "
                f"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,
        'model_type': 'D1Transformer',
        'ticker': TICKER,
        'timeframe': TIMEFRAME,
        'arch': {
            'input_size': input_size,
            'd_model': 64,
            'nhead': 4,
            'num_layers': 2,
            'dim_feedforward': 128,
            'dropout': 0.4,
            'activation': 'gelu',
            'total_params': total_params,
        },
        'training': {
            'batch_size': batch_size,
            'optimizer': 'AdamW',
            'lr': 5e-5,
            'weight_decay': 1e-3,
            'loss': 'WeightedFocalLoss(gamma=2.0)',
            'epochs_trained': 80,
            'patience': 25,
            'training_time_sec': elapsed,
            'data': 'x5_d1_augmented.h5 (без pre-train)',
        },
        '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),
            'threshold': float(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_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()
