"""
Модуль Transformer-архитектур для прогнозирования направления цены.

Содержит классы:
    PositionalEncoding — позиционное кодирование (синусоидальное)
    D1Transformer     — Transformer для D1 таймфрейма (альтернатива LSTM/GRU)
"""

import math

import torch
import torch.nn as nn


class PositionalEncoding(nn.Module):
    """
    Синусоидальное позиционное кодирование.

    Позволяет модели учитывать порядок следования временных шагов
    без рекуррентных связей.

    Args:
        d_model: размерность модели (embeddings).
        max_len: максимальная длина последовательности.
        dropout: вероятность dropout после добавления кодирования.
    """

    def __init__(self, d_model: int = 64, max_len: int = 120, dropout: float = 0.1):
        super().__init__()
        self.dropout = nn.Dropout(dropout)

        pe = torch.zeros(max_len, d_model)
        position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
        div_term = torch.exp(
            torch.arange(0, d_model, 2).float()
            * (-math.log(10000.0) / d_model)
        )
        pe[:, 0::2] = torch.sin(position * div_term)
        pe[:, 1::2] = torch.cos(position * div_term)
        pe = pe.unsqueeze(0)  # (1, max_len, d_model)
        self.register_buffer('pe', pe)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Args:
            x: тензор формы (batch_size, seq_len, d_model).

        Returns:
            x + позиционное кодирование (с dropout).
        """
        x = x + self.pe[:, :x.size(1), :]
        return self.dropout(x)


class D1Transformer(nn.Module):
    """
    Transformer-модель для предсказания направления цены (3 класса).

    Отличается от LSTM/GRU использованием механизма внимания вместо
    рекуррентных связей. Добавляет архитектурную диверсификацию в ансамбль.

    Architecture:
        - Linear проекция признаков → d_model
        - Positional Encoding
        - TransformerEncoder (Multi-Head Self-Attention + FFN)
        - Mean pooling по временной оси
        - Полносвязный классификатор (d_model → 64 → 3)

    Args:
        input_size: количество признаков на шаг (по умолчанию 61).
        d_model: размерность эмбеддингов (должна быть кратна nhead).
        nhead: количество голов внимания.
        num_layers: количество слоёв TransformerEncoder.
        dim_feedforward: размерность скрытого слоя в FFN.
        dropout: вероятность dropout.
        activation: функция активации ('relu' или 'gelu').
    """

    def __init__(
        self,
        input_size: int = 61,
        d_model: int = 64,
        nhead: int = 4,
        num_layers: int = 2,
        dim_feedforward: int = 128,
        dropout: float = 0.3,
        activation: str = 'gelu',
    ):
        super().__init__()

        self.input_size = input_size
        self.d_model = d_model
        self.nhead = nhead
        self.num_layers = num_layers

        # Проекция признаков в d_model
        self.input_proj = nn.Linear(input_size, d_model)

        # Позиционное кодирование
        self.pos_encoder = PositionalEncoding(d_model, dropout=dropout)

        # TransformerEncoder
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=d_model,
            nhead=nhead,
            dim_feedforward=dim_feedforward,
            dropout=dropout,
            activation=activation,
            batch_first=True,
            norm_first=True,  # Pre-norm: стабильнее обучение
        )
        self.transformer_encoder = nn.TransformerEncoder(
            encoder_layer,
            num_layers=num_layers,
        )

        # Layer norm после encoder
        self.norm = nn.LayerNorm(d_model)

        # Классификатор
        self.classifier = nn.Sequential(
            nn.Linear(d_model, 64),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(64, 3),
        )

        # Инициализация Xavier
        self._init_weights()

    def _init_weights(self):
        """Инициализация весов Xavier uniform."""
        for p in self.parameters():
            if p.dim() > 1:
                nn.init.xavier_uniform_(p)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Прямой проход модели.

        Args:
            x: тензор формы (batch_size, seq_len, input_size).

        Returns:
            Логиты формы (batch_size, 3) для классов (HOLD, BUY, SELL).
        """
        # Проекция в d_model
        x = self.input_proj(x)  # (batch, seq, d_model)

        # Позиционное кодирование
        x = self.pos_encoder(x)  # (batch, seq, d_model)

        # Transformer encoder
        x = self.transformer_encoder(x)  # (batch, seq, d_model)

        # Layer norm
        x = self.norm(x)

        # Mean pooling по временной оси
        x = x.mean(dim=1)  # (batch, d_model)

        # Классификация
        logits = self.classifier(x)  # (batch, 3)
        return logits
