"""
Модуль BiGRU с Attention для прогнозирования направления цены.

Отличия от GRUPredictor/LSTMPredictor:
  - Bidirectional GRU (обрабатывает последовательность в обе стороны)
  - Attention механизм (взвешивает важность временных шагов вместо последнего)
  - Добавляет архитектурную диверсификацию в ансамбль

Классы:
    Attention     — слой аддитивного (Bahdanau) внимания
    BiGRUAttention — BiGRU + Attention + классификатор
"""

import torch
import torch.nn as nn
import torch.nn.functional as F


class Attention(nn.Module):
    """
    Слой аддитивного (Bahdanau-style) внимания.

    Для каждого временного шага вычисляет оценку важности,
    затем взвешенно суммирует скрытые состояния.

    Args:
        hidden_size: размерность скрытого состояния GRU.
    """

    def __init__(self, hidden_size: int = 128):
        super().__init__()
        self.hidden_size = hidden_size
        self.attention_weights = nn.Linear(hidden_size, 1, bias=False)

    def forward(self, gru_outputs: torch.Tensor) -> torch.Tensor:
        """
        Args:
            gru_outputs: выход GRU формы (batch, seq_len, hidden_size).

        Returns:
            context: взвешенная сумма формы (batch, hidden_size).
        """
        # (batch, seq_len, 1)
        scores = self.attention_weights(gru_outputs)
        # (batch, seq_len)
        attention_weights = F.softmax(scores.squeeze(-1), dim=1)
        # (batch, hidden_size) — взвешенная сумма
        context = torch.bmm(attention_weights.unsqueeze(1), gru_outputs).squeeze(1)
        return context


class BiGRUAttention(nn.Module):
    """
    BiGRU + Attention для предсказания BUY/HOLD/SELL.

    Архитектура:
        - 2-layer Bidirectional GRU (hidden=64 на направление → 128 всего)
        - Attention взвешивание временных шагов
        - LayerNorm + Dropout
        - Классификатор (128 → 64 → 3)

    Args:
        input_size: количество признаков на шаг (по умолчанию 61).
        hidden_size: размер скрытого состояния GRU на направление (всего ×2).
        num_layers: количество stacked BiGRU слоёв.
        dropout: вероятность dropout (между GRU слоями и в классификаторе).
        attention_size: размерность слоя attention.
    """

    def __init__(
        self,
        input_size: int = 61,
        hidden_size: int = 64,
        num_layers: int = 2,
        dropout: float = 0.4,
    ):
        super().__init__()

        self.input_size = input_size
        self.hidden_size = hidden_size
        self.num_layers = num_layers
        self.dropout = dropout

        # BiGRU: bidirectional=True → выход hidden_size*2
        self.gru = nn.GRU(
            input_size=input_size,
            hidden_size=hidden_size,
            num_layers=num_layers,
            batch_first=True,
            dropout=dropout if num_layers > 1 else 0,
            bidirectional=True,
        )
        gru_out_dim = hidden_size * 2  # 128

        # Attention
        self.attention = Attention(hidden_size=gru_out_dim)

        # LayerNorm для стабильности
        self.norm = nn.LayerNorm(gru_out_dim)

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

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

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

        Returns:
            Логиты формы (batch_size, 3) для классов (HOLD, BUY, SELL).
        """
        # BiGRU
        gru_out, _ = self.gru(x)  # (batch, seq, hidden*2)

        # Attention
        context = self.attention(gru_out)  # (batch, hidden*2)

        # Norm + classifier
        context = self.norm(context)
        logits = self.classifier(context)  # (batch, 3)
        return logits
