#Rolensky Louis CNN
import numpy as np
import pandas as pd
from pathlib import Path

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import Dataset, DataLoader

from sklearn.impute import SimpleImputer
from sklearn.preprocessing import StandardScaler
from sklearn.model_selection import train_test_split
from sklearn.utils.class_weight import compute_class_weight



# Paths


TRAIN_PATH = "set_17/data/train.csv"
DEV_PATH = "set_17/data/dev.csv"
EVAL_PATH = "set_17/data/eval.csv"

OUTPUT_DIR = Path("final_hyp")
OUTPUT_DIR.mkdir(exist_ok=True)



# Settings
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"

BATCH_SIZE = 64
MAX_EPOCHS = 60
PATIENCE = 8
LR = 5e-4
NUM_CLASSES = 9



# Load data
def load_data(path):
    df = pd.read_csv(path, comment="#", header=None)

    y = df.iloc[:, 0].astype(int).values
    X = df.iloc[:, 1:].astype(float).values

    return X, y


def save_hyp(predictions, path):
    with open(path, "w") as f:
        f.write("label\n")
        for p in predictions:
            f.write(f"{int(p)}\n")



# Dataset
class DCTDataset(Dataset):
    def __init__(self, X, y):
        self.X = torch.tensor(X, dtype=torch.float32)
        self.y = torch.tensor(y, dtype=torch.long)

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

    def __getitem__(self, index):
        x = self.X[index]

        # No IDCT here.
        # Just reshape raw DCT vector:
        # 3072 = 3 * 32 * 32
        x = x.reshape(3, 32, 32)

        return x, self.y[index]



# CNN model
class RawDCTCNN(nn.Module):
    def __init__(self, num_classes=9):
        super().__init__()

        self.features = nn.Sequential(
            nn.Conv2d(3, 32, kernel_size=3, padding=1),
            nn.BatchNorm2d(32),
            nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Dropout2d(0.15),

            nn.Conv2d(32, 64, kernel_size=3, padding=1),
            nn.BatchNorm2d(64),
            nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Dropout2d(0.20),

            nn.Conv2d(64, 128, kernel_size=3, padding=1),
            nn.BatchNorm2d(128),
            nn.ReLU(),
            nn.MaxPool2d(2),
            nn.Dropout2d(0.25),
        )

        self.classifier = nn.Sequential(
            nn.Flatten(),
            nn.Linear(128 * 4 * 4, 128),
            nn.ReLU(),
            nn.Dropout(0.40),
            nn.Linear(128, num_classes)
        )

    def forward(self, x):
        x = self.features(x)
        x = self.classifier(x)
        return x



# Training
def train_model(model, train_loader, val_loader, loss_fn):
    model.to(DEVICE)

    optimizer = optim.AdamW(
        model.parameters(),
        lr=LR,
        weight_decay=1e-3
    )

    scheduler = optim.lr_scheduler.ReduceLROnPlateau(
        optimizer,
        mode="min",
        factor=0.5,
        patience=3
    )

    best_val_loss = float("inf")
    best_state = None
    patience_counter = 0

    for epoch in range(1, MAX_EPOCHS + 1):
        model.train()
        train_loss = 0.0

        for X_batch, y_batch in train_loader:
            X_batch = X_batch.to(DEVICE)
            y_batch = y_batch.to(DEVICE)

            optimizer.zero_grad()

            logits = model(X_batch)
            loss = loss_fn(logits, y_batch)

            loss.backward()
            optimizer.step()

            train_loss += loss.item()

        train_loss /= len(train_loader)

        model.eval()
        val_loss = 0.0
        correct = 0
        total = 0

        with torch.no_grad():
            for X_batch, y_batch in val_loader:
                X_batch = X_batch.to(DEVICE)
                y_batch = y_batch.to(DEVICE)

                logits = model(X_batch)
                loss = loss_fn(logits, y_batch)

                val_loss += loss.item()

                preds = torch.argmax(logits, dim=1)
                correct += (preds == y_batch).sum().item()
                total += len(y_batch)

        val_loss /= len(val_loader)
        val_acc = correct / total

        scheduler.step(val_loss)

        print(
            f"Epoch {epoch:02d} | "
            f"train loss: {train_loss:.4f} | "
            f"val loss: {val_loss:.4f} | "
            f"val acc: {val_acc:.4f}"
        )

        if val_loss < best_val_loss:
            best_val_loss = val_loss
            best_state = model.state_dict()
            patience_counter = 0
        else:
            patience_counter += 1

        if patience_counter >= PATIENCE:
            print("Early stopping.")
            break

    model.load_state_dict(best_state)
    return model



# Prediction
def predict_model(model, X):
    dummy_y = np.zeros(len(X), dtype=int)

    dataset = DCTDataset(X, dummy_y)

    loader = DataLoader(
        dataset,
        batch_size=BATCH_SIZE,
        shuffle=False
    )

    model.eval()
    predictions = []

    with torch.no_grad():
        for X_batch, _ in loader:
            X_batch = X_batch.to(DEVICE)

            logits = model(X_batch)
            preds = torch.argmax(logits, dim=1)

            predictions.extend(preds.cpu().numpy())

    return np.array(predictions)



# Main
if __name__ == "__main__":

    print("=" * 60)
    print("Raw DCT CNN")
    print("No IDCT: 3072 DCT features -> 3 x 32 x 32 -> CNN")
    print("=" * 60)

    print("Using device:", DEVICE)

    # Load data
    X_train_raw, y_train = load_data(TRAIN_PATH)
    X_dev_raw, y_dev = load_data(DEV_PATH)
    X_eval_raw, y_eval_fake = load_data(EVAL_PATH)

    print("\nRaw shapes:")
    print("Train:", X_train_raw.shape)
    print("Dev:", X_dev_raw.shape)
    print("Eval:", X_eval_raw.shape)

    print("\nTrain label counts:")
    print(pd.Series(y_train).value_counts().sort_index())

    # Impute missing values
    imputer = SimpleImputer(strategy="median")

    X_train_raw = imputer.fit_transform(X_train_raw)
    X_dev_raw = imputer.transform(X_dev_raw)
    X_eval_raw = imputer.transform(X_eval_raw)

    # Scale raw DCT features
    scaler = StandardScaler()

    X_train_raw = scaler.fit_transform(X_train_raw)
    X_dev_raw = scaler.transform(X_dev_raw)
    X_eval_raw = scaler.transform(X_eval_raw)

    # Internal train/validation split
    X_train, X_val, y_train_split, y_val = train_test_split(
        X_train_raw,
        y_train,
        test_size=0.15,
        random_state=42,
        stratify=y_train
    )

    # Class weights
    weights = compute_class_weight(
        class_weight="balanced",
        classes=np.arange(NUM_CLASSES),
        y=y_train_split
    )

    weights = torch.tensor(weights, dtype=torch.float32)

    # Mild boost for difficult scored classes
    weights[3] *= 1.5
    weights[5] *= 2.5
    weights[6] *= 1.5

    weights = weights.to(DEVICE)

    print("\nClass weights:")
    for i, w in enumerate(weights.cpu().numpy()):
        print(f"class {i}: {w:.4f}")

    # DataLoaders
    train_dataset = DCTDataset(X_train, y_train_split)
    val_dataset = DCTDataset(X_val, y_val)

    train_loader = DataLoader(
        train_dataset,
        batch_size=BATCH_SIZE,
        shuffle=True
    )

    val_loader = DataLoader(
        val_dataset,
        batch_size=BATCH_SIZE,
        shuffle=False
    )

    # Train model
    model = RawDCTCNN(num_classes=NUM_CLASSES)

    loss_fn = nn.CrossEntropyLoss(weight=weights)

    print("\nTraining Raw DCT CNN...")
    model = train_model(model, train_loader, val_loader, loss_fn)

    # Predict
    print("\nPredicting train/dev/eval...")

    train_pred = predict_model(model, X_train_raw)
    dev_pred = predict_model(model, X_dev_raw)
    eval_pred = predict_model(model, X_eval_raw)

    # Save hyp files
    save_hyp(train_pred, OUTPUT_DIR / "hyp_raw_dct_cnn_train.csv")
    save_hyp(dev_pred, OUTPUT_DIR / "hyp_raw_dct_cnn_dev.csv")
    save_hyp(eval_pred, OUTPUT_DIR / "hyp_raw_dct_cnn_eval.csv")

    print("\nFiles created in final_hyp:")
    print(OUTPUT_DIR / "hyp_raw_dct_cnn_train.csv")
    print(OUTPUT_DIR / "hyp_raw_dct_cnn_dev.csv")
    print(OUTPUT_DIR / "hyp_raw_dct_cnn_eval.csv")

    print("\nScore with:")
    print("python score.py set_17/data/train.csv final_hyp/hyp_raw_dct_cnn_train.csv")
    print("python score.py set_17/data/dev.csv final_hyp/hyp_raw_dct_cnn_dev.csv")