# -- Code Cell --
import pandas as pd

# -- Code Cell --
train = pd.read_csv("train_data.csv")
test = pd.read_csv("test_data.csv")

# -- Code Cell --
test_task1 = test[test['subtaskID']==1]
test_task2 = test[test['subtaskID']==2]

# -- Code Cell --
train['image'] = train['image'].apply(lambda x: "./images/" + str(x))
test_task1['img1'] = test_task1['img1'].apply(lambda x: "./images/" + str(x))
test_task1['img2'] = test_task1['img2'].apply(lambda x: "./images/" + str(x))
test_task1['img3'] = test_task1['img3'].apply(lambda x: "./images/" + str(x))
test_task2['img1'] = test_task2['img1'].apply(lambda x: "./images/" + str(x))
test_task1['img1']

# -- Code Cell --
from skimage.metrics import structural_similarity as ssim
from PIL import Image
import numpy as np

# -- Code Cell --
def compare_images_task1(row_img):
    img1 = np.array(Image.open(row_img['img1']).convert("L"))
    img2 = np.array(Image.open(row_img['img2']).convert("L"))
    img3 = np.array(Image.open(row_img['img3']).convert("L"))
    similarrity = np.array([ssim(img1,img2),ssim(img2,img3), ssim(img1,img3)])
    if np.argmax(similarrity) == 0:
        return 3
    if np.argmax(similarrity) == 1:
        return 1
    if np.argmax(similarrity) == 2:
        return 2

# -- Code Cell --
answer = []
for i in range(len(test_task1)):
    row = test_task1.iloc[i]
    answer.append(compare_images_task1(row))

# -- Code Cell --
import cv2 as cv

# -- Code Cell --
import torch
from torch.utils.data import DataLoader,Dataset,random_split
from torchvision.transforms import v2
import torch.optim as optim
import torch.nn as nn

# -- Code Cell --
train

# -- Code Cell --
class Datasetttt(Dataset):
    def __init__(self, df,tranfsorm):
        self.df = df
        self.tranform = tranfsorm
    def __len__(self):
        return len(self.df)
    def __getitem__(self, index):
        row = self.df.iloc[index]
        label = row['angle']
        img_path = row['image']
        img =Image.open(img_path)
        img = self.tranform(img)
        return img, label

# -- Code Cell --
transform = v2.Compose([
    v2.Resize((64,64)),
    v2.ToTensor(),
    v2.Normalize([0.5,0.5,0.5],[0.5,0.5,0.5])
])

# -- Code Cell --
full_ds = Datasetttt(train, tranfsorm=transform)
train_sub, val_sub = random_split(full_ds, [0.8,0.2])
train_dt = DataLoader(train_sub, batch_size=32, shuffle=True, num_workers =0)
val_dt = DataLoader(val_sub, batch_size=32, shuffle=True, num_workers =0)

# -- Code Cell --
from torchvision import models

# -- Code Cell --
class modelas(nn.Module):
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        self.backbone = models.resnet18(pretrained =False)
        self.in_feat = self.backbone.fc.in_features
        self.backbone.fc = nn.Sequential(
             nn.Linear(self.in_feat, 360),
            nn.ReLU(),
             nn.Linear(360, 1)
        )
    def forward(self, x):
        x = self.backbone(x)
        return x

# -- Code Cell --
def angular_error(pred, true):
    diff = torch.abs(pred - true) % 360
    return torch.minimum(diff, 360 - diff)

# -- Code Cell --
model = modelas()
optimizer = optim.Adam(model.parameters(), lr=1e-4)
criterion = nn.L1Loss()
scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max')
device ='cuda'
model = model.to(device)

# -- Code Cell --
best_loss = float('inf')

for epoch in range(100):
    model.train()
    total_loss = 0

    for img, label in train_dt:
        img = img.to(device)
        label = label.to(device).float()

        optimizer.zero_grad()
        output = model(img).squeeze(1)

        loss = criterion(output, label)
        loss.backward()
        optimizer.step()

        total_loss += loss.item()

    model.eval()
    total_loss_eval = 0
    total_ang_err = 0
    with torch.no_grad():
        for img, label in val_dt:
            img = img.to(device)
            label = label.to(device).float()

            output = model(img).squeeze(1)
            loss = criterion(output, label)
            total_loss_eval += loss.item()
            
            err = angular_error(output, label).mean()
            total_ang_err += err.item()


    val_loss = total_loss_eval / len(val_dt)
    mean_ang_err = total_ang_err / len(val_dt)
    if val_loss < best_loss:
        best_loss = val_loss
        torch.save(model.state_dict(), "model.pth")

    scheduler.step(mean_ang_err)
    print(f"val_loss = {val_loss}, mean angular error = {mean_ang_err}")

# -- Code Cell --
class TestDataset(Dataset):
    def __init__(self, df,tranfsorm):
        self.df = df
        self.tranform = tranfsorm
    def __len__(self):
        return len(self.df)
    def __getitem__(self, index):
        row = self.df.iloc[index]
        img_path = row['img1']
        img =Image.open(img_path)
        img = self.tranform(img)
        return img

# -- Code Cell --
test_ds = TestDataset(test_task2,transform)
test_loader = DataLoader(test_ds, batch_size=32, shuffle=False, num_workers=0)

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

# -- Code Cell --
sub1 = pd.DataFrame({
    "subtaskID":1,
    "datapointID":test_task1['datapointID'],
    "answer":[int(value) for value in ans]
})
sub2 = pd.DataFrame({
    "subtaskID":2,
    "datapointID":test_task2['datapointID'],
    "answer":preds.astype(int)
})
final = pd.concat([sub1,sub2]).to_csv("subs.csv",index=False)

# -- Code Cell --
