# -- Code Cell --
import pandas as pd
import numpy as np

# -- Code Cell --
df_train = pd.read_csv("train_data.csv")
df_test = pd.read_csv("test_data.csv")
df_train.head()

# -- Code Cell --
df_test.head()

# -- Code Cell --


# -- Code Cell --
df_train["image"] = df_train["image"].apply(lambda x: "./images/" + x)
df_test["img1"] = df_test["img1"].apply(lambda x: "./images/" + str(x))
df_test["img2"] = df_test["img2"].apply(lambda x: "./images/" + str(x))
df_test["img3"] = df_test["img3"].apply(lambda x: "./images/" + str(x))


# -- Code Cell --
df_test_1 = df_test[df_test["subtaskID"] == 1]

# -- Code Cell --
from PIL import Image
img_tr = Image.open(df_train["image"].iloc[0]).convert("L")
img_tr

# -- Code Cell --
img_tst_1 = Image.open(df_test["img1"].iloc[0]).convert("L")
img_tst_1

# -- Code Cell --
img_tst_3 = Image.open(df_test["img3"].iloc[0]).convert("L")
img_tst_3

# -- Code Cell --
img_tst_2 = Image.open(df_test["img2"].iloc[0]).convert("L")
img_tst_2

# -- Code Cell --
from skimage.metrics import structural_similarity as ssim

# -- Code Cell --
np.array(img_tst_1).shape

# -- Code Cell --
ssim(np.array(img_tr), np.array(img_tst_3))

# -- Code Cell --
def get_lowest_ssim(row_tst):
    img_tst_1 = np.array(Image.open(row_tst.img1).convert("L"))
    img_tst_2 = np.array(Image.open(row_tst.img2).convert("L"))
    img_tst_3 = np.array(Image.open(row_tst.img3).convert("L"))
    idk_bro = np.array([ssim(img_tst_1, img_tst_2), ssim(img_tst_1, img_tst_3), ssim(img_tst_2, img_tst_3)])
    if np.argmax(idk_bro) +1 == 1:
        return 3
    elif np.argmax(idk_bro) +1 == 2:
        return 2
    else:
        return 1

# -- Code Cell --
from tqdm.notebook import tqdm
ans = []
for i in tqdm(range(len(df_test_1))):
    ans.append(get_lowest_ssim(df_test_1.iloc[i]))

# -- Code Cell --
df_pred_1 = pd.DataFrame({
    "subtaskID" : 1,
    "datapointID":  df_test_1["datapointID"],
    "answer" : ans
})

# -- Code Cell --
df_pred_2 = pd.DataFrame({
    "subtaskID" : 0,
    "datapointID":  df_test[df_test["subtaskID"] == 2]["datapointID"],
    "answer" : 0
})

# -- Code Cell --
df_pred = pd.concat([df_pred_1, df_pred_2])

# -- Code Cell --
df_pred.to_csv("subi.csv", index= False)

# -- Code Cell --
df_test_2 = df_test[df_test["subtaskID"] == 2]

# -- Code Cell --
pd.unique(df_test["subtaskID"])

# -- Code Cell --
df_test_2.head()

# -- Code Cell --
df_train.head()

# -- Code Cell --
from torchvision.transforms import v2
transform = v2.Compose([
    v2.Resize((64, 64)),
    v2.ToTensor(),
    v2.Normalize(mean=[0.5 ,0.5, 0.5], std=[0.5] * 3)
])

# -- Code Cell --
from torch.utils.data import Dataset, DataLoader
import numpy as np
import torch

class data_idk(Dataset):
    def __init__(self, df, test = False):
        super().__init__()
        self.test = test
        if self.test == False:
            self.X_path = df["image"]
        else:
            self.X_path = df["img1"]
        if self.test == False:
            y = df["angle"]
            self.y = y.astype(int)
    def __len__(self):
        return len(self.X_path)
    
    def __getitem__(self, index):
        img = Image.open(self.X_path.iloc[index])
        img_trns = transform(img)
        if self.test == False:
                label = torch.tensor(int(self.y.iloc[index]), dtype=torch.long)
                return img_trns, label
        return img_trns


# -- Code Cell --
df_test_2 = df_test_2.reset_index(drop=True)

# -- Code Cell --
full_dt = data_idk(df_train)
test_dt = data_idk(df_test_2, test= True)

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

train_dt, valid_dt =random_split(full_dt, [0.8, 0.2])

# -- Code Cell --
from torch.utils.data import DataLoader
train_loader=  DataLoader(train_dt, batch_size= 32, shuffle= True)
val_loader = DataLoader(valid_dt, batch_size=32, shuffle= True)
test_loader = DataLoader(test_dt,  batch_size=32 )

# -- Code Cell --
test_dt[0].shape

# -- Code Cell --
#Da prost cu clasificare cu 72 sau 360 de clase(48 sau 46), cel mai bine a dat cu regresie(83)

# -- Code Cell --
import torch.nn as nn
import torchvision
class model_idk(nn.Module):
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        self.backbone = torchvision.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 --
import torch
model = model_idk()
NUM_EPOCHS = 100
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
model = model.to(DEVICE)
criterion = nn.L1Loss()
optimizer =torch.optim.Adam(model.parameters())

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

for epoch in range(1, NUM_EPOCHS + 1):
    print(f"\nEpoch {epoch}/{NUM_EPOCHS}")

    # --- train ---
    model.train()
    total_loss = 0.0
    total = 0
    for batch_idx, (imgs, labels) in enumerate(train_loader):
        imgs   = imgs.to(DEVICE)
        labels = labels.to(DEVICE).float()

        optimizer.zero_grad()
        preds = model(imgs).squeeze(1)
        loss  = criterion(preds, labels)

        loss.backward()
        optimizer.step()

        total_loss += loss.item() * len(imgs)
        total      += len(imgs)

    train_loss = total_loss / total
    print(f"  [Train] Loss: {train_loss:.4f}")

    # --- validate ---
    model.eval()
    total_loss = 0.0
    total = 0
    with torch.no_grad():
        for imgs, labels in val_loader:
            imgs   = imgs.to(DEVICE)
            labels = labels.to(DEVICE).float()

            preds      = model(imgs).squeeze(1)
            loss       = criterion(preds, labels)
            total_loss += loss.item() * len(imgs)
            total      += len(imgs)

    val_loss = total_loss / total
    print(f"  [Val]   Loss: {val_loss:.4f}")

    if val_loss < best_val_loss:
        best_val_loss = val_loss
        torch.save(model.state_dict(), "best_model.pth")
        print(f"  --> Best saved (loss={best_val_loss:.4f})")


# -- Code Cell --
ans_v2 = []
with torch.no_grad():
 model.eval()
 labels=[]
 for batch in test_loader:
    X = batch.to(DEVICE)
   # print(model(X).detach().cpu().numpy().tolist())
    ans_v2.extend(model(X).detach().cpu().numpy().tolist())



# -- Code Cell --
import math

ans_v2 = [math.degrees(math.asin(int(x[0]))) for x in ans_v2]

# -- Code Cell --
df_pred_2 = pd.DataFrame({
    "subtaskID" : 2,
    "datapointID":  df_test[df_test["subtaskID"] == 2]["datapointID"],
    "answer" : ans_v2
})

df_pred_2

# -- Code Cell --
submission = pd.concat([df_pred_1, df_pred_2])

submission.to_csv('submission_matter_of_perspective.csv', index = False)

# -- Code Cell --


# -- Code Cell --
