"""
Скрипт обучения LSTM модели для тикера X5 на таймфрейме D1.

Загружает датасет из HDF5, создаёт модель LSTMPredictor,
обучает с early stopping, выводит метрики и сохраняет результат.

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

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

import h5py
import numpy as np
import torch
from torch.utils.data import DataLoader, TensorDataset

# Настройка логирования
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,
)
from src.ml.train.metrics import (
    calculate_all_metrics,
    classification_report,
)


# Константы
DATA_PATH = 'src/ml/models/saved/x5_d1_dataset.h5'
SAVE_DIR = 'src/ml/models/saved/x5_lstm_v1'
TICKER = 'X5'
TIMEFRAME = 'D1'
MODEL_NAME = 'x5_lstm_v1'
SEED = 42

# Гиперпараметры
INPUT_SIZE = 61
HIDDEN_SIZE = 128
NUM_LAYERS = 2
DROPOUT = 0.3
BATCH_SIZE = 16
N_EPOCHS = 100
LEARNING_RATE = 1e-3
WEIGHT_DECAY = 1e-4
PATIENCE = 15
GRADIENT_CLIP = 1.0


def load_dataset(path: str) -> Dict[str, np.ndarray]:
    """
    Загрузить датасет из HDF5 файла.

    Args:
        path: путь к HDF5 файлу.

    Returns:
        Словарь с массивами: X_train, y_train, X_val, y_val, X_test, y_test,
        feature_names, scaler_mean, scaler_scale.
    """
    logger.info(f"Загрузка датасета из {path}")

    with h5py.File(path, 'r') as f:
        data = {
            'X_train': f['X_train'][:],
            'y_train': f['y_train'][:],
            'X_val': f['X_val'][:],
            'y_val': f['y_val'][:],
            'X_test': f['X_test'][:],
            'y_test': f['y_test'][:],
            'feature_names': f['feature_names'][:],
            'scaler_mean': f['scaler_mean'][:],
            'scaler_scale': f['scaler_scale'][:],
        }
        # Метаданные
        metadata = f['metadata'][()]
        if isinstance(metadata, bytes):
            metadata = metadata.decode()
        data['metadata'] = json.loads(metadata) if isinstance(metadata, str) else metadata

    logger.info(
        f"Загружено: Train {data['X_train'].shape}, "
        f"Val {data['X_val'].shape}, "
        f"Test {data['X_test'].shape}"
    )
    return data


def create_dataloaders(
    data: Dict[str, np.ndarray],
    batch_size: int = 16,
) -> Tuple[DataLoader, DataLoader, DataLoader]:
    """
    Создать DataLoader'ы для обучения, валидации и тестирования.

    Args:
        data: словарь с данными из load_dataset().
        batch_size: размер батча.

    Returns:
        Кортеж (train_loader, val_loader, test_loader).
    """
    # Тренировочный датасет — перемешиваем
    train_dataset = TensorDataset(
        torch.FloatTensor(data['X_train']),
        torch.LongTensor(data['y_train']),
    )
    train_loader = DataLoader(
        train_dataset,
        batch_size=batch_size,
        shuffle=True,
        drop_last=False,
    )

    # Валидационный датасет — без перемешивания
    val_dataset = TensorDataset(
        torch.FloatTensor(data['X_val']),
        torch.LongTensor(data['y_val']),
    )
    val_loader = DataLoader(
        val_dataset,
        batch_size=batch_size,
        shuffle=False,
    )

    # Тестовый датасет — без перемешивания
    test_dataset = TensorDataset(
        torch.FloatTensor(data['X_test']),
        torch.LongTensor(data['y_test']),
    )
    test_loader = DataLoader(
        test_dataset,
        batch_size=batch_size,
        shuffle=False,
    )

    logger.info(
        f"DataLoader'ы созданы: train={len(train_loader)} батчей, "
        f"val={len(val_loader)} батчей, "
        f"test={len(test_loader)} батчей"
    )
    return train_loader, val_loader, test_loader


def print_metrics_report(
    phase: str,
    y_true: np.ndarray,
    y_pred: np.ndarray,
    loss: float,
) -> None:
    """
    Вывести отчёт по метрикам для фазы обучения/валидации/теста.

    Args:
        phase: название фазы ('Train', 'Val', 'Test').
        y_true: истинные метки.
        y_pred: предсказанные метки.
        loss: значение функции потерь.
    """
    class_names = {0: 'HOLD', 1: 'BUY', 2: 'SELL'}

    logger.info(f"\n{'=' * 60}")
    logger.info(f"МЕТРИКИ: {phase}")
    logger.info(f"{'=' * 60}")
    logger.info(f"Loss: {loss:.6f}")

    metrics = calculate_all_metrics(y_true, y_pred, num_classes=3)
    logger.info(f"Accuracy: {metrics['accuracy']:.4f}")
    logger.info(f"Macro F1: {metrics['macro_f1']:.4f}")
    logger.info(f"Weighted F1: {metrics['weighted_f1']:.4f}")

    cm = metrics['confusion_matrix']
    logger.info("Confusion Matrix:")
    logger.info(f"{'':>10} {'HOLD':>8} {'BUY':>8} {'SELL':>8}")
    for i, name in [(0, 'HOLD'), (1, 'BUY'), (2, 'SELL')]:
        row = f"{name:<10}"
        for j in range(3):
            row += f"{cm[i][j]:>8}"
        logger.info(row)

    per_class = metrics['per_class']
    logger.info(f"\nPer-class metrics:")
    for c, name in [(0, 'HOLD'), (1, 'BUY'), (2, 'SELL')]:
        m = per_class[c]
        cnt = int((y_true == c).sum())
        logger.info(
            f"  {name:<6}: precision={m['precision']:.4f}, "
            f"recall={m['recall']:.4f}, "
            f"f1={m['f1']:.4f}, "
            f"support={cnt}"
        )


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

    Сохраняет:
        - model.pt: веса модели
        - config.json: конфигурация модели
        - normalizer.pkl: параметры нормализации
        - features.json: список признаков

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

    Returns:
        Путь к директории сохранения.
    """
    save_path = Path(save_dir)
    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': [n.decode() if isinstance(n, bytes) else str(n)
                         for n in data['feature_names']],
        'n_classes': 3,
        'class_names': {0: 'HOLD', 1: 'BUY', 2: 'SELL'},
        'metrics': metrics,
        'ticker': TICKER,
        'timeframe': TIMEFRAME,
    }

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

    # Сохраняем параметры нормализации
    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}")

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

    return str(save_path)


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

    Последовательность:
        1. Установка seed
        2. Загрузка датасета
        3. Создание DataLoader'ов
        4. Вычисление весов классов
        5. Создание модели
        6. Обучение
        7. Вывод метрик
        8. Сохранение модели
        9. Регистрация в ModelRegistry
    """
    logger.info("=" * 60)
    logger.info("ЗАПУСК ОБУЧЕНИЯ LSTM МОДЕЛИ ДЛЯ X5 D1")
    logger.info("=" * 60)

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

    # 2. Загружаем датасет
    data = load_dataset(DATA_PATH)

    # 3. Создаём DataLoader'ы
    train_loader, val_loader, test_loader = create_dataloaders(
        data, batch_size=BATCH_SIZE
    )

    # 4. Вычисляем веса классов
    class_weights = compute_class_weights(data['y_train'], num_classes=3)
    logger.info(f"Веса классов: HOLD={class_weights[0]:.4f}, "
                f"BUY={class_weights[1]:.4f}, SELL={class_weights[2]:.4f}")

    # 5. Создаём модель
    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"Модель создана: LSTMPredictor "
        f"(input={INPUT_SIZE}, hidden={HIDDEN_SIZE}, layers={NUM_LAYERS}, "
        f"dropout={DROPOUT})"
    )
    logger.info(f"Параметры: всего {total_params:,}, обучаемых {trainable_params:,}")

    # 6. Обучаем модель
    logger.info(f"\nНачало обучения (max {N_EPOCHS} эпох, patience={PATIENCE})...")
    logger.info(f"  Optimizer: AdamW (lr={LEARNING_RATE}, wd={WEIGHT_DECAY})")
    logger.info(f"  Scheduler: CosineAnnealingWarmRestarts")
    logger.info(f"  Batch size: {BATCH_SIZE}, Gradient clip: {GRADIENT_CLIP}")

    start_time = time.time()
    result = train_model(
        model=model,
        train_loader=train_loader,
        val_loader=val_loader,
        test_loader=test_loader,
        n_epochs=N_EPOCHS,
        learning_rate=LEARNING_RATE,
        weight_decay=WEIGHT_DECAY,
        patience=PATIENCE,
        gradient_clip=GRADIENT_CLIP,
        device='cpu',
        model_name=MODEL_NAME,
        save_dir=Path(SAVE_DIR).parent.as_posix(),
        class_weights=class_weights,
        verbose=True,
    )
    training_time = time.time() - start_time

    # 7. Выводим финальные метрики
    logger.info("\n" + "=" * 60)
    logger.info("РЕЗУЛЬТАТЫ ОБУЧЕНИЯ")
    logger.info("=" * 60)

    logger.info(f"\nОбщая информация:")
    logger.info(f"  Всего эпох: {result.total_epochs}")
    logger.info(f"  Лучшая эпоха: {result.best_epoch + 1}")
    logger.info(f"  Early stopping: {'Да' if result.total_epochs < N_EPOCHS else 'Нет'}")
    logger.info(f"  Время обучения: {training_time:.2f} сек ({training_time / 60:.2f} мин)")

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

    if result.test_metrics is not None:
        logger.info(f"\nФинальные метрики (тест):")
        test_preds = result.test_metrics['predictions']
        test_labels = result.test_metrics['labels']
        print_metrics_report(
            'TEST',
            test_labels,
            test_preds,
            result.test_metrics['loss'],
        )

    # 8. Сохраняем модель и артефакты
    best_params = {
        'input_size': INPUT_SIZE,
        'hidden_size': HIDDEN_SIZE,
        'num_layers': NUM_LAYERS,
        'dropout': DROPOUT,
        'bidirectional': False,
    }

    metrics_summary = {
        'val_loss': float(result.val_metrics['loss']),
        'val_accuracy': float(result.val_metrics['accuracy']),
        'val_confusion_matrix': result.val_metrics['confusion_matrix'].tolist(),
        'best_epoch': result.best_epoch + 1,
        'total_epochs': result.total_epochs,
        'training_time_sec': training_time,
    }

    if result.test_metrics is not None:
        metrics_summary['test_loss'] = float(result.test_metrics['loss'])
        metrics_summary['test_accuracy'] = float(result.test_metrics['accuracy'])
        metrics_summary['test_confusion_matrix'] = result.test_metrics['confusion_matrix'].tolist()

    model_path = save_model_artifacts(
        model=result.model,
        data=data,
        best_params=best_params,
        metrics=metrics_summary,
        save_dir=SAVE_DIR,
    )

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

    # Финальный отчёт
    logger.info("\n" + "=" * 60)
    logger.info("ИТОГОВЫЙ ОТЧЁТ")
    logger.info("=" * 60)
    logger.info(f"  Тикер: {TICKER}")
    logger.info(f"  Таймфрейм: {TIMEFRAME}")
    logger.info(f"  Архитектура: LSTMPredictor (hidden={HIDDEN_SIZE}, layers={NUM_LAYERS})")
    logger.info(f"  Датасет: {DATA_PATH}")
    logger.info(f"  Веса классов: HOLD={class_weights[0]:.4f}, "
                f"BUY={class_weights[1]:.4f}, SELL={class_weights[2]:.4f}")
    logger.info(f"  Эпох до early stopping: {result.total_epochs} / {N_EPOCHS}")
    logger.info(f"  Лучшая эпоха: {result.best_epoch + 1}")
    logger.info(f"  Val Loss: {result.val_metrics['loss']:.6f}")
    logger.info(f"  Val Accuracy: {result.val_metrics['accuracy']:.4f}")
    if result.test_metrics:
        logger.info(f"  Test Loss: {result.test_metrics['loss']:.6f}")
        logger.info(f"  Test Accuracy: {result.test_metrics['accuracy']:.4f}")
    logger.info(f"  Время обучения: {training_time:.2f} сек")
    logger.info(f"  Модель сохранена: {model_path}")
    logger.info("=" * 60)


if __name__ == '__main__':
    main()
