import copy
import numpy as np
import pandas as pd
import torch
import torch.nn as nn

from torch.utils.data import TensorDataset, DataLoader
from sklearn.metrics import confusion_matrix
from sklearn.utils.class_weight import compute_class_weight
from sklearn.model_selection import StratifiedKFold
from sklearn.base import clone
from scipy.fft import idctn
from scipy.ndimage import zoom
from lightgbm import LGBMClassifier


# ============================================================
# Paths
# ============================================================

train_path = r"C:\Users\Faezeh\OneDrive - Temple University\Semester 2\Intro Machine Learning\Project\Data\train.csv"
dev_path   = r"C:\Users\Faezeh\OneDrive - Temple University\Semester 2\Intro Machine Learning\Project\Data\dev.csv"
eval_path  = r"C:\Users\Faezeh\OneDrive - Temple University\Semester 2\Intro Machine Learning\Project\Data\eval.csv"


# ============================================================
# Labels
# ============================================================

valid_classes = [0, 2, 3, 5, 6, 8]
tissue_classes = [0, 2, 3, 5, 6]
background_class = 8

label_names = {
    0: "norm",
    2: "nneo",
    3: "infl",
    5: "dcis",
    6: "indc",
    8: "bckg"
}


# ============================================================
# Load data
# ============================================================

def load_data(path):
    df = pd.read_csv(path, sep=",", skipinitialspace=True)

    print(f"\nLoaded {path}")
    print("Raw shape:", df.shape)

    if df.shape[1] < 100:
        print("Detected collapsed columns. Splitting manually...")

        rows = []
        with open(path, "r") as f:
            next(f)
            for line in f:
                parts = line.strip().replace(",", " ").split()
                rows.append(parts)

        df = pd.DataFrame(rows).astype(float)

    print("Fixed shape:", df.shape)

    y = df.iloc[:, 0].values.astype(int)
    X = df.iloc[:, 1:].values.astype(np.float32)

    print("X shape:", X.shape)
    print("y shape:", y.shape)

    return X, y


X_train_full, y_train_full = load_data(train_path)
X_dev_full, y_dev_full = load_data(dev_path)
X_eval_full, y_eval_full = load_data(eval_path)


def filter_valid_classes(X, y):
    mask = np.isin(y, valid_classes)
    return X[mask], y[mask], mask


X_train, y_train, train_mask = filter_valid_classes(X_train_full, y_train_full)
X_dev, y_dev, dev_mask = filter_valid_classes(X_dev_full, y_dev_full)

print("\nAfter removing ignored classes:")
print("Train:", X_train.shape, y_train.shape)
print("Dev:  ", X_dev.shape, y_dev.shape)
print("Eval: ", X_eval_full.shape)


# ============================================================
# Scoring
# ============================================================

def dpath_score(y_true, y_pred, verbose=False):
    cm = confusion_matrix(y_true, y_pred, labels=list(range(9)))

    class_errors = {}

    for cls in tissue_classes:
        total = cm[cls].sum()
        correct = cm[cls, cls]
        error = 0.0 if total == 0 else 1.0 - correct / total
        class_errors[cls] = error

    avg_tissue_error = np.mean([class_errors[cls] for cls in tissue_classes])

    total_bg = cm[background_class].sum()
    correct_bg = cm[background_class, background_class]
    bg_error = 0.0 if total_bg == 0 else 1.0 - correct_bg / total_bg

    score = 0.90 * avg_tissue_error + 0.10 * bg_error

    if verbose:
        print("\nConfusion matrix:")
        print(cm)

        print("\nClass error rates:")
        for cls in tissue_classes:
            print(f"{cls} ({label_names[cls]}): {class_errors[cls] * 100:.4f}%")

        print(f"{background_class} ({label_names[background_class]}): {bg_error * 100:.4f}%")
        print(f"\nAverage tissue error: {avg_tissue_error * 100:.4f}%")
        print(f"Background error:     {bg_error * 100:.4f}%")
        print(f"DPATH SCORE:          {score * 100:.4f}%")

    return score


# ============================================================
# LightGBM sample weights
# ============================================================

def make_sample_weights(y):
    weights = np.ones_like(y, dtype=np.float32)

    counts = {cls: np.sum(y == cls) for cls in valid_classes}
    max_count = max(counts.values())

    for cls in valid_classes:
        w = max_count / counts[cls]

        if cls in tissue_classes:
            weights[y == cls] = w
        elif cls == background_class:
            weights[y == cls] = 0.5 * w

    weights[y == 5] *= 2.0
    weights[y == 6] *= 2.0

    return weights


# ============================================================
# Define LightGBM model
# ============================================================

lgbm_model = LGBMClassifier(
    objective="multiclass",
    num_class=9,
    learning_rate=0.1,
    n_estimators=300,
    num_leaves=63,
    max_depth=-1,
    random_state=42,
    n_jobs=-1,
    verbose=-1
)


# ============================================================
# GBM out-of-fold train predictions
# ============================================================

def get_oof_predictions_lgbm(model, X, y, n_splits=5):
    oof_pred = np.zeros_like(y)

    kf = StratifiedKFold(
        n_splits=n_splits,
        shuffle=True,
        random_state=42
    )

    for fold, (tr_idx, val_idx) in enumerate(kf.split(X, y), start=1):
        print(f"\nTraining GBM fold {fold}/{n_splits}")

        fold_model = clone(model)

        X_tr = X[tr_idx]
        y_tr = y[tr_idx]
        X_val = X[val_idx]

        fold_weights = make_sample_weights(y_tr)

        fold_model.fit(X_tr, y_tr, sample_weight=fold_weights)

        oof_pred[val_idx] = fold_model.predict(X_val)

    return oof_pred


print("\n================================")
print("LightGBM OUT-OF-FOLD TRAIN SCORE")
print("================================")

y_train_pred_lgbm_oof = get_oof_predictions_lgbm(
    lgbm_model,
    X_train,
    y_train,
    n_splits=5
)

lgbm_train_score = dpath_score(
    y_train,
    y_train_pred_lgbm_oof,
    verbose=True
)

print(f"\nLightGBM OOF TRAIN SCORE = {lgbm_train_score * 100:.4f}%")


# ============================================================
# Train final LightGBM on full training set
# ============================================================

print("\n================================")
print("Training final LightGBM on full train set")
print("================================")

sample_weights = make_sample_weights(y_train)
lgbm_model.fit(X_train, y_train, sample_weight=sample_weights)


# ============================================================
# LightGBM dev/eval predictions
# ============================================================

y_dev_pred_lgbm_full = lgbm_model.predict(X_dev_full)
y_eval_pred_lgbm_full = lgbm_model.predict(X_eval_full)

print("\nLightGBM DEV results:")
lgbm_dev_score = dpath_score(
    y_dev_full[dev_mask],
    y_dev_pred_lgbm_full[dev_mask],
    verbose=True
)

print(f"\nLightGBM TRAIN SCORE (OOF) = {lgbm_train_score * 100:.4f}%")
print(f"LightGBM DEV SCORE         = {lgbm_dev_score * 100:.4f}%")


# Create full train prediction file for GBM using OOF predictions
# For ignored rows, keep the original ignored label so row count stays correct.
y_train_pred_lgbm_full = np.copy(y_train_full)
y_train_pred_lgbm_full[train_mask] = y_train_pred_lgbm_oof

print("\nChecking LightGBM OOF train predictions:")
print("Identical to all train labels?",
      np.array_equal(y_train_pred_lgbm_full, y_train_full))
print("Identical to valid train labels?",
      np.array_equal(y_train_pred_lgbm_full[train_mask], y_train_full[train_mask]))


# ============================================================
# IDCT reconstruction for CNN
# ============================================================

def dct_to_reconstructed_images(X, full_size=256, output_size=64):
    N = X.shape[0]
    X_dct_small = X.reshape(N, 3, 32, 32)
    X_out = np.zeros((N, 3, output_size, output_size), dtype=np.float32)

    resize_factor = output_size / full_size

    for i in range(N):
        if i % 1000 == 0:
            print(f"Reconstructing image {i}/{N}")

        full_dct = np.zeros((3, full_size, full_size), dtype=np.float32)
        full_dct[:, :32, :32] = X_dct_small[i]

        img256 = idctn(full_dct, axes=(-2, -1), norm="ortho")

        img64 = zoom(
            img256,
            zoom=(1, resize_factor, resize_factor),
            order=1
        )

        X_out[i] = img64.astype(np.float32)

    return X_out


print("\n================================")
print("Reconstructing images for CNN")
print("================================")

X_train_img = dct_to_reconstructed_images(X_train)
X_dev_img = dct_to_reconstructed_images(X_dev)

X_train_full_img = dct_to_reconstructed_images(X_train_full)
X_dev_full_img = dct_to_reconstructed_images(X_dev_full)
X_eval_full_img = dct_to_reconstructed_images(X_eval_full)


def get_train_norm_stats(X):
    mean = X.mean(axis=(0, 2, 3), keepdims=True)
    std = X.std(axis=(0, 2, 3), keepdims=True) + 1e-6
    return mean, std


def apply_norm(X, mean, std):
    return ((X - mean) / std).astype(np.float32)


mean, std = get_train_norm_stats(X_train_img)

X_train_norm = apply_norm(X_train_img, mean, std)
X_dev_norm = apply_norm(X_dev_img, mean, std)

X_train_full_norm = apply_norm(X_train_full_img, mean, std)
X_dev_full_norm = apply_norm(X_dev_full_img, mean, std)
X_eval_full_norm = apply_norm(X_eval_full_img, mean, std)


# ============================================================
# CNN model
# ============================================================

class IDCTCNN64(nn.Module):
    def __init__(self):
        super().__init__()

        self.features = nn.Sequential(
            nn.Conv2d(3, 32, 3, padding=1),
            nn.BatchNorm2d(32),
            nn.ReLU(),
            nn.MaxPool2d(2),

            nn.Conv2d(32, 64, 3, padding=1),
            nn.BatchNorm2d(64),
            nn.ReLU(),
            nn.MaxPool2d(2),

            nn.Conv2d(64, 128, 3, padding=1),
            nn.BatchNorm2d(128),
            nn.ReLU(),
            nn.MaxPool2d(2),

            nn.Conv2d(128, 256, 3, padding=1),
            nn.BatchNorm2d(256),
            nn.ReLU(),
            nn.MaxPool2d(2),

            nn.Conv2d(256, 256, 3, padding=1),
            nn.BatchNorm2d(256),
            nn.ReLU(),
        )

        self.classifier = nn.Sequential(
            nn.AdaptiveAvgPool2d((1, 1)),
            nn.Flatten(),
            nn.Dropout(0.45),
            nn.Linear(256, 128),
            nn.ReLU(),
            nn.Dropout(0.30),
            nn.Linear(128, 9)
        )

    def forward(self, x):
        x = self.features(x)
        return self.classifier(x)


def make_class_weights(y):
    weights = np.zeros(9, dtype=np.float32)

    base_weights = compute_class_weight(
        class_weight="balanced",
        classes=np.array(valid_classes),
        y=y
    )

    for cls, w in zip(valid_classes, base_weights):
        weights[cls] = w

    weights[5] *= 2.0
    weights[6] *= 2.0
    weights[8] *= 0.6

    print("\nCNN class weights:")
    for cls in valid_classes:
        print(f"{cls} ({label_names[cls]}): {weights[cls]:.4f}")

    return torch.tensor(weights, dtype=torch.float32)


def predict_cnn(model, X, batch_size=256):
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    model.eval()
    preds = []

    with torch.no_grad():
        for start in range(0, len(X), batch_size):
            xb = torch.tensor(
                X[start:start + batch_size],
                dtype=torch.float32
            ).to(device)

            logits = model(xb)
            pred = torch.argmax(logits, dim=1).cpu().numpy()
            preds.extend(pred)

    return np.array(preds)


def train_cnn_with_early_stopping(
    X_tr,
    y_tr,
    X_val,
    y_val,
    epochs=70,
    batch_size=64,
    lr=1e-3,
    patience=12
):
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print("Using device:", device)

    train_ds = TensorDataset(
        torch.tensor(X_tr, dtype=torch.float32),
        torch.tensor(y_tr, dtype=torch.long)
    )

    train_loader = DataLoader(
        train_ds,
        batch_size=batch_size,
        shuffle=True
    )

    model = IDCTCNN64().to(device)

    class_weights = make_class_weights(y_tr).to(device)
    criterion = nn.CrossEntropyLoss(weight=class_weights)

    optimizer = torch.optim.AdamW(
        model.parameters(),
        lr=lr,
        weight_decay=1e-4
    )

    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
        optimizer,
        mode="min",
        factor=0.5,
        patience=4
    )

    best_score = np.inf
    best_epoch = 0
    best_state = None
    epochs_without_improvement = 0

    for epoch in range(1, epochs + 1):
        model.train()
        total_loss = 0.0

        for xb, yb in train_loader:
            xb = xb.to(device)
            yb = yb.to(device)

            optimizer.zero_grad()
            logits = model(xb)
            loss = criterion(logits, yb)
            loss.backward()
            optimizer.step()

            total_loss += loss.item()

        y_val_pred = predict_cnn(model, X_val)
        val_score = dpath_score(y_val, y_val_pred)

        scheduler.step(val_score)

        print(
            f"Epoch {epoch}/{epochs}, "
            f"loss={total_loss / len(train_loader):.4f}, "
            f"dev score={val_score * 100:.4f}%"
        )

        if val_score < best_score:
            best_score = val_score
            best_epoch = epoch
            best_state = copy.deepcopy(model.state_dict())
            epochs_without_improvement = 0
            print("  New best CNN saved.")
        else:
            epochs_without_improvement += 1

        if epochs_without_improvement >= patience:
            print(f"Early stopping at epoch {epoch}.")
            break

    model.load_state_dict(best_state)

    print(f"\nBest CNN epoch: {best_epoch}")
    print(f"Best CNN dev score: {best_score * 100:.4f}%")

    return model, best_score


# ============================================================
# Train CNN
# ============================================================

print("\n================================")
print("Training CNN")
print("================================")

cnn_model, best_cnn_score = train_cnn_with_early_stopping(
    X_train_norm,
    y_train,
    X_dev_norm,
    y_dev,
    epochs=70,
    batch_size=64,
    lr=1e-3,
    patience=12
)


# ============================================================
# CNN predictions and scores
# ============================================================

print("\n================================")
print("CNN Scores")
print("================================")

y_train_pred_cnn_full = predict_cnn(cnn_model, X_train_full_norm)
y_dev_pred_cnn_full = predict_cnn(cnn_model, X_dev_full_norm)
y_eval_pred_cnn_full = predict_cnn(cnn_model, X_eval_full_norm)

print("\nCNN TRAIN results:")
cnn_train_score = dpath_score(
    y_train_full[train_mask],
    y_train_pred_cnn_full[train_mask],
    verbose=True
)

print("\nCNN DEV results:")
cnn_dev_score = dpath_score(
    y_dev_full[dev_mask],
    y_dev_pred_cnn_full[dev_mask],
    verbose=True
)

print(f"\nCNN TRAIN SCORE = {cnn_train_score * 100:.4f}%")
print(f"CNN DEV SCORE   = {cnn_dev_score * 100:.4f}%")

print("\nChecking CNN train predictions:")
print("Identical to all train labels?",
      np.array_equal(y_train_pred_cnn_full, y_train_full))
print("Identical to valid train labels?",
      np.array_equal(y_train_pred_cnn_full[train_mask], y_train_full[train_mask]))


# ============================================================
# Final comparison
# ============================================================

print("\n================================")
print("FINAL MODEL COMPARISON")
print("================================")

print(f"LightGBM Train Score (OOF): {lgbm_train_score * 100:.4f}%")
print(f"LightGBM Dev Score:         {lgbm_dev_score * 100:.4f}%")
print(f"CNN Train Score:            {cnn_train_score * 100:.4f}%")
print(f"CNN Dev Score:              {cnn_dev_score * 100:.4f}%")


# ============================================================
# Save corrected prediction files
# ============================================================

pd.DataFrame({"label": y_train_pred_lgbm_full}).to_csv(
    "train_predictions_LightGBM.csv",
    index=False
)

pd.DataFrame({"label": y_dev_pred_lgbm_full}).to_csv(
    "dev_predictions_LightGBM.csv",
    index=False
)

pd.DataFrame({"label": y_eval_pred_lgbm_full}).to_csv(
    "eval_predictions_LightGBM.csv",
    index=False
)

pd.DataFrame({"label": y_train_pred_cnn_full}).to_csv(
    "train_predictions_CNN.csv",
    index=False
)

pd.DataFrame({"label": y_dev_pred_cnn_full}).to_csv(
    "dev_predictions_CNN.csv",
    index=False
)

pd.DataFrame({"label": y_eval_pred_cnn_full}).to_csv(
    "eval_predictions_CNN.csv",
    index=False
)

print("\nSaved corrected prediction files:")
print("train_predictions_LightGBM.csv")
print("dev_predictions_LightGBM.csv")
print("eval_predictions_LightGBM.csv")
print("train_predictions_CNN.csv")
print("dev_predictions_CNN.csv")
print("eval_predictions_CNN.csv")