#!/usr/bin/env python3
"""
Гибридная модель: LSTM-кодировщик (окно 20 свечей) → Latent features → Random Forest.
"""
import copy
import numpy as np
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
import warnings
warnings.filterwarnings('ignore')

import config  # DEVICE is shared with config — call config.set_device() before import

DEVICE = config.DEVICE  # alias for backward compatibility
WINDOW = 20
LATENT_DIM = 32
BATCH_SIZE = 256
EPOCHS = 50
PATIENCE = 6


class SequenceDataset(Dataset):
    def __init__(self, X: np.ndarray, y: np.ndarray | None = None):
        self.X = torch.FloatTensor(X)
        if y is not None:
            # y может быть (N,) или (N, H) — сохраняем размерность
            if y.ndim == 1:
                self.y = torch.FloatTensor(y).view(-1, 1)
            else:
                self.y = torch.FloatTensor(y)  # (N, H)
        else:
            self.y = None

    def __len__(self):
        return len(self.X)

    def __getitem__(self, idx):
        if self.y is not None:
            return self.X[idx], self.y[idx]
        return self.X[idx], 0.0


class LSTMEncoder(nn.Module):
    """LSTM-кодировщик: окно → скрытое состояние (latent features)."""
    def __init__(self, n_features: int, hidden_size: int = 32, num_layers: int = 2):
        super().__init__()
        self.hidden_size = hidden_size
        self.num_layers = num_layers
        self.lstm = nn.LSTM(
            input_size=n_features,
            hidden_size=hidden_size,
            num_layers=num_layers,
            batch_first=True,
            dropout=0.2 if num_layers > 1 else 0,
        )
        # Выходной слой для supervised pretraining (predict outcome)
        self.fc = nn.Linear(hidden_size, 1)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        lstm_out, (h_n, _) = self.lstm(x)
        # h_n shape: (num_layers, batch, hidden_size) -> берём последний слой
        h_last = h_n[-1]  # (batch, hidden_size)
        return self.fc(h_last)

    def encode(self, x: torch.Tensor) -> np.ndarray:
        """Извлекает latent features (последний скрытый слой)."""
        _, (h_n, _) = self.lstm(x)
        return h_n[-1].detach().cpu().numpy()


def create_sequences(features: np.ndarray, target: np.ndarray | None = None, window: int = WINDOW):
    """Создаёт последовательности для LSTM."""
    n = len(features)
    X, y = [], []
    for i in range(window, n):
        X.append(features[i - window : i])
        if target is not None:
            y.append(target[i])
    X_arr = np.array(X, dtype=np.float32)
    y_arr = np.array(y, dtype=np.float32) if target is not None else None
    return X_arr, y_arr


def train_encoder(
    X_seq: np.ndarray,
    y_seq: np.ndarray | None = None,
    n_features: int = 51,
    hidden_size: int = 32,
    num_layers: int = 2,
    lr: float = 5e-4,
    epochs: int = EPOCHS,
    batch_size: int = BATCH_SIZE,
    patience: int = PATIENCE,
    device: str = DEVICE,
    verbose: bool = True,
) -> LSTMEncoder:
    """Обучает LSTMEncoder supervised (prediction) или self-supervised."""
    model = LSTMEncoder(n_features, hidden_size, num_layers).to(device)
    optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-5)

    use_target = y_seq is not None
    if use_target:
        # Автоопределение: бинарный таргет (0/1) → BCE, непрерывный → MSE
        unique_vals = np.unique(y_seq)
        if len(unique_vals) <= 2 and set(unique_vals).issubset({0, 1}):
            pos_weight = (y_seq == 0).sum() / max((y_seq == 1).sum(), 1)
            criterion = nn.BCEWithLogitsLoss(pos_weight=torch.tensor(pos_weight).to(device))
        else:
            criterion = nn.MSELoss()
    else:
        criterion = nn.MSELoss()

    n = len(X_seq)
    val_size = int(n * 0.15)
    X_train, X_val = X_seq[:-val_size], X_seq[-val_size:]
    y_train = y_seq[:-val_size] if use_target else X_seq[:-val_size]
    y_val = y_seq[-val_size:] if use_target else X_seq[-val_size:]

    train_loader = DataLoader(SequenceDataset(X_train, y_train), batch_size=batch_size, shuffle=False)
    val_loader = DataLoader(SequenceDataset(X_val, y_val), batch_size=batch_size, shuffle=False)

    best_val_loss = float('inf')
    patience_counter = 0
    best_model = None

    for epoch in range(epochs):
        model.train()
        train_loss = 0.0
        for Xb, yb in train_loader:
            Xb, yb = Xb.to(device), yb.to(device)
            optimizer.zero_grad()
            out = model(Xb)
            if use_target:
                loss = criterion(out, yb)
            else:
                loss = criterion(out.squeeze(), yb.mean(dim=(1, 2)))  # self-supervised proxy
            loss.backward()
            nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            optimizer.step()
            train_loss += loss.item() * len(Xb)
        train_loss /= len(X_train)

        model.eval()
        val_loss = 0.0
        with torch.no_grad():
            for Xb, yb in val_loader:
                Xb, yb = Xb.to(device), yb.to(device)
                out = model(Xb)
                if use_target:
                    loss = criterion(out, yb)
                else:
                    loss = criterion(out.squeeze(), yb.mean(dim=(1, 2)))
                val_loss += loss.item() * len(Xb)
        val_loss /= len(X_val)

        if verbose and (epoch + 1) % 10 == 0:
            print(f'    LSTM Epoch {epoch+1}/{epochs} train_loss={train_loss:.6f} val_loss={val_loss:.6f}')

        if val_loss < best_val_loss:
            best_val_loss = val_loss
            patience_counter = 0
            # Сохраняем через deepcopy, чтобы избежать проблем с PyTorch state_dict
            best_model = copy.deepcopy(model)
        else:
            patience_counter += 1
            if patience_counter >= patience:
                if verbose:
                    print(f'    LSTM Early stopping at epoch {epoch+1}')
                break

    if best_model is not None:
        return best_model
    return model


def extract_lstm_features(
    df,
    feature_cols: list[str],
    window: int = WINDOW,
    hidden_size: int = LATENT_DIM,
    num_layers: int = 2,
    retrain: bool = True,
    model: LSTMEncoder | None = None,
    device: str = DEVICE,
    verbose: bool = True,
):
    """
    Извлекает LSTM-latent признаки и добавляет их в DataFrame.
    Возвращает (df_with_features, LSTM_model).
    """
    available = [c for c in feature_cols if c in df.columns]
    data = df[available].values.astype(np.float32)
    data = np.nan_to_num(data, nan=0.0, posinf=0.0, neginf=0.0)

    # Нормализация
    mean = data.mean(axis=0, keepdims=True)
    std = data.std(axis=0, keepdims=True).clip(min=1e-8)
    data_norm = (data - mean) / std

    # Создание последовательностей
    X_seq, _ = create_sequences(data_norm, None, window)

    if retrain or model is None:
        if verbose:
            print(f'    LSTM: {len(X_seq)} seq, window={window}, features={data.shape[1]} → latent={hidden_size}')
        # Обучаем self-supervised: предсказание среднего следующего шага
        y_pretrain = data_norm[window:].mean(axis=1)
        model = train_encoder(
            X_seq, y_pretrain,
            n_features=data.shape[1],
            hidden_size=hidden_size,
            num_layers=num_layers,
            device=device,
            verbose=verbose,
        )

    # Извлекаем латентные признаки
    model.eval()
    loader = DataLoader(SequenceDataset(X_seq), batch_size=BATCH_SIZE, shuffle=False)
    all_feats = []
    with torch.no_grad():
        for Xb, _ in loader:
            feats = model.encode(Xb.to(device))
            all_feats.append(feats)
    latent = np.concatenate(all_feats)

    # Добавляем в df (первые window строк будут NaN)
    df_out = df.copy()
    for i in range(latent.shape[1]):
        col = f'lstm_feat_{i}'
        df_out[col] = np.nan
        df_out.iloc[window:, df_out.columns.get_loc(col)] = latent[:, i]

    return df_out, model, available


LSTM_FEATURE_COLS = [f'lstm_feat_{i}' for i in range(LATENT_DIM)]
