# -- Code Cell --
import os, random
import numpy as np
import pandas as pd

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

import torchvision
from torchvision import transforms
from PIL import Image

from sklearn.model_selection import train_test_split
from tqdm.auto import tqdm


# -- Code Cell --
SEED = 42
random.seed(SEED)
np.random.seed(SEED)
torch.manual_seed(SEED)

DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
print("Device:", DEVICE)

DATA_DIR = "."          # folderul unde sunt csv-urile
IMAGES_DIR = "images"   # folderul cu imaginile

IMG_SIZE = 224
BATCH_SIZE = 32
EPOCHS = 8
LR = 3e-4
NUM_WORKERS = 2

# -- Code Cell --
train_df = pd.read_csv(os.path.join(DATA_DIR, "train.csv"))
test_df  = pd.read_csv(os.path.join(DATA_DIR, "test.csv"))

print(train_df.head())
print(train_df["label"].value_counts())

classes = sorted(train_df["label"].unique())
class_to_idx = {c:i for i,c in enumerate(classes)}
idx_to_class = {i:c for c,i in class_to_idx.items()}

print("Classes:", classes)

# -- Code Cell --
train_df, val_df = train_test_split(
    train_df,
    test_size=0.2,
    random_state=SEED,
    stratify=train_df["label"]
)

print(len(train_df), len(val_df))


# -- Code Cell --
train_tfms = transforms.Compose([
    transforms.Resize((IMG_SIZE, IMG_SIZE)),
    transforms.RandomHorizontalFlip(p=0.5),
    transforms.RandomRotation(degrees=10),
    transforms.ColorJitter(brightness=0.15, contrast=0.15, saturation=0.1),
    transforms.ToTensor(),
    transforms.Normalize(mean=(0.485,0.456,0.406), std=(0.229,0.224,0.225))
])

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


# -- Code Cell --
class ChessDataset(Dataset):
    def __init__(self, df, images_dir, transform=None, is_test=False):
        self.df = df.reset_index(drop=True)
        self.images_dir = images_dir
        self.transform = transform
        self.is_test = is_test

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

    def __getitem__(self, idx):
        row = self.df.loc[idx]
        img_path = os.path.join(self.images_dir, row["image_path"])
        img = Image.open(img_path).convert("RGB")

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

        if self.is_test:
            return img, row["id"]
        else:
            y = class_to_idx[row["label"]]
            return img, torch.tensor(y, dtype=torch.long)


# -- Code Cell --
train_ds = ChessDataset(train_df, IMAGES_DIR, transform=train_tfms, is_test=False)
val_ds   = ChessDataset(val_df,   IMAGES_DIR, transform=val_tfms,   is_test=False)
test_ds  = ChessDataset(test_df,  IMAGES_DIR, transform=val_tfms,   is_test=True)

NUM_WORKERS = 0

train_loader = DataLoader(
    train_ds,
    batch_size=BATCH_SIZE,
    shuffle=True,
    num_workers=NUM_WORKERS
)

val_loader = DataLoader(
    val_ds,
    batch_size=BATCH_SIZE,
    shuffle=False,
    num_workers=NUM_WORKERS
)

test_loader = DataLoader(
    test_ds,
    batch_size=BATCH_SIZE,
    shuffle=False,
    num_workers=NUM_WORKERS
)

x, y = next(iter(train_loader))
print(x.shape, y.shape)


# -- Code Cell --
NUM_CLASSES = len(classes)

weights = torchvision.models.ResNet18_Weights.DEFAULT
model = torchvision.models.resnet18(weights=weights)

model.fc = nn.Linear(model.fc.in_features, NUM_CLASSES)
model = model.to(DEVICE)

criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-4)

# -- Code Cell --
def acc(logits, y):
    return (logits.argmax(1) == y).float().mean().item()

@torch.no_grad()
def evaluate(model, loader):
    model.eval()
    loss_sum, acc_sum, n = 0.0, 0.0, 0
    for x, y in loader:
        x, y = x.to(DEVICE), y.to(DEVICE)
        logits = model(x)
        loss = criterion(logits, y)
        loss_sum += loss.item()
        acc_sum += acc(logits, y)
        n += 1
    return loss_sum / n, acc_sum / n

def train_one_epoch(model, loader):
    model.train()
    loss_sum, acc_sum, n = 0.0, 0.0, 0
    for x, y in tqdm(loader, leave=False):
        x, y = x.to(DEVICE), y.to(DEVICE)
        optimizer.zero_grad()
        logits = model(x)
        loss = criterion(logits, y)
        loss.backward()
        optimizer.step()
        loss_sum += loss.item()
        acc_sum += acc(logits, y)
        n += 1
    return loss_sum / n, acc_sum / n

# -- Code Cell --
best_val_acc = 0.0
best_state = None

for epoch in range(1, EPOCHS + 1):
    tr_loss, tr_acc = train_one_epoch(model, train_loader)
    va_loss, va_acc = evaluate(model, val_loader)

    print(f"Epoch {epoch}/{EPOCHS} | "
          f"train acc {tr_acc:.4f} | val acc {va_acc:.4f}")

    if va_acc > best_val_acc:
        best_val_acc = va_acc
        best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}

print("Best val acc:", best_val_acc)

# -- Code Cell --
if best_state is not None:
    model.load_state_dict(best_state)
model = model.to(DEVICE)

# -- Code Cell --
import pandas as pd

model.eval()
all_ids = []
all_preds = []

with torch.no_grad():
    for x, ids in test_loader:
        x = x.to(DEVICE)
        logits = model(x)
        preds = logits.argmax(1).cpu().numpy()

        all_ids.extend(list(ids))
        all_preds.extend([idx_to_class[i] for i in preds])

submission = pd.DataFrame({
    "id": all_ids,
    "label": all_preds
})

submission.to_csv("submission.csv", index=False)
submission.head()