# -- Code Cell --
import matplotlib.pyplot as plt 
import cv2 as cv
import numpy as np

# -- Code Cell --
img = cv.imread("./data/train/263.jpg")
img_rgb = cv.cvtColor(img, cv.COLOR_BGR2RGB)

h, w = img_rgb.shape[:2]
margin = 5
rect = (margin, margin, w - 2*margin, h - 2*margin)
mask = np.zeros((h, w), np.uint8)
bgdModel = np.zeros((1,65), np.float64)
fgdModel = np.zeros((1,65), np.float64)
cv.grabCut(img_rgb, mask, rect, bgdModel, fgdModel, 10, cv.GC_INIT_WITH_RECT)
mask2 = np.where((mask==2)|(mask==0), 0, 1).astype('uint8')
result = np.full_like(img_rgb, 128)
result[mask2 == 1] = img_rgb[mask2 == 1]

plt.imshow(result), plt.show()

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

# -- Code Cell --
train = pd.read_csv("./data/train.csv")
train.head()

# -- Code Cell --
class FullDataset(Dataset):
    def __init__(self,path,transform):
        super().__init__()
        self.df = pd.read_csv(path)
        self.transform = transform
    def __len__(self):
        return len(self.df)
    def __getitem__(self, index):
        row = self.df.iloc[index]
        img = cv.imread(f"./data/train/{index}.jpg")
        label = row['y']
        img_crop = img[20:204, 20:204]
        imgmod = Image.fromarray(img_crop)
        if self.transform:
            imgmod = self.transform(imgmod)
        return imgmod,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 --
fullds = FullDataset("./data/train.csv",transform)

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

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

# -- Code Cell --
train["y"].unique()

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

# -- Code Cell --
class modelmeu(nn.Module):
    def __init__(self):
        super().__init__()
        self.resnet = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)
        self.resnet.fc = nn.Sequential(
            nn.Linear(2048, 512),
            nn.ReLU(),
            nn.Dropout(0.3),
            nn.Linear(512, 2),
        )
    def forward(self,x):
        out = self.resnet(x)
        return out

# -- Code Cell --
model = modelmeu()
for param in model.resnet.parameters():
    param.requires_grad = False
for param in model.resnet.fc.parameters():
    param.requires_grad = True
model = model.to(device)
optimizer = optim.Adam(model.resnet.fc.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()
scheduler = optim.lr_scheduler.StepLR(optimizer,step_size=30, gamma=0.2)

# -- Code Cell --
decoy = 1

# -- Code Cell --
from tqdm import tqdm
for epoch in range(15):
    total_loss = 0
    for img, label in tqdm(train_loader):
        img=img.to(device)
        label=label.to(device)
        optimizer.zero_grad()
        output = model(img)
        loss = criterion(output,label)
        loss.backward()
        optimizer.step()
        total_loss+=loss.item()
    scheduler.step()
    if total_loss/len(train_loader)<decoy:
        decoy = total_loss/len(train_loader)
        torch.save(model.state_dict(), "model.pth")
    print(f"epoch{epoch+1} -- loss{total_loss/len(train_loader):04f}")

# -- Code Cell --
model = modelmeu()
model.load_state_dict(torch.load("model.pth"))
model = model.to(device)
model.eval()

# -- Code Cell --
class TestDataset(Dataset):
    def __init__(self,path,transform):
        super().__init__()
        self.df = pd.read_csv(path)
        self.transform = transform
    def __len__(self):
        return len(self.df)
    def __getitem__(self, index):
        img = cv.imread(f"./data/test/{index}.jpg")
        img_crop = img[20:204, 20:204]
        imgmod = Image.fromarray(img_crop)
        if self.transform:
            imgmod = self.transform(imgmod)
        return imgmod

# -- Code Cell --
test_ds = TestDataset("./data/test.csv",transform)
test_loader = DataLoader(test_ds, batch_size=64, num_workers=0, shuffle=False)

# -- Code Cell --
preds=[]
with torch.no_grad():
    for img in test_loader:
        img = img.to(device)
        output = model(img)
        pred = output.argmax(dim=1)
        preds.append(pred.cpu())
    preds_final = torch.cat(preds)
preds_final = preds_final.numpy()
preds_final

# -- Code Cell --
test = pd.read_csv('./data/test.csv')
test.head()

# -- Code Cell --
sub = pd.DataFrame({
    'subtaskID':1,
    'datapointID':test['datapointID'],
    'answer':preds_final
}).to_csv('subs.csv',index=False)

# -- Code Cell --
