"""
BiGRU + Attention — v2: без pre-train (слишком медленно на CPU).

Только fine-tune на X5 augmented датасете.
"""

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.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'
MODEL_NAME = 'x5_bigru_attn_v1'


def load_hdf5(path):
    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][:]
    return data


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

    torch.manual_seed(42)
    np.random.seed(42)
    device = torch.device('cpu')
    batch_size = 32

    # ─── Данные ─────────────────────────────────────────────────────────
    data = load_hdf5(FINETUNE_DATA)
    logger.info(f"Train: {data['X_train'].shape}, Val: {data['X_val'].shape}, Test: {data['X_test'].shape}")

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

    # ─── Модель ─────────────────────────────────────────────────────────
    model = BiGRUAttention(
        input_size=data['X_train'].shape[2],
        hidden_size=64,
        num_layers=2,
        dropout=0.4,
    )
    total_params = sum(p.numel() for p in model.parameters())
    logger.info(f"Параметров: {total_params:,}")

    # ─── Веса ────────────────────────────────────────────────────────────
    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
    weights = weights / weights.mean()
    class_weights = torch.FloatTensor(weights)
    logger.info(f"Class weights: 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)

    best_val_loss = float('inf')
    best_state = None
    patience_counter = 0
    start = time.time()

    for epoch in range(1, 81):
        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)

        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}/80 | Train Loss: {train_loss:.4f} | "
                        f"Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.4f} | patience: {patience_counter}")

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

    if best_state:
        model.load_state_dict(best_state)
    elapsed = time.time() - start
    logger.info(f"  ✅ Обучено за {elapsed:.0f}с. Лучшая val_loss: {best_val_loss:.4f}")

    # ─── Калибровка ──────────────────────────────────────────────────────
    logger.info(f"\n{'='*70}")
    logger.info("КАЛИБРОВКА ПОРОГА")
    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_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)):
            if preds[i] in [1, 2] and all_probs[i, preds[i]] < threshold:
                preds[i] = 0
        macro_f1 = calculate_macro_f1(all_labels, preds)
        if macro_f1 > best_f1:
            best_f1 = macro_f1
            best_t = threshold
    logger.info(f"  Best threshold: {best_t:.2f} (macro F1: {best_f1:.4f})")

    # ─── Оценка ──────────────────────────────────────────────────────────
    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)

    preds_raw = np.argmax(all_probs, axis=1)
    preds_cal = preds_raw.copy()
    for i in range(len(preds_cal)):
        if preds_cal[i] in [1, 2] and all_probs[i, preds_cal[i]] < best_t:
            preds_cal[i] = 0

    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"  Raw: HOLD={int((preds_raw==0).sum())} BUY={int((preds_raw==1).sum())} SELL={int((preds_raw==2).sum())}")
    logger.info(f"  Cal: HOLD={int((preds_cal==0).sum())} BUY={int((preds_cal==1).sum())} SELL={int((preds_cal==2).sum())}")
    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}%")
    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(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': 'BiGRUAttention',
        'ticker': 'X5',
        'timeframe': 'D1',
        'arch': {
            'input_size': data['X_train'].shape[2],
            'hidden_size': 64,
            'num_layers': 2,
            'dropout': 0.4,
            'total_params': total_params,
        },
        'training': {
            'optimizer': 'AdamW',
            'lr': 5e-5,
            'weight_decay': 1e-3,
            'loss': 'WeightedFocalLoss(gamma=2.0)',
            '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)),
            'confusion_matrix': cm.tolist(),
            'sell_recall': float(sell_recall),
            'sell_precision': float(sell_precision),
            'val_loss': float(best_val_loss),
            'threshold': float(best_t),
        },
    }

    (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='X5',
        timeframe='D1',
        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}%")
    if sell_recall > 0:
        logger.info(f"  ✅ Можно добавить в ансамбль")
    else:
        logger.warning(f"  ❌ SELL recall = 0% — не пригодна")
    logger.info("=" * 70)


if __name__ == '__main__':
    main()
