# -- 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
import pandas as pd
from tqdm import tqdm

# -- Code Cell --
import pandas as pd
train = pd.read_csv('train.csv')
train.head()

# -- Code Cell --
train['label'].unique().tolist()

# -- Code Cell --
map_lb= {
    'violin': 0, 'moon': 1, 'book': 2, 'tree': 3, 'flower': 4,
    'elephant': 5, 'fish': 6, 'star': 7, 'pencil': 8, 'sun': 9,
    'apple': 10, 'boat': 11, 'umbrella': 12, 'lion': 13, 'banana': 14,
    'dog': 15, 'kite': 16, 'cat': 17, 'house': 18, 'cake': 19
}

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

# -- Code Cell --
train_ds = TrainDataset('./train.csv', transform=transform)
train_loader = DataLoader(train_ds, 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

# -- Code Cell --
resnet.eval()
for epoch in range(5):
    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()
        # print("t")
    print(f"epoch {epoch+1}, loss: {loss:.4f}")

# -- Code Cell --
class TestDataset(Dataset):
    def __init__(self, csv_path, transform):
        self.df = pd.read_csv(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.csv', transform=transform)
test_loader = DataLoader(test_ds, batch_size=32, shuffle=False, num_workers=0)

# -- Code Cell --
resnet.eval()
pred = []
with torch.no_grad():                         
    for imgs in test_loader:
        imgs = imgs.to(device)
        output = resnet(imgs)
        pred_per_file = output.argmax(dim=1)    #TREBUIE LOGITS
        pred.append(pred_per_file.cpu())        
preds = torch.cat(pred)

# -- Code Cell --
preds

# -- Code Cell --
map_lb

# -- Code Cell --
reverse_map = {}
for key,number in map_lb.items():
    reverse_map[number] = key
    print(f"key: {key}, number: {number}")

# -- Code Cell --
rev_pred = [reverse_map[p.item()] for p in preds]

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

# -- Code Cell --
