"""
Скрипт двухэтапного обучения LSTM модели для тикера X5.

Этап A — Pre-train на 9 тикерах MOEX (15,669 семплов):
    - Загружает multi_ticker_d1_pretrain.h5
    - Обучает LSTM (input_size=61, hidden_size=128, num_layers=2, dropout=0.3)
    - AdamW (lr=1e-3, wd=1e-4), CosineAnnealingWarmRestarts
    - Max 50 эпох, early stopping patience=10
    - Сохраняет чекпоинт

Этап B — Fine-tune на X5 D1:
    - Загружает чекпоинт из Этапа A
    - Загружает x5_d1_augmented.h5 (900 семплов)
    - Fine-tune с lr=1e-4 (в 10 раз меньше), max 30 эпох
    - Сохраняет финальную модель

Использование:
    python src/ml/train_pretrain_finetune.py
"""

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

import numpy as np
import torch

# Настройка логирования
logging.basicConfig(
    level=logging.INFO,
    format='%(asctime)s [%(levelname)s] %(message)s',
    datefmt='%H:%M:%S',
)
logger = logging.getLogger(__name__)

# Добавляем корень проекта в PYTHONPATH
project_root = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(project_root.parent))

from src.ml.models.lstm import LSTMPredictor
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,
    print_metrics_report,
    pretrain_then_finetune,
)
from src.ml.train.metrics import (
    calculate_all_metrics,
    calculate_macro_f1,
    classification_report,
)


# Константы
PRETRAIN_DATA_PATH = 'src/ml/models/saved/multi_ticker_d1_pretrain.h5'
FINETUNE_DATA_PATH = 'src/ml/models/saved/x5_d1_augmented.h5'
SAVE_DIR = 'src/ml/models/saved'
CHECKPOINT_PATH = 'src/ml/models/saved/x5_lstm_v2_pretrained.pt'
MODEL_NAME = 'x5_lstm_v2'
TICKER = 'X5'
TIMEFRAME = 'D1'
SEED = 42

# Гиперпараметры (одинаковые для pre-train, кроме lr)
INPUT_SIZE = 61
HIDDEN_SIZE = 128
NUM_LAYERS = 2
DROPOUT = 0.3
BATCH_SIZE = 32
PRETRAIN_EPOCHS = 50
FINETUNE_EPOCHS = 30
PRETRAIN_LR = 1e-3
FINETUNE_LR = 1e-4
WEIGHT_DECAY = 1e-4
PRETRAIN_PATIENCE = 10
FINETUNE_PATIENCE = 15
GRADIENT_CLIP = 1.0


def save_model_artifacts(
    model: torch.nn.Module,
    data: Dict,
    best_params: Dict,
    metrics: Dict,
    save_dir: str,
    model_name: str,
) -> str:
    """
    Сохранить модель и сопутствующие артефакты.

    Args:
        model: обученная модель.
        data: словарь с данными (для normalizer и feature_names).
        best_params: гиперпараметры модели.
        metrics: метрики производительности.
        save_dir: директория для сохранения.
        model_name: имя модели.

    Returns:
        Путь к директории сохранения.
    """
    save_path = Path(save_dir) / model_name
    save_path.mkdir(parents=True, exist_ok=True)

    # Сохраняем веса модели
    model_path = save_path / 'model.pt'
    torch.save(model.state_dict(), model_path)
    logger.info(f"Модель сохранена: {model_path}")

    # Сохраняем конфигурацию
    config = {
        'model_type': 'LSTMPredictor',
        'params': best_params,
        'feature_cols': list(data.get('feature_names', [])),
        'n_classes': 3,
        'class_names': {0: 'HOLD', 1: 'BUY', 2: 'SELL'},
        'metrics': metrics,
        'ticker': TICKER,
        'timeframe': TIMEFRAME,
        'training_type': 'pretrain_then_finetune',
        'pretrain_data': PRETRAIN_DATA_PATH,
        'finetune_data': FINETUNE_DATA_PATH,
    }

    config_path = save_path / 'config.json'
    config_path.write_text(json.dumps(config, indent=2, ensure_ascii=False))
    logger.info(f"Конфигурация сохранена: {config_path}")

    # Сохраняем параметры нормализации (из fine-tune данных)
    if 'scaler_mean' in data and 'scaler_scale' in data:
        normalizer_data = {
            'mean_': data['scaler_mean'].tolist(),
            'std_': data['scaler_scale'].tolist(),
        }
        normalizer_path = save_path / 'normalizer.json'
        normalizer_path.write_text(json.dumps(normalizer_data, indent=2))
        logger.info(f"Параметры нормализации сохранены: {normalizer_path}")

    # Сохраняем список признаков
    if 'feature_names' in data:
        features_path = save_path / 'features.json'
        features_path.write_text(
            json.dumps(list(data['feature_names']), indent=2, ensure_ascii=False)
        )
        logger.info(f"Список признаков сохранён: {features_path}")

    return str(save_path)


def compute_metrics_from_loader(
    model: torch.nn.Module,
    loader: torch.utils.data.DataLoader,
    device: torch.device,
) -> Dict:
    """
    Вычислить метрики для модели на даталоадере.

    Args:
        model: модель PyTorch.
        loader: DataLoader.
        device: устройство.

    Returns:
        Словарь с метриками.
    """
    model.eval()
    all_preds = []
    all_labels = []
    total_loss = 0.0
    criterion = torch.nn.CrossEntropyLoss()

    with torch.no_grad():
        for batch_x, batch_y in loader:
            batch_x = batch_x.to(device, non_blocking=True)
            batch_y = batch_y.to(device, non_blocking=True)
            outputs = model(batch_x)
            loss = criterion(outputs, batch_y)
            total_loss += loss.item() * batch_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)
    loss = total_loss / len(loader.dataset)

    metrics = calculate_all_metrics(all_labels, all_preds, num_classes=3)
    metrics['loss'] = float(loss)

    return {
        'loss': float(loss),
        'accuracy': float(metrics['accuracy']),
        'macro_f1': float(metrics['macro_f1']),
        'weighted_f1': float(metrics['weighted_f1']),
        'confusion_matrix': metrics['confusion_matrix'],
        'per_class': {str(k): v for k, v in metrics['per_class'].items()},
        'predictions': all_preds.tolist(),
        'labels': all_labels.tolist(),
    }


def main() -> None:
    """
    Основная функция двухэтапного обучения.

    Последовательность:
        1. Установка seed
        2. Pre-train: загрузка multi-ticker данных, обучение, сохранение чекпоинта
        3. Fine-tune: загрузка X5 D1 данных, дообучение
        4. Вывод метрик
        5. Сохранение модели
        6. Регистрация в ModelRegistry
    """
    logger.info("=" * 70)
    logger.info("ДВУХЭТАПНОЕ ОБУЧЕНИЕ LSTM МОДЕЛИ ДЛЯ X5")
    logger.info("=" * 70)
    logger.info(f"  Pre-train data: {PRETRAIN_DATA_PATH}")
    logger.info(f"  Fine-tune data: {FINETUNE_DATA_PATH}")
    logger.info(f"  Модель: LSTMPredictor (input={INPUT_SIZE}, "
                f"hidden={HIDDEN_SIZE}, layers={NUM_LAYERS})")

    # 1. Устанавливаем seed
    set_seed(SEED)
    logger.info(f"Seed установлен: {SEED}")

    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    logger.info(f"Используется устройство: {device}")

    total_start_time = time.time()

    # 2. Загружаем pre-train датасет
    logger.info("\n" + "-" * 60)
    logger.info("ЗАГРУЗКА ПРЕДВАРИТЕЛЬНЫХ ДАННЫХ (PRE-TRAIN)")
    logger.info("-" * 60)
    pretrain_data = load_hdf5_dataset(PRETRAIN_DATA_PATH)

    # 3. Загружаем fine-tune датасет (X5)
    logger.info("\n" + "-" * 60)
    logger.info("ЗАГРУЗКА ЦЕЛЕВЫХ ДАННЫХ (FINE-TUNE)")
    logger.info("-" * 60)
    finetune_data = load_hdf5_dataset(FINETUNE_DATA_PATH)

    # 4. Создаём DataLoader'ы
    logger.info("\n" + "-" * 60)
    logger.info("СОЗДАНИЕ DataLoader'ОВ")
    logger.info("-" * 60)

    # Pre-train данные — split 90/10 для train/val
    PRETRAIN_VAL_SPLIT = 0.1
    pretrain_loader, pretrain_val_loader, _ = create_dataloaders(
        pretrain_data, batch_size=BATCH_SIZE, shuffle_train=True,
        val_split=PRETRAIN_VAL_SPLIT,
    )

    # Fine-tune данные — уже есть val/test
    finetune_loader, val_loader, test_loader = create_dataloaders(
        finetune_data, batch_size=BATCH_SIZE, shuffle_train=True
    )

    # 5. Вычисляем веса классов
    # Pre-train классы почти сбалансированы — не используем веса
    # Fine-tune — используем веса
    finetune_class_weights = compute_class_weights(
        finetune_data['y_train'], num_classes=3
    )
    logger.info(f"\nВеса классов (fine-tune): HOLD={finetune_class_weights[0]:.4f}, "
                f"BUY={finetune_class_weights[1]:.4f}, SELL={finetune_class_weights[2]:.4f}")

    # 6. Создаём модель
    model = LSTMPredictor(
        input_size=INPUT_SIZE,
        hidden_size=HIDDEN_SIZE,
        num_layers=NUM_LAYERS,
        dropout=DROPOUT,
        bidirectional=False,
    )

    total_params = sum(p.numel() for p in model.parameters())
    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
    logger.info(f"\nПараметры модели: всего {total_params:,}, обучаемых {trainable_params:,}")

    # 7. Двухэтапное обучение
    logger.info("\n" + "=" * 70)
    logger.info("ЗАПУСК ДВУХЭТАПНОГО ОБУЧЕНИЯ")
    logger.info("=" * 70)

    training_start = time.time()

    pretrain_result, finetune_result = pretrain_then_finetune(
        model=model,
        pretrain_loader=pretrain_loader,
        pretrain_val_loader=pretrain_val_loader,
        finetune_loader=finetune_loader,
        finetune_val_loader=val_loader,
        test_loader=test_loader,
        pretrain_epochs=PRETRAIN_EPOCHS,
        finetune_epochs=FINETUNE_EPOCHS,
        pretrain_lr=PRETRAIN_LR,
        finetune_lr=FINETUNE_LR,
        weight_decay=WEIGHT_DECAY,
        pretrain_patience=PRETRAIN_PATIENCE,
        finetune_patience=FINETUNE_PATIENCE,
        gradient_clip=GRADIENT_CLIP,
        device='cpu',
        model_name=MODEL_NAME,
        save_dir=SAVE_DIR,
        pretrain_class_weights=None,  # pre-train почти сбалансирован
        finetune_class_weights=finetune_class_weights,
        verbose=True,
        checkpoint_path=CHECKPOINT_PATH,
    )

    total_training_time = time.time() - training_start

    # 8. Вывод финальных метрик
    logger.info("\n" + "=" * 70)
    logger.info("ИТОГОВЫЕ МЕТРИКИ")
    logger.info("=" * 70)

    if pretrain_result is not None:
        logger.info(f"\n--- Этап A: Pre-train ---")
        logger.info(f"  Эпох: {pretrain_result.total_epochs}")
        logger.info(f"  Лучшая эпоха: {pretrain_result.best_epoch + 1}")
        logger.info(f"  Val Loss: {pretrain_result.best_val_loss:.6f}")
        logger.info(f"  Val Accuracy: {pretrain_result.val_metrics['accuracy']:.4f}")
        pretrain_epochs_done = pretrain_result.total_epochs
    else:
        logger.info(f"\n--- Этап A: Pre-train загружен из чекпоинта ---")
        pretrain_epochs_done = 'loaded_from_checkpoint'

    logger.info(f"\n--- Этап B: Fine-tune ---")
    logger.info(f"  Эпох: {finetune_result.total_epochs}")
    logger.info(f"  Лучшая эпоха: {finetune_result.best_epoch + 1}")
    logger.info(f"  Val Loss: {finetune_result.best_val_loss:.6f}")
    logger.info(f"  Val Accuracy: {finetune_result.val_metrics['accuracy']:.4f}")

    # Метрики на валидации (fine-tune)
    logger.info(f"\nФинальные метрики (валидация):")
    val_preds = finetune_result.val_metrics['predictions']
    val_labels = finetune_result.val_metrics['labels']
    print_metrics_report(
        'VALIDATION',
        val_labels,
        val_preds,
        finetune_result.val_metrics['loss'],
    )

    # Метрики на тесте
    test_metrics = None
    if finetune_result.test_metrics is not None:
        logger.info(f"\nФинальные метрики (тест):")
        test_preds = finetune_result.test_metrics['predictions']
        test_labels = finetune_result.test_metrics['labels']
        print_metrics_report(
            'TEST',
            test_labels,
            test_preds,
            finetune_result.test_metrics['loss'],
        )
        test_metrics = {
            'loss': float(finetune_result.test_metrics['loss']),
            'accuracy': float(finetune_result.test_metrics['accuracy']),
            'macro_f1': float(calculate_macro_f1(test_labels, test_preds)),
            'confusion_matrix': finetune_result.test_metrics['confusion_matrix'].tolist(),
        }

    # 9. Сохраняем модель и артефакты
    best_params = {
        'input_size': INPUT_SIZE,
        'hidden_size': HIDDEN_SIZE,
        'num_layers': NUM_LAYERS,
        'dropout': DROPOUT,
        'bidirectional': False,
        'pretrain_epochs': PRETRAIN_EPOCHS,
        'finetune_epochs': FINETUNE_EPOCHS,
        'pretrain_lr': PRETRAIN_LR,
        'finetune_lr': FINETUNE_LR,
        'batch_size': BATCH_SIZE,
        'weight_decay': WEIGHT_DECAY,
    }

    # Собираем все метрики
    val_metrics = {
        'val_loss': float(finetune_result.val_metrics['loss']),
        'val_accuracy': float(finetune_result.val_metrics['accuracy']),
        'val_macro_f1': float(calculate_macro_f1(val_labels, val_preds)),
        'val_confusion_matrix': finetune_result.val_metrics['confusion_matrix'].tolist(),
        'best_epoch': finetune_result.best_epoch + 1,
        'total_epochs': finetune_result.total_epochs,
        'pretrain_epochs_done': pretrain_result.total_epochs if pretrain_result else 0,
        'training_time_sec': total_training_time,
    }

    if test_metrics:
        val_metrics['test_loss'] = test_metrics['loss']
        val_metrics['test_accuracy'] = test_metrics['accuracy']
        val_metrics['test_macro_f1'] = test_metrics['macro_f1']
        val_metrics['test_confusion_matrix'] = test_metrics['confusion_matrix']

    model_path = save_model_artifacts(
        model=finetune_result.model,
        data=finetune_data,
        best_params=best_params,
        metrics=val_metrics,
        save_dir=SAVE_DIR,
        model_name=MODEL_NAME,
    )

    # 10. Регистрируем модель в ModelRegistry
    registry = ModelRegistry()
    registry.register(
        model_name=MODEL_NAME,
        model_type='LSTMPredictor',
        ticker=TICKER,
        timeframe=TIMEFRAME,
        params=best_params,
        metrics=val_metrics,
        model_path=model_path,
    )
    logger.info(f"Модель зарегистрирована в ModelRegistry как '{MODEL_NAME}'")

    # Финальный отчёт
    total_time = time.time() - total_start_time
    logger.info("\n" + "=" * 70)
    logger.info("ИТОГОВЫЙ ОТЧЁТ — ДВУХЭТАПНОЕ ОБУЧЕНИЕ")
    logger.info("=" * 70)
    logger.info(f"  Тикер: {TICKER}")
    logger.info(f"  Таймфрейм: {TIMEFRAME}")
    logger.info(f"  Архитектура: LSTMPredictor (hidden={HIDDEN_SIZE}, layers={NUM_LAYERS})")
    logger.info(f"  Pre-train данные: {PRETRAIN_DATA_PATH}")
    logger.info(f"  Pre-train семплов: {len(pretrain_data['X_train'])}")
    logger.info(f"  Pre-train эпох: {pretrain_result.total_epochs if pretrain_result else 'N/A (checkpoint)'}")
    pretrain_loss_str = f"{pretrain_result.best_val_loss:.6f}" if pretrain_result else "N/A"
    logger.info(f"  Pre-train best val loss: {pretrain_loss_str}")
    logger.info(f"  Fine-tune данные: {FINETUNE_DATA_PATH}")
    logger.info(f"  Fine-tune семплов: {len(finetune_data['X_train'])}")
    logger.info(f"  Fine-tune эпох: {finetune_result.total_epochs}")
    logger.info(f"  Fine-tune best val loss: {finetune_result.best_val_loss:.6f}")
    logger.info(f"  Val Accuracy: {finetune_result.val_metrics['accuracy']:.4f}")
    if test_metrics:
        logger.info(f"  Test Accuracy: {test_metrics['accuracy']:.4f}")
        logger.info(f"  Test Macro F1: {test_metrics['macro_f1']:.4f}")
    logger.info(f"  Общее время обучения: {total_time:.2f} сек ({total_time / 60:.2f} мин)")
    logger.info(f"  Модель сохранена: {model_path}")
    logger.info("=" * 70)


if __name__ == '__main__':
    main()
