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

Содержит классы:
    LSTMPredictor — LSTM для классификации BUY/HOLD/SELL
    GRUPredictor  — GRU (облегчённая альтернатива LSTM)
"""

import torch
import torch.nn as nn
from typing import Optional


class LSTMPredictor(nn.Module):
    """
    LSTM-модель для предсказания направления цены (3 класса: BUY/HOLD/SELL).

    Architecture:
        - Stacked LSTM слои
        - Полносвязный классификатор (hidden → 64 → 3)
        - Dropout для регуляризации

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

    def __init__(
        self,
        input_size: int = 61,
        hidden_size: int = 128,
        num_layers: int = 2,
        dropout: float = 0.2,
        bidirectional: bool = False,
    ) -> None:
        super().__init__()

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

        # LSTM слои
        self.lstm = nn.LSTM(
            input_size=input_size,
            hidden_size=hidden_size,
            num_layers=num_layers,
            batch_first=True,
            dropout=dropout if num_layers > 1 else 0,
            bidirectional=bidirectional,
        )

        # Размер выхода LSTM
        lstm_out_dim = hidden_size * (2 if bidirectional else 1)

        # Классификатор
        self.classifier = nn.Sequential(
            nn.Linear(lstm_out_dim, 64),
            nn.ReLU(),
            nn.Dropout(dropout),
            nn.Linear(64, 3),  # 3 класса: BUY(1), HOLD(0), SELL(2)
        )

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

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

        Returns:
            Логиты формы (batch_size, 3) для классов (HOLD, BUY, SELL).
        """
        # LSTM проход
        lstm_out, (hidden, cell) = self.lstm(x)
        # Берём выход последнего временного шага
        last_out = lstm_out[:, -1, :]
        # Классификация
        logits = self.classifier(last_out)
        return logits


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

    Облегчённая альтернатива LSTM с меньшим количеством параметров.
    Рекомендуется при ограниченном объёме данных (< 10000 семплов).

    Args:
        input_size: количество признаков на шаг.
        hidden_size: размер скрытого состояния GRU.
        num_layers: количество stacked GRU слоёв.
        dropout: вероятность dropout.
    """

    def __init__(
        self,
        input_size: int = 61,
        hidden_size: int = 128,
        num_layers: int = 2,
        dropout: float = 0.2,
    ) -> None:
        super().__init__()

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

        self.classifier = nn.Sequential(
            nn.Linear(hidden_size, 64),
            nn.ReLU(),
            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) для классов.
        """
        gru_out, _ = self.gru(x)
        last_out = gru_out[:, -1, :]
        logits = self.classifier(last_out)
        return logits
