{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":6799,"databundleVersionId":4225553,"sourceType":"competition"}],"dockerImageVersionId":30804,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\n\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\n\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torchvision.transforms import functional as F\nimport torchvision.transforms as T\n\nimport random\n\nfrom tqdm import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-18T17:14:02.506445Z","iopub.execute_input":"2024-12-18T17:14:02.506881Z","iopub.status.idle":"2024-12-18T17:14:07.447106Z","shell.execute_reply.started":"2024-12-18T17:14:02.506834Z","shell.execute_reply":"2024-12-18T17:14:07.446034Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 1834579291\n\nnp.random.seed(SEED)\nrandom.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed(SEED)\n    torch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-18T17:14:07.449368Z","iopub.execute_input":"2024-12-18T17:14:07.450326Z","iopub.status.idle":"2024-12-18T17:14:07.520651Z","shell.execute_reply.started":"2024-12-18T17:14:07.450281Z","shell.execute_reply":"2024-12-18T17:14:07.519622Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Conf:\n    train_dir = '/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train'\n    val_dir = '/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/val'\n    test_dir = '/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/test'\n\n    batch_size = 16\n\n    scaling_factor = 4\n    x_patch_size = 24\n    y_patch_size = scaling_factor * x_patch_size\n    \n    subset_size = 30000","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-18T17:14:07.521848Z","iopub.execute_input":"2024-12-18T17:14:07.522235Z","iopub.status.idle":"2024-12-18T17:14:07.528244Z","shell.execute_reply.started":"2024-12-18T17:14:07.522194Z","shell.execute_reply":"2024-12-18T17:14:07.527189Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ImageDataset(Dataset):\n    def __init__(self, image_dir, x_patch_size, y_patch_size, train_set=False):\n        self.image_dir = image_dir\n        self.x_patch_size = x_patch_size\n        self.y_patch_size = y_patch_size\n        self.train_set = train_set\n\n        self.to_tensor = T.ToTensor()\n        \n        if self.train_set:\n            subdirs = [os.path.join(image_dir, d) for d in os.listdir(image_dir) \n                       if os.path.isdir(os.path.join(image_dir, d))]\n\n            self.image_paths = []\n            for d in subdirs:\n                for fname in os.listdir(d):\n                    if fname.lower().endswith(('.png', '.jpg', '.jpeg')):\n                        self.image_paths.append(os.path.join(d, fname))\n        else:\n            self.image_paths = [os.path.join(image_dir, fname) for fname in os.listdir(image_dir)\n                                if fname.lower().endswith(('.png', '.jpg', '.jpeg'))]\n\n        # if self.train_set:\n        self.image_paths = self.image_paths[:Conf.subset_size]\n\n        self.image_paths.sort()\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        img_path = self.image_paths[idx]\n        image = Image.open(img_path).convert('RGB')\n        w, h = image.size\n\n        if (w < self.y_patch_size or h < self.y_patch_size):\n            image = image.resize((self.y_patch_size, self.y_patch_size), Image.BICUBIC)\n            w, h = image.size\n\n        if self.train_set:\n            x1 = random.randint(0, w - self.y_patch_size)\n            y1 = random.randint(0, h - self.y_patch_size)\n        else:\n            x1 = (w - self.y_patch_size) // 2\n            y1 = (h - self.y_patch_size) // 2\n\n        y_patch = image.crop((x1, y1, x1 + self.y_patch_size, y1 + self.y_patch_size))\n        x_patch = y_patch.resize((self.x_patch_size, self.x_patch_size), Image.BICUBIC)\n\n        x_tensor = self.to_tensor(x_patch)\n        y_tensor = self.to_tensor(y_patch)\n\n        y_tensor = y_tensor * 2.0 - 1.0\n        \n        return x_tensor, y_tensor        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-18T17:14:07.529396Z","iopub.execute_input":"2024-12-18T17:14:07.529713Z","iopub.status.idle":"2024-12-18T17:14:07.544399Z","shell.execute_reply.started":"2024-12-18T17:14:07.529685Z","shell.execute_reply":"2024-12-18T17:14:07.543547Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n\ntrain_dataset = ImageDataset(Conf.train_dir, Conf.x_patch_size, Conf.y_patch_size, True)\nval_dataset = ImageDataset(Conf.val_dir, Conf.x_patch_size, Conf.y_patch_size)\ntest_dataset = ImageDataset(Conf.val_dir, Conf.x_patch_size, Conf.y_patch_size)\n\ntrain_loader = DataLoader(train_dataset, batch_size = Conf.batch_size, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size = Conf.batch_size, shuffle=True)\ntest_loader = DataLoader(test_dataset, batch_size = Conf.batch_size, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-18T17:14:07.546578Z","iopub.execute_input":"2024-12-18T17:14:07.546903Z","iopub.status.idle":"2024-12-18T17:14:24.996348Z","shell.execute_reply.started":"2024-12-18T17:14:07.546875Z","shell.execute_reply":"2024-12-18T17:14:24.995038Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def display_samples(loader, num_samples=5):\n    samples_displayed = 0\n    plt.figure(figsize=(12, 6))\n    for x, y in loader:\n        for i in range(x.size(0)):\n            if samples_displayed >= num_samples:\n                plt.show()\n                return\n            \n            plt.subplot(2, num_samples, samples_displayed + 1)\n            plt.imshow(x[i].permute(1, 2, 0).numpy())\n            plt.title(\"LR Image\")\n            plt.axis('off')\n\n            plt.subplot(2, num_samples, samples_displayed + 1 + num_samples)\n            plt.imshow(((y + 1) / 2.0)[i].permute(1, 2, 0).numpy(), cmap='gray')\n            plt.title(\"HR Image\")\n            plt.axis('off')\n            \n            samples_displayed += 1\n\ndisplay_samples(test_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-18T17:14:24.997685Z","iopub.execute_input":"2024-12-18T17:14:24.998079Z","iopub.status.idle":"2024-12-18T17:14:26.284313Z","shell.execute_reply.started":"2024-12-18T17:14:24.998037Z","shell.execute_reply":"2024-12-18T17:14:26.283217Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ResidualBlock(nn.Module):\n    def __init__(self, channels):\n        super(ResidualBlock, self).__init__()\n        self.block = nn.Sequential(\n            nn.Conv2d(channels, channels, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(channels, channels, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(channels)\n        )\n    \n    def forward(self, x):\n        return x + self.block(x)\n\nclass SRResNet(nn.Module):\n    def __init__(self, num_channels=3, num_blocks=16, upscale_factor=4):\n        super(SRResNet, self).__init__()\n        self.conv1 = nn.Sequential(\n            nn.Conv2d(num_channels, 64, kernel_size=9, stride=1, padding=4),\n            nn.ReLU(inplace=True)\n        )\n\n        self.res_blocks = nn.Sequential(\n            *[ResidualBlock(64) for _ in range(num_blocks)]\n        )\n\n        self.conv2 = nn.Sequential(\n            nn.Conv2d(64, 64, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(64)\n        )\n\n        self.upscale = nn.Sequential(\n            nn.Conv2d(64, 64 * upscale_factor, kernel_size=3, stride=1, padding=1),\n            nn.PixelShuffle(int(np.sqrt(upscale_factor))),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(64, 64 * upscale_factor, kernel_size=3, stride=1, padding=1),\n            nn.PixelShuffle(int(np.sqrt(upscale_factor))),\n            nn.ReLU(inplace=True),\n        )\n\n        self.conv3 = nn.Sequential(\n            nn.Conv2d(64, num_channels, kernel_size=9, stride=1, padding=4),\n            nn.Tanh()\n        )\n\n    def forward(self, x):\n        x = self.conv1(x)\n        residual = self.res_blocks(x)\n        x = self.conv2(residual) + x\n        x = self.upscale(x)\n        x = self.conv3(x)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-18T17:14:26.285824Z","iopub.execute_input":"2024-12-18T17:14:26.286217Z","iopub.status.idle":"2024-12-18T17:14:26.303825Z","shell.execute_reply.started":"2024-12-18T17:14:26.286177Z","shell.execute_reply":"2024-12-18T17:14:26.30255Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = SRResNet(num_channels=3, num_blocks=16, upscale_factor=Conf.scaling_factor).to(device)\ncriterion = nn.MSELoss()\noptimizer = optim.Adam(model.parameters(), lr=1e-4)\nscheduler = CosineAnnealingLR(optimizer, T_max=50, eta_min=1e-7)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-18T17:14:26.305143Z","iopub.execute_input":"2024-12-18T17:14:26.305586Z","iopub.status.idle":"2024-12-18T17:14:26.515372Z","shell.execute_reply.started":"2024-12-18T17:14:26.305543Z","shell.execute_reply":"2024-12-18T17:14:26.514253Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model(model, dataloader, criterion, optimizer, device):\n    model.train()\n    total_loss = 0.0\n    for lr_images, hr_images in tqdm(dataloader, desc=\"Training\", leave=False):\n        lr_images, hr_images = lr_images.to(device), hr_images.to(device)\n        optimizer.zero_grad()\n        sr_images = model(lr_images)\n        loss = criterion(sr_images, hr_images)\n        loss.backward()\n        optimizer.step()\n        total_loss += loss.item()\n    return total_loss / len(dataloader)\n\ndef validate_model(model, dataloader, criterion, device):\n    model.eval()\n    total_loss = 0.0\n    with torch.no_grad():\n        for lr_images, hr_images in tqdm(dataloader, desc=\"Validating\", leave=False):\n            lr_images, hr_images = lr_images.to(device), hr_images.to(device)\n            sr_images = model(lr_images)\n            loss = criterion(sr_images, hr_images)\n            total_loss += loss.item()\n    return total_loss / len(dataloader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-18T17:14:26.516776Z","iopub.execute_input":"2024-12-18T17:14:26.517203Z","iopub.status.idle":"2024-12-18T17:14:26.525931Z","shell.execute_reply.started":"2024-12-18T17:14:26.517162Z","shell.execute_reply":"2024-12-18T17:14:26.524843Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"epochs = 25\nfor epoch in range(epochs):\n    train_loss = train_model(model, train_loader, criterion, optimizer, device)\n    val_loss = validate_model(model, val_loader, criterion, device)\n    scheduler.step()\n    print(f\"Epoch {epoch + 1}/{epochs}, Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}\")\n\n    model.eval()\n    with torch.no_grad():\n        lr_sample, hr_sample = next(iter(val_loader))\n        \n        lr_sample = lr_sample.to(device)\n        hr_sample = hr_sample.to(device)\n        \n        sr_sample = model(lr_sample)\n        \n        lr_img = lr_sample[0].cpu().numpy().transpose(1, 2, 0)\n        sr_img = sr_sample[0].cpu().numpy().transpose(1, 2, 0)\n        hr_img = hr_sample[0].cpu().numpy().transpose(1, 2, 0)\n        \n        fig, axs = plt.subplots(1, 3, figsize=(15, 5))\n        \n        axs[0].imshow(lr_img.clip(0,1))\n        axs[0].set_title(\"LR Image\")\n        axs[0].axis('off')\n        \n        axs[1].imshow(((sr_img.clip(-1,1)) + 1) / 2.0)\n        axs[1].set_title(\"SR Image\")\n        axs[1].axis('off')\n        \n        axs[2].imshow(((hr_img.clip(-1,1)) + 1) / 2.0)\n        axs[2].set_title(\"HR Image\")\n        axs[2].axis('off')\n        \n        plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-18T17:14:26.527336Z","iopub.execute_input":"2024-12-18T17:14:26.52795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(model.state_dict(), 'SRResNet.pth')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}