{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":130932,"databundleVersionId":15769099}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport glob\nimport cv2\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models\nfrom tqdm import tqdm\nimport zipfile\n\n# ==========================================\n# 1. CONFIGURATION\n# ==========================================\nCONFIG = {\n    'img_size': 384,\n    'batch_size': 16,\n    'epochs': 2,   # 5\n    'lr': 1e-3,\n    'device': 'cuda' if torch.cuda.is_available() else 'cpu',\n    'num_workers': 4,\n    'train_dir': '/kaggle/input/automatic-lens-correction/lens-correction-train-cleaned',\n    'test_dir': '/kaggle/input/automatic-lens-correction/test-originals',\n    'model_path': 'resnet_lens_model.pth',\n    'submission_path': 'submission_resnet.zip'\n}\n\n# ==========================================\n# 2. DATASET\n# ==========================================\nclass LensDataset(Dataset):\n    def __init__(self, root_dir, img_size, is_train=True):\n        self.img_size = img_size\n        self.is_train = is_train\n        self.pairs = []\n        if is_train:\n            originals = glob.glob(os.path.join(root_dir, \"*_original.jpg\"))\n            for org in originals:\n                gen = org.replace(\"_original.jpg\", \"_generated.jpg\")\n                if os.path.exists(gen):\n                    self.pairs.append((org, gen))\n        else:\n            self.pairs = glob.glob(os.path.join(root_dir, \"*.jpg\"))\n\n    def __len__(self):\n        return len(self.pairs)\n\n    def __getitem__(self, idx):\n        if self.is_train:\n            org_path, gen_path = self.pairs[idx]\n            img_in = cv2.cvtColor(cv2.imread(org_path), cv2.COLOR_BGR2RGB)\n            img_gt = cv2.cvtColor(cv2.imread(gen_path), cv2.COLOR_BGR2RGB)\n            img_in = cv2.resize(img_in, (self.img_size, self.img_size))\n            img_gt = cv2.resize(img_gt, (self.img_size, self.img_size))\n            img_in = torch.from_numpy(img_in).permute(2, 0, 1).float() / 255.0\n            img_gt = torch.from_numpy(img_gt).permute(2, 0, 1).float() / 255.0\n            normalize = transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n            return normalize(img_in), img_gt, img_in \n        else:\n            path = self.pairs[idx]\n            img_in = cv2.imread(path)\n            h, w = img_in.shape[:2]\n            img_in_rgb = cv2.cvtColor(img_in, cv2.COLOR_BGR2RGB)\n            img_resized = cv2.resize(img_in_rgb, (self.img_size, self.img_size))\n            img_tensor = torch.from_numpy(img_resized).permute(2, 0, 1).float() / 255.0\n            normalize = transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n            filename = os.path.basename(path)\n            return normalize(img_tensor), filename, h, w, path\n\n# ==========================================\n# 3. RESNET MODEL\n# ==========================================\nclass ResNetLensModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        \n        # Load Pretrained ResNet50\n        base_model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1)\n        \n        # Extract layers. We drop 'avgpool' and 'fc'\n        self.encoder = nn.Sequential(\n            base_model.conv1,\n            base_model.bn1,\n            base_model.relu,\n            base_model.maxpool,\n            base_model.layer1,\n            base_model.layer2,\n            base_model.layer3,\n            base_model.layer4\n        )\n        \n        # FREEZE THE ENCODER\n        for param in self.encoder.parameters():\n            param.requires_grad = False\n            \n        # ResNet50 Layer4 output is 2048 channels\n        self.decoder = nn.Sequential(\n            nn.ConvTranspose2d(2048, 1024, kernel_size=2, stride=2), \n            nn.BatchNorm2d(1024), nn.ReLU(),\n            \n            nn.ConvTranspose2d(1024, 512, kernel_size=2, stride=2),\n            nn.BatchNorm2d(512), nn.ReLU(),\n            \n            nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2),\n            nn.BatchNorm2d(256), nn.ReLU(),\n            \n            nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2),\n            nn.BatchNorm2d(128), nn.ReLU(),\n            \n            nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2),\n            nn.BatchNorm2d(64), nn.ReLU(),\n            \n            nn.Conv2d(64, 2, kernel_size=3, padding=1) # Flow output\n        )\n        \n        nn.init.constant_(self.decoder[-1].weight, 0)\n        nn.init.constant_(self.decoder[-1].bias, 0)\n\n    def forward(self, x):\n        x = self.encoder(x)\n        flow = self.decoder(x)\n        flow = F.interpolate(flow, size=(CONFIG['img_size'], CONFIG['img_size']), mode='bilinear', align_corners=False)\n        return flow\n\ndef warp_image(x, flow):\n    B, C, H, W = x.size()\n    xx = torch.linspace(-1.0, 1.0, W).view(1, 1, 1, W).expand(B, -1, H, -1)\n    yy = torch.linspace(-1.0, 1.0, H).view(1, 1, H, 1).expand(B, -1, -1, W)\n    grid = torch.cat([xx, yy], 1).to(x.device)\n    sampling_grid = (grid + flow).permute(0, 2, 3, 1)\n    return F.grid_sample(x, sampling_grid, align_corners=True, padding_mode='border')\n\ndef apply_high_res_warp(img_path, flow_tensor):\n    img_cv = cv2.imread(img_path)\n    h_orig, w_orig = img_cv.shape[:2]\n    flow_np = flow_tensor.detach().cpu().numpy()[0].transpose(1, 2, 0)\n    flow_resized = cv2.resize(flow_np, (w_orig, h_orig), interpolation=cv2.INTER_LINEAR)\n    grid_x, grid_y = np.meshgrid(np.arange(w_orig), np.arange(h_orig))\n    map_x = ((2.0 * grid_x / (w_orig - 1) - 1.0 + flow_resized[:, :, 0] + 1.0) / 2.0) * (w_orig - 1)\n    map_y = ((2.0 * grid_y / (h_orig - 1) - 1.0 + flow_resized[:, :, 1] + 1.0) / 2.0) * (h_orig - 1)\n    return cv2.remap(img_cv, map_x.astype(np.float32), map_y.astype(np.float32), cv2.INTER_LINEAR)\n\n# ==========================================\n# 4. TRAINING & INFERENCE\n# ==========================================\ndef train():\n    dataset = LensDataset(CONFIG['train_dir'], CONFIG['img_size'], is_train=True)\n    dataloader = DataLoader(dataset, batch_size=CONFIG['batch_size'], shuffle=True, num_workers=CONFIG['num_workers'])\n    \n    model = ResNetLensModel().to(CONFIG['device'])\n    optimizer = torch.optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=CONFIG['lr'])\n    criterion = nn.L1Loss()\n    \n    print(f\"Training with Frozen ResNet50 Backbone...\")\n    \n    for epoch in range(CONFIG['epochs']):\n        model.train()\n        loop = tqdm(dataloader)\n        for batch_idx, (img_norm, img_gt, img_raw) in enumerate(loop):\n            img_norm, img_gt, img_raw = img_norm.to(CONFIG['device']), img_gt.to(CONFIG['device']), img_raw.to(CONFIG['device'])\n            \n            optimizer.zero_grad()\n            flow = model(img_norm)\n            warped = warp_image(img_raw, flow)\n            \n            loss = criterion(warped, img_gt)\n            loss.backward()\n            optimizer.step()\n            \n            loop.set_description(f\"Epoch {epoch+1}\")\n            loop.set_postfix(loss=loss.item())\n            \n    # Save Model\n    torch.save(model.state_dict(), CONFIG['model_path'])\n    print(f\"Model saved to {CONFIG['model_path']}\")\n    return model\n\ndef inference(model):\n    model.eval()\n    test_dataset = LensDataset(CONFIG['test_dir'], CONFIG['img_size'], is_train=False)\n    test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\n    \n    print(\"Running Inference and writing directly to ZIP...\")\n    \n    # Open ZipFile to write results directly\n    with zipfile.ZipFile(CONFIG['submission_path'], 'w') as zipf:\n        with torch.no_grad():\n            for img_norm, filename, h, w, raw_path in tqdm(test_loader):\n                img_norm = img_norm.to(CONFIG['device'])\n                \n                # Get flow\n                flow = model(img_norm)\n                \n                # Apply high-res correction\n                corrected_img = apply_high_res_warp(raw_path[0], flow)\n                \n                # Encode image to buffer (JPEG)\n                success, buffer = cv2.imencode('.jpg', corrected_img)\n                \n                # Write buffer to zip file if encoding succeeded\n                if success:\n                    zipf.writestr(filename[0], buffer.tobytes())\n                    \n    print(f\"Saved {CONFIG['submission_path']}\")\n\nif __name__ == \"__main__\":\n    model = train()\n    inference(model)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-22T17:22:59.396025Z","iopub.execute_input":"2026-02-22T17:22:59.396390Z","execution_failed":"2026-02-22T17:23:44.301Z"}},"outputs":[],"execution_count":null}]}