# -- Code Cell --
import os
import pandas as pd
from PIL import Image

# -- Code Cell --
train = pd.read_csv('train.csv')

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

# -- Code Cell --
train['Effect'].unique()

# -- Code Cell --
imgs = []
for i in range(len(train)):
    path = os.path.join("./starting_kit/", f"{train['Path'][i]}")
    img = Image.open(path)
    imgs.append(img)

# -- Code Cell --
imgs[0]

# -- Code Cell --
import pandas as pd
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 os

# -- Code Cell --
map_effect = {
    '-4': 0,
    '-2': 1,
    '-1': 2,
    '0': 3,
    '1': 4,
    '4': 5,
    'A': 6,
    'B': 7
}

# -- 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]
        img = Image.open(os.path.join("./starting_kit/",f"{train['Path'][index]}")).convert("L")
        if self.transform:
            img = self.transform(img)
        label = map_effect[row['Effect']]
        return img,label

# -- Code Cell --
transform = transforms.Compose([
    transforms.Resize((64,64)),
    transforms.ColorJitter(brightness=0.2,contrast=0.2,saturation=0.2,hue=0.05),
    transforms.ToTensor(),
    transforms.Normalize([0.5], [0.5])
])

# -- Code Cell --
train_ds=TrainDataset("./train.csv", transform=transform)

# -- Code Cell --
from torch.utils.data import random_split
train_size = int(0.8* len(train_ds))
valid_size = len(train_ds) - train_size
train_subset, valid_subset = random_split(train_ds, [train_size,valid_size])

# -- Code Cell --
valid_loader=  DataLoader(valid_subset, batch_size=32, shuffle=False, num_workers=0)
train_loader = DataLoader(train_subset, batch_size=32, shuffle=True, num_workers=0)

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

# -- Code Cell --
class CONVBlock(nn.Module):
    def __init__(self, in_c, out_c):
        super().__init__()
        self.conv = nn.Sequential(
            nn.Conv2d(in_c, out_c, kernel_size=3,padding=1),
            nn.ReLU(inplace=True),
            nn.BatchNorm2d(out_c),
            nn.MaxPool2d(2)
        )
    def forward(self,x):
        con = self.conv(x)
        return con

# -- Code Cell --
class CustomModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = CONVBlock(1, 32)
        self.conv2 = CONVBlock(32, 64)
        self.conv3 = CONVBlock(64, 128)
        self.adap = nn.AdaptiveAvgPool2d((1,1))
        self.cls = nn.Sequential(
            nn.Flatten(),
            nn.Linear(128,128),
            nn.ReLU(inplace=True),
            nn.Dropout(0.3),
            nn.Linear(128,8),
        )
    def forward(self,x):
        c1 = self.conv1(x)
        c2 = self.conv2(c1)
        c3 = self.conv3(c2)
        adap = self.adap(c3)
        cls = self.cls(adap)
        return cls

# -- Code Cell --
model = CustomModel()
model = model.to(device)
optimizer = optim.Adam(model.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()
scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.2)

# -- Code Cell --
from sklearn.metrics import f1_score

# -- Code Cell --
for epoch in range(30):
    model.train()
    total_loss = 0
    for img,label in 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()
    model.eval()
    all_labels = []
    all_preds = []
    with torch.no_grad():
        for img, label in valid_loader:
            img = img.to(device)
            label = label.to(device)
            output = model(img)
            pred = output.argmax(dim=1)
            all_preds.append(pred.cpu())
            all_labels.append(label.cpu())
    all_preds = torch.cat(all_preds)
    all_labels = torch.cat(all_labels)
    all_labels = all_labels.numpy()
    all_preds = all_preds.numpy()
    f1 = f1_score(all_labels,all_preds,average='macro')
    scheduler.step()
    print(f"epoch {epoch+1}--loss {total_loss/len(train_loader)}--f1 {f1}")

# -- Code Cell --
imgs = []
for i in range(len(train)):
    path = os.path.join("./starting_kit/", f"{train['Path'][i]}")
    img = Image.open(path).convert("L")
    imgs.append(img)

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

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

# -- Code Cell --
import cv2

all_sequences = []

for i in range(len(test)):
    path = os.path.join("./starting_kit/", test['datapointID'][i])
    img = cv2.imread(path)
    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
    _, binary = cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY_INV + cv2.THRESH_OTSU)
    num_labels, labels, stats, centroids = cv2.connectedComponentsWithStats(binary, connectivity=8)
    boxes = []
    for t in range(1, num_labels):
        x = stats[t, cv2.CC_STAT_LEFT]
        y = stats[t, cv2.CC_STAT_TOP]
        w = stats[t, cv2.CC_STAT_WIDTH]
        h = stats[t, cv2.CC_STAT_HEIGHT]
        area = stats[t, cv2.CC_STAT_AREA]
        if area > 20:
            boxes.append((x, y, w, h))

    boxes = sorted(boxes, key=lambda b: b[0])

    sequence = []
    for (x, y, w, h) in boxes:
        crop = gray[y:y+h, x:x+w]
        sequence.append(crop)

    all_sequences.append(sequence)

# -- Code Cell --
all_sequences[0][0]

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

for i in range(len(all_sequences)):
    seq_preds = []
    for crop in all_sequences[i]:
        crop_rgb = cv2.cvtColor(crop, cv2.COLOR_BGR2RGB)
        crop_pil = Image.fromarray(crop_rgb).convert("L")

        x = transform(crop_pil)          
        x = x.unsqueeze(0).to(device)    

        with torch.no_grad():
            out = model(x)
            pred = torch.argmax(out, dim=1).item()

        seq_preds.append(pred)

    preds.append(seq_preds)

# -- Code Cell --
rev_map = {0:'-4', 1:'-2', 2:'-1', 3:'0', 4:'1', 5:'4', 6:'A', 7:'B'}
pred_labels = []

for seq in preds:
    pred_labels.append([rev_map[p] for p in seq])

# -- Code Cell --
pred_labels

# -- Code Cell --
answers = []

for seq in pred_labels:
    total = 0
    current_seq = []

    for j in seq:
        if j in ['A', 'B']:
            effect = 0
        else:
            effect = int(j)

        total += effect
        current_seq.append(str(total))

    answers.append("|".join(current_seq))

# -- Code Cell --
answers[100]

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

# -- Code Cell --


# -- Code Cell --
