"""
Модуль нейросетевых архитектур для определения фаз рынка по методу Вайкоффа.

Архитектуры:
    1. WyckoffLSTM         — LSTM с Attention для Wyckoff-фаз (основная)
    2. WyckoffCNN          — 1D CNN для Wyckoff-фаз (лёгкая, быстрая)
    3. WyckoffTransformer  — Transformer для Wyckoff-фаз (точная, сложная)
    4. WyckoffEnsemble     — Ансамбль из двух моделей

Выход:
    - 8 классов (полные фазы Вайкоффа)
    - 5 классов (упрощённые фазы)
    - confidence (0-100%)

Usage:
    from src.ml.models.wyckoff import WyckoffLSTM, WyckoffCNN, WyckoffTransformer

    model = WyckoffLSTM(input_size=80, num_classes=8)
    logits = model(x)  # (batch, 8) — логиты для 8 фаз
"""

import math
from typing import Optional

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


class AttentionLayer(nn.Module):
    """
    Аддитивное (Bahdanau) внимание.

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

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

    def forward(self, lstm_outputs: torch.Tensor) -> torch.Tensor:
        """
        Args:
            lstm_outputs: (batch, seq_len, hidden_size)

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


class WyckoffLSTM(nn.Module):
    """
    LSTM с Attention для классификации фаз Вайкоффа.

    Архитектура:
        - 2-layer Bidirectional LSTM (захват контекста в обе стороны)
        - Attention (взвешивание важных моментов последовательности)
        - Полносвязный классификатор с residual connections
        - Выход: логиты для num_classes фаз

    Особенности:
        - LayerNorm для стабильности
        - Dropout 0.3-0.5 для регуляризации
        - Xavier инициализация

    Args:
        input_size: количество признаков на шаг.
        hidden_size: размер скрытого состояния LSTM (на направление).
        num_layers: количество LSTM слоёв.
        num_classes: количество классов фаз (8 или 5).
        dropout: вероятность dropout.
    """

    def __init__(
        self,
        input_size: int = 80,
        hidden_size: int = 128,
        num_layers: int = 2,
        num_classes: int = 8,
        dropout: float = 0.4,
    ):
        super().__init__()

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

        # BatchNorm для входных данных (стабилизация)
        self.input_norm = nn.LayerNorm(input_size)

        # 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=True,
        )
        lstm_out_dim = hidden_size * 2  # bidirectional

        # Attention
        self.attention = AttentionLayer(lstm_out_dim)
        self.attn_norm = nn.LayerNorm(lstm_out_dim)

        # Классификатор с residual
        self.classifier = nn.Sequential(
            nn.Linear(lstm_out_dim, 128),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(128, 64),
            nn.GELU(),
            nn.Dropout(dropout * 0.75),
            nn.Linear(64, num_classes),
        )

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

    def _init_weights(self):
        """Xavier uniform инициализация для LSTM и Linear слоёв."""
        for name, param in self.lstm.named_parameters():
            if 'weight_ih' in name:
                nn.init.xavier_uniform_(param)
            elif 'weight_hh' in name:
                nn.init.orthogonal_(param)
            elif 'bias' in name:
                nn.init.zeros_(param)
        for m in self.classifier:
            if isinstance(m, nn.Linear):
                nn.init.xavier_uniform_(m.weight)
                if m.bias is not None:
                    nn.init.zeros_(m.bias)

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

        Args:
            x: тензор признаков (batch_size, seq_len, input_size).

        Returns:
            logits: (batch_size, num_classes) — логиты для Wyckoff-фаз.
        """
        # Замена NaN/Inf для стабильности
        x = torch.nan_to_num(x, nan=0.0, posinf=10.0, neginf=-10.0)

        # Нормализация входа
        x = self.input_norm(x)

        # LSTM
        lstm_out, _ = self.lstm(x)  # (batch, seq, hidden*2)

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

        # Классификация
        logits = self.classifier(context)

        # Клиппинг логитов для численной стабильности
        logits = torch.clamp(logits, min=-15.0, max=15.0)

        return logits


class WyckoffCNN(nn.Module):
    """
    1D CNN для определения фаз Вайкоффа (лёгкая, быстрая архитектура).

    Подходит для:
        - Быстрого инференса на большом количестве тикеров
        - Retail-устройств (CPU/мало памяти)
        - Первичного скрининга

    Архитектура:
        - 3 свёрточных слоя с увеличивающимся dilation
        - Residual соединения
        - AdaptiveAvgPool для объединения по seq_len
        - Классификатор

    Args:
        input_size: количество признаков на шаг.
        num_classes: количество классов фаз.
        seq_len: длина входной последовательности (для расчёта padding).
        dropout: вероятность dropout.
    """

    def __init__(
        self,
        input_size: int = 80,
        num_classes: int = 8,
        seq_len: int = 60,
        dropout: float = 0.3,
    ):
        super().__init__()

        self.input_size = input_size
        self.num_classes = num_classes
        self.seq_len = seq_len

        # Входная проекция
        self.input_proj = nn.Conv1d(input_size, 64, kernel_size=1)

        # Свёрточные блоки с dilation (расширение рецептивного поля)
        self.conv1 = nn.Conv1d(64, 128, kernel_size=5, dilation=1, padding=2)
        self.conv2 = nn.Conv1d(128, 128, kernel_size=5, dilation=2, padding=4)
        self.conv3 = nn.Conv1d(128, 256, kernel_size=5, dilation=4, padding=8)

        # Residual проекции
        self.residual1 = nn.Conv1d(64, 128, kernel_size=1) if input_size != 128 else nn.Identity()
        self.residual2 = nn.Conv1d(128, 256, kernel_size=1)

        self.norm1 = nn.BatchNorm1d(128)
        self.norm2 = nn.BatchNorm1d(128)
        self.norm3 = nn.BatchNorm1d(256)

        self.dropout = nn.Dropout(dropout)

        # Pooling + Classifier
        self.pool = nn.AdaptiveAvgPool1d(1)
        self.classifier = nn.Sequential(
            nn.Linear(256, 128),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(128, num_classes),
        )

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

        Args:
            x: (batch, seq_len, input_size).

        Returns:
            logits: (batch, num_classes).
        """
        x = torch.nan_to_num(x, nan=0.0, posinf=10.0, neginf=-10.0)

        # (batch, input_size, seq_len) — Conv1d expects channels first
        x = x.transpose(1, 2)

        # Input projection
        x0 = self.input_proj(x)  # (batch, 64, seq)
        x0 = F.gelu(x0)

        # Conv block 1
        identity = self.residual1(x0)
        x1 = self.conv1(x0)
        x1 = self.norm1(x1)
        x1 = F.gelu(x1 + identity)

        # Conv block 2
        x2 = self.conv2(x1)
        x2 = self.norm2(x2)
        x2 = F.gelu(x2)

        # Conv block 3
        identity = self.residual2(x2)
        x3 = self.conv3(x2)
        x3 = self.norm3(x3)
        x3 = F.gelu(x3 + identity)

        x3 = self.dropout(x3)

        # Pooling
        x_pool = self.pool(x3).squeeze(-1)  # (batch, 256)

        # Classifier
        logits = self.classifier(x_pool)
        logits = torch.clamp(logits, min=-15.0, max=15.0)

        return logits


class WyckoffTransformer(nn.Module):
    """
    Transformer для определения фаз Вайкоффа (наиболее точная архитектура).

    Использует multi-head self-attention для захвата сложных зависимостей
    между VSA-событиями и ценовыми паттернами на всей длине последовательности.

    Архитектура:
        - Linear проекция + Positional Encoding
        - 3-4 слоя TransformerEncoder с Pre-LN
        - Mean + Max pooling (объединение по seq_len)
        - Классификатор с Dropout

    Особенности:
        - Pre-LayerNorm (стабильнее Post-LN)
        - GELU активация в FFN
        - Xavier инициализация

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

    def __init__(
        self,
        input_size: int = 80,
        d_model: int = 128,
        nhead: int = 4,
        num_layers: int = 3,
        dim_feedforward: int = 256,
        num_classes: int = 8,
        max_seq_len: int = 120,
        dropout: float = 0.3,
    ):
        super().__init__()

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

        # Входная проекция
        self.input_proj = nn.Linear(input_size, d_model)
        self.input_norm = nn.LayerNorm(d_model)

        # Позиционное кодирование (синусоидальное)
        self.pos_encoding = PositionalEncoding(d_model, max_seq_len, dropout)

        # Transformer Encoder с Pre-LN
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=d_model,
            nhead=nhead,
            dim_feedforward=dim_feedforward,
            dropout=dropout,
            activation='gelu',
            batch_first=True,
            norm_first=True,  # Pre-LN: стабильнее
        )
        self.transformer = nn.TransformerEncoder(
            encoder_layer,
            num_layers=num_layers,
        )

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

        # Классификатор (mean + max pooling)
        self.classifier = nn.Sequential(
            nn.Linear(d_model * 2, 128),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(128, 64),
            nn.GELU(),
            nn.Dropout(dropout * 0.75),
            nn.Linear(64, num_classes),
        )

        self._init_weights()

    def _init_weights(self):
        """Xavier инициализация."""
        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:
            logits: (batch_size, num_classes).
        """
        x = torch.nan_to_num(x, nan=0.0, posinf=10.0, neginf=-10.0)

        # Проекция в d_model
        x = self.input_proj(x)  # (batch, seq, d_model)
        x = self.input_norm(x)

        # Позиционное кодирование
        x = self.pos_encoding(x)

        # Transformer
        x = self.transformer(x)  # (batch, seq, d_model)
        x = self.encoder_norm(x)

        # Mean + Max pooling
        mean_pool = x.mean(dim=1)  # (batch, d_model)
        max_pool, _ = x.max(dim=1)  # (batch, d_model)

        # Конкатенация
        pooled = torch.cat([mean_pool, max_pool], dim=-1)  # (batch, d_model*2)

        # Классификация
        logits = self.classifier(pooled)
        logits = torch.clamp(logits, min=-15.0, max=15.0)

        return logits


class WyckoffMultiTimeframe(nn.Module):
    """
    Multi-timeframe модель для определения фаз Вайкоффа.

    Объединяет H1, D1 и W1 таймфреймы для более точного определения
    контекстуальной фазы рынка.

    Архитектура:
        - Отдельные Transformer энкодеры для каждого ТФ
        - Cross-attention fusion
        - Общий классификатор

    Args:
        input_size: количество признаков на шаг (общее для всех ТФ).
        d_model: размерность эмбеддингов.
        nhead: количество голов внимания.
        num_classes: количество классов фаз.
        dropout: вероятность dropout.
    """

    def __init__(
        self,
        input_size: int = 80,
        d_model: int = 128,
        nhead: int = 4,
        num_classes: int = 8,
        dropout: float = 0.3,
    ):
        super().__init__()

        # Общая проекция
        self.input_proj = nn.Linear(input_size, d_model)

        # H1 Encoder
        self.h1_encoder = nn.TransformerEncoder(
            nn.TransformerEncoderLayer(d_model, nhead, d_model * 4, dropout,
                                       activation='gelu', batch_first=True, norm_first=True),
            num_layers=2,
        )

        # D1 Encoder
        self.d1_encoder = nn.TransformerEncoder(
            nn.TransformerEncoderLayer(d_model, nhead, d_model * 4, dropout,
                                       activation='gelu', batch_first=True, norm_first=True),
            num_layers=2,
        )

        # W1 Encoder
        self.w1_encoder = nn.TransformerEncoder(
            nn.TransformerEncoderLayer(d_model, nhead, d_model * 4, dropout,
                                       activation='gelu', batch_first=True, norm_first=True),
            num_layers=2,
        )

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

        # Feature fusion
        self.fusion_norm = nn.LayerNorm(d_model * 3)

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

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

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

        Returns:
            logits: (batch, num_classes)
        """
        # Замена NaN/Inf
        h1_x, d1_x, w1_x = [
            torch.nan_to_num(t, nan=0.0, posinf=10.0, neginf=-10.0)
            for t in (h1_x, d1_x, w1_x)
        ]

        # Проекция
        h1_emb = self.input_proj(h1_x)
        d1_emb = self.input_proj(d1_x)
        w1_emb = self.input_proj(w1_x)

        # Encoders
        h1_feat = self.h1_encoder(h1_emb)[:, -1, :].unsqueeze(1)  # (B, 1, d_model)
        d1_feat = self.d1_encoder(d1_emb)[:, -1, :].unsqueeze(1)
        w1_feat = self.w1_encoder(w1_emb)[:, -1, :].unsqueeze(1)

        # Cross-attention: каждый attends ко всем
        combined = torch.cat([h1_feat, d1_feat, w1_feat], dim=1)  # (B, 3, d_model)
        attended, _ = self.cross_attn(combined, combined, combined)

        # Flatten
        fused = attended.reshape(attended.size(0), -1)  # (B, d_model*3)
        fused = self.fusion_norm(fused)

        # Классификация
        logits = self.classifier(fused)
        logits = torch.clamp(logits, min=-15.0, max=15.0)

        return logits


class WyckoffEnsemble(nn.Module):
    """
    Ансамбль из двух моделей для определения фаз Вайкоффа.

    Объединяет WyckoffLSTM (контекст) и WyckoffCNN (паттерны) через
    взвешенное голосование или конкатенацию эмбеддингов.

    Args:
        lstm_model: экземпляр WyckoffLSTM.
        cnn_model: экземпляр WyckoffCNN.
        num_classes: количество классов.
        fusion: способ объединения ('concat' или 'weighted').
    """

    def __init__(
        self,
        lstm_model: WyckoffLSTM,
        cnn_model: WyckoffCNN,
        num_classes: int = 8,
        fusion: str = 'concat',
    ):
        super().__init__()

        self.lstm = lstm_model
        self.cnn = cnn_model
        self.fusion = fusion

        if fusion == 'concat':
            # Убираем классификаторы из обеих моделей
            self.lstm.classifier = nn.Identity()
            self.cnn.classifier = nn.Identity()

            lstm_out = lstm_model.hidden_size * 2  # bidirectional
            cnn_out = 256  # последний канал CNN

            self.fusion_classifier = nn.Sequential(
                nn.Linear(lstm_out + cnn_out, 128),
                nn.GELU(),
                nn.Dropout(0.3),
                nn.Linear(128, num_classes),
            )

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

        Args:
            x: (batch, seq_len, input_size).

        Returns:
            logits: (batch, num_classes).
        """
        if self.fusion == 'weighted':
            # Взвешенное голосование
            lstm_logits = self.lstm(x)
            cnn_logits = self.cnn(x)
            # LSTM получает больший вес (контекст важнее)
            logits = 0.6 * lstm_logits + 0.4 * cnn_logits
        elif self.fusion == 'concat':
            # Конкатенация эмбеддингов
            lstm_emb = self.lstm(x)
            cnn_emb = self.cnn(x)
            combined = torch.cat([lstm_emb, cnn_emb], dim=-1)
            logits = self.fusion_classifier(combined)
        else:
            logits = self.lstm(x)

        return torch.clamp(logits, min=-15.0, max=15.0)


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

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

    def __init__(self, d_model: int, 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:
        """x: (batch, seq_len, d_model)"""
        x = x + self.pe[:, :x.size(1), :]
        return self.dropout(x)


# Словарь для выбора архитектуры по имени
WYCKOFF_MODELS = {
    'lstm': WyckoffLSTM,
    'cnn': WyckoffCNN,
    'transformer': WyckoffTransformer,
    'mtf': WyckoffMultiTimeframe,
    'ensemble': WyckoffEnsemble,
}


def create_wyckoff_model(
    model_type: str = 'lstm',
    input_size: int = 80,
    num_classes: int = 8,
    **kwargs,
) -> nn.Module:
    """
    Фабричный метод для создания модели определения фаз Вайкоффа.

    Args:
        model_type: тип модели ('lstm', 'cnn', 'transformer', 'mtf', 'ensemble').
        input_size: количество признаков.
        num_classes: количество классов (8 или 5).
        **kwargs: дополнительные параметры для конкретной архитектуры.

    Returns:
        Экземпляр модели.

    Raises:
        ValueError: если model_type не поддерживается.
    """
    # Фильтруем kwargs по сигнатуре конкретной модели
    import inspect

    if model_type == 'lstm':
        sig = inspect.signature(WyckoffLSTM.__init__)
        lstm_kwargs = {k: v for k, v in kwargs.items() if k in sig.parameters}
        return WyckoffLSTM(input_size=input_size, num_classes=num_classes, **lstm_kwargs)
    elif model_type == 'cnn':
        sig = inspect.signature(WyckoffCNN.__init__)
        cnn_kwargs = {k: v for k, v in kwargs.items() if k in sig.parameters}
        return WyckoffCNN(input_size=input_size, num_classes=num_classes, **cnn_kwargs)
    elif model_type == 'transformer':
        sig = inspect.signature(WyckoffTransformer.__init__)
        tfm_kwargs = {k: v for k, v in kwargs.items() if k in sig.parameters}
        return WyckoffTransformer(input_size=input_size, num_classes=num_classes, **tfm_kwargs)
    elif model_type == 'ensemble':
        lstm = WyckoffLSTM(input_size=input_size, num_classes=num_classes)
        cnn = WyckoffCNN(input_size=input_size, num_classes=num_classes)
        return WyckoffEnsemble(lstm, cnn, num_classes=num_classes, **kwargs)
    else:
        raise ValueError(
            f"Неизвестный тип модели: '{model_type}'. "
            f"Доступные: {list(WYCKOFF_MODELS.keys())}"
        )
