import numpy as np
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader


WINDOW_SIZE = 50
BATCH_SIZE = 256
EPOCHS = 100
PATIENCE = 8


class SequenceDataset(Dataset):
    def __init__(self, X: np.ndarray, y: np.ndarray):
        self.X = torch.FloatTensor(X)
        self.y = torch.FloatTensor(y).view(-1, 1)

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

    def __getitem__(self, idx):
        return self.X[idx], self.y[idx]


def create_sequences(
    features: np.ndarray,
    target: np.ndarray,
    window: int = WINDOW_SIZE,
) -> tuple[np.ndarray, np.ndarray]:
    n = len(features)
    X, y = [], []
    for i in range(window, n):
        X.append(features[i - window : i])
        y.append(target[i])
    return np.array(X, dtype=np.float32), np.array(y, dtype=np.float32)


def standardize(X: np.ndarray, mean: np.ndarray | None = None, std: np.ndarray | None = None):
    # Handle 1D arrays by adding a dummy dimension
    was_1d = False
    if X.ndim == 1:
        X = X[:, np.newaxis]
        was_1d = True

    if mean is None:
        mean = X.mean(axis=(0, 1), keepdims=True)
    if std is None:
        std = X.std(axis=(0, 1), keepdims=True).clip(min=1e-8)

    result = (X - mean) / std

    # Restore 1D shape if input was 1D
    if was_1d:
        result = result[:, 0]

    return result, mean, std


class LSTMOneOutput(nn.Module):
    def __init__(self, n_features: int, hidden_size: int = 32):
        super().__init__()
        self.lstm = nn.LSTM(
            input_size=n_features,
            hidden_size=hidden_size,
            num_layers=1,
            batch_first=True,
        )
        self.fc = nn.Linear(hidden_size, 1)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        lstm_out, _ = self.lstm(x)
        h = lstm_out[:, -1, :]
        return torch.sigmoid(self.fc(h))


def train_lstm_single(
    X_train: np.ndarray,
    y_train: np.ndarray,
    X_val: np.ndarray,
    y_val: np.ndarray,
    n_features: int,
    pos_weight: float = 2.0,
    lr: float = 5e-4,
    epochs: int = EPOCHS,
    batch_size: int = BATCH_SIZE,
    patience: int = PATIENCE,
    device: str = 'cpu',
    verbose: bool = True,
):
    model = LSTMOneOutput(n_features).to(device)
    optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-5)
    criterion = nn.BCEWithLogitsLoss(pos_weight=torch.tensor(pos_weight).to(device))

    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

    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()
            logits = model.lstm(Xb)[0][:, -1, :]
            logits = model.fc(logits)
            loss = criterion(logits, yb)
            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)
                logits = model.lstm(Xb)[0][:, -1, :]
                logits = model.fc(logits)
                loss = criterion(logits, yb)
                val_loss += loss.item() * len(Xb)
        val_loss /= len(X_val)

        if verbose:
            print(f"Epoch {epoch+1:3d}/{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
            best_state = model.state_dict().copy()
        else:
            patience_counter += 1
            if patience_counter >= patience:
                if verbose:
                    print(f"Early stopping at epoch {epoch+1}")
                break

    model.load_state_dict(best_state)
    return model


@torch.no_grad()
def predict_lstm(
    model: nn.Module,
    X: np.ndarray,
    batch_size: int = BATCH_SIZE,
    device: str = 'cpu',
) -> np.ndarray:
    model.eval()
    loader = DataLoader(SequenceDataset(X, np.zeros(len(X))), batch_size=batch_size, shuffle=False)
    all_probs = []
    for Xb, _ in loader:
        logits = model.lstm(Xb.to(device))[0][:, -1, :]
        logits = model.fc(logits)
        probs = torch.sigmoid(logits).cpu().numpy()
        all_probs.append(probs)
    return np.concatenate(all_probs).ravel()
