# -- Code Cell --
import pandas as pd
train_df = pd.read_json("train.json")
test_df= pd.read_json('test.json')

# -- Code Cell --
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import Dataset, DataLoader
from torchvision import models, transforms
from PIL import Image

# -- Code Cell --
train_df['cartoon_class'].unique()

# -- Code Cell --
map ={"Pokemon":1,"Snow White":2, "Tarzan":3,"Winnie the Pooh":4}
rev_map = {1: 'Pokemon',2:"Snow White",3:"Tarzan",4:'Winnie the Pooh'}

# -- Code Cell --


# -- Code Cell --
class TrainDataset(Dataset):
    def __init__(self,csv_path, transform):
        self.df = pd.read_json(csv_path)
        self.transform = transform
    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)
        label = map[row['cartoon_class']]
        return img, label

# -- Code Cell --
transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),
    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])

# -- Markdown Cell --
# # AICI INCEPE 

# -- Code Cell --
train_ds = TrainDataset("./train.json",transform = transform)

# -- Code Cell --
from torch.utils.data import random_split

# -- Code Cell --
train_size = int(0.8 * len(train_ds))
val_size = len(train_ds) - train_size
train_subset, val_subset = random_split(train_ds,[train_size, val_size])

# -- Code Cell --


# -- Code Cell --
train_loader = DataLoader(train_subset, batch_size=32, shuffle=True, num_workers=0)
val_loader = DataLoader(val_subset, batch_size=32, shuffle=False, num_workers=0)

# -- Code Cell --
device = torch.device('cuda')

# -- Code Cell --
weights = models.ResNet50_Weights.IMAGENET1K_V1
resnet = models.resnet50(weights=weights).eval()
resnet.fc = nn.Linear(2048,20)
resnet = resnet.to(device)
optimizer = optim.Adam(resnet.fc.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()

# -- Code Cell --
for param in resnet.parameters():
    param.requires_grad = False

for param in resnet.fc.parameters():
    param.requires_grad = True

# -- Markdown Cell --
# # ACCURACY

# -- Code Cell --
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=15, gamma=0.5)

# -- Code Cell --
for epoch in range(5):
    resnet.train()
    train_loss = 0.0
    correct = 0
    total = 0
    for imgs, labels in train_loader:
        imgs = imgs.to(device)
        labels = labels.to(device)

        optimizer.zero_grad()
        outputs = resnet(imgs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()

        train_loss += loss.item()

        preds = outputs.argmax(dim=1)
        correct += (preds == labels).sum().item()
        total += labels.size(0)
    train_acc = correct / total
    resnet.eval()
    val_loss = 0.0
    correct = 0
    total = 0
    with torch.no_grad():
        for imgs, labels in val_loader:
            imgs = imgs.to(device)
            labels = labels.to(device)

            outputs = resnet(imgs)
            loss = criterion(outputs, labels)

            val_loss += loss.item()

            preds = outputs.argmax(dim=1)
            correct += (preds == labels).sum().item()
            total += labels.size(0)

    val_acc = correct / total
    
    scheduler.step(val_acc)
    
    print(f"Epoch {epoch+1} | "f"train_loss: {train_loss/len(train_loader):.4f} | "f"train_acc: {train_acc:.4f} | " f"val_loss: {val_loss/len(val_loader):.4f} | "f"val_acc: {val_acc:.4f}")

# -- Markdown Cell --
# # F1

# -- Code Cell --
from sklearn.metrics import f1_score

# -- Code Cell --
for epoch in range(5):
    resnet.train()
    train_loss = 0.0

    for imgs, labels in train_loader:
        imgs = imgs.to(device)
        labels = labels.to(device)

        optimizer.zero_grad()
        outputs = resnet(imgs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()

        train_loss += loss.item()
    resnet.eval()
    val_loss = 0.0
    all_preds = []
    all_labels = []

    with torch.no_grad():
        for imgs, labels in val_loader:
            imgs = imgs.to(device)
            labels = labels.to(device)

            outputs = resnet(imgs)
            loss = criterion(outputs, labels)
            val_loss += loss.item()

            preds = outputs.argmax(dim=1)

            all_preds.append(preds.cpu())
            all_labels.append(labels.cpu())
    all_preds = torch.cat(all_preds)
    all_labels = torch.cat(all_labels)
    f1 = f1_score(all_labels.numpy(), all_preds.numpy(), average="macro")
    
    scheduler.step(f1)

    print(f"Epoch {epoch+1} | "f"train_loss: {train_loss/len(train_loader):.4f} | "f"val_loss: {val_loss/len(val_loader):.4f} | "f"val_f1_macro: {f1:.4f}")

# -- Code Cell --
class TestDataset(Dataset):
    def __init__(self,csv_path, transform):
        self.df = pd.read_json(csv_path)
        self.transform = transform
    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)
        return img

# -- Code Cell --
test_ds =TestDataset("./test.json",transform = transform)
test = DataLoader(test_ds, batch_size=32, shuffle=False, num_workers=0)

# -- Code Cell --
pred = []
for img in test:
    img = img.to(device)
    with torch.no_grad():
        pred_per_image = resnet(img)
        _, pred_per_image = torch.max(pred_per_image, 1)
        pred.append(pred_per_image.cpu())
        preds = torch.cat(pred) 
    

# -- Code Cell --
preds

# -- Code Cell --
preds = preds.numpy()

# -- Code Cell --
preds_rev = [rev_map[elem] for elem in preds]
preds_rev

# -- Code Cell --
rows = []
for idx, row in test_df.iterrows():
    rows.append({"image_path":row['image_path'],"cartoon_class":preds_rev[idx]})
pd.DataFrame(rows).to_csv("subs.csv",index=False)

# -- Code Cell --
