"""
MultiTimeframeFusion — архитектура для объединения H1 и D1 таймфреймов.

Архитектура:
  1. H1 Encoder: GRU (seq_len=120, input_size=n_features) → вектор h1_emb
  2. D1 Encoder: GRU (seq_len=30,  input_size=n_features) → вектор d1_emb
  3. Fusion: Cross-attention или конкатенация → (h1_emb || d1_emb) → FC → 3 класса

Варианты:
  - MTFBasic: конкатенация эмбеддингов (проще, меньше переобучение)
  - MTFAdvanced: cross-attention между эмбеддингами (сложнее, больше expressivity)
"""

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


class MTFBasic(nn.Module):
    """
    MultiTimeframeFusion — базовая версия.

    H1 Encoder (GRU) → h1_emb
    D1 Encoder (GRU) → d1_emb
    Concat(h1_emb, d1_emb) → FC(128→64→3) → logits

    Args:
        input_size: количество признаков (общее для H1 и D1).
        h1_hidden: размер скрытого состояния H1 GRU.
        d1_hidden: размер скрытого состояния D1 GRU.
        h1_num_layers: количество GRU слоёв для H1.
        d1_num_layers: количество GRU слоёв для D1.
        dropout: вероятность dropout в GRU и классификаторе.
    """

    def __init__(
        self,
        input_size: int = 61,
        h1_hidden: int = 96,
        d1_hidden: int = 96,
        h1_num_layers: int = 2,
        d1_num_layers: int = 2,
        dropout: float = 0.4,
    ):
        super().__init__()

        # H1 Encoder (обрабатывает ~120 H1-баров)
        self.h1_gru = nn.GRU(
            input_size=input_size,
            hidden_size=h1_hidden,
            num_layers=h1_num_layers,
            batch_first=True,
            dropout=dropout if h1_num_layers > 1 else 0,
        )

        # D1 Encoder (обрабатывает ~30 D1-баров)
        self.d1_gru = nn.GRU(
            input_size=input_size,
            hidden_size=d1_hidden,
            num_layers=d1_num_layers,
            batch_first=True,
            dropout=dropout if d1_num_layers > 1 else 0,
        )

        # Fusion classifier
        fusion_size = h1_hidden + d1_hidden
        self.classifier = nn.Sequential(
            nn.Linear(fusion_size, 128),
            nn.ReLU(),
            nn.Dropout(dropout),
            nn.Linear(128, 64),
            nn.ReLU(),
            nn.Dropout(dropout),
            nn.Linear(64, 3),  # BUY / HOLD / SELL
        )

        # LayerNorm для стабильности
        self.h1_norm = nn.LayerNorm(h1_hidden)
        self.d1_norm = nn.LayerNorm(d1_hidden)

    def forward(
        self,
        h1_x: torch.Tensor,
        d1_x: torch.Tensor,
    ) -> torch.Tensor:
        """
        Прямой проход.

        Args:
            h1_x: (batch, h1_seq_len, input_size) — H1 последовательность.
            d1_x: (batch, d1_seq_len, input_size) — D1 последовательность.

        Returns:
            logits: (batch, 3) — логиты для BUY/HOLD/SELL.
        """
        # Замена NaN/Inf на 0 для стабильности
        h1_x = torch.nan_to_num(h1_x, nan=0.0, posinf=10.0, neginf=-10.0)
        d1_x = torch.nan_to_num(d1_x, nan=0.0, posinf=10.0, neginf=-10.0)

        # H1 Encoding
        h1_out, _ = self.h1_gru(h1_x)
        h1_emb = self.h1_norm(h1_out[:, -1, :])  # последний шаг H1

        # D1 Encoding
        d1_out, _ = self.d1_gru(d1_x)
        d1_emb = self.d1_norm(d1_out[:, -1, :])  # последний шаг D1

        # Fusion: конкатенация
        fused = torch.cat([h1_emb, d1_emb], dim=-1)

        # Классификация
        logits = self.classifier(fused)
        # Клиппинг логитов для численной стабильности Focal Loss
        logits = torch.clamp(logits, min=-15.0, max=15.0)
        return logits


class MTFAdvanced(nn.Module):
    """
    MultiTimeframeFusion — продвинутая версия с cross-attention.

    H1 Encoder (Transformer) → h1_emb
    D1 Encoder (Transformer) → d1_emb
    Cross-attention между эмбеддингами → fused
    FC(128→3) → logits
    """

    def __init__(
        self,
        input_size: int = 61,
        d_model: int = 128,
        nhead: int = 4,
        num_encoder_layers: int = 2,
        dropout: float = 0.3,
    ):
        super().__init__()

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

        # H1 Transformer Encoder
        h1_enc_layer = nn.TransformerEncoderLayer(
            d_model, nhead, dim_feedforward=d_model * 4,
            dropout=dropout, batch_first=True,
        )
        self.h1_encoder = nn.TransformerEncoder(h1_enc_layer, num_encoder_layers)

        # D1 Transformer Encoder
        d1_enc_layer = nn.TransformerEncoderLayer(
            d_model, nhead, dim_feedforward=d_model * 4,
            dropout=dropout, batch_first=True,
        )
        self.d1_encoder = nn.TransformerEncoder(d1_enc_layer, num_encoder_layers)

        # Cross-attention fusion
        self.cross_attn = nn.MultiheadAttention(
            d_model, nhead, dropout=dropout, batch_first=True,
        )

        # Classifier
        self.classifier = nn.Sequential(
            nn.Linear(d_model * 2, 128),
            nn.ReLU(),
            nn.Dropout(dropout),
            nn.Linear(128, 3),
        )

    def forward(
        self,
        h1_x: torch.Tensor,
        d1_x: torch.Tensor,
    ) -> torch.Tensor:
        """
        Прямой проход.

        Args:
            h1_x: (batch, h1_seq_len, input_size)
            d1_x: (batch, d1_seq_len, input_size)

        Returns:
            logits: (batch, 3)
        """
        B = h1_x.size(0)

        # Проекция + позиционное кодирование
        h1_proj = self.pos_encoding(self.input_proj(h1_x))
        d1_proj = self.pos_encoding(self.input_proj(d1_x))

        # Encoders
        h1_encoded = self.h1_encoder(h1_proj)  # (B, h1_len, d_model)
        d1_encoded = self.d1_encoder(d1_proj)  # (B, d1_len, d_model)

        # Берём последний шаг каждого
        h1_feat = h1_encoded[:, -1, :].unsqueeze(1)  # (B, 1, d_model)
        d1_feat = d1_encoded[:, -1, :].unsqueeze(1)  # (B, 1, d_model)

        # Cross-attention: H1 attends to D1, D1 attends to H1
        h1_attended, _ = self.cross_attn(h1_feat, d1_encoded, d1_encoded)
        d1_attended, _ = self.cross_attn(d1_feat, h1_encoded, h1_encoded)

        # Fusion
        fused = torch.cat([h1_attended.squeeze(1), d1_attended.squeeze(1)], dim=-1)

        return self.classifier(fused)


class PositionalEncoding(nn.Module):
    """Sinusoidal positional encoding."""
    def __init__(self, d_model: int, max_len: int = 500):
        super().__init__()
        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() *
            (-torch.log(torch.tensor(10000.0)) / d_model)
        )
        pe[:, 0::2] = torch.sin(position * div_term)
        pe[:, 1::2] = torch.cos(position * div_term)
        self.register_buffer('pe', pe.unsqueeze(0))

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return x + self.pe[:, :x.size(1), :]
