# -- Code Cell --
import pandas as pd
train = pd.read_csv("./starter_kit/data/train_data.csv")
test =pd.read_csv("./starter_kit/data/test_data.csv")

# -- Code Cell --
train

# -- Code Cell --
import cv2 as cv
import numpy as np

# -- Code Cell --
result = []
for i in range(len(test)):
    img=cv.imread(f"./starter_kit/data/images/test_data/image_{i}.jpg")
    mean_odd = []
    mean_even = []
    for r in range(img.shape[0]):
        for c in range(img.shape[1]):
            if r%2 != 0:
                mean_odd.append(img[r][c])
            if r%2 == 0:
                mean_even.append(img[r][c])
    mean_odd = np.mean(mean_odd)
    mean_even = np.mean(mean_even)
    result.append(mean_even-mean_odd)


# -- Code Cell --
task1 = np.mean(result)

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

# -- Code Cell --
labels = train['label'].unique().tolist()
label_map = {}
rev_map = {}
for idx, element in enumerate(labels):
    label_map[element] = idx
    rev_map[idx] = element

# -- Code Cell --
def mod_img(image):
    # image = image[:, 5:-5]
    image_color = image.copy()
    image = cv.cvtColor(image, cv.COLOR_RGB2GRAY)
    b, g, r = cv.split(image_color)
    def shift_channel(channel, shift):
        h, w = channel.shape
        M = np.float32([[1, 0, shift], [0, 1, 0]])
        return cv.warpAffine(channel, M, (w, h))
    b_fixed = shift_channel(b, -5) 
    r_fixed = shift_channel(r, 5)   
    fixed = cv.merge([b_fixed, g, r_fixed])
    t,reconstructed = cv.threshold(image,1,255,cv.THRESH_BINARY_INV)
    mask = reconstructed
    #blur
    blur = cv.GaussianBlur(fixed,(5,5),0)
    #vhs
    img = np.zeros((image.shape[0],image.shape[1]), np.uint8)
    for r in range(image.shape[0]):
            if r%2 != 0:
                cv.line(img,(0,r),(256,r),(255,255,255),1)
    mask2 = img
    #combinat
    mask3 = cv.bitwise_or(mask, mask2)
    #bright
    bright = cv.convertScaleAbs(blur, alpha=1.0, beta=30)
    dstpt2 = cv.inpaint(bright,mask3,1,cv.INPAINT_TELEA)     
    return dstpt2

# -- Code Cell --
class TrainDataset(Dataset):
    def __init__(self, path, transform):
        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]
        image = cv.imread(f"./starter_kit/{row['path']}")
        image = cv.cvtColor(image, cv.COLOR_BGR2RGB)
        image = Image.fromarray(mod_img(image))
        image = self.transform(image)
        label = label_map[row['label']]
        label = torch.tensor(label, dtype=torch.long)
        return image, label

# -- Code Cell --
tranform = 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("./starter_kit/data/train_data.csv", tranform)
train_subset, val_subset = random_split(train_ds, [0.8,0.2])
train_loader = DataLoader(train_subset, batch_size=32, shuffle=True, num_workers=0)
valid_loader = DataLoader(val_subset, batch_size=32, shuffle=True, num_workers=0)

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

# -- Code Cell --
model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)
model.fc =nn.Linear(2048, 7)
for param in model.parameters():
    param.requires_grad = True
for param in model.fc.parameters():
    param.requires_grad = True
optimizer = optim.Adam(model.parameters(), lr=1e-4)
criterion = nn.CrossEntropyLoss()
scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max')
from sklearn.metrics import accuracy_score
model = model.to("cuda")

# -- Code Cell --
from tqdm import tqdm
decoy = -1
for epoch in range(20):
    model.train()
    for img,label in tqdm(train_loader):
        img = img.to("cuda")
        label = label.to("cuda")
        optimizer.zero_grad()
        output = model(img)
        loss = criterion(output,label)
        loss.backward()
        optimizer.step()
    model.eval()
    all_preds = []
    all_labels = []
    with torch.no_grad():
        for img,label in valid_loader:
            img = img.to("cuda")
            label = label.to("cuda")
            output = model(img)
            pred = output.argmax(dim=1)
            all_preds.append(pred.cpu())
            all_labels.append(label.cpu())
    all_labels = torch.cat(all_labels).numpy()
    all_preds = torch.cat(all_preds).numpy()
    acc = accuracy_score(all_labels,all_preds)
    if acc>decoy:
        decoy = acc
        torch.save(model.state_dict(), "model_full.pt")
    scheduler.step(acc)
    print(f"epoch {epoch}, acc {acc}")

# -- Code Cell --
class TestDataset(Dataset):
    def __init__(self, path, transform):
        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]
        image = cv.imread(f"./starter_kit/{row['path']}")
        image = cv.cvtColor(image, cv.COLOR_BGR2RGB)
        image = Image.fromarray(mod_img(image))
        image = self.transform(image)
        return image

# -- Code Cell --
test_ds = TestDataset("./starter_kit/data/test_data.csv", tranform)
test_loader = DataLoader(test_ds, num_workers=0, shuffle=False, batch_size=32)

# -- Code Cell --
model.load_state_dict(torch.load("./model_full.pt"))
model.eval()
preds=[]
with torch.no_grad():
    for img in test_loader:
        img = img.to("cuda")
        output = model(img)
        pred = output.argmax(dim=1)
        preds.append(pred.cpu())
preds = torch.cat(preds).numpy()
preds

# -- Code Cell --
preds_final = []
for i in preds:
    preds_final.append(rev_map[i])

# -- Code Cell --
preds_final

# -- Code Cell --
sub1 = pd.DataFrame({
    "subtaskID":[1],
    "datapointID":[1],
    'answer':[task1]
})
sub2 = pd.DataFrame({
    "subtaskID":2,
    "datapointID":test['ID'],
    'answer':preds_final
})
final = pd.concat([sub1,sub2]).to_csv("subs.csv",index=False)

# -- Code Cell --
