# -- Code Cell --
import copy
import numpy as np
import pandas as pd
from PIL import Image

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import Dataset, DataLoader, WeightedRandomSampler
from torchvision import models, transforms

from sklearn.model_selection import train_test_split
from sklearn.metrics import f1_score

# =========================
# CONFIG
# =========================
CSV_PATH = "train.csv"
TEST_CSV_PATH = "test.csv"
BATCH_SIZE = 32
EPOCHS = 10
IMG_SIZE = 224
SEED = 42

torch.manual_seed(SEED)
np.random.seed(SEED)

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# =========================
# DATA
# =========================
df = pd.read_csv(CSV_PATH)

train_df, val_df = train_test_split(
    df,
    test_size=0.2,
    random_state=SEED,
    stratify=df["label"]
)

train_df = train_df.reset_index(drop=True)
val_df = val_df.reset_index(drop=True)

print("Train size:", len(train_df))
print("Val size:", len(val_df))
print("Train label distribution:")
print(train_df["label"].value_counts())
print("Val label distribution:")
print(val_df["label"].value_counts())

train_transform = transforms.Compose([
    transforms.Resize((IMG_SIZE, IMG_SIZE)),
    transforms.RandomHorizontalFlip(p=0.5),
    transforms.RandomRotation(10),
    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),
    transforms.ToTensor(),
    transforms.Normalize([0.485, 0.456, 0.406],
                         [0.229, 0.224, 0.225])
])

val_transform = transforms.Compose([
    transforms.Resize((IMG_SIZE, IMG_SIZE)),
    transforms.ToTensor(),
    transforms.Normalize([0.485, 0.456, 0.406],
                         [0.229, 0.224, 0.225])
])

# =========================
# DATASET
# =========================
class SantaDataset(Dataset):
    def __init__(self, dataframe, transform=None, has_labels=True):
        self.df = dataframe.reset_index(drop=True)
        self.transform = transform
        self.has_labels = has_labels

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

    def __getitem__(self, idx):
        row = self.df.iloc[idx]
        img = Image.open(row["image_path"]).convert("RGB")

        if self.transform:
            img = self.transform(img)

        if self.has_labels:
            label = torch.tensor(int(row["label"]), dtype=torch.long)
            return img, label
        return img

train_ds = SantaDataset(train_df, transform=train_transform, has_labels=True)
val_ds = SantaDataset(val_df, transform=val_transform, has_labels=True)

# =========================
# OPTIONAL: balanced sampling
# =========================
class_counts = train_df["label"].value_counts().sort_index().values
class_weights = 1.0 / class_counts
sample_weights = train_df["label"].map(lambda x: class_weights[x]).values
sample_weights = torch.DoubleTensor(sample_weights)

sampler = WeightedRandomSampler(
    weights=sample_weights,
    num_samples=len(sample_weights),
    replacement=True
)

train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler, num_workers=0)
val_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=0)

# =========================
# MODEL
# =========================
weights = models.ResNet50_Weights.IMAGENET1K_V2
model = models.resnet50(weights=weights)
model.fc = nn.Linear(model.fc.in_features, 2)

# Înghețăm puțin la început
for param in model.parameters():
    param.requires_grad = False

# Dezghețăm mai mult decât înainte
for param in model.layer3.parameters():
    param.requires_grad = True
for param in model.layer4.parameters():
    param.requires_grad = True
for param in model.fc.parameters():
    param.requires_grad = True

model = model.to(device)

criterion = nn.CrossEntropyLoss()

optimizer = optim.AdamW([
    {"params": model.layer3.parameters(), "lr": 1e-5},
    {"params": model.layer4.parameters(), "lr": 3e-5},
    {"params": model.fc.parameters(), "lr": 1e-4}
], weight_decay=1e-4)

scheduler = optim.lr_scheduler.ReduceLROnPlateau(
    optimizer, mode="max", factor=0.5, patience=2
)

# =========================
# FUNCTIONS
# =========================
def train_one_epoch(model, loader, optimizer, criterion, device):
    model.train()
    running_loss = 0.0

    for images, labels in loader:
        images = images.to(device)
        labels = labels.to(device)

        optimizer.zero_grad()
        outputs = model(images)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()

        running_loss += loss.item()

    return running_loss / len(loader)

def evaluate_f1(model, loader, device):
    model.eval()
    all_probs = []
    all_labels = []

    with torch.no_grad():
        for images, labels in loader:
            images = images.to(device)
            outputs = model(images)

            probs = torch.softmax(outputs, dim=1)[:, 1]
            all_probs.extend(probs.cpu().numpy())
            all_labels.extend(labels.numpy())

    all_probs = np.array(all_probs)
    all_labels = np.array(all_labels)

    best_f1 = 0.0
    best_threshold = 0.5

    for threshold in np.arange(0.1, 0.91, 0.02):
        preds = (all_probs >= threshold).astype(int)
        f1 = f1_score(all_labels, preds)
        if f1 > best_f1:
            best_f1 = f1
            best_threshold = threshold

    return best_f1, best_threshold

# =========================
# TRAIN LOOP
# =========================
best_f1 = 0.0
best_threshold = 0.5
best_model_wts = copy.deepcopy(model.state_dict())

for epoch in range(EPOCHS):
    train_loss = train_one_epoch(model, train_loader, optimizer, criterion, device)
    val_f1, val_threshold = evaluate_f1(model, val_loader, device)

    scheduler.step(val_f1)

    print(f"Epoch {epoch+1}/{EPOCHS} | Train Loss: {train_loss:.4f} | Val F1: {val_f1:.4f} | Best Thr: {val_threshold:.2f}")

    if val_f1 > best_f1:
        best_f1 = val_f1
        best_threshold = val_threshold
        best_model_wts = copy.deepcopy(model.state_dict())

print(f"\nBest Val F1: {best_f1:.4f}")
print(f"Best Threshold: {best_threshold:.2f}")

model.load_state_dict(best_model_wts)

# =========================
# TEST PREDICTION
# =========================
test_df = pd.read_csv(TEST_CSV_PATH)
test_ds = SantaDataset(test_df, transform=val_transform, has_labels=False)
test_loader = DataLoader(test_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=0)

model.eval()
test_probs = []

with torch.no_grad():
    for images in test_loader:
        images = images.to(device)
        outputs = model(images)
        probs = torch.softmax(outputs, dim=1)[:, 1]
        test_probs.extend(probs.cpu().numpy())

test_probs = np.array(test_probs)
test_preds = (test_probs >= best_threshold).astype(int)

submission = pd.DataFrame({
    "image_path": test_df["image_path"],
    "label": test_preds
})

submission.to_csv("subs.csv", index=False)
print("Saved subs.csv")

# -- Code Cell --
