# -- 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
from torchvision.transforms import functional as TF
import cv2 as cv
import numpy as np

# -- Code Cell --
import pandas as pd
split_pictures = pd.read_csv('./validation_data/validation/validation_data.csv')

# -- Code Cell --
import os

# -- Code Cell --
os.listdir("./")

# -- Code Cell --
int(img_combined[1,1,0])

# -- Code Cell --
np.ndarray([int(img_combined[1,1,0]),int(img_combined[1,1,1]), int(img_combined[1,1,2])])

# -- Code Cell --
tile1 = cv.imread(f"./validation_data/validation/validation_tiles/105878.jpg")
tile2 = cv.imread(f"./validation_data/validation/validation_tiles/144060.jpg")
img_combined = np.zeros((128, 256, 3), dtype=np.uint8)
img_combined[0:128, 0:128] = tile1
img_combined[0:128, 128:256] = tile2
img_rgb = cv.cvtColor(img_combined, cv.COLOR_BGR2RGB)
plt.imshow(img_rgb)
plt.show()

# -- 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 --
img = cv.imread("./training_data/train_img_01.jpg")
import matplotlib.pyplot as plt
clipped_img = []
for i in range(3):
    for j in range(3):
        roi = img[128*i:128*(i+1), 128*j:128*(j+1)]
        clipped_img.append(roi)
for i in range(len(clipped_img)):
    plt.imshow(clipped_img[i])
    plt.show()

# -- Code Cell --
dt = {}
clipped_imgs = np.zeros((60, 16, 128, 128, 3), dtype=np.uint8)

for i in range(len(os.listdir("./training_data/"))):
    clipped_img = []
    img = cv.imread(f"./training_data/train_img_{i+1:02d}.jpg")
    count = 0
    for x in range(4):
        for y in range(4):
            roi = img[128*y:128*(y+1), 128*x:128*(x+1)]
            clipped_imgs[i][count] = roi
            count +=1

# for i in range(59):
#     dt[f"0_{i}"] = {"img1": clipped_imgs[i][0], "img2": clipped_imgs[i+1][1]}
 
#     dt[f"1_{i}"] = {"img1": clipped_imgs[i][0], "img2": clipped_imgs[i][1]}

#     dt[f"2_{i}"] = {"img1": clipped_imgs[i][0], "img2": clipped_imgs[i][4]

#     dt[f"3_{i}"] = {"img1": clipped_imgs[i][1], "img2": clipped_imgs[i][0]}
    
#     dt[f"4_{i}"] = {"img1": clipped_imgs[i][4], "img2": clipped_imgs[i][0]}

#     dt[f"5_{i}"] = {"img1": clipped_imgs[i][4], "img2": clipped_imgs[i][7]}
def get_corner(tile,pos):
    if pos=="tl": return tile[0:32,0:32]
    elif pos=="tr": return tile[0:32,-32:]
    elif pos=="bl": return tile[-32:,0:32]
    elif pos=="br": return tile[-32:,-32:]

for i in range(59):
    t0=clipped_imgs[i][0]
    t1=clipped_imgs[i][1]
    t4=clipped_imgs[i][4]
    t7=clipped_imgs[i][7]
    t_next=clipped_imgs[i+1][1]
    dt[f"0_{i}"]={"img1":get_corner(t0,"tr"),"img2":get_corner(t_next,"tl")}
    dt[f"1_{i}"]={"img1":get_corner(t0,"tr"),"img2":get_corner(t1,"tl")}
    dt[f"2_{i}"]={"img1":get_corner(t0,"bl"),"img2":get_corner(t4,"tl")}
    dt[f"3_{i}"]={"img1":get_corner(t1,"tl"),"img2":get_corner(t0,"tr")}
    dt[f"4_{i}"]={"img1":get_corner(t4,"tl"),"img2":get_corner(t0,"bl")}
    dt[f"5_{i}"]={"img1":get_corner(t4,"tr"),"img2":get_corner(t7,"bl")}

# -- Code Cell --
rows = []
for key, value in dt.items():
    label = int(key.split("_")[0])
    rows.append({ "tile1": value["img1"],"tile2": value["img2"],"label": label})
df_train_clipped = pd.DataFrame(rows)

# -- Code Cell --
class ClippedDataset(Dataset):
    def __init__(self, df, transform):
        self.df = df
        self.transform = transform

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

    def __getitem__(self, index):
        row = self.df.iloc[index]
        tile1 = row["tile1"]
        tile2 = row["tile2"]
        tile1 = cv.cvtColor(tile1, cv.COLOR_BGR2RGB)
        tile2 = cv.cvtColor(tile2, cv.COLOR_BGR2RGB)
        tile1 = Image.fromarray(tile1)
        tile2 = Image.fromarray(tile2)
        if self.transform:
            tile1 = self.transform(tile1)
            tile2 = self.transform(tile2)
        label = row["label"]
        return tile1, tile2, label

# -- Code Cell --
class ValidationDataset(Dataset):
    def __init__(self, transform, path=None):
        self.transform = transform
        if path:
            self.df = pd.read_csv(path)
    def __len__(self):
        return len(self.df)
    def __getitem__(self,index):
        row = self.df.iloc[index]
        tile1 = cv.imread(f"./validation_data/validation/validation_tiles/{row['tile1']}.jpg")
        tile2 = cv.imread(f"./validation_data/validation/validation_tiles/{row['tile2']}.jpg")
        
        img_combined = np.zeros((128, 256, 3), dtype=np.uint8)
        img_combined[0:128, 0:128] = tile1
        img_combined[0:128, 128:256] = tile2
        img_rgb = cv.cvtColor(img_combined, cv.COLOR_BGR2RGB)
        img_pil = Image.fromarray(img_rgb)
        label = row['answer']
        if self.transform:
            img_rgb = self.transform(img_pil)
        return img_rgb,label

# -- Code Cell --
valid_ds = ValidationDataset(path="./validation_data/validation/validation_data.csv", transform=transform)
valid_loader =DataLoader(valid_ds, batch_size=32, num_workers=0, shuffle=True)

# -- Code Cell --
clipped_ds = ClippedDataset(df_train_clipped, transform=transform)
clipped_loader = DataLoader(clipped_ds, batch_size=32, num_workers=0, shuffle=True)

# -- Code Cell --
weight = models.ResNet50_Weights.IMAGENET1K_V1

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

# -- Code Cell --
class Modelu(nn.Module):
    def __init__(self, num_classes=6):
        super().__init__()
        backbone = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)
        backbone.fc = nn.Identity()  
        self.encoder = backbone
        self.classifier = nn.Sequential(
            nn.Linear(512 * 3, 256),
            nn.ReLU(),
            nn.Linear(256, num_classes)
        )
    def forward(self, x1, x2):
        e1 = self.encoder(x1)
        e2 = self.encoder(x2)

        x = torch.cat([e1, e2, torch.abs(e1 - e2)], dim=1)

        return self.classifier(x)

# -- Code Cell --
model = Modelu().to(device)
optimizer = optim.Adam(model.parameters(), lr=1e-4)
criterion = nn.CrossEntropyLoss()
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=15, gamma=0.5)

# -- Code Cell --
for epoch in range(25):
    model.train()
    total_loss = 0

    for img1, img2, label in clipped_loader:
        img1 = img1.to(device)
        img2 = img2.to(device)
        label = label.to(device)

        out = model(img1, img2)
        loss = criterion(out, label)

        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

        total_loss += loss.item()

    scheduler.step()
    print(f"Epoch {epoch+1} | Loss {total_loss/len(clipped_loader):.4f}")

# -- Code Cell --
# for epoch in range(25):
#     model.train()       
#     total_loss = 0
#     correct = 0
#     total = 0

#     for combined, label in valid_loader:
#         combined, label = combined.to(device), label.to(device)  
#         out = model(combined)                  
#         loss = criterion(out, label)        

#         optimizer.zero_grad()           
#         loss.backward()                
#         optimizer.step()                
 
#         total_loss += loss.item()
#         correct += (out.argmax(1) == label).sum().item() 
#         total += label.size(0)                             
#     for combined, label in clipped_loader:
#         combined, label = combined.to(device), label.to(device)
#         out = model(combined)
#         loss = criterion(out, label)

#         optimizer.zero_grad()
#         loss.backward()
#         optimizer.step()

#         total_loss += loss.item()
#         correct += (out.argmax(1) == label).sum().item()
#         total += label.size(0)
#     scheduler.step()  
#     print(f"Epoc {epoch+1} -- Loss {total_loss/(len(valid_loader)+len(clipped_loader))}")
    

# -- Code Cell --
class TestDataset(Dataset):
    def __init__(self, transform, path):
        self.transform = transform
        self.df = pd.read_csv(path)

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

    def __getitem__(self, index):
        row = self.df.iloc[index]

        tile1 = cv.imread(f"./test/test_tiles/{row['tile1']}.jpg")
        tile2 = cv.imread(f"./test/test_tiles/{row['tile2']}.jpg")

        tile1 = cv.cvtColor(tile1, cv.COLOR_BGR2RGB)
        tile2 = cv.cvtColor(tile2, cv.COLOR_BGR2RGB)

        tile1 = Image.fromarray(tile1)
        tile2 = Image.fromarray(tile2)

        if self.transform:
            tile1 = self.transform(tile1)
            tile2 = self.transform(tile2)

        return tile1, tile2

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

# -- Code Cell --
model.eval()
preds = []

with torch.no_grad():
    for img1, img2 in test_loader:
        img1 = img1.to(device)
        img2 = img2.to(device)

        pred = model(img1, img2)
        pred = torch.argmax(pred, dim=1)

        preds.append(pred.cpu())

preds_final = torch.cat(preds).numpy()
preds_final

# -- Code Cell --
test_pairs = pd.read_csv("./test/test_pairs.csv")

# -- Code Cell --
subs1_adica_tot_2 = pd.DataFrame({
    'subtaskID': 1,
    'datapointID': test_pairs['datapointID'],
    'answer': preds_final
})
subs2 = pd.DataFrame({
    'subtaskID': 2,
    'datapointID': test_pairs['datapointID'],
    'answer': preds_final
})
sub_final = pd.concat([subs1_adica_tot_2, subs2], ignore_index=True)
sub_final.to_csv('subs.csv', index=False)

# -- Code Cell --
