{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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"},{"sourceId":213732047,"sourceType":"kernelVersion"}],"dockerImageVersionId":30822,"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\nimport torchvision.models as models\nfrom torchvision.utils import make_grid\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-20T08:03:02.470429Z","iopub.execute_input":"2024-12-20T08:03:02.470811Z","iopub.status.idle":"2024-12-20T08:03:06.719769Z","shell.execute_reply.started":"2024-12-20T08:03:02.470784Z","shell.execute_reply":"2024-12-20T08:03:06.718887Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 1834579290\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-20T08:03:06.720855Z","iopub.execute_input":"2024-12-20T08:03:06.721226Z","iopub.status.idle":"2024-12-20T08:03:06.77729Z","shell.execute_reply.started":"2024-12-20T08:03:06.721203Z","shell.execute_reply":"2024-12-20T08:03:06.776621Z"}},"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-20T08:03:06.778726Z","iopub.execute_input":"2024-12-20T08:03:06.778956Z","iopub.status.idle":"2024-12-20T08:03:06.782591Z","shell.execute_reply.started":"2024-12-20T08:03:06.778936Z","shell.execute_reply":"2024-12-20T08:03:06.781854Z"}},"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        self.image_paths = self.image_paths[:Conf.subset_size]\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-20T08:03:06.783586Z","iopub.execute_input":"2024-12-20T08:03:06.783879Z","iopub.status.idle":"2024-12-20T08:03:06.802019Z","shell.execute_reply.started":"2024-12-20T08:03:06.783848Z","shell.execute_reply":"2024-12-20T08:03:06.80108Z"}},"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-20T08:03:06.802582Z","iopub.execute_input":"2024-12-20T08:03:06.802838Z","iopub.status.idle":"2024-12-20T08:04:18.145618Z","shell.execute_reply.started":"2024-12-20T08:03:06.802817Z","shell.execute_reply":"2024-12-20T08:04:18.144715Z"}},"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(train_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T08:04:18.146373Z","iopub.execute_input":"2024-12-20T08:04:18.146611Z","iopub.status.idle":"2024-12-20T08:04:19.224582Z","shell.execute_reply.started":"2024-12-20T08:04:18.14659Z","shell.execute_reply":"2024-12-20T08:04:19.223639Z"}},"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-20T08:04:19.225559Z","iopub.execute_input":"2024-12-20T08:04:19.225904Z","iopub.status.idle":"2024-12-20T08:04:19.236182Z","shell.execute_reply.started":"2024-12-20T08:04:19.225871Z","shell.execute_reply":"2024-12-20T08:04:19.235503Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ConvBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel_size=3):\n        super(ConvBlock, self).__init__()\n        self.block = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=1, padding=1),\n            nn.LeakyReLU(0.2, inplace=True),\n            nn.Conv2d(out_channels, out_channels, kernel_size=kernel_size, stride=2, padding=1),\n            nn.LeakyReLU(0.2, inplace=True)\n        )\n\n    def forward(self, x):\n        return self.block(x)\n\nclass Discriminator(nn.Module):\n    def __init__(self, in_channels=3):\n        super(Discriminator, self).__init__()\n        \n        self.layers = nn.Sequential(\n            ConvBlock(in_channels, 64, kernel_size=3),\n            ConvBlock(64, 128, kernel_size=3),\n            ConvBlock(128, 256, kernel_size=3),\n            ConvBlock(256, 512, kernel_size=3),\n            ConvBlock(512, 512, kernel_size=3),\n            \n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(),\n            nn.Linear(512, 1024),\n            nn.LeakyReLU(0.2, inplace=True),\n            nn.Linear(1024, 1),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        return self.layers(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T08:04:19.237084Z","iopub.execute_input":"2024-12-20T08:04:19.237381Z","iopub.status.idle":"2024-12-20T08:04:19.25432Z","shell.execute_reply.started":"2024-12-20T08:04:19.237352Z","shell.execute_reply":"2024-12-20T08:04:19.253451Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class VGG19(nn.Module):\n    def __init__(self, feature_layer='relu5_4'):\n        super(VGG19, self).__init__()\n        self.vgg = self._load_vgg()\n        self.feature_layer = self._get_layer_index(feature_layer)\n        self.criterion = nn.MSELoss()\n\n        for param in self.vgg.parameters():\n            param.requires_grad = False\n\n    def _load_vgg(self):\n        vgg = models.vgg19(pretrained=True).features\n        return vgg\n\n    def _get_layer_index(self, layer_name):\n        layer_mapping = {\n            'relu1_1': 1, 'relu1_2': 3,\n            'relu2_1': 6, 'relu2_2': 8,\n            'relu3_1': 11, 'relu3_2': 13, 'relu3_4': 17,\n            'relu4_1': 20, 'relu4_2': 22, 'relu4_4': 26,\n            'relu5_1': 29, 'relu5_2': 31, 'relu5_4': 35\n        }\n        return layer_mapping[layer_name]\n\n    def forward(self, input_image, target):\n        input_image = self._vgg_preprocess(input_image)\n        target = self._vgg_preprocess(target)\n\n        input_features = self._extract_features(input_image)\n        target_features = self._extract_features(target)\n        \n        loss = self.criterion(input_features, target_features) / 12.75\n        return loss\n\n    def _extract_features(self, x):\n        for idx, layer in enumerate(self.vgg):\n            if idx == self.feature_layer: # Taking features before ReLU5_4 (idea from ESRGAN paper)\n                break\n            x = layer(x)\n        return x\n\n    def _vgg_preprocess(self, x):\n        mean = torch.tensor([0.485, 0.456, 0.406], device=x.device).view(1, 3, 1, 1)\n        std = torch.tensor([0.229, 0.224, 0.225], device=x.device).view(1, 3, 1, 1)\n        return (x - mean) / std","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T08:04:19.256355Z","iopub.execute_input":"2024-12-20T08:04:19.256612Z","iopub.status.idle":"2024-12-20T08:04:19.271122Z","shell.execute_reply.started":"2024-12-20T08:04:19.25659Z","shell.execute_reply":"2024-12-20T08:04:19.27039Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AdversarialLoss(nn.Module):\n    def __init__(self):\n        super(AdversarialLoss, self).__init__()\n\n    def forward(self, discriminator_output):\n        loss = -torch.log(discriminator_output + 1e-8).mean()\n        return loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T08:04:19.272283Z","iopub.execute_input":"2024-12-20T08:04:19.27258Z","iopub.status.idle":"2024-12-20T08:04:19.287607Z","shell.execute_reply.started":"2024-12-20T08:04:19.272538Z","shell.execute_reply":"2024-12-20T08:04:19.286951Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\ngenerator = SRResNet(num_channels=3, num_blocks=16, upscale_factor=Conf.scaling_factor).to(device)\ndiscriminator = Discriminator(in_channels=3).to(device)\n\nvgg_loss = VGG19(feature_layer='relu5_4').to(device)\nadversarial_loss = AdversarialLoss()\n\nbce_loss = nn.BCELoss()\n\noptimizer_G = optim.Adam(generator.parameters(), lr=1e-4, betas=(0.9, 0.999))\noptimizer_D = optim.Adam(discriminator.parameters(), lr=1e-4, betas=(0.9, 0.999))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T08:04:19.288438Z","iopub.execute_input":"2024-12-20T08:04:19.288634Z","iopub.status.idle":"2024-12-20T08:04:24.204527Z","shell.execute_reply.started":"2024-12-20T08:04:19.288616Z","shell.execute_reply":"2024-12-20T08:04:24.203838Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"generator.load_state_dict(torch.load('/kaggle/input/srresnet/SRResNet.pth', map_location=device))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T08:04:24.205374Z","iopub.execute_input":"2024-12-20T08:04:24.205708Z","iopub.status.idle":"2024-12-20T08:04:24.357869Z","shell.execute_reply.started":"2024-12-20T08:04:24.205661Z","shell.execute_reply":"2024-12-20T08:04:24.356895Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def display_images(low_res, super_res, high_res, epoch):\n    low_res = low_res[0].permute(1, 2, 0).cpu().clamp(0, 1)\n    super_res = ((super_res[0].permute(1, 2, 0).cpu().clamp(-1, 1) + 1) / 2.0)\n    high_res = ((high_res[0].permute(1, 2, 0).cpu().clamp(-1, 1) + 1) / 2.0)\n\n    plt.figure(figsize=(12, 4))\n    plt.subplot(1, 3, 1)\n    plt.title(\"Low-Resolution\")\n    plt.imshow(low_res)\n    plt.axis('off')\n\n    plt.subplot(1, 3, 2)\n    plt.title(\"Super-Resolution (Generated)\")\n    plt.imshow(super_res)\n    plt.axis('off')\n\n    plt.subplot(1, 3, 3)\n    plt.title(\"High-Resolution (Ground Truth)\")\n    plt.imshow(high_res)\n    plt.axis('off')\n\n    plt.suptitle(f\"Validation Results - Epoch {epoch+1}\")\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T08:04:24.358668Z","iopub.execute_input":"2024-12-20T08:04:24.35895Z","iopub.status.idle":"2024-12-20T08:04:24.364555Z","shell.execute_reply.started":"2024-12-20T08:04:24.358928Z","shell.execute_reply":"2024-12-20T08:04:24.363768Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_epochs = 50\nfor epoch in range(num_epochs):\n    generator.train()\n    discriminator.train()\n\n    train_d_loss, train_g_loss, train_content_loss, train_adv_loss = 0, 0, 0, 0\n    for x_tensor, y_tensor in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs} - Training\"):\n        x_tensor = x_tensor.to(device)\n        y_tensor = y_tensor.to(device)\n\n        ############################\n        # Train Discriminator\n        ############################\n        \n        optimizer_D.zero_grad()\n\n        real_output = discriminator(y_tensor)\n        real_labels = torch.ones_like(real_output).to(device)\n        d_real_loss = bce_loss(real_output, real_labels)\n\n        fake_images = generator(x_tensor).detach()\n        fake_output = discriminator(fake_images)\n        fake_labels = torch.zeros_like(fake_output).to(device)\n        d_fake_loss = bce_loss(fake_output, fake_labels)\n\n        d_loss = (d_real_loss + d_fake_loss) / 2\n        d_loss.backward()\n        optimizer_D.step()\n\n        ############################\n        # Train Generator\n        ############################\n        \n        optimizer_G.zero_grad()\n\n        fake_images = generator(x_tensor)\n        fake_output = discriminator(fake_images)\n\n        g_adv_loss = adversarial_loss(fake_output)\n        g_content_loss = vgg_loss(fake_images, y_tensor)\n\n        g_loss = g_content_loss + g_adv_loss\n        g_loss.backward()\n        optimizer_G.step()\n\n        train_d_loss += d_loss.item()\n        train_g_loss += g_loss.item()\n        train_content_loss += g_content_loss.item()\n        train_adv_loss += g_adv_loss.item()\n\n    ############################\n    # Validate Model\n    ############################\n    \n    generator.eval()\n    val_d_loss, val_content_loss, val_adv_loss = 0, 0, 0\n    sample_lr, sample_sr, sample_hr = None, None, None\n    with torch.no_grad():\n        for x_tensor, y_tensor in tqdm(val_loader, desc=\"Validation\"):\n            x_tensor = x_tensor.to(device)\n            y_tensor = y_tensor.to(device)\n\n            real_output = discriminator(y_tensor)\n            d_real_loss = bce_loss(real_output, torch.ones_like(real_output).to(device))\n\n            fake_images = generator(x_tensor)\n            fake_output = discriminator(fake_images)\n            d_fake_loss = bce_loss(fake_output, torch.zeros_like(fake_output).to(device))\n            d_loss = (d_real_loss + d_fake_loss) / 2\n\n            g_adv_loss = adversarial_loss(fake_output)\n            g_content_loss = vgg_loss(fake_images, y_tensor)\n\n            val_d_loss += d_loss.item()\n            val_content_loss += g_content_loss.item()\n            val_adv_loss += g_adv_loss.item()\n\n            if sample_lr is None:\n                sample_lr = x_tensor.cpu()\n                sample_sr = fake_images.cpu()\n                sample_hr = y_tensor.cpu()\n\n    train_d_loss /= len(train_loader)\n    train_g_loss /= len(train_loader)\n    train_content_loss /= len(train_loader)\n    train_adv_loss /= len(train_loader)\n\n    val_d_loss /= len(val_loader)\n    val_content_loss /= len(val_loader)\n    val_adv_loss /= len(val_loader)\n\n    ############################\n    # Logging\n    ############################\n    \n    print(f\"Epoch [{epoch+1}/{num_epochs}]\")\n    print(f\"  Train D Loss: {train_d_loss:.4f}, G Loss: {train_g_loss:.4f}, Content Loss: {train_content_loss:.4f}, Adv Loss: {train_adv_loss:.4f}\")\n    print(f\"  Val   D Loss: {val_d_loss:.4f}, Content Loss: {val_content_loss:.4f}, Adv Loss: {val_adv_loss:.4f}\")\n\n    ############################\n    # Display Images After Validation\n    ############################\n    \n    display_images(sample_lr, sample_sr, sample_hr, epoch)\n\n    ############################\n    # Save Models\n    ############################\n    \n    if (epoch + 1) % 2 == 0:\n        torch.save(generator.state_dict(), f\"generator_epoch_{epoch+1}.pth\")\n        torch.save(discriminator.state_dict(), f\"discriminator_epoch_{epoch+1}.pth\")\n        print(f\"Saved checkpoints at epoch {epoch+1}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T08:04:24.365547Z","iopub.execute_input":"2024-12-20T08:04:24.365856Z","iopub.status.idle":"2024-12-20T08:04:46.477076Z","shell.execute_reply.started":"2024-12-20T08:04:24.365833Z","shell.execute_reply":"2024-12-20T08:04:46.475925Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(generator.state_dict(), 'SRGAN.pth')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T08:04:46.47754Z","iopub.status.idle":"2024-12-20T08:04:46.477803Z","shell.execute_reply":"2024-12-20T08:04:46.477694Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}