{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":99552,"databundleVersionId":13441085}],"dockerImageVersionId":31090,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# On Kaggle, torch/numpy/skimage/tqdm/PIL are preinstalled. lpips and timm are not -- install them.\n# timm is needed for SwinIR's DropPath / trunc_normal_ / to_2tuple utilities.\nimport subprocess, sys\ndef pip_install(pkg):\n    subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", pkg], check=False)\n\npip_install(\"lpips\")\npip_install(\"timm\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T04:55:47.427745Z","iopub.execute_input":"2026-08-12T04:55:47.428648Z","iopub.status.idle":"2026-08-12T04:56:52.27063Z","shell.execute_reply.started":"2026-08-12T04:55:47.42862Z","shell.execute_reply":"2026-08-12T04:56:52.269689Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport glob\nimport json\nimport csv\nimport time\nimport random\nimport math\n\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.utils.checkpoint as checkpoint\nfrom torch.utils.data import Dataset, DataLoader\nfrom timm.models.layers import DropPath, to_2tuple, trunc_normal_\n\nprint(\"torch:\", torch.__version__, \"| cuda available:\", torch.cuda.is_available())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T04:56:52.271801Z","iopub.execute_input":"2026-08-12T04:56:52.27211Z","iopub.status.idle":"2026-08-12T04:57:04.799495Z","shell.execute_reply.started":"2026-08-12T04:56:52.27209Z","shell.execute_reply":"2026-08-12T04:57:04.798579Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CONFIG = {\n    \"data\": {\n        \"root\": \"/kaggle/input/datasets/abhinav1609/degradation-dataset\",  # <-- EDIT to your Kaggle dataset path\n        \"gt_subdir\": \"GT\",\n        \"lr_subdir\": \"NoisyLR\",\n        \"crop_size\": 128,        # LR-space crop; GT crop = crop_size * scale. SwinIR pads internally to a\n                                 # multiple of window_size at forward time (check_image_size), so this does\n                                 # NOT need to be a multiple of 8 -- 64 just matches the official SwinIR\n                                 # classical-SR training recipe (64x64 LR patch for x2).\n        \"scale\": 2,              # GT:NoisyLR resolution ratio -- 128 -> 256 = 2x, matches the task\n        \"val_fraction\": 0.1,\n        \"seed\": 42,\n    },\n    \"model\": {\n        \"arch\": \"swinir\",\n        \"variant\": \"classical\",  # \"light\" (~0.9M params @1ch, fastest) | \"classical\" (~11.7M params @1ch)\n        \"in_channels\": 1,        # grayscale\n        \"window_size\": 8,\n    },\n    \"train\": {\n        \"epochs\": 100,\n        \"batch_size\": 2,        # drop this if you hit OOM on \"classical\" (embed_dim=180); \"light\" can go higher\n        \"lr\": 2.0e-4,\n        \"min_lr\": 1.0e-6,\n        \"weight_decay\": 0.0,\n        \"betas\": (0.9, 0.99),     # AdamW betas from the official SwinIR training recipe\n        \"num_workers\": 2,        # Kaggle notebooks: keep this low (2-4)\n        \"amp\": False,\n        \"huber_delta\": 0.02,      # Huber loss threshold (pixel range ~[0,1]) -- robust core/PSNR term\n        \"huber_weight\": 1.0,\n        \"ssim_weight\": 0.2,       # structural term, targets the SSIM metric\n        \"perceptual_weight\": 0.02,  # VGG perceptual term weight, targets LPIPS (only used if use_perceptual=True)\n        \"use_perceptual\": False,  # set True to enable -- downloads torchvision's ImageNet VGG16 weights on\n                                   # first run (internet required); VGG is used ONLY as a frozen loss feature\n                                   # extractor, never as an architectural component or checkpoint to load into\n                                   # the model itself\n        \"val_every\": 5,\n        \"ckpt_every\": 5,\n        \"log_every\": 50,\n        \"seed\": 42,\n    },\n    \"eval\": {\n        # 4-way rotation self-ensemble (\"+\"-style TTA). Used for final validation/inference ONLY --\n        # never inside the training loop, since it costs 4x the forward passes per call.\n        \"self_ensemble\": True,\n    },\n    \"paths\": {\n        \"out_dir\": \"/kaggle/working/results\",\n        \"weights_dir\": \"/kaggle/working/weights\",\n        \"ckpt_name\": \"swinir_kla.pt\",\n    },\n}\n\nos.makedirs(CONFIG[\"paths\"][\"out_dir\"], exist_ok=True)\nos.makedirs(CONFIG[\"paths\"][\"weights_dir\"], exist_ok=True)\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(\"device:\", DEVICE)\n\ndef set_seed(seed: int):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nset_seed(CONFIG[\"train\"][\"seed\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T04:57:04.800406Z","iopub.execute_input":"2026-08-12T04:57:04.800729Z","iopub.status.idle":"2026-08-12T04:57:04.81494Z","shell.execute_reply.started":"2026-08-12T04:57:04.800702Z","shell.execute_reply":"2026-08-12T04:57:04.814133Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -----------------------------------------------------------------------------------\n# SwinIR: Image Restoration Using Swin Transformer, https://arxiv.org/abs/2108.10257\n# Originally Written by Ze Liu, Modified by Jingyun Liang.\n# -----------------------------------------------------------------------------------\n\nimport math\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.utils.checkpoint as checkpoint\nfrom timm.models.layers import DropPath, to_2tuple, trunc_normal_\n\n\nclass Mlp(nn.Module):\n    def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):\n        super().__init__()\n        out_features = out_features or in_features\n        hidden_features = hidden_features or in_features\n        self.fc1 = nn.Linear(in_features, hidden_features)\n        self.act = act_layer()\n        self.fc2 = nn.Linear(hidden_features, out_features)\n        self.drop = nn.Dropout(drop)\n\n    def forward(self, x):\n        x = self.fc1(x)\n        x = self.act(x)\n        x = self.drop(x)\n        x = self.fc2(x)\n        x = self.drop(x)\n        return x\n\n\ndef window_partition(x, window_size):\n    \"\"\"\n    Args:\n        x: (B, H, W, C)\n        window_size (int): window size\n\n    Returns:\n        windows: (num_windows*B, window_size, window_size, C)\n    \"\"\"\n    B, H, W, C = x.shape\n    x = x.view(B, H // window_size, window_size, W // window_size, window_size, C)\n    windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)\n    return windows\n\n\ndef window_reverse(windows, window_size, H, W):\n    \"\"\"\n    Args:\n        windows: (num_windows*B, window_size, window_size, C)\n        window_size (int): Window size\n        H (int): Height of image\n        W (int): Width of image\n\n    Returns:\n        x: (B, H, W, C)\n    \"\"\"\n    B = int(windows.shape[0] / (H * W / window_size / window_size))\n    x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1)\n    x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1)\n    return x\n\n\nclass WindowAttention(nn.Module):\n    r\"\"\" Window based multi-head self attention (W-MSA) module with relative position bias.\n    It supports both of shifted and non-shifted window.\n\n    Args:\n        dim (int): Number of input channels.\n        window_size (tuple[int]): The height and width of the window.\n        num_heads (int): Number of attention heads.\n        qkv_bias (bool, optional):  If True, add a learnable bias to query, key, value. Default: True\n        qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set\n        attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.0\n        proj_drop (float, optional): Dropout ratio of output. Default: 0.0\n    \"\"\"\n\n    def __init__(self, dim, window_size, num_heads, qkv_bias=True, qk_scale=None, attn_drop=0., proj_drop=0.):\n\n        super().__init__()\n        self.dim = dim\n        self.window_size = window_size  # Wh, Ww\n        self.num_heads = num_heads\n        head_dim = dim // num_heads\n        self.scale = qk_scale or head_dim ** -0.5\n\n        # define a parameter table of relative position bias\n        self.relative_position_bias_table = nn.Parameter(\n            torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads))  # 2*Wh-1 * 2*Ww-1, nH\n\n        # get pair-wise relative position index for each token inside the window\n        coords_h = torch.arange(self.window_size[0])\n        coords_w = torch.arange(self.window_size[1])\n        coords = torch.stack(torch.meshgrid([coords_h, coords_w]))  # 2, Wh, Ww\n        coords_flatten = torch.flatten(coords, 1)  # 2, Wh*Ww\n        relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :]  # 2, Wh*Ww, Wh*Ww\n        relative_coords = relative_coords.permute(1, 2, 0).contiguous()  # Wh*Ww, Wh*Ww, 2\n        relative_coords[:, :, 0] += self.window_size[0] - 1  # shift to start from 0\n        relative_coords[:, :, 1] += self.window_size[1] - 1\n        relative_coords[:, :, 0] *= 2 * self.window_size[1] - 1\n        relative_position_index = relative_coords.sum(-1)  # Wh*Ww, Wh*Ww\n        self.register_buffer(\"relative_position_index\", relative_position_index)\n\n        self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)\n        self.attn_drop = nn.Dropout(attn_drop)\n        self.proj = nn.Linear(dim, dim)\n\n        self.proj_drop = nn.Dropout(proj_drop)\n\n        trunc_normal_(self.relative_position_bias_table, std=.02)\n        self.softmax = nn.Softmax(dim=-1)\n\n    def forward(self, x, mask=None):\n        \"\"\"\n        Args:\n            x: input features with shape of (num_windows*B, N, C)\n            mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None\n        \"\"\"\n        B_, N, C = x.shape\n        qkv = self.qkv(x).reshape(B_, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)\n        q, k, v = qkv[0], qkv[1], qkv[2]  # make torchscript happy (cannot use tensor as tuple)\n\n        q = q * self.scale\n        attn = (q @ k.transpose(-2, -1))\n\n        relative_position_bias = self.relative_position_bias_table[self.relative_position_index.view(-1)].view(\n            self.window_size[0] * self.window_size[1], self.window_size[0] * self.window_size[1], -1)  # Wh*Ww,Wh*Ww,nH\n        relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous()  # nH, Wh*Ww, Wh*Ww\n        attn = attn + relative_position_bias.unsqueeze(0)\n\n        if mask is not None:\n            nW = mask.shape[0]\n            attn = attn.view(B_ // nW, nW, self.num_heads, N, N) + mask.unsqueeze(1).unsqueeze(0)\n            attn = attn.view(-1, self.num_heads, N, N)\n            attn = self.softmax(attn)\n        else:\n            attn = self.softmax(attn)\n\n        attn = self.attn_drop(attn)\n\n        x = (attn @ v).transpose(1, 2).reshape(B_, N, C)\n        x = self.proj(x)\n        x = self.proj_drop(x)\n        return x\n\n    def extra_repr(self) -> str:\n        return f'dim={self.dim}, window_size={self.window_size}, num_heads={self.num_heads}'\n\n    def flops(self, N):\n        # calculate flops for 1 window with token length of N\n        flops = 0\n        # qkv = self.qkv(x)\n        flops += N * self.dim * 3 * self.dim\n        # attn = (q @ k.transpose(-2, -1))\n        flops += self.num_heads * N * (self.dim // self.num_heads) * N\n        #  x = (attn @ v)\n        flops += self.num_heads * N * N * (self.dim // self.num_heads)\n        # x = self.proj(x)\n        flops += N * self.dim * self.dim\n        return flops\n\n\nclass SwinTransformerBlock(nn.Module):\n    r\"\"\" Swin Transformer Block.\n\n    Args:\n        dim (int): Number of input channels.\n        input_resolution (tuple[int]): Input resulotion.\n        num_heads (int): Number of attention heads.\n        window_size (int): Window size.\n        shift_size (int): Shift size for SW-MSA.\n        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.\n        qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True\n        qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.\n        drop (float, optional): Dropout rate. Default: 0.0\n        attn_drop (float, optional): Attention dropout rate. Default: 0.0\n        drop_path (float, optional): Stochastic depth rate. Default: 0.0\n        act_layer (nn.Module, optional): Activation layer. Default: nn.GELU\n        norm_layer (nn.Module, optional): Normalization layer.  Default: nn.LayerNorm\n    \"\"\"\n\n    def __init__(self, dim, input_resolution, num_heads, window_size=7, shift_size=0,\n                 mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0., drop_path=0.,\n                 act_layer=nn.GELU, norm_layer=nn.LayerNorm):\n        super().__init__()\n        self.dim = dim\n        self.input_resolution = input_resolution\n        self.num_heads = num_heads\n        self.window_size = window_size\n        self.shift_size = shift_size\n        self.mlp_ratio = mlp_ratio\n        if min(self.input_resolution) <= self.window_size:\n            # if window size is larger than input resolution, we don't partition windows\n            self.shift_size = 0\n            self.window_size = min(self.input_resolution)\n        assert 0 <= self.shift_size < self.window_size, \"shift_size must in 0-window_size\"\n\n        self.norm1 = norm_layer(dim)\n        self.attn = WindowAttention(\n            dim, window_size=to_2tuple(self.window_size), num_heads=num_heads,\n            qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop)\n\n        self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()\n        self.norm2 = norm_layer(dim)\n        mlp_hidden_dim = int(dim * mlp_ratio)\n        self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)\n\n        if self.shift_size > 0:\n            attn_mask = self.calculate_mask(self.input_resolution)\n        else:\n            attn_mask = None\n\n        self.register_buffer(\"attn_mask\", attn_mask)\n\n    def calculate_mask(self, x_size):\n        # calculate attention mask for SW-MSA\n        H, W = x_size\n        img_mask = torch.zeros((1, H, W, 1))  # 1 H W 1\n        h_slices = (slice(0, -self.window_size),\n                    slice(-self.window_size, -self.shift_size),\n                    slice(-self.shift_size, None))\n        w_slices = (slice(0, -self.window_size),\n                    slice(-self.window_size, -self.shift_size),\n                    slice(-self.shift_size, None))\n        cnt = 0\n        for h in h_slices:\n            for w in w_slices:\n                img_mask[:, h, w, :] = cnt\n                cnt += 1\n\n        mask_windows = window_partition(img_mask, self.window_size)  # nW, window_size, window_size, 1\n        mask_windows = mask_windows.view(-1, self.window_size * self.window_size)\n        attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)\n        attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))\n\n        return attn_mask\n\n    def forward(self, x, x_size):\n        H, W = x_size\n        B, L, C = x.shape\n        # assert L == H * W, \"input feature has wrong size\"\n\n        shortcut = x\n        x = self.norm1(x)\n        x = x.view(B, H, W, C)\n\n        # cyclic shift\n        if self.shift_size > 0:\n            shifted_x = torch.roll(x, shifts=(-self.shift_size, -self.shift_size), dims=(1, 2))\n        else:\n            shifted_x = x\n\n        # partition windows\n        x_windows = window_partition(shifted_x, self.window_size)  # nW*B, window_size, window_size, C\n        x_windows = x_windows.view(-1, self.window_size * self.window_size, C)  # nW*B, window_size*window_size, C\n\n        # W-MSA/SW-MSA (to be compatible for testing on images whose shapes are the multiple of window size\n        if self.input_resolution == x_size:\n            attn_windows = self.attn(x_windows, mask=self.attn_mask)  # nW*B, window_size*window_size, C\n        else:\n            attn_windows = self.attn(x_windows, mask=self.calculate_mask(x_size).to(x.device))\n\n        # merge windows\n        attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C)\n        shifted_x = window_reverse(attn_windows, self.window_size, H, W)  # B H' W' C\n\n        # reverse cyclic shift\n        if self.shift_size > 0:\n            x = torch.roll(shifted_x, shifts=(self.shift_size, self.shift_size), dims=(1, 2))\n        else:\n            x = shifted_x\n        x = x.view(B, H * W, C)\n\n        # FFN\n        x = shortcut + self.drop_path(x)\n        x = x + self.drop_path(self.mlp(self.norm2(x)))\n\n        return x\n\n    def extra_repr(self) -> str:\n        return f\"dim={self.dim}, input_resolution={self.input_resolution}, num_heads={self.num_heads}, \" \\\n               f\"window_size={self.window_size}, shift_size={self.shift_size}, mlp_ratio={self.mlp_ratio}\"\n\n    def flops(self):\n        flops = 0\n        H, W = self.input_resolution\n        # norm1\n        flops += self.dim * H * W\n        # W-MSA/SW-MSA\n        nW = H * W / self.window_size / self.window_size\n        flops += nW * self.attn.flops(self.window_size * self.window_size)\n        # mlp\n        flops += 2 * H * W * self.dim * self.dim * self.mlp_ratio\n        # norm2\n        flops += self.dim * H * W\n        return flops\n\n\nclass PatchMerging(nn.Module):\n    r\"\"\" Patch Merging Layer.\n\n    Args:\n        input_resolution (tuple[int]): Resolution of input feature.\n        dim (int): Number of input channels.\n        norm_layer (nn.Module, optional): Normalization layer.  Default: nn.LayerNorm\n    \"\"\"\n\n    def __init__(self, input_resolution, dim, norm_layer=nn.LayerNorm):\n        super().__init__()\n        self.input_resolution = input_resolution\n        self.dim = dim\n        self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)\n        self.norm = norm_layer(4 * dim)\n\n    def forward(self, x):\n        \"\"\"\n        x: B, H*W, C\n        \"\"\"\n        H, W = self.input_resolution\n        B, L, C = x.shape\n        assert L == H * W, \"input feature has wrong size\"\n        assert H % 2 == 0 and W % 2 == 0, f\"x size ({H}*{W}) are not even.\"\n\n        x = x.view(B, H, W, C)\n\n        x0 = x[:, 0::2, 0::2, :]  # B H/2 W/2 C\n        x1 = x[:, 1::2, 0::2, :]  # B H/2 W/2 C\n        x2 = x[:, 0::2, 1::2, :]  # B H/2 W/2 C\n        x3 = x[:, 1::2, 1::2, :]  # B H/2 W/2 C\n        x = torch.cat([x0, x1, x2, x3], -1)  # B H/2 W/2 4*C\n        x = x.view(B, -1, 4 * C)  # B H/2*W/2 4*C\n\n        x = self.norm(x)\n        x = self.reduction(x)\n\n        return x\n\n    def extra_repr(self) -> str:\n        return f\"input_resolution={self.input_resolution}, dim={self.dim}\"\n\n    def flops(self):\n        H, W = self.input_resolution\n        flops = H * W * self.dim\n        flops += (H // 2) * (W // 2) * 4 * self.dim * 2 * self.dim\n        return flops\n\n\nclass BasicLayer(nn.Module):\n    \"\"\" A basic Swin Transformer layer for one stage.\n\n    Args:\n        dim (int): Number of input channels.\n        input_resolution (tuple[int]): Input resolution.\n        depth (int): Number of blocks.\n        num_heads (int): Number of attention heads.\n        window_size (int): Local window size.\n        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.\n        qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True\n        qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.\n        drop (float, optional): Dropout rate. Default: 0.0\n        attn_drop (float, optional): Attention dropout rate. Default: 0.0\n        drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0\n        norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm\n        downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None\n        use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.\n    \"\"\"\n\n    def __init__(self, dim, input_resolution, depth, num_heads, window_size,\n                 mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0.,\n                 drop_path=0., norm_layer=nn.LayerNorm, downsample=None, use_checkpoint=False):\n\n        super().__init__()\n        self.dim = dim\n        self.input_resolution = input_resolution\n        self.depth = depth\n        self.use_checkpoint = use_checkpoint\n\n        # build blocks\n        self.blocks = nn.ModuleList([\n            SwinTransformerBlock(dim=dim, input_resolution=input_resolution,\n                                 num_heads=num_heads, window_size=window_size,\n                                 shift_size=0 if (i % 2 == 0) else window_size // 2,\n                                 mlp_ratio=mlp_ratio,\n                                 qkv_bias=qkv_bias, qk_scale=qk_scale,\n                                 drop=drop, attn_drop=attn_drop,\n                                 drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path,\n                                 norm_layer=norm_layer)\n            for i in range(depth)])\n\n        # patch merging layer\n        if downsample is not None:\n            self.downsample = downsample(input_resolution, dim=dim, norm_layer=norm_layer)\n        else:\n            self.downsample = None\n\n    def forward(self, x, x_size):\n        for blk in self.blocks:\n            if self.use_checkpoint:\n                x = checkpoint.checkpoint(blk, x, x_size)\n            else:\n                x = blk(x, x_size)\n        if self.downsample is not None:\n            x = self.downsample(x)\n        return x\n\n    def extra_repr(self) -> str:\n        return f\"dim={self.dim}, input_resolution={self.input_resolution}, depth={self.depth}\"\n\n    def flops(self):\n        flops = 0\n        for blk in self.blocks:\n            flops += blk.flops()\n        if self.downsample is not None:\n            flops += self.downsample.flops()\n        return flops\n\n\nclass RSTB(nn.Module):\n    \"\"\"Residual Swin Transformer Block (RSTB).\n\n    Args:\n        dim (int): Number of input channels.\n        input_resolution (tuple[int]): Input resolution.\n        depth (int): Number of blocks.\n        num_heads (int): Number of attention heads.\n        window_size (int): Local window size.\n        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.\n        qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True\n        qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.\n        drop (float, optional): Dropout rate. Default: 0.0\n        attn_drop (float, optional): Attention dropout rate. Default: 0.0\n        drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0\n        norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm\n        downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None\n        use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.\n        img_size: Input image size.\n        patch_size: Patch size.\n        resi_connection: The convolutional block before residual connection.\n    \"\"\"\n\n    def __init__(self, dim, input_resolution, depth, num_heads, window_size,\n                 mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0.,\n                 drop_path=0., norm_layer=nn.LayerNorm, downsample=None, use_checkpoint=False,\n                 img_size=224, patch_size=4, resi_connection='1conv'):\n        super(RSTB, self).__init__()\n\n        self.dim = dim\n        self.input_resolution = input_resolution\n\n        self.residual_group = BasicLayer(dim=dim,\n                                         input_resolution=input_resolution,\n                                         depth=depth,\n                                         num_heads=num_heads,\n                                         window_size=window_size,\n                                         mlp_ratio=mlp_ratio,\n                                         qkv_bias=qkv_bias, qk_scale=qk_scale,\n                                         drop=drop, attn_drop=attn_drop,\n                                         drop_path=drop_path,\n                                         norm_layer=norm_layer,\n                                         downsample=downsample,\n                                         use_checkpoint=use_checkpoint)\n\n        if resi_connection == '1conv':\n            self.conv = nn.Conv2d(dim, dim, 3, 1, 1)\n        elif resi_connection == '3conv':\n            # to save parameters and memory\n            self.conv = nn.Sequential(nn.Conv2d(dim, dim // 4, 3, 1, 1), nn.LeakyReLU(negative_slope=0.2, inplace=True),\n                                      nn.Conv2d(dim // 4, dim // 4, 1, 1, 0),\n                                      nn.LeakyReLU(negative_slope=0.2, inplace=True),\n                                      nn.Conv2d(dim // 4, dim, 3, 1, 1))\n\n        self.patch_embed = PatchEmbed(\n            img_size=img_size, patch_size=patch_size, in_chans=0, embed_dim=dim,\n            norm_layer=None)\n\n        self.patch_unembed = PatchUnEmbed(\n            img_size=img_size, patch_size=patch_size, in_chans=0, embed_dim=dim,\n            norm_layer=None)\n\n    def forward(self, x, x_size):\n        return self.patch_embed(self.conv(self.patch_unembed(self.residual_group(x, x_size), x_size))) + x\n\n    def flops(self):\n        flops = 0\n        flops += self.residual_group.flops()\n        H, W = self.input_resolution\n        flops += H * W * self.dim * self.dim * 9\n        flops += self.patch_embed.flops()\n        flops += self.patch_unembed.flops()\n\n        return flops\n\n\nclass PatchEmbed(nn.Module):\n    r\"\"\" Image to Patch Embedding\n\n    Args:\n        img_size (int): Image size.  Default: 224.\n        patch_size (int): Patch token size. Default: 4.\n        in_chans (int): Number of input image channels. Default: 3.\n        embed_dim (int): Number of linear projection output channels. Default: 96.\n        norm_layer (nn.Module, optional): Normalization layer. Default: None\n    \"\"\"\n\n    def __init__(self, img_size=224, patch_size=4, in_chans=3, embed_dim=96, norm_layer=None):\n        super().__init__()\n        img_size = to_2tuple(img_size)\n        patch_size = to_2tuple(patch_size)\n        patches_resolution = [img_size[0] // patch_size[0], img_size[1] // patch_size[1]]\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.patches_resolution = patches_resolution\n        self.num_patches = patches_resolution[0] * patches_resolution[1]\n\n        self.in_chans = in_chans\n        self.embed_dim = embed_dim\n\n        if norm_layer is not None:\n            self.norm = norm_layer(embed_dim)\n        else:\n            self.norm = None\n\n    def forward(self, x):\n        x = x.flatten(2).transpose(1, 2)  # B Ph*Pw C\n        if self.norm is not None:\n            x = self.norm(x)\n        return x\n\n    def flops(self):\n        flops = 0\n        H, W = self.img_size\n        if self.norm is not None:\n            flops += H * W * self.embed_dim\n        return flops\n\n\nclass PatchUnEmbed(nn.Module):\n    r\"\"\" Image to Patch Unembedding\n\n    Args:\n        img_size (int): Image size.  Default: 224.\n        patch_size (int): Patch token size. Default: 4.\n        in_chans (int): Number of input image channels. Default: 3.\n        embed_dim (int): Number of linear projection output channels. Default: 96.\n        norm_layer (nn.Module, optional): Normalization layer. Default: None\n    \"\"\"\n\n    def __init__(self, img_size=224, patch_size=4, in_chans=3, embed_dim=96, norm_layer=None):\n        super().__init__()\n        img_size = to_2tuple(img_size)\n        patch_size = to_2tuple(patch_size)\n        patches_resolution = [img_size[0] // patch_size[0], img_size[1] // patch_size[1]]\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.patches_resolution = patches_resolution\n        self.num_patches = patches_resolution[0] * patches_resolution[1]\n\n        self.in_chans = in_chans\n        self.embed_dim = embed_dim\n\n    def forward(self, x, x_size):\n        B, HW, C = x.shape\n        x = x.transpose(1, 2).view(B, self.embed_dim, x_size[0], x_size[1])  # B Ph*Pw C\n        return x\n\n    def flops(self):\n        flops = 0\n        return flops\n\n\nclass Upsample(nn.Sequential):\n    \"\"\"Upsample module.\n\n    Args:\n        scale (int): Scale factor. Supported scales: 2^n and 3.\n        num_feat (int): Channel number of intermediate features.\n    \"\"\"\n\n    def __init__(self, scale, num_feat):\n        m = []\n        if (scale & (scale - 1)) == 0:  # scale = 2^n\n            for _ in range(int(math.log(scale, 2))):\n                m.append(nn.Conv2d(num_feat, 4 * num_feat, 3, 1, 1))\n                m.append(nn.PixelShuffle(2))\n        elif scale == 3:\n            m.append(nn.Conv2d(num_feat, 9 * num_feat, 3, 1, 1))\n            m.append(nn.PixelShuffle(3))\n        else:\n            raise ValueError(f'scale {scale} is not supported. ' 'Supported scales: 2^n and 3.')\n        super(Upsample, self).__init__(*m)\n\n\nclass UpsampleOneStep(nn.Sequential):\n    \"\"\"UpsampleOneStep module (the difference with Upsample is that it always only has 1conv + 1pixelshuffle)\n       Used in lightweight SR to save parameters.\n\n    Args:\n        scale (int): Scale factor. Supported scales: 2^n and 3.\n        num_feat (int): Channel number of intermediate features.\n\n    \"\"\"\n\n    def __init__(self, scale, num_feat, num_out_ch, input_resolution=None):\n        self.num_feat = num_feat\n        self.input_resolution = input_resolution\n        m = []\n        m.append(nn.Conv2d(num_feat, (scale ** 2) * num_out_ch, 3, 1, 1))\n        m.append(nn.PixelShuffle(scale))\n        super(UpsampleOneStep, self).__init__(*m)\n\n    def flops(self):\n        H, W = self.input_resolution\n        flops = H * W * self.num_feat * 3 * 9\n        return flops\n\n\nclass SwinIR(nn.Module):\n    r\"\"\" SwinIR\n        A PyTorch impl of : `SwinIR: Image Restoration Using Swin Transformer`, based on Swin Transformer.\n\n    Args:\n        img_size (int | tuple(int)): Input image size. Default 64\n        patch_size (int | tuple(int)): Patch size. Default: 1\n        in_chans (int): Number of input image channels. Default: 3\n        embed_dim (int): Patch embedding dimension. Default: 96\n        depths (tuple(int)): Depth of each Swin Transformer layer.\n        num_heads (tuple(int)): Number of attention heads in different layers.\n        window_size (int): Window size. Default: 7\n        mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4\n        qkv_bias (bool): If True, add a learnable bias to query, key, value. Default: True\n        qk_scale (float): Override default qk scale of head_dim ** -0.5 if set. Default: None\n        drop_rate (float): Dropout rate. Default: 0\n        attn_drop_rate (float): Attention dropout rate. Default: 0\n        drop_path_rate (float): Stochastic depth rate. Default: 0.1\n        norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.\n        ape (bool): If True, add absolute position embedding to the patch embedding. Default: False\n        patch_norm (bool): If True, add normalization after patch embedding. Default: True\n        use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False\n        upscale: Upscale factor. 2/3/4/8 for image SR, 1 for denoising and compress artifact reduction\n        img_range: Image range. 1. or 255.\n        upsampler: The reconstruction reconstruction module. 'pixelshuffle'/'pixelshuffledirect'/'nearest+conv'/None\n        resi_connection: The convolutional block before residual connection. '1conv'/'3conv'\n    \"\"\"\n\n    def __init__(self, img_size=64, patch_size=1, in_chans=3,\n                 embed_dim=96, depths=[6, 6, 6, 6], num_heads=[6, 6, 6, 6],\n                 window_size=7, mlp_ratio=4., qkv_bias=True, qk_scale=None,\n                 drop_rate=0., attn_drop_rate=0., drop_path_rate=0.1,\n                 norm_layer=nn.LayerNorm, ape=False, patch_norm=True,\n                 use_checkpoint=False, upscale=2, img_range=1., upsampler='', resi_connection='1conv',\n                 **kwargs):\n        super(SwinIR, self).__init__()\n        num_in_ch = in_chans\n        num_out_ch = in_chans\n        num_feat = 64\n        self.img_range = img_range\n        if in_chans == 3:\n            rgb_mean = (0.4488, 0.4371, 0.4040)\n            self.mean = torch.Tensor(rgb_mean).view(1, 3, 1, 1)\n        else:\n            self.mean = torch.zeros(1, 1, 1, 1)\n        self.upscale = upscale\n        self.upsampler = upsampler\n        self.window_size = window_size\n\n        #####################################################################################################\n        ################################### 1, shallow feature extraction ###################################\n        self.conv_first = nn.Conv2d(num_in_ch, embed_dim, 3, 1, 1)\n\n        #####################################################################################################\n        ################################### 2, deep feature extraction ######################################\n        self.num_layers = len(depths)\n        self.embed_dim = embed_dim\n        self.ape = ape\n        self.patch_norm = patch_norm\n        self.num_features = embed_dim\n        self.mlp_ratio = mlp_ratio\n\n        # split image into non-overlapping patches\n        self.patch_embed = PatchEmbed(\n            img_size=img_size, patch_size=patch_size, in_chans=embed_dim, embed_dim=embed_dim,\n            norm_layer=norm_layer if self.patch_norm else None)\n        num_patches = self.patch_embed.num_patches\n        patches_resolution = self.patch_embed.patches_resolution\n        self.patches_resolution = patches_resolution\n\n        # merge non-overlapping patches into image\n        self.patch_unembed = PatchUnEmbed(\n            img_size=img_size, patch_size=patch_size, in_chans=embed_dim, embed_dim=embed_dim,\n            norm_layer=norm_layer if self.patch_norm else None)\n\n        # absolute position embedding\n        if self.ape:\n            self.absolute_pos_embed = nn.Parameter(torch.zeros(1, num_patches, embed_dim))\n            trunc_normal_(self.absolute_pos_embed, std=.02)\n\n        self.pos_drop = nn.Dropout(p=drop_rate)\n\n        # stochastic depth\n        dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))]  # stochastic depth decay rule\n\n        # build Residual Swin Transformer blocks (RSTB)\n        self.layers = nn.ModuleList()\n        for i_layer in range(self.num_layers):\n            layer = RSTB(dim=embed_dim,\n                         input_resolution=(patches_resolution[0],\n                                           patches_resolution[1]),\n                         depth=depths[i_layer],\n                         num_heads=num_heads[i_layer],\n                         window_size=window_size,\n                         mlp_ratio=self.mlp_ratio,\n                         qkv_bias=qkv_bias, qk_scale=qk_scale,\n                         drop=drop_rate, attn_drop=attn_drop_rate,\n                         drop_path=dpr[sum(depths[:i_layer]):sum(depths[:i_layer + 1])],  # no impact on SR results\n                         norm_layer=norm_layer,\n                         downsample=None,\n                         use_checkpoint=use_checkpoint,\n                         img_size=img_size,\n                         patch_size=patch_size,\n                         resi_connection=resi_connection\n\n                         )\n            self.layers.append(layer)\n        self.norm = norm_layer(self.num_features)\n\n        # build the last conv layer in deep feature extraction\n        if resi_connection == '1conv':\n            self.conv_after_body = nn.Conv2d(embed_dim, embed_dim, 3, 1, 1)\n        elif resi_connection == '3conv':\n            # to save parameters and memory\n            self.conv_after_body = nn.Sequential(nn.Conv2d(embed_dim, embed_dim // 4, 3, 1, 1),\n                                                 nn.LeakyReLU(negative_slope=0.2, inplace=True),\n                                                 nn.Conv2d(embed_dim // 4, embed_dim // 4, 1, 1, 0),\n                                                 nn.LeakyReLU(negative_slope=0.2, inplace=True),\n                                                 nn.Conv2d(embed_dim // 4, embed_dim, 3, 1, 1))\n\n        #####################################################################################################\n        ################################ 3, high quality image reconstruction ################################\n        if self.upsampler == 'pixelshuffle':\n            # for classical SR\n            self.conv_before_upsample = nn.Sequential(nn.Conv2d(embed_dim, num_feat, 3, 1, 1),\n                                                      nn.LeakyReLU(inplace=True))\n            self.upsample = Upsample(upscale, num_feat)\n            self.conv_last = nn.Conv2d(num_feat, num_out_ch, 3, 1, 1)\n        elif self.upsampler == 'pixelshuffledirect':\n            # for lightweight SR (to save parameters)\n            self.upsample = UpsampleOneStep(upscale, embed_dim, num_out_ch,\n                                            (patches_resolution[0], patches_resolution[1]))\n        elif self.upsampler == 'nearest+conv':\n            # for real-world SR (less artifacts)\n            self.conv_before_upsample = nn.Sequential(nn.Conv2d(embed_dim, num_feat, 3, 1, 1),\n                                                      nn.LeakyReLU(inplace=True))\n            self.conv_up1 = nn.Conv2d(num_feat, num_feat, 3, 1, 1)\n            if self.upscale == 4:\n                self.conv_up2 = nn.Conv2d(num_feat, num_feat, 3, 1, 1)\n            self.conv_hr = nn.Conv2d(num_feat, num_feat, 3, 1, 1)\n            self.conv_last = nn.Conv2d(num_feat, num_out_ch, 3, 1, 1)\n            self.lrelu = nn.LeakyReLU(negative_slope=0.2, inplace=True)\n        else:\n            # for image denoising and JPEG compression artifact reduction\n            self.conv_last = nn.Conv2d(embed_dim, num_out_ch, 3, 1, 1)\n\n        self.apply(self._init_weights)\n\n    def _init_weights(self, m):\n        if isinstance(m, nn.Linear):\n            trunc_normal_(m.weight, std=.02)\n            if isinstance(m, nn.Linear) and m.bias is not None:\n                nn.init.constant_(m.bias, 0)\n        elif isinstance(m, nn.LayerNorm):\n            nn.init.constant_(m.bias, 0)\n            nn.init.constant_(m.weight, 1.0)\n\n    @torch.jit.ignore\n    def no_weight_decay(self):\n        return {'absolute_pos_embed'}\n\n    @torch.jit.ignore\n    def no_weight_decay_keywords(self):\n        return {'relative_position_bias_table'}\n\n    def check_image_size(self, x):\n        _, _, h, w = x.size()\n        mod_pad_h = (self.window_size - h % self.window_size) % self.window_size\n        mod_pad_w = (self.window_size - w % self.window_size) % self.window_size\n        x = F.pad(x, (0, mod_pad_w, 0, mod_pad_h), 'reflect')\n        return x\n\n    def forward_features(self, x):\n        x_size = (x.shape[2], x.shape[3])\n        x = self.patch_embed(x)\n        if self.ape:\n            x = x + self.absolute_pos_embed\n        x = self.pos_drop(x)\n\n        for layer in self.layers:\n            x = layer(x, x_size)\n\n        x = self.norm(x)  # B L C\n        x = self.patch_unembed(x, x_size)\n\n        return x\n\n    def forward(self, x):\n        H, W = x.shape[2:]\n        x = self.check_image_size(x)\n        \n        self.mean = self.mean.type_as(x)\n        x = (x - self.mean) * self.img_range\n\n        if self.upsampler == 'pixelshuffle':\n            # for classical SR\n            x = self.conv_first(x)\n            x = self.conv_after_body(self.forward_features(x)) + x\n            x = self.conv_before_upsample(x)\n            x = self.conv_last(self.upsample(x))\n        elif self.upsampler == 'pixelshuffledirect':\n            # for lightweight SR\n            x = self.conv_first(x)\n            x = self.conv_after_body(self.forward_features(x)) + x\n            x = self.upsample(x)\n        elif self.upsampler == 'nearest+conv':\n            # for real-world SR\n            x = self.conv_first(x)\n            x = self.conv_after_body(self.forward_features(x)) + x\n            x = self.conv_before_upsample(x)\n            x = self.lrelu(self.conv_up1(torch.nn.functional.interpolate(x, scale_factor=2, mode='nearest')))\n            if self.upscale == 4:\n                x = self.lrelu(self.conv_up2(torch.nn.functional.interpolate(x, scale_factor=2, mode='nearest')))\n            x = self.conv_last(self.lrelu(self.conv_hr(x)))\n        else:\n            # for image denoising and JPEG compression artifact reduction\n            x_first = self.conv_first(x)\n            res = self.conv_after_body(self.forward_features(x_first)) + x_first\n            x = x + self.conv_last(res)\n\n        x = x / self.img_range + self.mean\n\n        return x[:, :, :H*self.upscale, :W*self.upscale]\n\n    def flops(self):\n        flops = 0\n        H, W = self.patches_resolution\n        flops += H * W * 3 * self.embed_dim * 9\n        flops += self.patch_embed.flops()\n        for i, layer in enumerate(self.layers):\n            flops += layer.flops()\n        flops += H * W * 3 * self.embed_dim * self.embed_dim\n        flops += self.upsample.flops()\n        return flops\n\n\n\n# --------------------------------------------------------------------------------------\n# Wrappers used by the rest of this notebook (build_model / self_ensemble_predict take\n# the place of the old ShiftLUT build_model / rotation_ensemble_predict).\n# --------------------------------------------------------------------------------------\n\ndef build_model(variant: str = \"classical\", scale: int = 2, in_channels: int = 1, window_size: int = 8):\n    \"\"\"\n    variant:\n      \"light\"     -> SwinIR-light: embed_dim=60, 4x RSTB(depth 6), pixelshuffledirect head\n                     (~0.9M params @1ch) -- fastest, best default under a speed-weighted metric.\n      \"classical\" -> SwinIR classical SR: embed_dim=180, 6x RSTB(depth 6), pixelshuffle head\n                     (~11.7M params @1ch) -- higher ceiling, slower.\n    Both are trained fully from scratch here -- no pretrained checkpoint is loaded.\n\n    Note: `img_size` below only sizes SwinIR's (unused here, ape=False) absolute position embedding\n    and its flops() helper. The network itself handles arbitrary H/W at forward time via\n    check_image_size() reflect-padding to a multiple of window_size, so training crop size and test\n    image size never need to match img_size or be a multiple of window_size themselves.\n    \"\"\"\n    variants = {\n        \"light\": dict(embed_dim=60, depths=[6, 6, 6, 6], num_heads=[6, 6, 6, 6],\n                      mlp_ratio=2, upsampler=\"pixelshuffledirect\", resi_connection=\"1conv\"),\n        \"classical\": dict(embed_dim=180, depths=[6, 6, 6, 6, 6, 6], num_heads=[6, 6, 6, 6, 6, 6],\n                          mlp_ratio=2, upsampler=\"pixelshuffle\", resi_connection=\"1conv\"),\n    }\n    if variant not in variants:\n        raise ValueError(f\"Unknown variant '{variant}', choose from {list(variants)}\")\n    cfg = variants[variant]\n    return SwinIR(\n        img_size=64, patch_size=1, in_chans=in_channels,\n        window_size=window_size, img_range=1., upscale=scale,\n        **cfg,\n    )\n\n\ndef self_ensemble_predict(model, x):\n    \"\"\"4-way rotation self-ensemble: rotate input, predict, rotate prediction back, average.\n    Use for final validation/inference only -- NOT inside the training loop (4x the forward-pass\n    cost per call). This is a lighter version of SwinIR's official 8-way flip+rotate \"+\" self-ensemble\n    (this notebook only rotates, matching what was already implemented previously).\"\"\"\n    preds = []\n    for k in range(4):\n        x_rot = torch.rot90(x, k, dims=[2, 3])\n        y_rot = model(x_rot)\n        y = torch.rot90(y_rot, -k, dims=[2, 3])\n        preds.append(y)\n    return torch.stack(preds, dim=0).mean(dim=0)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T04:57:04.816842Z","iopub.execute_input":"2026-08-12T04:57:04.817121Z","iopub.status.idle":"2026-08-12T04:57:04.895144Z","shell.execute_reply.started":"2026-08-12T04:57:04.817103Z","shell.execute_reply":"2026-08-12T04:57:04.894517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# quick sanity check\n_m = build_model(CONFIG[\"model\"][\"variant\"], scale=CONFIG[\"data\"][\"scale\"],\n                  in_channels=CONFIG[\"model\"][\"in_channels\"], window_size=CONFIG[\"model\"][\"window_size\"])\n_x = torch.randn(2, CONFIG[\"model\"][\"in_channels\"], 128, 128)  # matches the task's 128x128 LR input\n_y = _m(_x)\nexpected = (128 * CONFIG[\"data\"][\"scale\"], 128 * CONFIG[\"data\"][\"scale\"])\nassert tuple(_y.shape[-2:]) == expected, f\"expected {expected}, got {tuple(_y.shape[-2:])}\"\nprint(\"sanity check -> input:\", _x.shape, \"output:\", _y.shape,\n      \"| params:\", f\"{sum(p.numel() for p in _m.parameters())/1e6:.2f}M\")\ndel _m, _x, _y","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T04:57:04.896014Z","iopub.execute_input":"2026-08-12T04:57:04.896264Z","iopub.status.idle":"2026-08-12T04:57:23.511001Z","shell.execute_reply.started":"2026-08-12T04:57:04.896242Z","shell.execute_reply":"2026-08-12T04:57:23.510294Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Unchanged from the original notebook: load_array() never clips the LR input (out-of-range pixel\n# values carry real degradation-strength information and SwinIR has no sigmoid/clamp, so they pass\n# through untouched), and the augmentation (flip + 90-degree rotation) already matches the standard\n# SwinIR training recipe. No changes needed here for the model swap.\n\nIMG_EXTS = (\".npy\", \".png\", \".jpg\", \".jpeg\", \".tif\", \".tiff\", \".bmp\")\n\n\ndef load_array(path: str) -> np.ndarray:\n    ext = os.path.splitext(path)[1].lower()\n    if ext == \".npy\":\n        arr = np.load(path).astype(np.float32)\n    else:\n        from PIL import Image\n        arr = np.array(Image.open(path).convert(\"RGB\")).astype(np.float32) / 255.0\n\n    if arr.ndim == 2:\n        arr = arr[..., None]\n    if arr.shape[-1] not in (1, 3):\n        arr = np.transpose(arr, (1, 2, 0))\n    return arr\n\n\ndef save_array(path: str, arr: np.ndarray):\n    ext = os.path.splitext(path)[1].lower()\n    if ext == \".npy\":\n        np.save(path, arr.astype(np.float32))\n    else:\n        from PIL import Image\n        img = np.clip(arr, 0.0, 1.0)\n        img = (img * 255.0).round().astype(np.uint8)\n        if img.shape[-1] == 1:\n            img = img[..., 0]\n        Image.fromarray(img).save(path)\n\n\nclass PairedRestorationDataset(Dataset):\n    def __init__(self, root, gt_subdir=\"GT\", lr_subdir=\"NoisyLR\",\n                 crop_size=96, scale=2, train=True, file_list=None):\n        self.gt_dir = os.path.join(root, gt_subdir)\n        self.lr_dir = os.path.join(root, lr_subdir)\n        self.crop_size = crop_size\n        self.scale = scale\n        self.train = train\n\n        if file_list is None:\n            gt_files = set(os.path.basename(f) for f in glob.glob(os.path.join(self.gt_dir, \"*\")))\n            lr_files = set(os.path.basename(f) for f in glob.glob(os.path.join(self.lr_dir, \"*\")))\n            common = sorted(gt_files & lr_files)\n            if not common:\n                raise FileNotFoundError(f\"No matching filenames between {self.gt_dir} and {self.lr_dir}\")\n            self.files = common\n        else:\n            self.files = file_list\n\n    def __len__(self):\n        return len(self.files)\n\n    def _augment(self, lr, gt):\n        if np.random.rand() < 0.5:\n            lr, gt = lr[:, ::-1, :].copy(), gt[:, ::-1, :].copy()\n        if np.random.rand() < 0.5:\n            lr, gt = lr[::-1, :, :].copy(), gt[::-1, :, :].copy()\n        k = np.random.randint(0, 4)\n        if k:\n            lr, gt = np.rot90(lr, k).copy(), np.rot90(gt, k).copy()\n        return lr, gt\n\n    def _random_crop(self, lr, gt):\n        h, w = lr.shape[:2]\n        cs = self.crop_size\n        if h < cs or w < cs:\n            pad_h, pad_w = max(0, cs - h), max(0, cs - w)\n            lr = np.pad(lr, ((0, pad_h), (0, pad_w), (0, 0)), mode=\"reflect\")\n            gt = np.pad(gt, ((0, pad_h * self.scale), (0, pad_w * self.scale), (0, 0)), mode=\"reflect\")\n            h, w = lr.shape[:2]\n        top = np.random.randint(0, h - cs + 1)\n        left = np.random.randint(0, w - cs + 1)\n        lr_c = lr[top:top + cs, left:left + cs, :]\n        gt_c = gt[top * self.scale:(top + cs) * self.scale, left * self.scale:(left + cs) * self.scale, :]\n        return lr_c, gt_c\n\n    def __getitem__(self, idx):\n        fname = self.files[idx]\n        lr = load_array(os.path.join(self.lr_dir, fname))\n        gt = load_array(os.path.join(self.gt_dir, fname))\n\n        if self.train:\n            lr, gt = self._random_crop(lr, gt)\n            lr, gt = self._augment(lr, gt)\n\n        lr_t = torch.from_numpy(np.ascontiguousarray(lr.transpose(2, 0, 1))).float()\n        gt_t = torch.from_numpy(np.ascontiguousarray(gt.transpose(2, 0, 1))).float()\n        return lr_t, gt_t, fname\n\n\ndef make_train_val_split(root, gt_subdir=\"GT\", lr_subdir=\"NoisyLR\", val_fraction=0.1, seed=42):\n    gt_dir = os.path.join(root, gt_subdir)\n    lr_dir = os.path.join(root, lr_subdir)\n    gt_files = set(os.path.basename(f) for f in glob.glob(os.path.join(gt_dir, \"*\")))\n    lr_files = set(os.path.basename(f) for f in glob.glob(os.path.join(lr_dir, \"*\")))\n    common = sorted(gt_files & lr_files)\n\n    rng = np.random.RandomState(seed)\n    idx = rng.permutation(len(common))\n    n_val = max(1, int(len(common) * val_fraction))\n    val_idx, train_idx = set(idx[:n_val].tolist()), set(idx[n_val:].tolist())\n\n    train_files = [common[i] for i in sorted(train_idx)]\n    val_files = [common[i] for i in sorted(val_idx)]\n    return train_files, val_files\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T04:57:23.512281Z","iopub.execute_input":"2026-08-12T04:57:23.512624Z","iopub.status.idle":"2026-08-12T04:57:23.530332Z","shell.execute_reply.started":"2026-08-12T04:57:23.512605Z","shell.execute_reply":"2026-08-12T04:57:23.529788Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def _gaussian_kernel(window_size, sigma, channels, device, dtype):\n    coords = torch.arange(window_size, device=device, dtype=dtype) - window_size // 2\n    g = torch.exp(-(coords ** 2) / (2 * sigma ** 2))\n    g = g / g.sum()\n    kernel_2d = g[:, None] @ g[None, :]\n    kernel = kernel_2d.expand(channels, 1, window_size, window_size).contiguous()\n    return kernel\n\n\nclass SSIMLoss(nn.Module):\n    \"\"\"SSIM used as a training loss (1 - SSIM). See loss-fn survey sec 8.1.3 (SSIM as loss) /\n    9.2.2 (SSIM as metric); L=1.0 for [0,1]-range images, not 255.\"\"\"\n\n    def __init__(self, window_size=11, sigma=1.5, channels=1):\n        super().__init__()\n        self.window_size = window_size\n        self.sigma = sigma\n        self.channels = channels\n        self.register_buffer(\"window\", torch.empty(0), persistent=False)\n\n    def _get_window(self, device, dtype):\n        if self.window.numel() == 0 or self.window.device != device:\n            self.window = _gaussian_kernel(self.window_size, self.sigma, self.channels, device, dtype)\n        return self.window\n\n    def forward(self, pred, target):\n        window = self._get_window(pred.device, pred.dtype)\n        pad = self.window_size // 2\n        c1, c2 = 0.01 ** 2, 0.03 ** 2\n\n        mu_p = F.conv2d(pred, window, padding=pad, groups=self.channels)\n        mu_t = F.conv2d(target, window, padding=pad, groups=self.channels)\n        mu_p2, mu_t2, mu_pt = mu_p * mu_p, mu_t * mu_t, mu_p * mu_t\n\n        sigma_p2 = F.conv2d(pred * pred, window, padding=pad, groups=self.channels) - mu_p2\n        sigma_t2 = F.conv2d(target * target, window, padding=pad, groups=self.channels) - mu_t2\n        sigma_pt = F.conv2d(pred * target, window, padding=pad, groups=self.channels) - mu_pt\n\n        ssim_map = ((2 * mu_pt + c1) * (2 * sigma_pt + c2)) / (\n            (mu_p2 + mu_t2 + c1) * (sigma_p2 + sigma_t2 + c2)\n        )\n        return 1.0 - ssim_map.mean()\n\n\nclass VGGPerceptualLoss(nn.Module):\n    \"\"\"Perceptual loss (loss-fn survey sec 9.1.3, Eq. 172): normalized L2 distance between frozen\n    VGG16 feature activations of prediction vs. target. VGG16 is used ONLY as a fixed feature\n    extractor for this loss term -- never fine-tuned, never used as an architectural component of\n    the restoration model itself. This is the one pretrained-weights exception allowed for this task.\n\n    Requires internet access on first run to download torchvision's ImageNet-pretrained VGG16 weights.\n    \"\"\"\n\n    # relu2_2=8, relu3_3=15, relu4_3=22 in torchvision's vgg16.features indexing\n    def __init__(self, layers=(\"8\", \"15\", \"22\"), device=\"cuda\"):\n        super().__init__()\n        from torchvision.models import vgg16, VGG16_Weights\n        vgg = vgg16(weights=VGG16_Weights.IMAGENET1K_V1).features.to(device).eval()\n        for p in vgg.parameters():\n            p.requires_grad_(False)\n        self.vgg = vgg\n        self.layers = set(layers)\n        self.last_layer = max(int(n) for n in self.layers)\n        self.register_buffer(\"mean\", torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1))\n        self.register_buffer(\"std\", torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1))\n\n    def _preprocess(self, x):\n        if x.shape[1] == 1:\n            x = x.repeat(1, 3, 1, 1)  # grayscale -> 3ch, VGG16 is RGB/ImageNet-trained\n        x = x.clamp(0, 1)\n        return (x - self.mean) / self.std\n\n    def forward(self, pred, target):\n        p, t = self._preprocess(pred), self._preprocess(target)\n        loss = 0.0\n        with torch.no_grad():\n            t_feats = {}\n            tt = t\n            for name, layer in self.vgg._modules.items():\n                tt = layer(tt)\n                if name in self.layers:\n                    t_feats[name] = tt\n                if int(name) >= self.last_layer:\n                    break\n        pp = p\n        for name, layer in self.vgg._modules.items():\n            pp = layer(pp)\n            if name in self.layers:\n                loss = loss + F.mse_loss(pp, t_feats[name])\n            if int(name) >= self.last_layer:\n                break\n        return loss\n\n\nclass RestorationLoss(nn.Module):\n    \"\"\"Huber (robust core/PSNR term) + (1-SSIM) (structural term) + optional VGG perceptual (LPIPS-\n    aligned term). See the loss-fn survey discussion: Huber sec 2.1.3, SSIM sec 8.1.3/9.2.2,\n    Perceptual sec 9.1.3 Eq. 172.\"\"\"\n\n    def __init__(self, huber_delta=0.02, huber_weight=1.0, ssim_weight=0.2,\n                 perceptual_weight=0.0, use_perceptual=False, device=\"cuda\", in_channels=1):\n        super().__init__()\n        self.huber = nn.SmoothL1Loss(beta=huber_delta)\n        self.ssim = SSIMLoss(channels=in_channels)\n        self.huber_weight = huber_weight\n        self.ssim_weight = ssim_weight\n        self.perceptual_weight = perceptual_weight\n        self.perceptual_fn = None\n        if use_perceptual and perceptual_weight > 0:\n            self.perceptual_fn = VGGPerceptualLoss(device=device)\n\n    def forward(self, pred, target):\n        # Huber runs on the RAW (unclamped) prediction: it is what supervises pixels the network\n        # pushes outside [0,1], since SSIM/perceptual clamp first and get zero gradient out of range.\n        huber = self.huber(pred, target)\n\n        pred_c = pred.clamp(0, 1)\n        target_c = target.clamp(0, 1)\n        ssim = self.ssim(pred_c, target_c)\n\n        total = self.huber_weight * huber + self.ssim_weight * ssim\n        logs = {\"huber\": huber.item(), \"ssim_loss\": ssim.item()}\n\n        if self.perceptual_fn is not None:\n            perc = self.perceptual_fn(pred_c, target_c)\n            total = total + self.perceptual_weight * perc\n            logs[\"perceptual\"] = perc.item()\n\n        logs[\"total\"] = total.item()\n        return total, logs\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T04:57:23.531211Z","iopub.execute_input":"2026-08-12T04:57:23.531484Z","iopub.status.idle":"2026-08-12T04:57:23.550852Z","shell.execute_reply.started":"2026-08-12T04:57:23.53141Z","shell.execute_reply":"2026-08-12T04:57:23.550008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"_lpips_model = None\n\ndef psnr(pred, target, data_range=1.0):\n    pred = np.clip(pred, 0, data_range)\n    target = np.clip(target, 0, data_range)\n    mse = np.mean((pred.astype(np.float64) - target.astype(np.float64)) ** 2)\n    if mse == 0:\n        return float(\"inf\")\n    return 10.0 * np.log10((data_range ** 2) / mse)\n\n\ndef ssim_metric(pred, target, data_range=1.0):\n    from skimage.metrics import structural_similarity as sk_ssim\n    pred = np.clip(pred, 0, data_range)\n    target = np.clip(target, 0, data_range)\n    channel_axis = -1 if pred.ndim == 3 and pred.shape[-1] in (1, 3) else None\n    return float(sk_ssim(target, pred, data_range=data_range, channel_axis=channel_axis))\n\n\ndef lpips_score(pred, target, device=\"cuda\"):\n    global _lpips_model\n    import lpips\n    if _lpips_model is None:\n        _lpips_model = lpips.LPIPS(net=\"alex\").to(device)\n        _lpips_model.eval()\n\n    def to_tensor(x):\n        x = np.clip(x, 0, 1).astype(np.float32)\n        t = torch.from_numpy(x.transpose(2, 0, 1)).unsqueeze(0).to(device)\n        return t * 2 - 1\n\n    with torch.no_grad():\n        d = _lpips_model(to_tensor(pred), to_tensor(target))\n    return float(d.item())\n\n\ndef compute_all_metrics(pred, target, device=\"cuda\"):\n    metrics = {\"psnr\": psnr(pred, target), \"ssim\": ssim_metric(pred, target)}\n    try:\n        metrics[\"lpips\"] = lpips_score(pred, target, device=device)\n    except Exception as e:\n        # first-call weight download can hiccup on flaky connections -- don't\n        # kill the whole eval run over it, just flag it and keep going.\n        print(f\"[warn] LPIPS unavailable ({e}); reporting NaN for this metric\")\n        metrics[\"lpips\"] = float(\"nan\")\n    return metrics\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T04:57:23.55177Z","iopub.execute_input":"2026-08-12T04:57:23.552013Z","iopub.status.idle":"2026-08-12T04:57:23.566964Z","shell.execute_reply.started":"2026-08-12T04:57:23.551988Z","shell.execute_reply":"2026-08-12T04:57:23.566198Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"split_path = os.path.join(CONFIG[\"paths\"][\"out_dir\"], \"val_split.json\")\nif os.path.exists(split_path):\n    with open(split_path) as f:\n        split = json.load(f)\n    train_files, val_files = split[\"train\"], split[\"val\"]\nelse:\n    train_files, val_files = make_train_val_split(\n        CONFIG[\"data\"][\"root\"], CONFIG[\"data\"][\"gt_subdir\"], CONFIG[\"data\"][\"lr_subdir\"],\n        val_fraction=CONFIG[\"data\"][\"val_fraction\"], seed=CONFIG[\"data\"][\"seed\"],\n    )\n    with open(split_path, \"w\") as f:\n        json.dump({\"train\": train_files, \"val\": val_files}, f, indent=2)\n\nprint(f\"train files: {len(train_files)} | val files: {len(val_files)}\")\n\ntrain_ds = PairedRestorationDataset(\n    CONFIG[\"data\"][\"root\"], CONFIG[\"data\"][\"gt_subdir\"], CONFIG[\"data\"][\"lr_subdir\"],\n    crop_size=CONFIG[\"data\"][\"crop_size\"], scale=CONFIG[\"data\"][\"scale\"],\n    train=True, file_list=train_files,\n)\nval_ds = PairedRestorationDataset(\n    CONFIG[\"data\"][\"root\"], CONFIG[\"data\"][\"gt_subdir\"], CONFIG[\"data\"][\"lr_subdir\"],\n    scale=CONFIG[\"data\"][\"scale\"], train=False, file_list=val_files,\n)\ntrain_loader = DataLoader(\n    train_ds, batch_size=CONFIG[\"train\"][\"batch_size\"], shuffle=True,\n    num_workers=CONFIG[\"train\"][\"num_workers\"], pin_memory=True, drop_last=True,\n)\nval_loader = DataLoader(val_ds, batch_size=1, shuffle=False, num_workers=2)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T04:57:23.56873Z","iopub.execute_input":"2026-08-12T04:57:23.568923Z","iopub.status.idle":"2026-08-12T04:57:23.722437Z","shell.execute_reply.started":"2026-08-12T04:57:23.568909Z","shell.execute_reply":"2026-08-12T04:57:23.721522Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = build_model(\n    CONFIG[\"model\"][\"variant\"], scale=CONFIG[\"data\"][\"scale\"],\n    in_channels=CONFIG[\"model\"][\"in_channels\"], window_size=CONFIG[\"model\"][\"window_size\"],\n).to(DEVICE)\nn_params = sum(p.numel() for p in model.parameters())\nprint(f\"model=SwinIR-{CONFIG['model']['variant']} params={n_params/1e6:.2f}M\")\n\ncriterion = RestorationLoss(\n    huber_delta=CONFIG[\"train\"][\"huber_delta\"],\n    huber_weight=CONFIG[\"train\"][\"huber_weight\"],\n    ssim_weight=CONFIG[\"train\"][\"ssim_weight\"],\n    perceptual_weight=CONFIG[\"train\"][\"perceptual_weight\"],\n    use_perceptual=CONFIG[\"train\"][\"use_perceptual\"],\n    device=DEVICE,\n    in_channels=CONFIG[\"model\"][\"in_channels\"],\n).to(DEVICE)\n\n# AdamW with SwinIR's standard betas=(0.9, 0.99). Equivalent to Adam here since weight_decay=0.\noptimizer = torch.optim.AdamW(\n    model.parameters(), lr=CONFIG[\"train\"][\"lr\"], betas=CONFIG[\"train\"][\"betas\"],\n    weight_decay=CONFIG[\"train\"][\"weight_decay\"],\n)\n# Cosine schedule kept (simpler than SwinIR's official multi-step schedule; fine for a from-scratch\n# run on a fixed epoch budget).\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer, T_max=CONFIG[\"train\"][\"epochs\"], eta_min=CONFIG[\"train\"][\"min_lr\"]\n)\nscaler = torch.amp.GradScaler(enabled=CONFIG[\"train\"][\"amp\"] and DEVICE == \"cuda\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T04:57:23.724436Z","iopub.execute_input":"2026-08-12T04:57:23.725164Z","iopub.status.idle":"2026-08-12T04:57:24.254862Z","shell.execute_reply.started":"2026-08-12T04:57:23.725136Z","shell.execute_reply":"2026-08-12T04:57:24.25403Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef validate(model, loader, device, self_ensemble=False):\n    model.eval()\n    psnrs, ssims = [], []\n    for lr_img, gt_img, _ in loader:\n        lr_img, gt_img = lr_img.to(device), gt_img.to(device)\n        pred = self_ensemble_predict(model, lr_img) if self_ensemble else model(lr_img)\n        pred = pred.clamp(0, 1)\n        pred_np = pred[0].permute(1, 2, 0).cpu().numpy()\n        gt_np = gt_img[0].permute(1, 2, 0).cpu().numpy()\n        psnrs.append(psnr(pred_np, gt_np))\n        ssims.append(ssim_metric(pred_np, gt_np))\n    model.train()\n    return float(np.mean(psnrs)), float(np.mean(ssims))\n\n\nRESUME_FROM = None  # e.g. os.path.join(CONFIG[\"paths\"][\"weights_dir\"], CONFIG[\"paths\"][\"ckpt_name\"])\n\nstart_epoch = 0\nbest_psnr = -1.0\nif RESUME_FROM and os.path.exists(RESUME_FROM):\n    ckpt = torch.load(RESUME_FROM, map_location=DEVICE)\n    model.load_state_dict(ckpt[\"model\"])\n    optimizer.load_state_dict(ckpt[\"optimizer\"])\n    start_epoch = ckpt[\"epoch\"] + 1\n    best_psnr = ckpt.get(\"best_psnr\", -1.0)\n    print(f\"resumed from {RESUME_FROM} at epoch {start_epoch}\")\n\nlog_path = os.path.join(CONFIG[\"paths\"][\"out_dir\"], \"train_log.csv\")\nwrite_header = not os.path.exists(log_path)\nlog_file = open(log_path, \"a\", newline=\"\")\nlog_writer = csv.writer(log_file)\nif write_header:\n    log_writer.writerow([\"epoch\", \"loss\", \"huber\", \"ssim_loss\", \"perceptual\", \"val_psnr\", \"val_ssim\", \"lr\", \"epoch_time_s\"])\n\nfor epoch in range(start_epoch, CONFIG[\"train\"][\"epochs\"]):\n    model.train()\n    t0 = time.time()\n    running = {}\n    for step, (lr_img, gt_img, _) in enumerate(train_loader):\n        lr_img, gt_img = lr_img.to(DEVICE, non_blocking=True), gt_img.to(DEVICE, non_blocking=True)\n\n        optimizer.zero_grad(set_to_none=True)\n        with torch.amp.autocast(device_type=\"cuda\", enabled=CONFIG[\"train\"][\"amp\"] and DEVICE == \"cuda\"):\n            # Plain single forward pass. The old notebook trained against a 4-way rotation\n            # self-ensemble every step (4x the forward-pass cost per batch) -- self-ensembling is a\n            # test-time-only trick (see self_ensemble_predict / CONFIG[\"eval\"][\"self_ensemble\"]),\n            # it buys nothing during training and just burns compute.\n            pred = model(lr_img)\n            loss, logs = criterion(pred, gt_img)\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        for k, v in logs.items():\n            running[k] = running.get(k, 0.0) + v\n        if step % CONFIG[\"train\"][\"log_every\"] == 0:\n            avg = {k: v / (step + 1) for k, v in running.items()}\n            print(f\"epoch {epoch} step {step}/{len(train_loader)}  \" +\n                  \"  \".join(f\"{k}={v:.4f}\" for k, v in avg.items()))\n\n    scheduler.step()\n    epoch_time = time.time() - t0\n    avg_logs = {k: v / len(train_loader) for k, v in running.items()}\n\n    val_psnr, val_ssim = (None, None)\n    if (epoch + 1) % CONFIG[\"train\"][\"val_every\"] == 0 or epoch == CONFIG[\"train\"][\"epochs\"] - 1:\n        # Fast per-epoch validation: plain forward, no self-ensemble (keep the epoch loop quick).\n        val_psnr, val_ssim = validate(model, val_loader, DEVICE, self_ensemble=False)\n        print(f\"[epoch {epoch}] val PSNR={val_psnr:.3f} dB  SSIM={val_ssim:.4f}  ({epoch_time:.1f}s)\")\n        if val_psnr > best_psnr:\n            best_psnr = val_psnr\n            torch.save(\n                {\"model\": model.state_dict(), \"epoch\": epoch, \"best_psnr\": best_psnr, \"config\": CONFIG},\n                os.path.join(CONFIG[\"paths\"][\"weights_dir\"], CONFIG[\"paths\"][\"ckpt_name\"].replace(\".pt\", \"_best.pt\")),\n            )\n\n    log_writer.writerow([\n        epoch, avg_logs.get(\"total\"), avg_logs.get(\"huber\"), avg_logs.get(\"ssim_loss\"), avg_logs.get(\"perceptual\"),\n        val_psnr, val_ssim, optimizer.param_groups[0][\"lr\"], epoch_time,\n    ])\n    log_file.flush()\n\n    if (epoch + 1) % CONFIG[\"train\"][\"ckpt_every\"] == 0 or epoch == CONFIG[\"train\"][\"epochs\"] - 1:\n        torch.save(\n            {\"model\": model.state_dict(), \"optimizer\": optimizer.state_dict(),\n             \"epoch\": epoch, \"best_psnr\": best_psnr, \"config\": CONFIG},\n            os.path.join(CONFIG[\"paths\"][\"weights_dir\"], CONFIG[\"paths\"][\"ckpt_name\"]),\n        )\n\nlog_file.close()\nprint(f\"training complete. best val PSNR = {best_psnr:.3f} dB\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T04:57:24.256082Z","iopub.execute_input":"2026-08-12T04:57:24.256399Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def bicubic_baseline(lr_img, scale):\n    return F.interpolate(lr_img, scale_factor=scale, mode=\"bicubic\", align_corners=False).clamp(0, 1)\n\n\ndef run_evaluation(checkpoint_path, num_examples=8, self_ensemble=None):\n    if self_ensemble is None:\n        self_ensemble = CONFIG[\"eval\"][\"self_ensemble\"]\n    ckpt = torch.load(checkpoint_path, map_location=DEVICE)\n    mcfg = ckpt[\"config\"][\"model\"]\n    scale = ckpt[\"config\"][\"data\"][\"scale\"]\n    eval_model = build_model(\n        mcfg[\"variant\"], scale=scale, in_channels=mcfg[\"in_channels\"], window_size=mcfg[\"window_size\"],\n    ).to(DEVICE)\n    eval_model.load_state_dict(ckpt[\"model\"])\n    eval_model.eval()\n\n    model_metrics, baseline_metrics = [], []\n    examples_dir = os.path.join(CONFIG[\"paths\"][\"out_dir\"], \"examples\")\n    os.makedirs(examples_dir, exist_ok=True)\n\n    with torch.no_grad():\n        for i in range(len(val_ds)):\n            lr_img, gt_img, fname = val_ds[i]\n            lr_img, gt_img = lr_img.unsqueeze(0).to(DEVICE), gt_img.unsqueeze(0).to(DEVICE)\n\n            pred_fn = self_ensemble_predict if self_ensemble else (lambda m, x: m(x))\n            pred = pred_fn(eval_model, lr_img).clamp(0, 1)\n            base = bicubic_baseline(lr_img, scale)\n\n            gt_np = gt_img[0].permute(1, 2, 0).cpu().numpy()\n            pred_np = pred[0].permute(1, 2, 0).cpu().numpy()\n            base_np = base[0].permute(1, 2, 0).cpu().numpy()\n\n            model_metrics.append(compute_all_metrics(pred_np, gt_np, device=DEVICE))\n            baseline_metrics.append(compute_all_metrics(base_np, gt_np, device=DEVICE))\n\n            if i < num_examples:\n                from PIL import Image\n                def to_uint8(x):\n                    x = (np.clip(x, 0, 1) * 255).astype(np.uint8)\n                    if x.shape[-1] == 1:\n                        x = x[..., 0]   # (H,W,1) -> (H,W) so PIL treats it as grayscale\n                    return x\n                side_by_side = np.concatenate([to_uint8(base_np), to_uint8(pred_np), to_uint8(gt_np)], axis=1)\n                Image.fromarray(side_by_side).save(\n                    os.path.join(examples_dir, f\"{os.path.splitext(fname)[0]}_bicubic_pred_gt.png\")\n                )\n\n    def agg(metrics_list):\n        keys = metrics_list[0].keys()\n        return {k: float(np.mean([m[k] for m in metrics_list if np.isfinite(m[k])])) for k in keys}\n\n    report = {\n        \"num_val_images\": len(val_ds),\n        \"self_ensemble\": self_ensemble,\n        \"model\": agg(model_metrics),\n        \"baseline_bicubic\": agg(baseline_metrics),\n    }\n    report_path = os.path.join(CONFIG[\"paths\"][\"out_dir\"], \"eval_report.json\")\n    with open(report_path, \"w\") as f:\n        json.dump(report, f, indent=2)\n\n    print(json.dumps(report, indent=2))\n    print(f\"\\nSaved report to {report_path}, examples to {examples_dir}\")\n    return report\n\n\nbest_ckpt = os.path.join(CONFIG[\"paths\"][\"weights_dir\"], CONFIG[\"paths\"][\"ckpt_name\"].replace(\".pt\", \"_best.pt\"))\nif os.path.exists(best_ckpt):\n    run_evaluation(best_ckpt)\nelse:\n    print(f\"no checkpoint found at {best_ckpt} yet -- run the training cell first\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T04:54:53.752728Z","iopub.status.idle":"2026-08-12T04:54:53.753036Z","shell.execute_reply.started":"2026-08-12T04:54:53.752906Z","shell.execute_reply":"2026-08-12T04:54:53.752919Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def pad_to_multiple(t, multiple=8):\n    _, _, h, w = t.shape\n    pad_h = (multiple - h % multiple) % multiple\n    pad_w = (multiple - w % multiple) % multiple\n    if pad_h or pad_w:\n        t = F.pad(t, (0, pad_w, 0, pad_h), mode=\"reflect\")\n    return t, pad_h, pad_w\n\n\ndef run_inference(input_dir, output_dir, checkpoint_path, batch_size=4, clip_output=True,\n                   device=None, self_ensemble=None):\n    # self_ensemble=True (default from CONFIG) roughly 4x's per-image latency -- if the eval score\n    # weights inference speed, benchmark both settings before choosing which to submit.\n    device = device or DEVICE\n    if self_ensemble is None:\n        self_ensemble = CONFIG[\"eval\"][\"self_ensemble\"]\n    os.makedirs(output_dir, exist_ok=True)\n\n    t_start = time.time()\n    ckpt = torch.load(checkpoint_path, map_location=device)\n    mcfg = ckpt[\"config\"][\"model\"]\n    scale = ckpt[\"config\"][\"data\"][\"scale\"]\n    # NOTE: the original notebook built this model without in_channels/window_size, silently\n    # defaulting to a 3-channel model that would fail to load a 1-channel checkpoint. Fixed by\n    # rebuilding from the full saved model config, same as run_evaluation.\n    inf_model = build_model(\n        mcfg[\"variant\"], scale=scale, in_channels=mcfg[\"in_channels\"], window_size=mcfg[\"window_size\"],\n    ).to(device)\n    inf_model.load_state_dict(ckpt[\"model\"])\n    inf_model.eval()\n\n    files = sorted(f for f in os.listdir(input_dir) if os.path.splitext(f)[1].lower() in IMG_EXTS)\n    if not files:\n        raise FileNotFoundError(f\"No supported image files found in {input_dir}\")\n    print(f\"loaded model on {device}, scale={scale}, self_ensemble={self_ensemble}, {len(files)} images to restore\")\n\n    with torch.no_grad():\n        for i in range(0, len(files), batch_size):\n            batch_files = files[i:i + batch_size]\n            arrs = [load_array(os.path.join(input_dir, f)) for f in batch_files]\n            shapes = [a.shape[:2] for a in arrs]\n\n            same_shape = len(set(shapes)) == 1\n            groups = [batch_files] if same_shape else [[f] for f in batch_files]\n            arr_groups = [arrs] if same_shape else [[a] for a in arrs]\n\n            for g_files, g_arrs in zip(groups, arr_groups):\n                tin = torch.from_numpy(\n                    np.stack([a.transpose(2, 0, 1) for a in g_arrs]).astype(np.float32)\n                ).to(device, non_blocking=True)\n\n                tin_padded, pad_h, pad_w = pad_to_multiple(tin, multiple=8)\n                pred_fn = self_ensemble_predict if self_ensemble else (lambda m, x: m(x))\n                pred = pred_fn(inf_model, tin_padded)\n\n                if pad_h or pad_w:\n                    pred = pred[..., :pred.shape[-2] - pad_h * scale, :pred.shape[-1] - pad_w * scale]\n                if clip_output:\n                    pred = pred.clamp(0, 1)\n\n                pred_np = pred.permute(0, 2, 3, 1).cpu().numpy()\n                for fname, out_arr in zip(g_files, pred_np):\n                    save_array(os.path.join(output_dir, fname), out_arr)\n\n    total_time = time.time() - t_start\n    timing = {\n        \"num_images\": len(files), \"total_time_s\": total_time,\n        \"avg_time_per_image_s\": total_time / len(files), \"device\": device,\n        \"batch_size\": batch_size, \"self_ensemble\": self_ensemble,\n    }\n    with open(os.path.join(output_dir, \"_timing.json\"), \"w\") as f:\n        json.dump(timing, f, indent=2)\n\n    print(f\"done. {len(files)} images restored in {total_time:.2f}s \"\n          f\"({total_time/len(files)*1000:.1f} ms/image avg). Wrote outputs + _timing.json to {output_dir}\")\n\n\n# Example call -- edit paths and uncomment to run:\n# run_inference(\n#     input_dir=\"/kaggle/input/degradation-dataset/NoisyLR_test\",\n#     output_dir=\"/kaggle/working/restored_output\",\n#     checkpoint_path=best_ckpt,\n#     batch_size=8,\n# )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-12T04:54:53.754516Z","iopub.status.idle":"2026-08-12T04:54:53.754916Z","shell.execute_reply.started":"2026-08-12T04:54:53.754785Z","shell.execute_reply":"2026-08-12T04:54:53.754799Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}