{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.0"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":130932,"databundleVersionId":15769099,"isSourceIdPinned":false}],"dockerImageVersionId":31286,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Automatic Lens Correction — Hybrid Model\n1. Parametric radial warp (k1, k2, cx, cy) corrects bulk barrel distortion\n2. Residual dense flow handles asymmetric quirks\nLoss: gradient magnitude + direction + SSIM + smoothness","metadata":{}},{"cell_type":"code","source":"import os, json, glob, re, random, zipfile\nimport numpy as np\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms.functional as TF\n\ndef ssim(x, y, data_range=1.0, size_average=True):\n    C1 = (0.01 * data_range) ** 2\n    C2 = (0.03 * data_range) ** 2\n    mu_x = F.avg_pool2d(x, 3, 1, 1)\n    mu_y = F.avg_pool2d(y, 3, 1, 1)\n    mu_x_sq, mu_y_sq, mu_xy = mu_x**2, mu_y**2, mu_x*mu_y\n    sig_x  = F.avg_pool2d(x*x, 3, 1, 1) - mu_x_sq\n    sig_y  = F.avg_pool2d(y*y, 3, 1, 1) - mu_y_sq\n    sig_xy = F.avg_pool2d(x*y, 3, 1, 1) - mu_xy\n    num = (2*mu_xy + C1) * (2*sig_xy + C2)\n    den = (mu_x_sq + mu_y_sq + C1) * (sig_x + sig_y + C2)\n    return (num/den).mean() if size_average else num/den\n\nprint(\"Imports done. Device:\", \"cuda\" if torch.cuda.is_available() else \"cpu\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T04:07:36.456417Z","iopub.execute_input":"2026-02-22T04:07:36.457208Z","iopub.status.idle":"2026-02-22T04:07:43.817106Z","shell.execute_reply.started":"2026-02-22T04:07:36.457176Z","shell.execute_reply":"2026-02-22T04:07:43.816219Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CFG = {\n    \"train_root\": \"/kaggle/input/competitions/automatic-lens-correction/lens-correction-train-cleaned\",\n    \"test_root\":  \"/kaggle/input/competitions/automatic-lens-correction/test-originals\",\n    \"out_root\":   \"/kaggle/working\",\n    \"img_size\":   384,\n    \"batch_size\": 8,\n    \"epochs\":     10,\n    \"max_pairs\":  10000,   # set to None to use all 23k\n    \"lr\":         1e-4,\n    \"val_split\":  0.05,\n    \"seed\":       42,\n    \"w_grad_mag\": 1.0,\n    \"w_grad_dir\": 0.5,\n    \"w_ssim\":     0.2,\n    \"w_l1\":       0.05,\n    \"w_smooth\":   0.2,\n    \"flow_tanh_scale\": 0.08,\n    \"grad_clip\":  1.0,\n    \"patience\":   2,\n}\nCFG[\"device\"] = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nrandom.seed(CFG[\"seed\"])\nnp.random.seed(CFG[\"seed\"])\ntorch.manual_seed(CFG[\"seed\"])\ntorch.cuda.manual_seed_all(CFG[\"seed\"])\nprint(\"Device:\", CFG[\"device\"])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T04:07:43.818655Z","iopub.execute_input":"2026-02-22T04:07:43.819076Z","iopub.status.idle":"2026-02-22T04:07:43.831761Z","shell.execute_reply.started":"2026-02-22T04:07:43.819050Z","shell.execute_reply":"2026-02-22T04:07:43.831026Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def sobel(x):\n    kx = torch.tensor([[-1,0,1],[-2,0,2],[-1,0,1]], dtype=x.dtype, device=x.device).view(1,1,3,3)\n    ky = torch.tensor([[-1,-2,-1],[0,0,0],[1,2,1]],  dtype=x.dtype, device=x.device).view(1,1,3,3)\n    B, C, H, W = x.shape\n    xf = x.view(B*C, 1, H, W)\n    gx = F.conv2d(xf, kx, padding=1).view(B, C, H, W)\n    gy = F.conv2d(xf, ky, padding=1).view(B, C, H, W)\n    return gx, gy\n\ndef gradient_losses(pred, target):\n    pgx, pgy = sobel(pred);  tgx, tgy = sobel(target)\n    p_mag = torch.sqrt(pgx**2 + pgy**2 + 1e-6)\n    t_mag = torch.sqrt(tgx**2 + tgy**2 + 1e-6)\n    mag_loss = F.l1_loss(p_mag, t_mag)\n    p_nx = pgx/(p_mag+1e-6); p_ny = pgy/(p_mag+1e-6)\n    t_nx = tgx/(t_mag+1e-6); t_ny = tgy/(t_mag+1e-6)\n    dir_loss = (1 - (p_nx*t_nx + p_ny*t_ny)).mean()\n    return mag_loss, dir_loss\n\ndef smoothness_loss(flow):\n    dy = flow[:,:,1:,:] - flow[:,:,:-1,:]\n    dx = flow[:,:,:,1:] - flow[:,:,:,:-1]\n    return dx.abs().mean() + dy.abs().mean()\n\ndef total_loss(pred, target, flow, cfg):\n    mag, dirn = gradient_losses(pred, target)\n    ss   = 1 - ssim(pred, target, data_range=1.0)\n    l1   = F.l1_loss(pred, target)\n    sm   = smoothness_loss(flow)\n    loss = (cfg[\"w_grad_mag\"]*mag + cfg[\"w_grad_dir\"]*dirn +\n            cfg[\"w_ssim\"]*ss + cfg[\"w_l1\"]*l1 + cfg[\"w_smooth\"]*sm)\n    return loss, {\"mag\": mag.item(), \"dir\": dirn.item(), \"ssim\": ss.item(),\n                  \"l1\": l1.item(), \"smooth\": sm.item()}\n\nprint(\"Losses defined.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T04:07:43.832827Z","iopub.execute_input":"2026-02-22T04:07:43.833117Z","iopub.status.idle":"2026-02-22T04:07:43.844808Z","shell.execute_reply.started":"2026-02-22T04:07:43.833085Z","shell.execute_reply":"2026-02-22T04:07:43.843986Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_pairs(train_root):\n    exts = [\"jpg\", \"jpeg\", \"png\"]\n    pairs = []\n    for ext in exts:\n        for gf in glob.glob(os.path.join(train_root, f'*_generated.{ext}')):\n            of = gf.replace(f'_generated.{ext}', f'_original.{ext}')\n            if os.path.exists(of):\n                pairs.append((gf, of))\n    print(f\"Found {len(pairs)} pairs in {train_root}\")\n    return sorted(pairs)\n\nclass LensPairsDataset(Dataset):\n    def __init__(self, pairs, img_size, augment=False):\n        self.pairs = pairs\n        self.img_size = img_size\n        self.augment = augment\n\n    def __len__(self):\n        return len(self.pairs)\n\n    def __getitem__(self, idx):\n        dist_path, corr_path = self.pairs[idx]\n        x = Image.open(dist_path).convert('RGB').resize((self.img_size, self.img_size), Image.BILINEAR)\n        y = Image.open(corr_path).convert('RGB').resize((self.img_size, self.img_size), Image.BILINEAR)\n        x, y = TF.to_tensor(x), TF.to_tensor(y)\n        if self.augment:\n            if torch.rand(1) > 0.5:\n                f = 0.8 + 0.4*torch.rand(1).item()\n                x = TF.adjust_brightness(x, f); y = TF.adjust_brightness(y, f)\n            if torch.rand(1) > 0.5:\n                f = 0.8 + 0.4*torch.rand(1).item()\n                x = TF.adjust_contrast(x, f); y = TF.adjust_contrast(y, f)\n            if torch.rand(1) > 0.5:\n                f = 0.8 + 0.4*torch.rand(1).item()\n                x = TF.adjust_saturation(x, f); y = TF.adjust_saturation(y, f)\n        return x, y\n\nall_pairs = build_pairs(CFG['train_root'])\nrandom.shuffle(all_pairs)\n\nif CFG.get('max_pairs') and len(all_pairs) > CFG['max_pairs']:\n    all_pairs = all_pairs[:CFG['max_pairs']]\n    print(f\"Capped to {CFG['max_pairs']} pairs for speed\")\n\nn_val = max(1, int(len(all_pairs) * CFG['val_split']))\nval_pairs   = all_pairs[:n_val]\ntrain_pairs = all_pairs[n_val:]\n\ntrain_ds = LensPairsDataset(train_pairs, CFG['img_size'], augment=True)\nval_ds   = LensPairsDataset(val_pairs,   CFG['img_size'], augment=False)\ntrain_dl = DataLoader(train_ds, batch_size=CFG['batch_size'], shuffle=True,\n                      num_workers=2, pin_memory=True, persistent_workers=True)\nval_dl   = DataLoader(val_ds,   batch_size=CFG['batch_size'], shuffle=False,\n                      num_workers=2, pin_memory=True, persistent_workers=True)\nprint(f\"Train: {len(train_ds)}  Val: {len(val_ds)}\")\nprint(f\"Train batches: {len(train_dl)}  Val batches: {len(val_dl)}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T04:07:43.845845Z","iopub.execute_input":"2026-02-22T04:07:43.846129Z","iopub.status.idle":"2026-02-22T04:08:16.305794Z","shell.execute_reply.started":"2026-02-22T04:07:43.846107Z","shell.execute_reply":"2026-02-22T04:08:16.305051Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_base_grid(B, H, W, device, dtype):\n    ys, xs = torch.meshgrid(\n        torch.linspace(-1, 1, H, device=device, dtype=dtype),\n        torch.linspace(-1, 1, W, device=device, dtype=dtype),\n        indexing='ij'\n    )\n    return torch.stack([xs, ys], dim=-1).unsqueeze(0).repeat(B, 1, 1, 1)\n\ndef radial_warp_grid(base_grid, k1, k2, cx, cy):\n    x = base_grid[..., 0] - cx.view(-1,1,1)\n    y = base_grid[..., 1] - cy.view(-1,1,1)\n    r2 = x*x + y*y\n    factor = 1.0 + k1.view(-1,1,1)*r2 + k2.view(-1,1,1)*(r2*r2)\n    x_d = x*factor + cx.view(-1,1,1)\n    y_d = y*factor + cy.view(-1,1,1)\n    return torch.stack([x_d, y_d], dim=-1)\n\nclass DoubleConv(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True),\n        )\n    def forward(self, x): return self.net(x)\n\nclass Down(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.pool = nn.MaxPool2d(2)\n        self.conv = DoubleConv(in_ch, out_ch)\n    def forward(self, x): return self.conv(self.pool(x))\n\nclass Up(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super().__init__()\n        self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=False)\n        self.conv = DoubleConv(in_ch, out_ch)\n    def forward(self, x, skip):\n        x = self.up(x)\n        if x.shape[-2:] != skip.shape[-2:]:\n            x = F.interpolate(x, size=skip.shape[-2:], mode='bilinear', align_corners=False)\n        return self.conv(torch.cat([skip, x], dim=1))\n\nclass HybridWarpNet(nn.Module):\n    def __init__(self, base_ch=32, flow_tanh_scale=0.08):\n        super().__init__()\n        self.flow_tanh_scale = flow_tanh_scale\n        self.enc1 = DoubleConv(3, base_ch)\n        self.enc2 = Down(base_ch,   base_ch*2)\n        self.enc3 = Down(base_ch*2, base_ch*4)\n        self.enc4 = Down(base_ch*4, base_ch*8)\n        self.bottleneck = Down(base_ch*8, base_ch*16)\n        self.up1 = Up(base_ch*16 + base_ch*8, base_ch*8)\n        self.up2 = Up(base_ch*8  + base_ch*4, base_ch*4)\n        self.up3 = Up(base_ch*4  + base_ch*2, base_ch*2)\n        self.up4 = Up(base_ch*2  + base_ch,   base_ch)\n        self.flow_head = nn.Conv2d(base_ch, 2, 1)\n        nn.init.zeros_(self.flow_head.weight)\n        nn.init.zeros_(self.flow_head.bias)\n        self.param_fc = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1), nn.Flatten(),\n            nn.Linear(base_ch*16, 128), nn.ReLU(inplace=True),\n            nn.Linear(128, 4)\n        )\n        nn.init.zeros_(self.param_fc[-1].weight)\n        nn.init.zeros_(self.param_fc[-1].bias)\n\n    def forward(self, x):\n        s1 = self.enc1(x)\n        s2 = self.enc2(s1)\n        s3 = self.enc3(s2)\n        s4 = self.enc4(s3)\n        b  = self.bottleneck(s4)\n        raw = self.param_fc(b)\n        k1 = 0.35 * torch.tanh(raw[:, 0])\n        k2 = 0.35 * torch.tanh(raw[:, 1])\n        cx = 0.30 * torch.tanh(raw[:, 2])\n        cy = 0.30 * torch.tanh(raw[:, 3])\n        x1 = self.up1(b, s4)\n        x2 = self.up2(x1, s3)\n        x3 = self.up3(x2, s2)\n        x4 = self.up4(x3, s1)\n        flow = self.flow_tanh_scale * torch.tanh(self.flow_head(x4))\n        return (k1, k2, cx, cy), flow\n\n    def warp(self, x, params, flow):\n        B, C, H, W = x.shape\n        base = make_base_grid(B, H, W, device=x.device, dtype=x.dtype)\n        k1, k2, cx, cy = params\n        grid_radial = radial_warp_grid(base, k1, k2, cx, cy)\n        x1 = F.grid_sample(x, grid_radial, mode='bilinear', padding_mode='border', align_corners=True)\n        flow_grid = base + flow.permute(0, 2, 3, 1)\n        return F.grid_sample(x1, flow_grid, mode='bilinear', padding_mode='border', align_corners=True)\n\nprint(\"HybridWarpNet defined.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T04:08:16.307363Z","iopub.execute_input":"2026-02-22T04:08:16.307602Z","iopub.status.idle":"2026-02-22T04:08:16.327355Z","shell.execute_reply.started":"2026-02-22T04:08:16.307578Z","shell.execute_reply":"2026-02-22T04:08:16.326509Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = CFG['device']\nmodel = HybridWarpNet(base_ch=32, flow_tanh_scale=CFG['flow_tanh_scale']).to(device)\noptimizer = torch.optim.AdamW(model.parameters(), lr=CFG['lr'])\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG['epochs'])\nscaler = torch.cuda.amp.GradScaler()\n\nbest_val = 1e9\nbad_epochs = 0\nckpt_path = os.path.join(CFG['out_root'], 'best_hybrid.pth')\n\nfor epoch in range(CFG['epochs']):\n    model.train()\n    tr = 0.0\n    for i, (dist, corr) in enumerate(train_dl):\n        dist = dist.to(device, non_blocking=True)\n        corr = corr.to(device, non_blocking=True)\n        optimizer.zero_grad(set_to_none=True)\n\n        with torch.cuda.amp.autocast():\n            params, flow = model(dist)\n            pred = model.warp(dist, params, flow)\n            loss, stats = total_loss(pred, corr, flow, CFG)\n\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), CFG['grad_clip'])\n        scaler.step(optimizer)\n        scaler.update()\n\n        tr += loss.item()\n        if (i+1) % 50 == 0:\n            print(f\"  Epoch {epoch+1}/{CFG['epochs']} step {i+1}/{len(train_dl)} \"\n                  f\"loss {loss.item():.4f} mag {stats['mag']:.4f} \"\n                  f\"dir {stats['dir']:.4f} smooth {stats['smooth']:.4f}\")\n\n    scheduler.step()\n    tr /= max(1, len(train_dl))\n\n    model.eval()\n    va = 0.0\n    with torch.no_grad():\n        for dist, corr in val_dl:\n            dist = dist.to(device, non_blocking=True)\n            corr = corr.to(device, non_blocking=True)\n            with torch.cuda.amp.autocast():\n                params, flow = model(dist)\n                pred = model.warp(dist, params, flow)\n                loss, _ = total_loss(pred, corr, flow, CFG)\n            va += loss.item()\n    va /= max(1, len(val_dl))\n    print(f\"Epoch {epoch+1}/{CFG['epochs']} train {tr:.4f} val {va:.4f}\")\n\n    if va < best_val - 1e-4:\n        best_val = va\n        bad_epochs = 0\n        torch.save(model.state_dict(), ckpt_path)\n        print(f\"  ✓ Saved best: val {best_val:.4f}\")\n    else:\n        bad_epochs += 1\n        print(f\"  No improvement ({bad_epochs}/{CFG['patience']})\")\n        if bad_epochs >= CFG['patience']:\n            print(\"Early stopping triggered.\")\n            break\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T04:08:16.328456Z","iopub.execute_input":"2026-02-22T04:08:16.328876Z","iopub.status.idle":"2026-02-22T05:40:08.321623Z","shell.execute_reply.started":"2026-02-22T04:08:16.328842Z","shell.execute_reply":"2026-02-22T05:40:08.320553Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_state_dict(torch.load(ckpt_path, map_location=device))\nmodel.eval()\n\ntest_files = sorted(\n    glob.glob(os.path.join(CFG['test_root'], '*.jpg')) +\n    glob.glob(os.path.join(CFG['test_root'], '*.png'))\n)\nprint(f\"Test files: {len(test_files)}\")\n\nout_dir = os.path.join(CFG['out_root'], 'corrected_test')\nos.makedirs(out_dir, exist_ok=True)\n\n@torch.no_grad()\ndef correct_one(path):\n    img = Image.open(path).convert('RGB')\n    W0, H0 = img.size\n    inp = TF.to_tensor(\n        img.resize((CFG['img_size'], CFG['img_size']), Image.BILINEAR)\n    ).unsqueeze(0).to(device)\n    with torch.cuda.amp.autocast():\n        params, flow = model(inp)\n        pred = model.warp(inp, params, flow).clamp(0, 1)\n    out = TF.to_pil_image(pred.squeeze(0).cpu())\n    return out.resize((W0, H0), Image.BILINEAR)\n\nfor p in test_files:\n    correct_one(p).save(os.path.join(out_dir, os.path.basename(p)), quality=95)\n\nzip_path = os.path.join(CFG['out_root'], 'submission.zip')\nwith zipfile.ZipFile(zip_path, 'w', compression=zipfile.ZIP_DEFLATED) as z:\n    for f in sorted(glob.glob(os.path.join(out_dir, '*.jpg')) +\n                    glob.glob(os.path.join(out_dir, '*.png'))):\n        z.write(f, arcname=os.path.basename(f))\n\nprint(f\"Done. Submission zip: {zip_path}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-22T05:40:08.323823Z","iopub.execute_input":"2026-02-22T05:40:08.324156Z","iopub.status.idle":"2026-02-22T05:42:29.424405Z","shell.execute_reply.started":"2026-02-22T05:40:08.324096Z","shell.execute_reply":"2026-02-22T05:42:29.423791Z"}},"outputs":[],"execution_count":null}]}