"""
MultiTimeframeFusion — Pre-train на мульти-тикерах + Fine-tune X5.

Этапы:
  1. Для каждого из 9 MOEX тикеров создаём MTF датасет
  2. Объединяем train-сплиты → Pre-train на ~2000+ сэмплах
  3. Pre-train MTF модели (общие паттерны H1×D1)
  4. Fine-tune на X5 (специфика тикера)
  5. Оценка на X5 test

Usage:
    python src/ml/train_mtf_pretrain.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, TensorDataset

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,
)

# 9 MOEX tickers (те же, что в D1 pre-train)
PRETRAIN_TICKERS = ['SBER', 'LKOH', 'NVTK', 'MTSS', 'PHOR', 'VTBR', 'MOEX', 'GAZP', 'PLZL']
TICKER = 'X5'
MODEL_NAME = 'x5_mtf_v2'
DEVICE = torch.device('cpu')


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:
    classes, counts = np.unique(y, return_counts=True)
    total = len(y)
    weights = total / (len(classes) * counts.astype(float))
    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:
    model.train()
    total_loss = 0.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)
        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)
    return total_loss / len(dataloader.dataset)


@torch.no_grad()
def validate_mtf(
    model: nn.Module,
    dataloader: DataLoader,
    criterion: nn.Module,
    device: torch.device,
) -> Dict:
    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),
    }


def calibrate_thresholds_mtf(
    model: nn.Module,
    dataloader: DataLoader,
    device: torch.device,
) -> Tuple[float, Dict[int, float]]:
    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 = [], []
    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_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 load_ticker_mtf_data(ticker: str, h1_limit: int = 5000, d1_limit: int = 1500,
                          feature_cols: Optional[list] = None) -> Optional[Dict]:
    """
    Загрузить и подготовить MTF данные для одного тикера.
    Возвращает {'h1': arr, 'd1': arr, 'y': arr, 'feature_cols': list} или None.
    """
    try:
        logger.info(f"  [{ticker}] Загрузка H1 (лимит={h1_limit}) + D1 (лимит={d1_limit})...")
        ds = MultiTimeframeFusionDataset(
            ticker=ticker,
            h1_limit=h1_limit,
            d1_limit=d1_limit,
            h1_seq_len=120,  # увеличили 40→120 (15 торговых дней)
            d1_seq_len=30,
            forecast_horizon=1,
            threshold_pct=0.5,
            val_split=0.0,
            test_split=0.0,
            feature_cols=feature_cols,
            top_k_features=30,  # feature selection: 61→30
        )
        train = ds.data['train']
        logger.info(f"  [{ticker}] Готово: {len(train['y'])} сэмплов, "
                     f"распред: HOLD={int((train['y']==0).sum())} "
                     f"BUY={int((train['y']==1).sum())} "
                     f"SELL={int((train['y']==2).sum())}")
        return {
            'h1': train['h1'],
            'd1': train['d1'],
            'y': train['y'],
            'feature_cols': ds.get_feature_cols(),
        }
    except Exception as e:
        logger.warning(f"  [{ticker}] Ошибка загрузки: {e}")
        return None


def main():
    logger.info("=" * 70)
    logger.info("MTF PRE-TRAIN (9 MOEX) + FINE-TUNE X5")
    logger.info("=" * 70)

    set_seed(42)
    batch_size = 32
    feature_cols = None

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

    all_h1, all_d1, all_y = [], [], []

    for ticker in PRETRAIN_TICKERS:
        data = load_ticker_mtf_data(ticker, h1_limit=5000, d1_limit=1500,
                                     feature_cols=feature_cols)
        if data is not None:
            all_h1.append(data['h1'])
            all_d1.append(data['d1'])
            all_y.append(data['y'])
            if feature_cols is None:
                feature_cols = data.get('feature_cols')

    if not all_h1:
        logger.error("Не удалось загрузить данные ни для одного тикера!")
        return

    if feature_cols is None:
        logger.error("feature_cols не получены!")
        return

    X_h1_pretrain = np.concatenate(all_h1, axis=0)
    X_d1_pretrain = np.concatenate(all_d1, axis=0)
    y_pretrain = np.concatenate(all_y, axis=0)
    input_size = len(feature_cols)

    logger.info(f"\n  Pre-train датасет: {len(X_h1_pretrain)} сэмплов")
    logger.info(f"  Input size: {input_size}")
    dist = {0: int((y_pretrain == 0).sum()), 1: int((y_pretrain == 1).sum()), 2: int((y_pretrain == 2).sum())}
    logger.info(f"  Распределение: HOLD={dist[0]} BUY={dist[1]} SELL={dist[2]}")

    pretrain_dataset = TensorDataset(
        torch.FloatTensor(X_h1_pretrain),
        torch.FloatTensor(X_d1_pretrain),
        torch.LongTensor(y_pretrain),
    )
    pretrain_loader = DataLoader(pretrain_dataset, batch_size=batch_size, shuffle=True)

    # ── ЭТАП 2: Pre-train обучение ──────────────────────────────
    logger.info("\n" + "=" * 70)
    logger.info("ЭТАП 2: Pre-train обучение")
    logger.info("=" * 70)

    model = MTFBasic(
        input_size=input_size,
        h1_hidden=64,
        d1_hidden=64,
        h1_num_layers=2,
        d1_num_layers=1,
        dropout=0.3,  # чуть ниже dropout для pre-train (больше данных)
    )
    logger.info(f"  Параметров: {sum(p.numel() for p in model.parameters()):,}")

    class_weights = compute_class_weights(y_pretrain, sell_multiplier=2.0)
    logger.info(f"  Class weights: {class_weights.numpy().round(3)}")

    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=15, T_mult=2)

    model = model.to(DEVICE)
    best_val_loss = float('inf')
    patience = 10
    patience_counter = 0
    max_epochs = 50
    best_state = None

    # Используем 10% pre-train данных как val
    n_pretrain = len(y_pretrain)
    n_val = max(1, n_pretrain // 10)

    for epoch in range(1, max_epochs + 1):
        train_loss = train_epoch_mtf(model, pretrain_loader, optimizer, criterion, DEVICE)
        scheduler.step()

        # Оценка на небольшой val-выборке
        val_h1 = torch.FloatTensor(X_h1_pretrain[-n_val:]).to(DEVICE)
        val_d1 = torch.FloatTensor(X_d1_pretrain[-n_val:]).to(DEVICE)
        val_y = torch.LongTensor(y_pretrain[-n_val:]).to(DEVICE)

        model.eval()
        with torch.no_grad():
            val_out = model(val_h1, val_d1)
            val_loss = criterion(val_out, val_y).item()
            val_preds = torch.argmax(val_out, dim=1)
            val_acc = (val_preds == val_y).float().mean().item()
        model.train()

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

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

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

    # ── ЭТАП 3: Fine-tune на X5 ─────────────────────────────────
    logger.info("\n" + "=" * 70)
    logger.info("ЭТАП 3: Fine-tune на X5")
    logger.info("=" * 70)

    x5_ds = MultiTimeframeFusionDataset(
        ticker=TICKER,
        h1_limit=5000,
        d1_limit=2000,
        h1_seq_len=120,  # увеличили 40→120
        d1_seq_len=30,
        forecast_horizon=1,
        threshold_pct=0.5,
        val_split=0.15,
        test_split=0.15,
        feature_cols=feature_cols,  # те же признаки, что и pre-train
        top_k_features=30,  # feature selection
    )

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

    # Fine-tune: снижаем dropout, уменьшаем lr
    model.dropout = 0.4  # увеличиваем dropout для fine-tune (меньше данных)

    # Пересоздаём классификатор с новым dropout
    model.classifier = nn.Sequential(
        nn.Linear(128, 128),
        nn.ReLU(),
        nn.Dropout(0.4),
        nn.Linear(128, 64),
        nn.ReLU(),
        nn.Dropout(0.4),
        nn.Linear(64, 3),
    )

    x5_weights = compute_class_weights(x5_ds.data['train']['y'], sell_multiplier=2.0)
    logger.info(f"  X5 class weights: {x5_weights.numpy().round(3)}")
    ft_criterion = WeightedFocalLoss(gamma=2.0, class_weights=x5_weights)
    ft_optimizer = AdamW(model.parameters(), lr=5e-5, weight_decay=1e-3)

    model = model.to(DEVICE)
    best_ft_loss = float('inf')
    best_ft_state = None
    ft_patience = 15
    ft_counter = 0
    ft_epochs = 50

    for epoch in range(1, ft_epochs + 1):
        train_loss = train_epoch_mtf(model, train_loader, ft_optimizer, ft_criterion, DEVICE)
        val_result = validate_mtf(model, val_loader, ft_criterion, DEVICE)

        if val_result['loss'] < best_ft_loss:
            best_ft_loss = val_result['loss']
            best_ft_state = model.state_dict().copy()
            ft_counter = 0
        else:
            ft_counter += 1

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

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

    if best_ft_state:
        model.load_state_dict(best_ft_state)

    # ── ЭТАП 4: Калибровка + Оценка ─────────────────────────────
    logger.info("\n" + "=" * 70)
    logger.info("ЭТАП 4: Калибровка + Оценка на тесте")
    logger.info("=" * 70)

    best_threshold, thresholds = calibrate_thresholds_mtf(model, val_loader, DEVICE)
    logger.info(f"  Best threshold: {best_threshold:.2f}")

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

    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',
            'pretrain_tickers': PRETRAIN_TICKERS,
            'pretrain_samples': len(X_h1_pretrain),
        },
        '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_ft_loss),
        },
        'calibrated_thresholds': thresholds,
        'feature_cols': feature_cols,
        'model_path': str(saved_dir),
    }

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

    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("\n" + "=" * 70)
    logger.info(f"ИТОГ: {MODEL_NAME}")
    logger.info(f"  Pre-train: {len(PRETRAIN_TICKERS)} tickers, {len(X_h1_pretrain)} samples")
    logger.info(f"  Fine-tune: {TICKER} ({len(x5_ds.data['train']['y'])} train samples)")
    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("=" * 70)


if __name__ == '__main__':
    main()
