{"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,              # Higher LR because we are only training the head\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    # 'output_dir' is no longer needed for file storage, \n    # but we keep the key if you want to reference working directory\n    'working_dir': '/kaggle/working/', \n    'model_path': '/kaggle/working/effnet_lens_model.pth'\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            \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            \n            # EfficientNet expects normalized data roughly in this range\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            \n            # Normalize for ImageNet stats\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 # Return raw input for warping\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            \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            \n            filename = os.path.basename(path)\n            return normalize(img_tensor), filename, h, w, path\n\n# ==========================================\n# 3. EFFICIENTNET MODEL\n# ==========================================\nclass EfficientNetLensModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        \n        # Load Pretrained EfficientNet\n        # We use the 'features' part as the Encoder\n        base_model = models.efficientnet_b0(weights=models.EfficientNet_B0_Weights.IMAGENET1K_V1)\n        self.encoder = base_model.features\n        \n        # FREEZE THE ENCODER\n        for param in self.encoder.parameters():\n            param.requires_grad = False\n            \n        # EfficientNet B0 output channels is 1280 at the final layer\n        self.decoder = nn.Sequential(\n            # Input: 1280 x 12 x 12 (at 384 input)\n            nn.ConvTranspose2d(1280, 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.ConvTranspose2d(64, 32, kernel_size=2, stride=2),\n            nn.BatchNorm2d(32), nn.ReLU(),\n            \n            nn.Conv2d(32, 2, kernel_size=3, padding=1) # Output Flow\n        )\n        \n        # Initialize last layer to zero for identity mapping\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        # Pass through Frozen Encoder\n        x = self.encoder(x)\n        \n        # Pass through Trainable Decoder\n        flow = self.decoder(x)\n        \n        # Final upsample to ensure exact match\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    \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    \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 = EfficientNetLensModel().to(CONFIG['device'])\n    # Only optimize parameters that require grad (The Decoder)\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 EfficientNet 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 the 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    zip_output_path = os.path.join(CONFIG['working_dir'], 'submission_effnet.zip')\n    print(f\"Running Inference and streaming to {zip_output_path}...\")\n    \n    # Open Zip file directly. Images are written to memory and then to zip,\n    # ensuring no individual image files are left on the disk.\n    with zipfile.ZipFile(zip_output_path, 'w', zipfile.ZIP_DEFLATED) 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 correction\n                corrected_img = apply_high_res_warp(raw_path[0], flow)\n                \n                # Encode image to memory buffer (JPG)\n                success, buffer = cv2.imencode('.jpg', corrected_img)\n                \n                if success:\n                    # Write buffer directly to zip file\n                    # filename is a tuple from dataloader, so use filename[0]\n                    zipf.writestr(filename[0], buffer.tobytes())\n\n    print(\"Inference complete.\")\n    print(f\"Files saved: {CONFIG['model_path']} and {zip_output_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-22T16:22:17.406757Z","iopub.execute_input":"2026-02-22T16:22:17.407507Z","execution_failed":"2026-02-22T16:24:53.730Z"}},"outputs":[],"execution_count":null}]}