{"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":31286,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-22T22:48:04.178502Z","iopub.execute_input":"2026-02-22T22:48:04.178743Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 1 — Imports, Paths, Config\n# ============================================================================\nimport os, sys, glob, time, random, zipfile, gc, math, csv\nfrom pathlib import Path\nfrom collections import defaultdict\nimport warnings; warnings.filterwarnings('ignore')\nimport numpy as np\nimport cv2\nfrom PIL import Image\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport torchvision.models as tv_models\ntry:\n    from pytorch_msssim import SSIM\nexcept ImportError:\n    os.system('pip install -q pytorch-msssim')\n    from pytorch_msssim import SSIM\n\nDATA_ROOT  = '/kaggle/input/automatic-lens-correction'\nTRAIN_DIR  = os.path.join(DATA_ROOT, 'lens-correction-train-cleaned')\nTEST_DIR   = os.path.join(DATA_ROOT, 'test-originals')\nOUTPUT_DIR = '/kaggle/working/corrected'\nCKPT_DIR   = '/kaggle/working/checkpoints'\nos.makedirs(OUTPUT_DIR, exist_ok=True)\nos.makedirs(CKPT_DIR, exist_ok=True)\n\nprint('TRAIN_DIR:', TRAIN_DIR, 'exists:', os.path.isdir(TRAIN_DIR))\nprint('TEST_DIR :', TEST_DIR,  'exists:', os.path.isdir(TEST_DIR))\n\nclass CFG:\n    seed            = 42\n    device          = 'cuda' if torch.cuda.is_available() else 'cpu'\n\n    # Architecture — WIDER bounds, HIGHER flow res\n    backbone        = 'resnet18'\n    num_k           = 3\n    predict_tang    = True\n    predict_center  = True\n    predict_zoom    = True\n    residual_res    = 64           # was 32 — 4x more flow detail\n    residual_lambda = 0.15         # was 0.1 — let flow contribute more\n    reflect_pad     = 64\n\n    # WIDER parameter bounds\n    k_bounds        = [1.0, 0.5, 0.25]   # was [0.6, 0.3, 0.15]\n    tang_bound      = 0.05               # was 0.02\n    center_bound    = 0.10               # was 0.05\n    zoom_min        = 0.90               # was 1.02\n    zoom_range      = 0.35               # s = zoom_min + sigmoid * zoom_range\n\n    # Progressive training phases\n    # Phase A: 128px, frozen backbone, parametric only\n    phaseA_size     = 128\n    phaseA_epochs   = 6\n    phaseA_batch    = 64\n    phaseA_lr       = 8e-4\n    phaseA_workers  = 4\n    phaseA_flow     = False\n\n    # Phase B: 256px, full backbone, parametric + flow\n    phaseB_size     = 256\n    phaseB_epochs   = 6\n    phaseB_batch    = 16\n    phaseB_lr       = 2e-4\n    phaseB_workers  = 4\n    phaseB_flow     = True\n\n    # Phase C: 320px, fine-tune everything\n    phaseC_size     = 320\n    phaseC_epochs   = 3\n    phaseC_batch    = 8\n    phaseC_lr       = 5e-5\n    phaseC_workers  = 4\n    phaseC_flow     = True\n\n    # Loss weights — COMPETITION ALIGNED\n    # Scoring: edge 40%, line 22%, gradient 18%, SSIM 15%, pixel 5%\n    w_charb         = 0.5          # pixel (5% weight in scoring)\n    w_ssim          = 1.0          # SSIM (15% weight)\n    w_grad          = 1.2          # gradient orientation (18% weight)\n    w_edge          = 2.0          # edge similarity (40% weight) — DOMINANT\n    w_border        = 0.01\n    w_tv            = 0.005\n    w_mag           = 0.003\n    w_jacobian      = 0.001\n    w_line          = 0.8          # NEW: line straightness loss (22% weight)\n\n    # Validation\n    val_ratio       = 0.02\n    max_train_pairs = 23000\n\n    # Inference\n    encoder_size    = 320\n    tta_scales      = [256, 320]\n    infer_interp    = 'bicubic'\n\ndef seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = False\n    torch.backends.cudnn.benchmark = True\n\nseed_everything(CFG.seed)\nprint(f'Device: {CFG.device}')\nprint('=== CELL 1 OK ===')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 2 — Pair Discovery\n# ============================================================================\n\ndef discover_train_pairs(train_dir, max_pairs=None):\n    train_dir = Path(train_dir)\n    pairs = []\n    for o in sorted(train_dir.glob('*_original.jpg')):\n        g = Path(str(o).replace('_original.jpg', '_generated.jpg'))\n        if g.exists():\n            pairs.append((str(o), str(g)))\n            if max_pairs and len(pairs) >= max_pairs:\n                break\n    random.shuffle(pairs)\n    return pairs\n\nn_total = len(list(Path(TRAIN_DIR).glob('*_original.jpg')))\nprint(f'Total pairs: {n_total}')\nprint('=== CELL 2 OK ===')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 3 — LensDataset + Helpers\n# ============================================================================\n\ndef resize_preserve_aspect(img, longest_side):\n    h, w = img.shape[:2]\n    if h >= w:\n        new_h = longest_side\n        new_w = int(round(w * longest_side / h))\n    else:\n        new_w = longest_side\n        new_h = int(round(h * longest_side / w))\n    return cv2.resize(img, (max(new_w, 2), max(new_h, 2)), interpolation=cv2.INTER_AREA)\n\ndef pad_to_square(img, target_size):\n    h, w = img.shape[:2]\n    pad_top   = (target_size - h) // 2\n    pad_bot   = target_size - h - pad_top\n    pad_left  = (target_size - w) // 2\n    pad_right = target_size - w - pad_left\n    padded = cv2.copyMakeBorder(img, pad_top, pad_bot, pad_left, pad_right,\n                                 cv2.BORDER_REFLECT_101)\n    return padded, pad_top, pad_left\n\nclass LensDataset(Dataset):\n    def __init__(self, pairs, size=256, is_train=True):\n        self.pairs = pairs\n        self.size = size\n        self.is_train = is_train\n\n    def __len__(self):\n        return len(self.pairs)\n\n    def _aug(self, img):\n        if not self.is_train: return img\n        img = img.astype(np.float32)\n        if random.random() < 0.5:\n            img = np.clip(img + random.uniform(-30, 30), 0, 255)\n        if random.random() < 0.5:\n            factor = random.uniform(0.7, 1.3)\n            img = np.clip((img - img.mean()) * factor + img.mean(), 0, 255)\n        if random.random() < 0.3:\n            gamma = random.uniform(0.8, 1.2)\n            img = np.clip(255.0 * (img / 255.0) ** gamma, 0, 255)\n        return img.astype(np.uint8)\n\n    def __getitem__(self, idx):\n        dist_path, gt_path = self.pairs[idx]\n        dist_bgr = cv2.imread(dist_path, cv2.IMREAD_COLOR)\n        gt_bgr   = cv2.imread(gt_path,   cv2.IMREAD_COLOR)\n        if dist_bgr is None or gt_bgr is None:\n            s = self.size\n            return torch.zeros(3, s, s), torch.zeros(3, s, s)\n\n        dist = cv2.cvtColor(dist_bgr, cv2.COLOR_BGR2RGB)\n        gt   = cv2.cvtColor(gt_bgr,   cv2.COLOR_BGR2RGB)\n        dist = resize_preserve_aspect(dist, self.size)\n        gt   = resize_preserve_aspect(gt,   self.size)\n\n        if self.is_train:\n            py_s = random.getstate(); np_s = np.random.get_state()\n            dist = self._aug(dist)\n            random.setstate(py_s); np.random.set_state(np_s)\n            gt = self._aug(gt)\n            # Random horizontal flip (both images)\n            if random.random() < 0.5:\n                dist = np.ascontiguousarray(dist[:, ::-1])\n                gt   = np.ascontiguousarray(gt[:, ::-1])\n\n        dist, _, _ = pad_to_square(dist, self.size)\n        gt,   _, _ = pad_to_square(gt,   self.size)\n        dist_t = torch.from_numpy(dist.astype(np.float32) / 255.0).permute(2, 0, 1)\n        gt_t   = torch.from_numpy(gt.astype(np.float32)   / 255.0).permute(2, 0, 1)\n        return dist_t, gt_t\n\n# Smoke test\n_tp = discover_train_pairs(TRAIN_DIR, max_pairs=2)\nif _tp:\n    _ds = LensDataset(_tp, size=128, is_train=False)\n    _d, _g = _ds[0]\n    print(f'Smoke: {_d.shape} range=[{_d.min():.2f}, {_d.max():.2f}]')\n    del _ds, _d, _g, _tp\nprint('=== CELL 3 OK ===')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 4 — Parametric Grid Builder\n# ============================================================================\n\ndef build_parametric_grid(H, W, k1, k2, k3, p1, p2, cx_off, cy_off, s, device=None):\n    B = k1.shape[0]\n    if device is None: device = k1.device\n\n    ys = torch.linspace(-1.0, 1.0, H, device=device)\n    xs = torch.linspace(-1.0, 1.0, W, device=device)\n    grid_y, grid_x = torch.meshgrid(ys, xs, indexing='ij')\n    gx = grid_x.unsqueeze(0).expand(B, -1, -1)\n    gy = grid_y.unsqueeze(0).expand(B, -1, -1)\n\n    s_ = s.view(B, 1, 1)\n    gx = gx / (s_ + 1e-8)\n    gy = gy / (s_ + 1e-8)\n\n    cx_ = cx_off.view(B, 1, 1); cy_ = cy_off.view(B, 1, 1)\n    gx_c = gx - cx_; gy_c = gy - cy_\n    r2 = gx_c**2 + gy_c**2\n\n    k1_ = k1.view(B,1,1); k2_ = k2.view(B,1,1); k3_ = k3.view(B,1,1)\n    p1_ = p1.view(B,1,1); p2_ = p2.view(B,1,1)\n\n    radial = 1.0 + k1_*r2 + k2_*(r2**2) + k3_*(r2**3)\n    gx_d = gx_c*radial + 2.0*p1_*gx_c*gy_c + p2_*(r2 + 2.0*gx_c**2)\n    gy_d = gy_c*radial + p1_*(r2 + 2.0*gy_c**2) + 2.0*p2_*gx_c*gy_c\n    gx_d = gx_d + cx_; gy_d = gy_d + cy_\n\n    return torch.stack([gx_d, gy_d], dim=-1)\n\ndef adjust_grid_for_padding(grid, H_orig, W_orig, pad, device=None):\n    if device is None: device = grid.device\n    H_pad = H_orig + 2*pad; W_pad = W_orig + 2*pad\n    pix_x = (grid[...,0]+1.0)*(W_orig-1)/2.0 + pad\n    norm_x = 2.0*pix_x/(W_pad-1) - 1.0\n    pix_y = (grid[...,1]+1.0)*(H_orig-1)/2.0 + pad\n    norm_y = 2.0*pix_y/(H_pad-1) - 1.0\n    return torch.stack([norm_x, norm_y], dim=-1)\n\nprint('=== CELL 4 OK ===')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 5 — WarpNet with wider bounds + higher-res flow\n# ============================================================================\n\nclass WarpNet(nn.Module):\n    def __init__(self, cfg=CFG):\n        super().__init__()\n        self.cfg = cfg\n\n        base = tv_models.resnet18(weights=tv_models.ResNet18_Weights.DEFAULT)\n        self.layer0 = nn.Sequential(base.conv1, base.bn1, base.relu, base.maxpool)\n        self.layer1 = base.layer1\n        self.layer2 = base.layer2\n        self.layer3 = base.layer3\n        self.layer4 = base.layer4\n        self.pool   = nn.Sequential(base.avgpool, nn.Flatten())\n        feat_dim = 512\n\n        n_params = cfg.num_k\n        if cfg.predict_tang:   n_params += 2\n        if cfg.predict_center: n_params += 2\n        if cfg.predict_zoom:   n_params += 1\n        self.n_params = n_params\n\n        # Deeper param head for better parameter prediction\n        self.param_head = nn.Sequential(\n            nn.Linear(feat_dim, 256), nn.GELU(), nn.Dropout(0.1),\n            nn.Linear(256, 128), nn.GELU(), nn.Dropout(0.05),\n            nn.Linear(128, 64), nn.GELU(),\n            nn.Linear(64, n_params),\n        )\n        nn.init.zeros_(self.param_head[-1].weight)\n        nn.init.zeros_(self.param_head[-1].bias)\n\n        # Higher-res flow head with conv decoder\n        R = cfg.residual_res  # 64\n        self.flow_head = nn.Sequential(\n            nn.Linear(feat_dim, 1024), nn.GELU(), nn.Dropout(0.1),\n            nn.Linear(1024, 2 * R * R),\n        )\n        nn.init.zeros_(self.flow_head[-1].weight)\n        nn.init.zeros_(self.flow_head[-1].bias)\n\n        self.residual_lambda = cfg.residual_lambda\n        self._flow_enabled = False\n\n    def encode(self, x):\n        x = self.layer0(x)\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n        return self.pool(x)\n\n    def freeze_early(self):\n        for m in [self.layer0, self.layer1, self.layer2]:\n            for p in m.parameters(): p.requires_grad = False\n        print('Frozen: layer0-2')\n\n    def unfreeze_all(self):\n        for p in self.parameters(): p.requires_grad = True\n        if not self._flow_enabled:\n            for p in self.flow_head.parameters(): p.requires_grad = False\n        print('Unfrozen all')\n\n    def enable_flow(self):\n        self._flow_enabled = True\n        for p in self.flow_head.parameters(): p.requires_grad = True\n\n    def disable_flow(self):\n        self._flow_enabled = False\n        for p in self.flow_head.parameters(): p.requires_grad = False\n\n    def parse_params(self, raw):\n        cfg = self.cfg; idx = 0; B = raw.shape[0]; dev = raw.device\n        ks = []\n        for i in range(cfg.num_k):\n            ks.append(torch.tanh(raw[:, idx]) * cfg.k_bounds[i]); idx += 1\n        while len(ks) < 3: ks.append(torch.zeros(B, device=dev))\n\n        if cfg.predict_tang:\n            p1 = torch.tanh(raw[:, idx]) * cfg.tang_bound; idx += 1\n            p2 = torch.tanh(raw[:, idx]) * cfg.tang_bound; idx += 1\n        else:\n            p1 = p2 = torch.zeros(B, device=dev)\n\n        if cfg.predict_center:\n            cx = torch.tanh(raw[:, idx]) * cfg.center_bound; idx += 1\n            cy = torch.tanh(raw[:, idx]) * cfg.center_bound; idx += 1\n        else:\n            cx = cy = torch.zeros(B, device=dev)\n\n        if cfg.predict_zoom:\n            s = cfg.zoom_min + torch.sigmoid(raw[:, idx]) * cfg.zoom_range; idx += 1\n        else:\n            s = torch.ones(B, device=dev)\n\n        return ks[0], ks[1], ks[2], p1, p2, cx, cy, s\n\n    def forward_params(self, x):\n        feat = self.encode(x)\n        raw = self.param_head(feat)\n        k1, k2, k3, p1, p2, cx, cy, s = self.parse_params(raw)\n        flow_lr = None\n        if self._flow_enabled:\n            R = self.cfg.residual_res\n            flow_lr = self.flow_head(feat).view(-1, 2, R, R)\n        return k1, k2, k3, p1, p2, cx, cy, s, flow_lr\n\n    def forward(self, x, apply_warp=True):\n        B, C, H, W = x.shape\n        pad = self.cfg.reflect_pad\n        k1, k2, k3, p1, p2, cx, cy, s, flow_lr = self.forward_params(x)\n\n        if not apply_warp:\n            return None, {'k1':k1,'k2':k2,'k3':k3,'p1':p1,'p2':p2,\n                          'cx':cx,'cy':cy,'s':s,'flow_lr':flow_lr}\n\n        grid = build_parametric_grid(H, W, k1, k2, k3, p1, p2, cx, cy, s, x.device)\n\n        if self._flow_enabled and flow_lr is not None:\n            flow_hr = F.interpolate(flow_lr, size=(H, W), mode='bilinear', align_corners=True)\n            flow_hr = flow_hr.permute(0, 2, 3, 1)\n            flow_hr = torch.clamp(flow_hr, -0.20, 0.20)\n            grid = grid + self.residual_lambda * flow_hr\n\n        x_padded = F.pad(x, [pad]*4, mode='reflect')\n        grid_padded = adjust_grid_for_padding(grid, H, W, pad, x.device)\n        output = F.grid_sample(x_padded, grid_padded, mode='bilinear',\n                               padding_mode='border', align_corners=True)\n\n        info = {'k1':k1,'k2':k2,'k3':k3,'p1':p1,'p2':p2,'cx':cx,'cy':cy,\n                's':s,'flow_lr':flow_lr,'grid':grid}\n        return output, info\n\n# Smoke test\n_m = WarpNet(CFG).to(CFG.device)\nwith torch.no_grad():\n    _x = torch.rand(2, 3, 128, 128, device=CFG.device)\n    _o, _i = _m(_x)\n    print(f'WarpNet: in={_x.shape} out={_o.shape} k1={_i[\"k1\"][0].item():.4f} s={_i[\"s\"][0].item():.4f}')\n    print(f'Param bounds: k1∈±{CFG.k_bounds[0]}, tang∈±{CFG.tang_bound}, '\n          f'center∈±{CFG.center_bound}, zoom∈[{CFG.zoom_min:.2f}, {CFG.zoom_min+CFG.zoom_range:.2f}]')\ndel _m, _x, _o, _i; torch.cuda.empty_cache()\nprint('=== CELL 5 OK ===')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 6 — Competition-Aligned Loss Functions\n# ============================================================================\n# Scoring: Edge 40%, Line 22%, Gradient 18%, SSIM 15%, Pixel 5%\n\nclass CharbonnierLoss(nn.Module):\n    def __init__(self, eps=1e-3):\n        super().__init__()\n        self.eps_sq = eps ** 2\n    def forward(self, pred, target):\n        return torch.sqrt((pred - target)**2 + self.eps_sq).mean()\n\nclass GradientLoss(nn.Module):\n    \"\"\"L1 on Sobel gradients — targets gradient orientation metric (18%).\"\"\"\n    def __init__(self):\n        super().__init__()\n        sx = torch.tensor([[-1,0,1],[-2,0,2],[-1,0,1]], dtype=torch.float32).view(1,1,3,3)\n        sy = torch.tensor([[-1,-2,-1],[0,0,0],[1,2,1]], dtype=torch.float32).view(1,1,3,3)\n        self.register_buffer('sx', sx)\n        self.register_buffer('sy', sy)\n    def _sobel(self, img):\n        B,C,H,W = img.shape\n        flat = img.reshape(B*C,1,H,W)\n        return (F.conv2d(flat, self.sx, padding=1).reshape(B,C,H,W),\n                F.conv2d(flat, self.sy, padding=1).reshape(B,C,H,W))\n    def forward(self, pred, target):\n        px,py = self._sobel(pred); tx,ty = self._sobel(target)\n        return F.l1_loss(px,tx) + F.l1_loss(py,ty)\n\nclass MultiScaleEdgeLoss(nn.Module):\n    \"\"\"Edge mag + orientation at multiple scales — targets edge similarity (40%).\"\"\"\n    def __init__(self, scales=(1.0, 0.5)):\n        super().__init__()\n        self.scales = scales\n        sx = torch.tensor([[-1,0,1],[-2,0,2],[-1,0,1]], dtype=torch.float32).view(1,1,3,3)\n        sy = torch.tensor([[-1,-2,-1],[0,0,0],[1,2,1]], dtype=torch.float32).view(1,1,3,3)\n        self.register_buffer('sx', sx)\n        self.register_buffer('sy', sy)\n    def _edge(self, img):\n        gray = img.mean(dim=1, keepdim=True)\n        gx = F.conv2d(gray, self.sx, padding=1)\n        gy = F.conv2d(gray, self.sy, padding=1)\n        return torch.sqrt(gx**2 + gy**2 + 1e-8), gx, gy\n    def forward(self, pred, target):\n        total = 0.0\n        for sc in self.scales:\n            if sc < 1.0:\n                p = F.interpolate(pred,   scale_factor=sc, mode='bilinear', align_corners=True)\n                t = F.interpolate(target, scale_factor=sc, mode='bilinear', align_corners=True)\n            else:\n                p, t = pred, target\n            pm,pgx,pgy = self._edge(p); tm,tgx,tgy = self._edge(t)\n            total += F.l1_loss(pm, tm)\n            # Orientation cosine sim on strong edges\n            mask = (tm > tm.mean()).float()\n            dot = pgx*tgx + pgy*tgy\n            pn = torch.sqrt(pgx**2 + pgy**2 + 1e-8)\n            tn = torch.sqrt(tgx**2 + tgy**2 + 1e-8)\n            total += (mask * (1.0 - dot/(pn*tn+1e-8))).mean() * 0.5\n        return total / len(self.scales)\n\nclass LineStraightnessLoss(nn.Module):\n    \"\"\"Penalizes curvature in horizontal/vertical edge responses — targets line straightness (22%).\"\"\"\n    def __init__(self):\n        super().__init__()\n        # Second-order derivative filters to detect curvature\n        # Horizontal curvature: d²I/dy² along strong horizontal edges\n        laplacian = torch.tensor([[0,1,0],[1,-4,1],[0,1,0]], dtype=torch.float32).view(1,1,3,3)\n        self.register_buffer('laplacian', laplacian)\n        sx = torch.tensor([[-1,0,1],[-2,0,2],[-1,0,1]], dtype=torch.float32).view(1,1,3,3)\n        sy = torch.tensor([[-1,-2,-1],[0,0,0],[1,2,1]], dtype=torch.float32).view(1,1,3,3)\n        self.register_buffer('sx', sx)\n        self.register_buffer('sy', sy)\n\n    def forward(self, pred, target):\n        # Work on grayscale\n        pred_g = pred.mean(dim=1, keepdim=True)\n        tgt_g  = target.mean(dim=1, keepdim=True)\n\n        # Detect strong edges in target\n        tgx = F.conv2d(tgt_g, self.sx, padding=1)\n        tgy = F.conv2d(tgt_g, self.sy, padding=1)\n        t_mag = torch.sqrt(tgx**2 + tgy**2 + 1e-8)\n        edge_mask = (t_mag > t_mag.mean() * 1.5).float()\n\n        # Laplacian measures local curvature\n        pred_lap = F.conv2d(pred_g, self.laplacian, padding=1)\n        tgt_lap  = F.conv2d(tgt_g, self.laplacian, padding=1)\n\n        # Match curvature on edges\n        return (edge_mask * (pred_lap - tgt_lap).abs()).mean()\n\nclass TVLoss(nn.Module):\n    def forward(self, flow):\n        return ((flow[:,:,1:,:] - flow[:,:,:-1,:]).abs().mean() +\n                (flow[:,:,:,1:] - flow[:,:,:,:-1]).abs().mean())\n\nclass JacobianPenalty(nn.Module):\n    def forward(self, grid):\n        dxdu = grid[:,:,1:,0] - grid[:,:,:-1,0]\n        dxdv = grid[:,1:,:,0] - grid[:,:-1,:,0]\n        dydu = grid[:,:,1:,1] - grid[:,:,:-1,1]\n        dydv = grid[:,1:,:,1] - grid[:,:-1,:,1]\n        H = min(dxdv.shape[1], dxdu.shape[1])\n        W = min(dxdu.shape[2], dxdv.shape[2])\n        det = dxdu[:,:H,:W]*dydv[:,:H,:W] - dxdv[:,:H,:W]*dydu[:,:H,:W]\n        return F.relu(0.1 - det).mean()\n\nclass BorderPenalty(nn.Module):\n    def __init__(self, bw=12, threshold=0.02):\n        super().__init__()\n        self.bw = bw; self.thr = threshold\n    def forward(self, img):\n        bw = self.bw\n        penalty = 0.0\n        for s in [img[:,:,:bw,:], img[:,:,-bw:,:], img[:,:,:,:bw], img[:,:,:,-bw:]]:\n            penalty += F.relu(self.thr - s.var(dim=[2,3])).mean()\n        return penalty\n\nclass CombinedLoss(nn.Module):\n    def __init__(self, cfg=CFG, use_flow_losses=False):\n        super().__init__()\n        self.cfg = cfg\n        self.use_flow = use_flow_losses\n        self.charb  = CharbonnierLoss()\n        self.ssim   = SSIM(data_range=1.0, size_average=True, channel=3)\n        self.grad   = GradientLoss()\n        self.edge   = MultiScaleEdgeLoss()\n        self.line   = LineStraightnessLoss()\n        self.border = BorderPenalty()\n        self.tv     = TVLoss()\n        self.jac    = JacobianPenalty()\n\n    def forward(self, pred, target, info=None):\n        c = self.cfg; losses = {}\n        losses['charb'] = self.charb(pred, target)\n        losses['ssim']  = 1.0 - self.ssim(pred.clamp(0,1), target)\n        losses['grad']  = self.grad(pred, target)\n        losses['edge']  = self.edge(pred, target)\n        losses['line']  = self.line(pred, target)\n        losses['border'] = self.border(pred)\n\n        total = (c.w_charb  * losses['charb'] +\n                 c.w_ssim   * losses['ssim']  +\n                 c.w_grad   * losses['grad']  +\n                 c.w_edge   * losses['edge']  +\n                 c.w_line   * losses['line']  +\n                 c.w_border * losses['border'])\n\n        if self.use_flow and info is not None:\n            if info.get('flow_lr') is not None:\n                losses['tv']  = self.tv(info['flow_lr'])\n                losses['mag'] = info['flow_lr'].abs().mean()\n                total += c.w_tv * losses['tv'] + c.w_mag * losses['mag']\n            if info.get('grid') is not None:\n                losses['jac'] = self.jac(info['grid'])\n                total += c.w_jacobian * losses['jac']\n\n        losses['total'] = total\n        return total, losses\n\nprint('=== CELL 6 OK ===')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 7 — Progressive Training (A: 128px → B: 256px → C: 320px)\n# ============================================================================\n\ndef train_one_epoch(model, loader, criterion, optimizer, scaler, device,\n                    scheduler=None, log_every=400):\n    model.train()\n    running = defaultdict(float); count = 0\n    for bi, (dist, gt) in enumerate(loader):\n        dist = dist.to(device, non_blocking=True)\n        gt   = gt.to(device, non_blocking=True)\n        optimizer.zero_grad(set_to_none=True)\n        with autocast():\n            output, info = model(dist)\n            loss, losses = criterion(output, gt, info)\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0)\n        scaler.step(optimizer)\n        scaler.update()\n        if scheduler: scheduler.step()\n        bs = dist.shape[0]; count += bs\n        for k, v in losses.items():\n            running[k] += float(v.item()) * bs\n        if bi % log_every == 0:\n            print(f'  batch {bi}: loss={running[\"total\"]/max(count,1):.5f}', flush=True)\n    return {k: v/max(count,1) for k, v in running.items()}\n\n@torch.no_grad()\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running = defaultdict(float); count = 0\n    for dist, gt in loader:\n        dist = dist.to(device, non_blocking=True)\n        gt   = gt.to(device, non_blocking=True)\n        with autocast():\n            output, info = model(dist)\n            _, losses = criterion(output, gt, info)\n        bs = dist.shape[0]; count += bs\n        for k, v in losses.items():\n            running[k] += float(v.item()) * bs\n    return {k: v/max(count,1) for k, v in running.items()}\n\ndef run_phase(model, train_pairs, val_pairs, size, epochs, batch, lr, workers,\n              phase_name, ckpt_name, device, use_flow=False, cfg=CFG):\n    print(f'\\n{\"=\"*60}')\n    print(f'{phase_name}: {size}px, {epochs}ep, bs={batch}, lr={lr}, flow={use_flow}')\n    print(f'{\"=\"*60}')\n\n    train_ds = LensDataset(train_pairs, size=size, is_train=True)\n    val_ds   = LensDataset(val_pairs,   size=size, is_train=False)\n    train_dl = DataLoader(train_ds, batch_size=batch, shuffle=True,\n                          num_workers=workers, pin_memory=True, drop_last=True)\n    val_dl   = DataLoader(val_ds, batch_size=max(batch, 8), shuffle=False,\n                          num_workers=workers, pin_memory=True)\n\n    criterion = CombinedLoss(cfg, use_flow_losses=use_flow).to(device)\n    scaler = GradScaler()\n    trainable = [p for p in model.parameters() if p.requires_grad]\n    n_trainable = sum(p.numel() for p in trainable)\n    print(f'Trainable: {n_trainable:,} params, batches/epoch: {len(train_dl)}')\n\n    optimizer = torch.optim.AdamW(trainable, lr=lr, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(\n        optimizer, max_lr=lr, epochs=epochs,\n        steps_per_epoch=len(train_dl), pct_start=0.3)\n\n    best_val = float('inf')\n    ckpt_path = os.path.join(CKPT_DIR, ckpt_name)\n\n    for epoch in range(epochs):\n        t0 = time.time()\n        tm = train_one_epoch(model, train_dl, criterion, optimizer, scaler, device,\n                             scheduler=scheduler)\n        # Validate every epoch (fast at low res)\n        vm = validate(model, val_dl, criterion, device)\n        val_str = f'val={vm[\"total\"]:.5f}'\n        if vm['total'] < best_val:\n            best_val = vm['total']\n            torch.save(model.state_dict(), ckpt_path)\n            val_str += ' *BEST*'\n\n        print(f'Ep {epoch+1}/{epochs} ({time.time()-t0:.0f}s) | '\n              f'train={tm[\"total\"]:.5f} {val_str} '\n              f'ssim={vm.get(\"ssim\",0):.4f} edge={vm.get(\"edge\",0):.4f}')\n\n    if os.path.exists(ckpt_path):\n        model.load_state_dict(torch.load(ckpt_path, map_location=device))\n    print(f'{phase_name} done. Best val={best_val:.5f}')\n    torch.cuda.empty_cache()\n    return model\n\n\ndef run_training():\n    all_pairs = discover_train_pairs(TRAIN_DIR, max_pairs=CFG.max_train_pairs)\n    if not all_pairs:\n        raise RuntimeError(f'No pairs in {TRAIN_DIR}')\n\n    n_val = max(int(len(all_pairs) * CFG.val_ratio), 50)\n    val_pairs   = all_pairs[:n_val]\n    train_pairs = all_pairs[n_val:]\n    print(f'Train: {len(train_pairs)}, Val: {len(val_pairs)}')\n\n    model = WarpNet(CFG).to(CFG.device)\n\n    # Resume if checkpoint exists\n    for p in ['best_phaseC.pth', 'best_phaseB.pth', 'best_phaseA.pth']:\n        ckpt = os.path.join(CKPT_DIR, p)\n        if os.path.exists(ckpt):\n            print(f'Resuming from {ckpt}')\n            model.load_state_dict(torch.load(ckpt, map_location=CFG.device))\n            break\n\n    # --- PHASE A: 128px, frozen backbone, parametric only ---\n    model.disable_flow()\n    model.freeze_early()\n    model = run_phase(model, train_pairs, val_pairs,\n        size=CFG.phaseA_size, epochs=CFG.phaseA_epochs,\n        batch=CFG.phaseA_batch, lr=CFG.phaseA_lr,\n        workers=CFG.phaseA_workers,\n        phase_name='PHASE A (128px parametric)', ckpt_name='best_phaseA.pth',\n        device=CFG.device, use_flow=False)\n\n    # --- PHASE B: 256px, all layers, parametric + flow ---\n    model.unfreeze_all()\n    model.enable_flow()\n    model = run_phase(model, train_pairs, val_pairs,\n        size=CFG.phaseB_size, epochs=CFG.phaseB_epochs,\n        batch=CFG.phaseB_batch, lr=CFG.phaseB_lr,\n        workers=CFG.phaseB_workers,\n        phase_name='PHASE B (256px + flow)', ckpt_name='best_phaseB.pth',\n        device=CFG.device, use_flow=True)\n\n    # --- PHASE C: 320px, fine-tune ---\n    model = run_phase(model, train_pairs, val_pairs,\n        size=CFG.phaseC_size, epochs=CFG.phaseC_epochs,\n        batch=CFG.phaseC_batch, lr=CFG.phaseC_lr,\n        workers=CFG.phaseC_workers,\n        phase_name='PHASE C (320px fine-tune)', ckpt_name='best_phaseC.pth',\n        device=CFG.device, use_flow=True)\n\n    torch.save(model.state_dict(), os.path.join(CKPT_DIR, 'final.pth'))\n    print('\\nTraining complete.')\n    return model\n\nprint('=== CELL 7 OK ===')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 8 — Full-Resolution Inference with TTA\n# ============================================================================\n\ndef preprocess_for_encoder(img_rgb, target_size):\n    resized = resize_preserve_aspect(img_rgb, target_size)\n    padded, _, _ = pad_to_square(resized, target_size)\n    return torch.from_numpy(padded.astype(np.float32)/255.0).permute(2,0,1).unsqueeze(0)\n\n@torch.no_grad()\ndef infer_single_image(model, img_bgr, device, cfg=CFG):\n    img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)\n    orig_h, orig_w = img_rgb.shape[:2]\n\n    all_params = {k: [] for k in ['k1','k2','k3','p1','p2','cx','cy','s']}\n    all_flow = []\n\n    # Multi-scale + horizontal flip TTA\n    for scale in cfg.tta_scales:\n        for flip in [False, True]:\n            img_in = img_rgb[:, ::-1].copy() if flip else img_rgb\n            enc_input = preprocess_for_encoder(img_in, scale).to(device)\n            with autocast():\n                k1, k2, k3, p1, p2, cx, cy, s, flow_lr = model.forward_params(enc_input)\n\n            if flip:\n                # Mirror cx, p2 signs for horizontal flip\n                cx = -cx\n                p2 = -p2\n                if flow_lr is not None:\n                    flow_lr = torch.flip(flow_lr, dims=[3])  # flip W\n                    flow_lr[:, 0] = -flow_lr[:, 0]  # negate x component\n\n            for name, val in zip(['k1','k2','k3','p1','p2','cx','cy','s'],\n                                 [k1, k2, k3, p1, p2, cx, cy, s]):\n                all_params[name].append(val)\n            if flow_lr is not None:\n                all_flow.append(flow_lr)\n\n    params = {k: torch.stack(v).mean(0) for k, v in all_params.items()}\n\n    grid = build_parametric_grid(\n        orig_h, orig_w,\n        params['k1'], params['k2'], params['k3'],\n        params['p1'], params['p2'],\n        params['cx'], params['cy'], params['s'], device)\n\n    if model._flow_enabled and all_flow:\n        avg_flow = torch.stack(all_flow).mean(0)\n        flow_hr = F.interpolate(avg_flow, size=(orig_h, orig_w),\n                                mode='bilinear', align_corners=True)\n        flow_hr = flow_hr.permute(0, 2, 3, 1)\n        flow_hr = torch.clamp(flow_hr, -0.20, 0.20)\n        grid = grid + model.residual_lambda * flow_hr\n\n    pad = max(32, int(cfg.reflect_pad * max(orig_h, orig_w) / cfg.encoder_size))\n    pad = min(pad, 160)\n\n    img_t = torch.from_numpy(img_rgb.astype(np.float32)/255.0)\\\n                 .permute(2,0,1).unsqueeze(0).to(device)\n    img_padded = F.pad(img_t, [pad]*4, mode='reflect')\n    grid_padded = adjust_grid_for_padding(grid, orig_h, orig_w, pad, device)\n\n    output = F.grid_sample(img_padded, grid_padded, mode=cfg.infer_interp,\n                           padding_mode='border', align_corners=True)\n\n    out_np = (output[0].clamp(0,1)*255).byte().permute(1,2,0).cpu().numpy()\n    out_bgr = cv2.cvtColor(out_np, cv2.COLOR_RGB2BGR)\n    assert out_bgr.shape[:2] == (orig_h, orig_w)\n    return out_bgr\n\ndef run_inference(model, test_dir, output_dir, device, cfg=CFG):\n    model.eval()\n    os.makedirs(output_dir, exist_ok=True)\n    test_files = sorted([f for f in os.listdir(test_dir)\n                         if f.lower().endswith(('.jpg','.jpeg','.png'))])\n    print(f'\\nProcessing {len(test_files)} test images...')\n    success = fallback = 0\n    for i, fname in enumerate(test_files):\n        fpath = os.path.join(test_dir, fname)\n        out_path = os.path.join(output_dir, fname)\n        try:\n            img_bgr = cv2.imread(fpath, cv2.IMREAD_COLOR)\n            if img_bgr is None: raise ValueError(f'Cannot read {fpath}')\n            corrected = infer_single_image(model, img_bgr, device, cfg)\n            cv2.imwrite(out_path, corrected, [cv2.IMWRITE_JPEG_QUALITY, 95])\n            success += 1\n        except Exception as e:\n            print(f'  FALLBACK {fname}: {e}')\n            img_bgr = cv2.imread(fpath, cv2.IMREAD_COLOR)\n            if img_bgr is not None:\n                cv2.imwrite(out_path, img_bgr, [cv2.IMWRITE_JPEG_QUALITY, 95])\n            fallback += 1\n        if (i+1) % 100 == 0:\n            print(f'  {i+1}/{len(test_files)} ({success} ok, {fallback} fallback)')\n            torch.cuda.empty_cache()\n    print(f'Done: {success} corrected, {fallback} fallbacks')\n    return success, fallback\n\nprint('=== CELL 8 OK ===')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 9 — Sanity + Zip\n# ============================================================================\n\ndef sanity_check(output_dir, test_dir, expected=1000):\n    out_files  = sorted(os.listdir(output_dir))\n    test_files = sorted([f for f in os.listdir(test_dir)\n                         if f.lower().endswith(('.jpg','.jpeg','.png'))])\n    print(f'Output: {len(out_files)} (expected: {expected})')\n    if len(out_files) < expected:\n        missing = set(test_files) - set(out_files)\n        print(f'  MISSING {len(missing)} files!')\n    dim_bad = black_bad = 0\n    for fname in out_files:\n        out = cv2.imread(os.path.join(output_dir, fname))\n        ref = cv2.imread(os.path.join(test_dir, fname))\n        if out is None or ref is None: continue\n        if out.shape[:2] != ref.shape[:2]: dim_bad += 1\n        for strip in [out[:8,:], out[-8:,:], out[:,:8], out[:,-8:]]:\n            if strip.mean() < 3.0: black_bad += 1; break\n    print(f'Dim mismatches: {dim_bad}, Black border warnings: {black_bad}')\n    return len(out_files) >= expected and dim_bad == 0\n\ndef create_zip(output_dir, zip_path='/kaggle/working/submission.zip'):\n    with zipfile.ZipFile(zip_path, 'w', zipfile.ZIP_DEFLATED) as zf:\n        for fname in sorted(os.listdir(output_dir)):\n            zf.write(os.path.join(output_dir, fname), fname)\n    print(f'Created {zip_path} ({os.path.getsize(zip_path)/1024/1024:.1f} MB)')\n\nprint('=== CELL 9 OK ===')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 10 — RUN ALL\n# ============================================================================\n\nprint('=' * 60)\nprint('LENS CORRECTION — Progressive 128→256→320px + Flow')\nprint('=' * 60)\n\nn_orig = len(glob.glob(os.path.join(TRAIN_DIR, '*_original.jpg')))\nn_test = len([f for f in os.listdir(TEST_DIR) if f.lower().endswith('.jpg')])\nprint(f'Training pairs: {n_orig}, Test images: {n_test}')\nif n_orig == 0: raise RuntimeError('No training data!')\nif n_test == 0: raise RuntimeError('No test data!')\n\nt_start = time.time()\nmodel = run_training()\nt_train = time.time() - t_start\nprint(f'\\nTraining: {t_train/60:.1f} min')\n\nt0 = time.time()\nrun_inference(model, TEST_DIR, OUTPUT_DIR, CFG.device, CFG)\nprint(f'Inference: {(time.time()-t0)/60:.1f} min')\n\nprint('\\n' + '=' * 60)\nsanity_check(OUTPUT_DIR, TEST_DIR, expected=n_test)\ncreate_zip(OUTPUT_DIR)\n\nprint(f'\\nTotal: {(time.time()-t_start)/60:.1f} min')\nprint('Upload submission.zip to bounty.autohdr.com → download CSV → submit to Kaggle')","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}