# -- 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 torch.nn.functional as F

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

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

# -- Code Cell --
import ast

# -- Code Cell --
train_df['answer'] = train_df['answer'].apply(ast.literal_eval)

# -- Code Cell --
(rgb,depth) = (train_df['answer'].iloc[0][0][0][0][0],train_df['answer'].iloc[0][0][0][0][1])
(rgb,depth)

# -- Code Cell --
rgb = []
depth = []
for i in range(len(train_df)):
    rgb.append(train_df['answer'].iloc[i])
    depth.append(train_df['answer'].iloc[i])

# -- Code Cell --
depth[0]

# -- Code Cell --
class DoubleConv(nn.Module):
    def __init__(self, in_ch, out_ch):
        super().__init__()
        self.conv = nn.Sequential(
            nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),
            nn.BatchNorm2d(out_ch),
            nn.ReLU(inplace=True),
            nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1),
            nn.BatchNorm2d(out_ch),
            nn.ReLU(inplace=True),
        )

    def forward(self, x):
        return self.conv(x)


class UNet(nn.Module):
    def __init__(self, in_channels=24, out_channels=4):
        super().__init__()
        
        self.enc1 = DoubleConv(in_channels, 64)
        self.pool1 = nn.MaxPool2d(2)

        self.enc2 = DoubleConv(64, 128)
        self.pool2 = nn.MaxPool2d(2)

        self.enc3 = DoubleConv(128, 256)
        self.pool3 = nn.MaxPool2d(2)

        self.bottleneck = DoubleConv(256, 512)

        self.up3 = nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2)
        self.dec3 = DoubleConv(512, 256)

        self.up2 = nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2)
        self.dec2 = DoubleConv(256, 128)

        self.up1 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2)
        self.dec1 = DoubleConv(128, 64)

        self.out_conv = nn.Conv2d(64, out_channels, kernel_size=1)

    def forward(self, x):
        c1 = self.enc1(x)
        p1 = self.pool1(c1)

        c2 = self.enc2(p1)
        p2 = self.pool2(c2)

        c3 = self.enc3(p2)
        p3 = self.pool3(c3)

        bn = self.bottleneck(p3)

        u3 = F.interpolate(self.up3(bn), size=c3.shape[2:])
        d3 = self.dec3(torch.cat([u3, c3], dim=1))

        u2 = F.interpolate(self.up2(d3), size=c2.shape[2:])
        d2 = self.dec2(torch.cat([u2, c2], dim=1))

        u1 = F.interpolate(self.up1(d2), size=c1.shape[2:])
        d1 = self.dec1(torch.cat([u1, c1], dim=1))

        return self.out_conv(d1)


# -- Code Cell --
import numpy as np

# -- Code Cell --
rgb = np.array(rgb, dtype=np.float32)
depth = np.array(depth, dtype=np.float32)

# -- Code Cell --
data = np.array(rgb)  
rgb = data[:, 0]    
depth = data[:, 1] 

# -- Code Cell --
rgb_mean = rgb.mean()
rgb_std = rgb.std()
rgb = (rgb - rgb_mean) / rgb_std

depth_mean = depth.mean()
depth_std = depth.std()
depth = (depth - depth_mean) / depth_std

# -- Code Cell --
class FrameDataset(Dataset):
    def __init__(self, rgb, depth, window=3):
        self.rgb = rgb
        self.depth = depth
        self.window = window
    
    def __len__(self):
        return len(self.rgb) - self.window  
    
    def __getitem__(self, idx):
        
        rgb_in = self.rgb[idx:idx+self.window]      
        depth_in = self.depth[idx:idx+self.window]   
        
        x = np.concatenate([
            rgb_in.reshape(-1, 90, 160),   
            depth_in.reshape(-1, 90, 160),  
        ], axis=0)  
        
        y = self.rgb[idx + self.window]  
        
        return torch.tensor(x), torch.tensor(y)

# -- Code Cell --
dataset = FrameDataset(rgb, depth, window=3)

# -- Code Cell --
from torch.utils.data import random_split
train_size = int(0.8*len(dataset))
val_size = len(dataset) - train_size
train_subset,val_subset = random_split(dataset, [train_size,val_size])

# -- Code Cell --
train_loader = DataLoader(train_subset, batch_size=32, shuffle=True, num_workers=0)
val_loader = DataLoader(val_subset, batch_size=32, shuffle=False, num_workers=0)
dataloader = DataLoader(dataset, shuffle=True, num_workers=0, batch_size=32)

# -- Code Cell --
model = UNet()
model = model.to(torch.device("cuda"))
optimizer = optim.Adam(model.parameters(), lr=1e-4)
criterion = nn.MSELoss()
scheduler = optim.lr_scheduler.StepLR(optimizer,step_size=30, gamma=0.3)

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

# -- Code Cell --
for epoch in range(300):
    model.train()
    total_loss = 0
    for x,y in dataloader:
        x = x.to(torch.device("cuda"))
        y = y.to(torch.device("cuda"))
        optimizer.zero_grad()
        output = model(x)
        loss = criterion(output,y)
        loss.backward()
        optimizer.step()
    print(f"ep{epoch+1}-los{loss:04f}")

# -- Code Cell --
model.eval()
last_rgb = rgb[-3:]
last_depth = depth[-3:]

predictions = []
last_datapoint_id = 40

with torch.no_grad():
    for i in range(5):
        x = np.concatenate([
            last_rgb.reshape(-1, 90, 160),
            last_depth.reshape(-1, 90, 160),
        ], axis=0)
        x = torch.tensor(x).unsqueeze(0).to(torch.device("cuda"))

        pred = model(x).cpu().numpy()[0]
        pred_denorm = pred * rgb_std + rgb_mean
        predictions.append(pred_denorm)

        last_rgb = np.concatenate([last_rgb[1:], pred[np.newaxis]], axis=0)
        last_depth = np.concatenate([last_depth[1:], last_depth[-1:]], axis=0)

# -- Code Cell --
arr[0][0]

# -- Code Cell --
rows = []
for i, pred in enumerate(predictions):
    arr = pred[np.newaxis]  
    rows.append({
        'subtaskID': 1,
        'datapointID': 41 + i,
        'answer': arr[0][0].tolist()
    })
 
submission = pd.DataFrame(rows)
submission.to_csv('submission.csv', index=False)

# -- Code Cell --
