# -- Code Cell --
import numpy as np
import pandas as pd
import ast
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
import gc

device = 'cuda' if torch.cuda.is_available() else 'cpu'
print(f"Device: {device}")

# ===================== LOAD DATA =====================
print("Loading data...")
df = pd.read_csv('train_data.csv')
latents = []
for i in range(len(df)):
    ans = np.array(ast.literal_eval(df.iloc[i]['answer']), dtype=np.float32)
    latents.append(ans)
latents = np.array(latents)  # (41, 2, 4, 90, 160)
print(f"Latents: {latents.shape}")

all_frames = latents.reshape(41, 8, 90, 160)  # RGB(4ch) + Depth(4ch)
rgb_frames = latents[:, 0]  # (41, 4, 90, 160)
del df, latents; gc.collect()

# ===================== DATASET =====================
WINDOW = 5

class FrameDataset(Dataset):
    def __init__(self, frames_8ch, targets_4ch, window):
        self.X, self.Y = [], []
        for i in range(len(frames_8ch) - window):
            self.X.append(frames_8ch[i:i+window])
            self.Y.append(targets_4ch[i+window])
        self.X = torch.tensor(np.array(self.X))
        self.Y = torch.tensor(np.array(self.Y))
    def __len__(self): return len(self.X)
    def __getitem__(self, i): return self.X[i], self.Y[i]

ds = FrameDataset(all_frames, rgb_frames, WINDOW)
dl = DataLoader(ds, batch_size=len(ds), shuffle=True, num_workers=0)
print(f"Dataset: {len(ds)} samples")

# ===================== MODEL =====================
class FramePredictor(nn.Module):
    def __init__(self):
        super().__init__()
        # Conv3D encoder: (B, 8, 5, 90, 160) -> temporal collapses to 1
        self.enc = nn.Sequential(
            nn.Conv3d(8, 32, kernel_size=(3,3,3), padding=(0,1,1)),
            nn.GELU(),
            nn.Conv3d(32, 64, kernel_size=(3,3,3), padding=(0,1,1)),
            nn.GELU(),
        )
        self.dec = nn.Sequential(
            nn.Conv2d(64, 64, 3, padding=1),
            nn.GELU(),
            nn.Conv2d(64, 32, 3, padding=1),
            nn.GELU(),
            nn.Conv2d(32, 4, 3, padding=1),
        )
        self.skip = nn.Conv2d(8, 4, 1)  # residual from last input frame

    def forward(self, x):
        # x: (B, T, 8, H, W)
        B, T, C, H, W = x.shape
        h = self.enc(x.permute(0, 2, 1, 3, 4))  # (B, 64, 1, H, W)
        h = h.squeeze(2)                          # (B, 64, H, W)
        out = self.dec(h)                          # (B, 4, H, W)
        return out + self.skip(x[:, -1])           # residual connection

model = FramePredictor().to(device)
print(f"Params: {sum(p.numel() for p in model.parameters()):,}")

# ===================== TRAINING =====================
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=3000)
loss_fn = nn.L1Loss()

model.train()
for epoch in range(3000):
    for X, Y in dl:
        X, Y = X.to(device), Y.to(device)
        pred = model(X)
        loss = loss_fn(pred, Y)
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()
    scheduler.step()
    if (epoch + 1) % 300 == 0:
        print(f"  Epoch {epoch+1}: loss={loss.item():.6f}, lr={scheduler.get_last_lr()[0]:.6f}")

# ===================== INFERENCE (autoregressive) =====================
print("Predicting next 5 frames...")
model.eval()

last_depth = all_frames[-1, 4:]  # (4, 90, 160) - reuse last depth
window = all_frames[-WINDOW:].copy()  # (5, 8, 90, 160)
predictions = []

with torch.no_grad():
    for step in range(5):
        inp = torch.tensor(window, dtype=torch.float32).unsqueeze(0).to(device)
        pred_rgb = model(inp).cpu().numpy()[0]  # (4, 90, 160)
        predictions.append(pred_rgb)
        # Shift window
        new_frame = np.concatenate([pred_rgb, last_depth], axis=0)[np.newaxis]
        window = np.concatenate([window[1:], new_frame], axis=0)
        print(f"  Step {step+1}: range=[{pred_rgb.min():.2f}, {pred_rgb.max():.2f}]")

# ===================== WRITE CSV =====================
print("Writing submission.csv...")
rows = []
for i, pred in enumerate(predictions):
    pred_out = pred[np.newaxis]  # (1, 4, 90, 160)
    rows.append({
        'subtaskID': 1,
        'datapointID': 41 + i,
        'answer': repr(pred_out.tolist())
    })

out_df = pd.DataFrame(rows)
out_df.to_csv('submission.csv', index=False, quoting=1)

# Verify
verify = pd.read_csv('submission.csv')
t = np.array(ast.literal_eval(verify.iloc[0]['answer']))
print(f"Verify: shape={t.shape}, IDs={verify['datapointID'].tolist()}")
print("Done!")

# -- Code Cell --
