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

Загружает x5_h1_dataset.h5 (5,098 + 1,093 + 1,093 семплов).
Использует class weights для балансировки (HOLD=77% + BUY=13% + SELL=13%).
Обучает LSTM (input_size=63, hidden_size=128, num_layers=2, dropout=0.3).
100 эпох, patience=15.
Сохраняет как src/ml/models/saved/x5_h1_lstm_v1/.

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

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

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


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

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


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

    Args:
        model: обученная модель.
        data: словарь с данными.
        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': list(data.get('feature_names', [])),
        'n_classes': 3,
        'class_names': {0: 'HOLD', 1: 'BUY', 2: 'SELL'},
        'metrics': metrics,
        'ticker': TICKER,
        'timeframe': TIMEFRAME,
        'training_type': 'single_stage_h1',
    }

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

    # Сохраняем параметры нормализации
    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 main() -> None:
    """
    Основная функция обучения на H1.

    Последовательность:
        1. Установка seed
        2. Загрузка датасета
        3. Создание DataLoader'ов
        4. Вычисление весов классов (для балансировки HOLD)
        5. Создание модели
        6. Обучение
        7. Вывод метрик
        8. Сохранение модели
        9. Регистрация в ModelRegistry
    """
    logger.info("=" * 70)
    logger.info("ОБУЧЕНИЕ LSTM МОДЕЛИ ДЛЯ X5 H1")
    logger.info("=" * 70)
    logger.info(f"  Данные: {DATA_PATH}")
    logger.info(f"  Модель: LSTMPredictor (input={INPUT_SIZE}, "
                f"hidden={HIDDEN_SIZE}, layers={NUM_LAYERS})")
    logger.info(f"  H1 seq_len=120, features={INPUT_SIZE}")

    # 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. Загружаем датасет
    logger.info("\n" + "-" * 60)
    logger.info("ЗАГРУЗКА ДАННЫХ")
    logger.info("-" * 60)
    data = load_hdf5_dataset(DATA_PATH)

    # 3. Создаём DataLoader'ы
    logger.info("\n" + "-" * 60)
    logger.info("СОЗДАНИЕ DataLoader'ОВ")
    logger.info("-" * 60)
    train_loader, val_loader, test_loader = create_dataloaders(
        data, batch_size=BATCH_SIZE, shuffle_train=True
    )

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

    # 5. Создаём модель (input_size=63 для H1)
    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Модель создана: 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}")
    logger.info(f"  Class weights: применены для балансировки")

    training_start = 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() - training_start

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

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

    test_metrics = None
    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'],
        )
        test_metrics = {
            'loss': float(result.test_metrics['loss']),
            'accuracy': float(result.test_metrics['accuracy']),
            'macro_f1': float(calculate_macro_f1(test_labels, test_preds)),
            'confusion_matrix': result.test_metrics['confusion_matrix'].tolist(),
        }

    # 8. Сохраняем модель и артефакты
    best_params = {
        'input_size': INPUT_SIZE,
        'hidden_size': HIDDEN_SIZE,
        'num_layers': NUM_LAYERS,
        'dropout': DROPOUT,
        'bidirectional': False,
        'seq_len': 120,  # H1 использует 120 баров
        'batch_size': BATCH_SIZE,
        'learning_rate': LEARNING_RATE,
        'weight_decay': WEIGHT_DECAY,
    }

    val_metrics = {
        'val_loss': float(result.val_metrics['loss']),
        'val_accuracy': float(result.val_metrics['accuracy']),
        'val_macro_f1': float(calculate_macro_f1(val_labels, val_preds)),
        '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 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=result.model,
        data=data,
        best_params=best_params,
        metrics=val_metrics,
        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=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("ИТОГОВЫЙ ОТЧЁТ — ОБУЧЕНИЕ НА H1")
    logger.info("=" * 70)
    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"  Train семплов: {len(data['X_train'])}")
    logger.info(f"  Val семплов: {len(data['X_val'])}")
    logger.info(f"  Test семплов: {len(data['X_test'])}")
    logger.info(f"  Seq len: 120, Features: {INPUT_SIZE}")
    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 test_metrics:
        logger.info(f"  Test Loss: {test_metrics['loss']:.6f}")
        logger.info(f"  Test Accuracy: {test_metrics['accuracy']:.4f}")
        logger.info(f"  Test Macro F1: {test_metrics['macro_f1']:.4f}")
    logger.info(f"  Время обучения: {training_time:.2f} сек ({training_time / 60:.2f} мин)")
    logger.info(f"  Модель сохранена: {model_path}")
    logger.info("=" * 70)


if __name__ == '__main__':
    main()
