{"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":"none","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":6654895,"sourceType":"datasetVersion","datasetId":3840540},{"sourceId":6654988,"sourceType":"datasetVersion","datasetId":3840578},{"sourceId":6655032,"sourceType":"datasetVersion","datasetId":3840603},{"sourceId":6655122,"sourceType":"datasetVersion","datasetId":3840660},{"sourceId":6655159,"sourceType":"datasetVersion","datasetId":3840681},{"sourceId":7116619,"sourceType":"datasetVersion","datasetId":4104230},{"sourceId":150637573,"sourceType":"kernelVersion"}],"dockerImageVersionId":30588,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"[Tellez *et al.* (2020)](https://arxiv.org/abs/1811.02840) proposed two-stage method to build deep neural networks for gigapixel distopathology image analysis solely using image-level labels.\n\n<img src=\"https://storage.googleapis.com/ubco/Screenshot%202023-11-16%20at%2011.24.50.png\">\n\nThe encoder compress gigapixel image to much smaller feature space and the classifier is trained feeding the features created by the encoder.\nThey tested three different encoding methods: variational autoencoder, contrastive training, and bidirectional GAN.","metadata":{}},{"cell_type":"markdown","source":"My previous notebooks trained [VAE](https://www.kaggle.com/code/emiz6413/neural-image-compression-with-vae?scriptVersionId=151666102) and [contrastive learning](https://www.kaggle.com/code/emiz6413/neural-image-compression-with-barlowtwins) models.\n\nIn this notebook, a bidirectional GAN (BiGAN) is trained. BiGAN was proposed in the paper [\"Adversarial Feature Learning\" (Donahue *et al,*, 2016)](https://arxiv.org/abs/1605.09782v7) at ICLR 2017.\n\nGAN trains *generator* and *discriminator* to compete against each other by plaing a zero-sum game. The generator takes an input sampled from a fixed distribution (e.g. normal distribution) and generate as realistic data as possible. Whereas the discriminator tries to predict whether an input is fake or real.\n\n<img src=\"https://storage.googleapis.com/ubco/Generative_Adversarial_Network_illustration.svg\">\n\nAlthough typical GANs can learn the forward mapping of latent representation to \ndata, they have no means for inverse mapping - projecting data back into the latent space.\nBiGAN introduces additional *encoder* to encode a data to latent space, and the discriminator takes not only data but also latent vector.\n\n<img src=\"https://storage.googleapis.com/ubco/BiGAN_diagram.png\">\n\nOnce they are trained, a joint of real data and its latent vector encoded by the encoder is indistinguishable from a joint of generated data and true noise. This way, the encoder can learn inverse mapping.\n\nThe encoder takes a patch of 64x64 pixels and map to 128 dimension latent vector, therefore, the compression rate of this model is $\\frac{3 \\times 64 \\times 64}{128} = 96$","metadata":{}},{"cell_type":"code","source":"from itertools import chain\nfrom pathlib import Path\nfrom typing import Iterator, Literal, TypeVar\nimport random\nimport os\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch import nn, optim, Tensor\nfrom torch.nn import functional as F\nfrom torch.utils.data import IterableDataset, DataLoader\nfrom torch.nn.utils.parametrizations import spectral_norm as _spectral_norm\nfrom torchvision import io, transforms\nfrom torchvision.utils import make_grid\nfrom tqdm.auto import tqdm\nimport wandb\nimport matplotlib.pyplot as plt\nfrom kaggle_secrets import UserSecretsClient","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-04T02:26:50.048274Z","iopub.execute_input":"2023-12-04T02:26:50.048717Z","iopub.status.idle":"2023-12-04T02:26:50.471917Z","shell.execute_reply.started":"2023-12-04T02:26:50.048677Z","shell.execute_reply":"2023-12-04T02:26:50.470507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMG_SIZE = 64\nINITIAL_EPOCH = 0\nLAST_EPOCH = 0  # I trained the model for 150 epochs\nBATCH_SIZE = 256\nLOSS = \"Hinge\"  # Hinge loss converges better than BCE\nDISC_SN_ENABLED = True  # Spectral Normalization in the discriminator\nDISC_BN_ENABLED = False  # Batch Normalization in the discriminator\nAMP = torch.cuda.is_available()\nDEVICE = torch.device(\"cuda\") if torch.cuda.is_available() else torch.device(\"cpu\")","metadata":{"execution":{"iopub.status.busy":"2023-12-04T02:12:59.149548Z","iopub.execute_input":"2023-12-04T02:12:59.149961Z","iopub.status.idle":"2023-12-04T02:12:59.157261Z","shell.execute_reply.started":"2023-12-04T02:12:59.14992Z","shell.execute_reply":"2023-12-04T02:12:59.156154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Utils","metadata":{}},{"cell_type":"code","source":"class AverageMeter:\n    def __init__(self) -> None:\n        self.reset()\n\n    def reset(self) -> None:\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val: float, n: int = 1) -> None:\n        self.sum += val * n\n        self.count += n\n\n    @property\n    def average(self) -> float:\n        if self.count == 0:\n            return 0.0\n        return self.sum / self.count\n\n\nclass HingeLoss(nn.Module):\n    def __init__(self, for_discriminator: bool) -> None:\n        super().__init__()\n        self.for_discriminator = for_discriminator\n\n    def forward(self, input: torch.Tensor, target: torch.Tensor) -> torch.Tensor:\n        if self.for_discriminator:\n            return self.compute_loss_for_discriminator(input, target)\n        else:\n            return self.compute_loss_for_generator(input, target)\n\n    @staticmethod\n    def compute_loss_for_discriminator(input: torch.Tensor, target: torch.Tensor) -> torch.Tensor:\n        if all(target == 1):\n            return torch.mean(F.relu(1.0 - input))\n        elif all(target == 0):\n            return torch.mean(F.relu(1.0 + input))\n        else:\n            raise ValueError(\"invalid target value\")\n\n    @staticmethod\n    def compute_loss_for_generator(input: torch.Tensor, target: torch.Tensor) -> torch.Tensor:\n        if all(target == 1):\n            return -input.mean()\n        elif all(target == 0):\n            return input.mean()\n        else:\n            raise ValueError(\"invalid target value\")\n\n\n\nT = TypeVar(\"T\")\n\n\ndef spectral_norm(module: T, enabled: bool = True) -> T:\n    if enabled:\n        module = _spectral_norm(module)\n    return module\n\n\nLOSS_TYPE = Literal[\"BCE\", \"Hinge\"]\n\n\ndef image_colorfulness(image: torch.Tensor) -> torch.Tensor:\n    \"\"\"\n    Args:\n        image: Tensor RGB order CHW format\n\n    Returns:\n        float: colorfulness\n        \n    .. _For details, please refer:\n        https://www.kaggle.com/code/emiz6413/image-colorfulness-to-filter-uninformative-patches\n    \"\"\"\n    r, g, b = image\n    rg = torch.abs(r - g)\n    yb = torch.abs(0.5 * (r + g) - b)\n    rb_mean = torch.mean(rg)\n    rb_std = torch.std(rg)\n    yb_mean = torch.mean(yb)\n    yb_std = torch.std(yb)\n    std_root = torch.sqrt((rb_std ** 2) + (yb_std ** 2))\n    mean_root = torch.sqrt((rb_mean ** 2) + (yb_mean ** 2))\n    return std_root + (0.3 * mean_root)","metadata":{"execution":{"iopub.status.busy":"2023-12-04T02:26:52.487275Z","iopub.execute_input":"2023-12-04T02:26:52.487713Z","iopub.status.idle":"2023-12-04T02:26:52.505638Z","shell.execute_reply.started":"2023-12-04T02:26:52.487679Z","shell.execute_reply":"2023-12-04T02:26:52.504105Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Modules","metadata":{}},{"cell_type":"code","source":"class EncoderBlock(nn.Module):\n    def __init__(\n        self,\n        in_channels: int,\n        out_channels: int,\n        kernel_size: int = 4,\n        stride: int = 2,\n        padding: int = 1,\n        bias: bool = False,\n    ) -> None:\n        super().__init__()\n        self.conv = nn.Conv2d(\n            in_channels=in_channels,\n            out_channels=out_channels,\n            kernel_size=kernel_size,\n            stride=stride,\n            padding=padding,\n            bias=bias,\n        )\n        self.bn = nn.BatchNorm2d(num_features=out_channels)\n        self.activation = nn.ReLU(inplace=True)\n\n    def forward(self, x: Tensor) -> Tensor:\n        x = self.conv(x)\n        x = self.bn(x)\n        x = self.activation(x)\n        return x\n\n\nclass Encoder64(nn.Module):\n    def __init__(self, latent_dim: int = 128, in_channels: int = 3) -> None:\n        super().__init__()\n        self.latent_dim = latent_dim  # for compatibility\n        self.layers = nn.Sequential(\n            EncoderBlock(in_channels=in_channels, out_channels=64),  # 64 -> 32\n            EncoderBlock(in_channels=64, out_channels=128),  # 32 -> 16\n            EncoderBlock(in_channels=128, out_channels=256),  # 16 -> 8\n            EncoderBlock(in_channels=256, out_channels=512),  # 8 -> 4\n            EncoderBlock(in_channels=512, out_channels=1024, stride=1, padding=0),  # 4 -> 1\n            nn.Conv2d(in_channels=1024, out_channels=latent_dim, kernel_size=1),\n        )\n\n    def forward(self, x: Tensor) -> Tensor:\n        return self.layers(x)\n\n\nclass GeneratorBlock(nn.Module):\n    def __init__(\n        self,\n        in_channles: int,\n        out_channels: int,\n        kernel_size: int = 4,\n        stride: int = 2,\n        padding: int = 1,\n        bias: bool = False,\n    ) -> None:\n        super().__init__()\n        self.conv_t = nn.ConvTranspose2d(\n            in_channels=in_channles,\n            out_channels=out_channels,\n            kernel_size=kernel_size,\n            stride=stride,\n            padding=padding,\n            bias=bias,\n        )\n        self.bn = nn.BatchNorm2d(num_features=out_channels)\n        self.activation = nn.ReLU(inplace=True)\n\n    def forward(self, x: Tensor) -> Tensor:\n        x = self.conv_t(x)\n        x = self.bn(x)\n        x = self.activation(x)\n        return x\n\n\nclass Generator64(nn.Module):\n    def __init__(self, latent_dim: int = 128, out_channels: int = 3) -> None:\n        super().__init__()\n        self.latent_dim = latent_dim\n        self.layers = nn.Sequential(\n            GeneratorBlock(in_channles=latent_dim, out_channels=1024, stride=1, padding=0),  # 1 -> 4\n            GeneratorBlock(in_channles=1024, out_channels=512),  # 4 -> 8\n            GeneratorBlock(in_channles=512, out_channels=256),  # 8 -> 16\n            GeneratorBlock(in_channles=256, out_channels=128),  # 16 -> 32\n            nn.ConvTranspose2d(\n                in_channels=128, out_channels=out_channels, kernel_size=4, stride=2, padding=1, bias=False\n            ),  # 32 -> 64\n            nn.Tanh(),\n        )\n\n    def forward(self, x: Tensor) -> Tensor:\n        return self.layers(x)\n\n\nclass DiscriminatorBlock(nn.Module):\n    def __init__(\n        self,\n        in_channles: int,\n        out_channels: int,\n        kernel_size: int = 4,\n        stride: int = 2,\n        padding: int = 1,\n        bias: bool = True,\n        sn_enabled: bool = False,\n        bn_enabled: bool = True,\n    ) -> None:\n        super().__init__()\n        self.conv = spectral_norm(\n            nn.Conv2d(\n                in_channels=in_channles,\n                out_channels=out_channels,\n                kernel_size=kernel_size,\n                stride=stride,\n                padding=padding,\n                bias=bias,\n            ),\n            enabled=sn_enabled,\n        )\n        self.bn = nn.BatchNorm2d(num_features=out_channels) if bn_enabled else nn.Identity()\n        self.activation = nn.LeakyReLU(negative_slope=0.2, inplace=True)\n\n    def forward(self, x: Tensor) -> Tensor:\n        x = self.conv(x)\n        x = self.bn(x)\n        x = self.activation(x)\n        return x\n\n\nclass Discriminator64(nn.Module):\n    def __init__(\n        self,\n        in_channels: int = 3,\n        latent_dim: int = 128,\n        sn_enabled: bool = False,\n        bn_enabled: bool = True,\n    ) -> None:\n        super().__init__()\n        self.latent_dim = latent_dim\n        self.x_mapping = nn.Sequential(\n            DiscriminatorBlock(\n                in_channles=in_channels, out_channels=64, sn_enabled=sn_enabled, bn_enabled=bn_enabled\n            ),  # 64 -> 32\n            DiscriminatorBlock(\n                in_channles=64, out_channels=128, sn_enabled=sn_enabled, bn_enabled=bn_enabled\n            ),  # 32 -> 16\n            DiscriminatorBlock(\n                in_channles=128, out_channels=256, sn_enabled=sn_enabled, bn_enabled=bn_enabled\n            ),  # 16 -> 8\n            DiscriminatorBlock(\n                in_channles=256, out_channels=512, sn_enabled=sn_enabled, bn_enabled=bn_enabled\n            ),  # 8 -> 4\n            DiscriminatorBlock(\n                in_channles=512, out_channels=1024, stride=1, padding=0, sn_enabled=sn_enabled, bn_enabled=bn_enabled\n            ),  # 4 -> 1\n        )\n\n        self.z_mapping = nn.Sequential(\n            DiscriminatorBlock(\n                in_channles=latent_dim,\n                out_channels=512,\n                kernel_size=1,\n                stride=1,\n                padding=0,\n                sn_enabled=sn_enabled,\n                bn_enabled=bn_enabled,\n            ),\n            DiscriminatorBlock(\n                in_channles=512,\n                out_channels=512,\n                kernel_size=1,\n                stride=1,\n                padding=0,\n                sn_enabled=sn_enabled,\n                bn_enabled=bn_enabled,\n            ),\n        )\n\n        self.joint_mapping = nn.Sequential(\n            DiscriminatorBlock(\n                in_channles=1024 + 512,\n                out_channels=2048,\n                kernel_size=1,\n                stride=1,\n                padding=0,\n                sn_enabled=sn_enabled,\n                bn_enabled=bn_enabled,\n            ),\n            DiscriminatorBlock(\n                in_channles=2048,\n                out_channels=2048,\n                kernel_size=1,\n                stride=1,\n                padding=0,\n                sn_enabled=sn_enabled,\n                bn_enabled=bn_enabled,\n            ),\n            nn.Conv2d(in_channels=2048, out_channels=1, kernel_size=1),\n        )\n\n    def forward(self, x: Tensor, z: Tensor) -> Tensor:\n        x = self.x_mapping(x)\n        z = self.z_mapping(z)\n        joint = torch.concat((x, z), dim=1)\n        joint = self.joint_mapping(joint)\n        return joint\n","metadata":{"execution":{"iopub.status.busy":"2023-12-04T02:20:07.590063Z","iopub.execute_input":"2023-12-04T02:20:07.590536Z","iopub.status.idle":"2023-12-04T02:20:07.618146Z","shell.execute_reply.started":"2023-12-04T02:20:07.590503Z","shell.execute_reply":"2023-12-04T02:20:07.616999Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### BiGAN","metadata":{}},{"cell_type":"code","source":"class BiGAN(nn.Module):\n    def __init__(\n        self,\n        encoder: nn.Module,\n        generator: nn.Module,\n        discriminator: nn.Module,\n        device: torch.device = torch.device(\"cpu\"),\n        amp: bool = False,\n        eval_amp: bool = False,\n        loss_type: Literal[\"BCE\", \"Hinge\"] = \"Hinge\",\n        disc_iters: int = 2,\n        ge_iters: int = 1,\n    ) -> None:\n        super().__init__()\n        self.encoder = encoder\n        self.generator = generator\n        self.discriminator = discriminator\n        self.latent_dim = self.encoder.latent_dim\n        self.loss_type = loss_type\n        self.ge_criterion = self.create_eg_criterion()\n        self.d_criterion = self.create_d_criterion()\n        self.criterion = nn.BCEWithLogitsLoss()\n        self.ge_optimizer = self.create_eg_optimizer()\n        self.d_optimizer = self.create_d_optimizer()\n        self.device = device\n        self.scaler = torch.cuda.amp.grad_scaler.GradScaler(enabled=amp)\n        self.eval_amp = eval_amp\n        self.disc_iters = disc_iters\n        self.ge_iters = ge_iters\n        self.init_parameters()\n\n    def init_parameters(self) -> None:\n        for m in self.modules():\n            if isinstance(m, (nn.Conv2d, nn.ConvTranspose2d, nn.Linear)):\n                nn.init.normal_(m.weight.data, 0.0, 0.02)\n                if hasattr(m, \"bias\") and m.bias is not None:\n                    nn.init.constant_(m.bias, 0.0)\n            elif isinstance(m, (nn.BatchNorm1d, nn.BatchNorm2d)):\n                nn.init.normal_(m.weight.data, 1.0, 0.02)\n                nn.init.constant_(m.bias.data, 0.0)\n\n    def create_eg_optimizer(self, lr: float = 5e-5, betas: tuple[float, float] = (0.0, 0.999)) -> optim.Optimizer:\n        self.ge_optimizer = optim.Adam(\n            chain(self.encoder.parameters(), self.generator.parameters()), lr=lr, betas=betas\n        )\n        return self.ge_optimizer\n\n    def create_d_optimizer(self, lr: float = 2e-4, betas: tuple[float, float] = (0.0, 0.999)) -> optim.Optimizer:\n        \"\"\"\n        Note:\n            Setting Discriminator's learning rate larger converges faster\n\n        .. _GANs Trained by a Two Time-Scale Update Rule Converge to a Local Nash Equilibrium (TTUR):\n            Heusel et al. (2017)\n            https://arxiv.org/abs/1706.08500\n        \"\"\"\n        self.d_optimizer = optim.Adam(self.discriminator.parameters(), lr=lr, betas=betas)\n        return self.d_optimizer\n\n    def create_eg_criterion(self) -> nn.Module:\n        if self.loss_type == \"BCE\":\n            return nn.BCEWithLogitsLoss()\n        if self.loss_type == \"Hinge\":\n            return HingeLoss(for_discriminator=False)\n\n    def create_d_criterion(self) -> nn.Module:\n        if self.loss_type == \"BCE\":\n            return nn.BCEWithLogitsLoss()\n        if self.loss_type == \"Hinge\":\n            return HingeLoss(for_discriminator=True)\n\n    def encode(self, x: torch.Tensor) -> torch.Tensor:\n        return self.encoder(x)\n\n    def generate(self, z: torch.Tensor) -> torch.Tensor:\n        return self.generator(z)\n\n    def reconstruct(self, x: torch.Tensor) -> torch.Tensor:\n        return self.generate(self.encode(x))\n\n    def discriminate(\n        self, x: torch.Tensor, z_hat: torch.Tensor, x_tilde: torch.Tensor, z: torch.Tensor\n    ) -> tuple[torch.Tensor, torch.Tensor]:\n        x = torch.cat([x, x_tilde], dim=0)\n        z = torch.cat([z_hat, z], dim=0)\n        output = self.discriminator(x, z)\n        real_preds, tilde_preds = torch.tensor_split(output, 2, dim=0)\n        return real_preds, tilde_preds\n\n    @torch.no_grad()\n    def evaluate(self, eval_loader: DataLoader) -> float:\n        rec_loss_meter = AverageMeter()\n        pbar = tqdm(total=len(eval_loader), leave=False)\n        self.eval()\n        for x in eval_loader:\n            x = x.to(self.device)\n            with torch.cuda.amp.autocast(enabled=self.eval_amp):\n                reconstructed = self.reconstruct(x)\n            mse = nn.functional.mse_loss(input=reconstructed, target=x)\n            rec_loss_meter.update(mse.item())\n            pbar.set_description(f\"reconstruction loss: {rec_loss_meter.average:.3f}\")\n            pbar.update()\n        pbar.close()\n        return rec_loss_meter.average\n\n    def train_single_epoch(self, train_loader: DataLoader) -> tuple[float, float]:\n        ge_loss_meter = AverageMeter()\n        d_loss_meter = AverageMeter()\n        pbar = tqdm(total=len(train_loader), leave=False)\n        self.train()\n        d_iter = 0\n        ge_iter = 0\n        for x in train_loader:\n            x = x.to(self.device)\n            if d_iter < self.disc_iters:\n                d_loss = self.train_disc(x)\n                d_loss_meter.update(d_loss.item())\n                d_iter += 1\n            else:\n                ge_loss = self.train_ge(x)\n                ge_loss_meter.update(ge_loss.item())\n                ge_iter += 1\n\n                if ge_iter == self.ge_iters:\n                    d_iter = 0\n                    ge_iter = 0\n\n            pbar.set_description(f\"GE loss: {ge_loss_meter.average:.3f}. D loss: {d_loss_meter.average:.3f}\")\n            pbar.update()\n        pbar.close()\n\n        return ge_loss_meter.average, d_loss_meter.average\n\n    def train_disc(self, x_real: torch.Tensor) -> torch.Tensor:\n        \"\"\"Train the discriminator\"\"\"\n        self.d_optimizer.zero_grad()\n\n        y_real = torch.ones((x_real.size(0), 1), device=x_real.device)\n        y_fake = torch.zeros_like(y_real)\n        z_fake = torch.randn(x_real.size(0), self.latent_dim, 1, 1, device=x_real.device)\n        with torch.cuda.amp.autocast(enabled=self.scaler.is_enabled()):\n            with torch.no_grad():\n                z_real = self.encode(x_real)\n                x_fake = self.generate(z_fake)\n\n            real_preds, fake_preds = self.discriminate(x_real, z_real, x_fake, z_fake)\n            d_loss: torch.Tensor = self.d_criterion(real_preds.view(-1, 1), y_real) + self.d_criterion(\n                fake_preds.view(-1, 1), y_fake\n            )\n\n        self.scaler.scale(d_loss).backward()\n        self.scaler.step(self.d_optimizer)\n        self.scaler.update()\n        return d_loss\n\n    def train_ge(self, x_real: torch.Tensor) -> torch.Tensor:\n        \"\"\"Train the generator and encoder\"\"\"\n        self.ge_optimizer.zero_grad()\n\n        y_real = torch.ones((x_real.size(0), 1), device=x_real.device)\n        y_fake = torch.zeros_like(y_real)\n        z_fake = torch.randn(x_real.size(0), self.latent_dim, 1, 1, device=x_real.device)\n\n        with torch.cuda.amp.autocast(enabled=self.scaler.is_enabled()):\n            z_real = self.encode(x_real)\n            x_fake = self.generate(z_fake)\n\n            real_preds, fake_preds = self.discriminate(x_real, z_real, x_fake, z_fake)\n            ge_loss: torch.Tensor = self.ge_criterion(fake_preds.view(-1, 1), y_real) + self.ge_criterion(\n                real_preds.view(-1, 1), y_fake\n            )\n\n        self.scaler.scale(ge_loss).backward()\n        self.scaler.step(self.ge_optimizer)\n        self.scaler.update()\n        return ge_loss","metadata":{"execution":{"iopub.status.busy":"2023-12-04T02:21:43.048971Z","iopub.execute_input":"2023-12-04T02:21:43.049384Z","iopub.status.idle":"2023-12-04T02:21:43.081135Z","shell.execute_reply.started":"2023-12-04T02:21:43.049349Z","shell.execute_reply":"2023-12-04T02:21:43.080032Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataset and transforms","metadata":{}},{"cell_type":"code","source":"class UBCODataset(IterableDataset):\n    normalize = transforms.Normalize(mean=[0.5]*3, std=[0.5]*3)\n    denormalize = transforms.Normalize(mean=[-1]*3, std=[2,]*3)\n\n    def __init__(\n        self,\n        image_dirs: list[Path],\n        is_train: bool,\n        shuffle: bool,\n        image_size: int,\n        colorfulness_thresh: float = 5.\n    ) -> None:\n        self.image_paths = list(chain.from_iterable([d.glob(\"*.png\") for d in tqdm(image_dirs)]))\n        self.colorfulness_thresh = colorfulness_thresh\n        self.transforms = transforms.RandomCrop(size=image_size) if is_train else transforms.CenterCrop(size=image_size)\n        self.shuffle = shuffle\n\n    def __len__(self) -> int:\n        \"\"\"\n        Notes:\n            This is just a workaround for tqdm to work.\n            Black and white input will be skipped, thus StopIteration may be raised before reaching this number.\n        \"\"\"\n        return len(self.image_paths)\n\n    def __iter__(self) -> Iterator[torch.Tensor]:\n        worker_info = torch.utils.data.get_worker_info()\n        if worker_info is None:\n            image_paths = self.image_paths\n        else:\n            w_id = worker_info.id\n            n_workers = worker_info.num_workers\n            image_paths = self.image_paths[w_id::n_workers]\n        if self.shuffle:\n            random.shuffle(image_paths)\n        for img_path in image_paths:\n            img = self.read_image(img_path)\n            cropped = self.transforms(img)\n            if image_colorfulness(cropped) < self.colorfulness_thresh:\n                # if the cropped image is black and white, discard it\n                continue\n            yield self.normalize(cropped / 255.)\n\n    def read_image(self, image_path: Path) -> torch.FloatTensor:\n        return io.read_image(str(image_path)).float()","metadata":{"execution":{"iopub.status.busy":"2023-12-04T02:07:21.98373Z","iopub.execute_input":"2023-12-04T02:07:21.98414Z","iopub.status.idle":"2023-12-04T02:07:21.996492Z","shell.execute_reply.started":"2023-12-04T02:07:21.984108Z","shell.execute_reply":"2023-12-04T02:07:21.995153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Instantiate datasets","metadata":{}},{"cell_type":"code","source":"# For details of tile creation, please refer pjmathematician's notebook: \n# https://www.kaggle.com/code/pjmathematician/ucbo-tilemaker\n# For details of the split, please refer: \n# https://www.kaggle.com/code/emiz6413/stratified-shuffle-split/notebook\nmeta_df = pd.read_csv(\"/kaggle/input/save-stratified-split-to-csv/train_split.csv\")\ntrain_image_ids = meta_df[meta_df[\"is_train\"]][\"image_id\"].values\neval_image_ids = meta_df[~meta_df[\"is_train\"]][\"image_id\"].values\ntrain_image_dirs = [next(Path(\"/kaggle/input/\").glob(f\"ucbo-tiles-256-*/256_{img_id}\")) for img_id in train_image_ids]\neval_image_dirs = [next(Path(\"/kaggle/input/\").glob(f\"ucbo-tiles-256-*/256_{img_id}\")) for img_id in eval_image_ids]","metadata":{"execution":{"iopub.status.busy":"2023-12-04T02:07:23.344041Z","iopub.execute_input":"2023-12-04T02:07:23.344441Z","iopub.status.idle":"2023-12-04T02:07:25.904674Z","shell.execute_reply.started":"2023-12-04T02:07:23.344411Z","shell.execute_reply":"2023-12-04T02:07:25.90328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = UBCODataset(train_image_dirs, is_train=True, shuffle=True, image_size=IMG_SIZE)\neval_ds = UBCODataset(eval_image_dirs, is_train=False, shuffle=False, image_size=IMG_SIZE)\ntrain_loader = DataLoader(dataset=train_ds, batch_size=BATCH_SIZE, num_workers=os.cpu_count(), drop_last=True)\neval_loader = DataLoader(dataset=eval_ds, batch_size=BATCH_SIZE*2, num_workers=os.cpu_count())","metadata":{"execution":{"iopub.status.busy":"2023-12-04T02:07:25.906761Z","iopub.execute_input":"2023-12-04T02:07:25.907297Z","iopub.status.idle":"2023-12-04T02:09:09.69854Z","shell.execute_reply.started":"2023-12-04T02:07:25.907255Z","shell.execute_reply":"2023-12-04T02:09:09.69746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Visualize the first few samples","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 2, figsize=(10, 5))\n\ntrain_imgs = next(iter(train_loader))[:64]\nval_imgs = next(iter(eval_loader))[:64]\n\nfor ax, img, title in zip(axes, [train_imgs, val_imgs], [\"train\", \"eval\"]):\n    ax.imshow(transforms.ToPILImage()(make_grid(UBCODataset.denormalize(img), nrow=8)))\n    ax.axis(\"off\")\n    ax.set_title(title)\n\nfig.tight_layout()\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-04T02:09:09.70031Z","iopub.execute_input":"2023-12-04T02:09:09.700666Z","iopub.status.idle":"2023-12-04T02:09:32.17289Z","shell.execute_reply.started":"2023-12-04T02:09:09.700635Z","shell.execute_reply":"2023-12-04T02:09:32.170277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Instantiate a model","metadata":{}},{"cell_type":"code","source":"encoder = Encoder64()\ngenerator = Generator64()\ndiscriminator = Discriminator64(sn_enabled=DISC_SN_ENABLED, bn_enabled=DISC_BN_ENABLED)\nmodel = BiGAN(encoder=encoder, generator=generator, discriminator=discriminator, device=DEVICE, amp=AMP)\nmodel = model.to(DEVICE)\nmodel.load_state_dict(\n    torch.load(\"/kaggle/input/bidirectionalgan-ubco/bigan_ubco_64.pth\", map_location=\"cpu\")\n)","metadata":{"execution":{"iopub.status.busy":"2023-12-04T02:22:21.498514Z","iopub.execute_input":"2023-12-04T02:22:21.499243Z","iopub.status.idle":"2023-12-04T02:22:22.627901Z","shell.execute_reply.started":"2023-12-04T02:22:21.4992Z","shell.execute_reply":"2023-12-04T02:22:22.626872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Train","metadata":{}},{"cell_type":"code","source":"pbar = tqdm(total=LAST_EPOCH - INITIAL_EPOCH)\nsample_img = next(iter(eval_loader))[:64]\nsample_img = sample_img.to(model.device)\nbest_loss = np.inf\nbest_epoch = -1\n\nfor epoch in range(INITIAL_EPOCH, LAST_EPOCH):\n    ge_loss, disc_loss = model.train_single_epoch(train_loader=train_loader)\n    rec_loss = model.evaluate(eval_loader=eval_loader)\n\n    with torch.no_grad():\n        gen_img = model.generate(torch.randn(64, model.latent_dim, 1, 1, device=model.device))\n        rec_img = model.reconstruct(sample_img)\n\n    fig, axes = plt.subplots(1, 3)\n\n    for ax, img, title in zip(axes, [sample_img, rec_img, gen_img], [\"original\", \"reconstructed\", \"generated\"]):\n        ax.imshow(transforms.ToPILImage()(UBCODataset.denormalize(make_grid(img, nrow=8))))\n        ax.axis(\"off\")\n        ax.set_title(title)\n\n    fig.tight_layout()\n    fig.suptitle(f\"epoch: {epoch}\")\n    \n    torch.save(model.state_dict(), \"bigan_ubco_last.pth\")\n    if rec_loss < best_loss:\n        torch.save(model.state_dict(), \"bigan_ubco_best.pth\")\n        best_loss = rec_loss\n        best_epoch = epoch\n\n    pbar.update()\n\npbar.close()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-12-04T02:22:24.227328Z","iopub.execute_input":"2023-12-04T02:22:24.228848Z","iopub.status.idle":"2023-12-04T02:22:30.524691Z","shell.execute_reply.started":"2023-12-04T02:22:24.228802Z","shell.execute_reply":"2023-12-04T02:22:30.523372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Visualize the results","metadata":{}},{"cell_type":"code","source":"model = model.eval()\n\nwith torch.no_grad():\n    rec_img = model.reconstruct(sample_img)\n    z_p = torch.randn(64, model.latent_dim, 1, 1)\n    gen_img = model.generate(z_p)\n    \nfig, axes = plt.subplots(1, 3)\n\nfor ax, img, title in zip(axes, [sample_img, rec_img, gen_img], [\"original\", \"reconstructed\", \"generated\"]):\n    ax.imshow(transforms.ToPILImage()(UBCODataset.denormalize(make_grid(img, nrow=8))))\n    ax.axis(\"off\")\n    ax.set_title(title)\n\nfig.tight_layout()\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-04T02:30:43.423208Z","iopub.execute_input":"2023-12-04T02:30:43.423616Z","iopub.status.idle":"2023-12-04T02:30:44.876938Z","shell.execute_reply.started":"2023-12-04T02:30:43.423582Z","shell.execute_reply":"2023-12-04T02:30:44.875961Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}