{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9187072,"sourceType":"datasetVersion","datasetId":5504483}],"dockerImageVersionId":30005,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Code from Here [https://www.kaggle.com/code/nachiket273/cyclegan-pytorch]\n* I tried to see if I could convert CT images to Sagittal images using CycleGAN.","metadata":{}},{"cell_type":"code","source":"%reload_ext autoreload\n%autoreload 2\n%matplotlib inline","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-10-28T05:57:28.532033Z","iopub.execute_input":"2024-10-28T05:57:28.532451Z","iopub.status.idle":"2024-10-28T05:57:28.573093Z","shell.execute_reply.started":"2024-10-28T05:57:28.532412Z","shell.execute_reply":"2024-10-28T05:57:28.572229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from glob import glob\nimport itertools\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport os\nimport pandas as pd\nimport PIL\nfrom PIL import Image\nimport random\nimport shutil\nfrom sklearn.model_selection import GroupKFold\nfrom sklearn.metrics import roc_curve\nfrom sklearn import metrics\nimport time\nfrom tqdm.notebook import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.nn.init as init\nfrom torch.utils.data import Dataset, random_split, DataLoader\n\nimport torchvision.models as models\nimport torchvision.transforms as transforms","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","execution":{"iopub.status.busy":"2024-10-28T05:57:28.574887Z","iopub.execute_input":"2024-10-28T05:57:28.575154Z","iopub.status.idle":"2024-10-28T05:57:30.876614Z","shell.execute_reply.started":"2024-10-28T05:57:28.575127Z","shell.execute_reply":"2024-10-28T05:57:30.875807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Seed","metadata":{}},{"cell_type":"code","source":"def set_seed(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:30.877739Z","iopub.execute_input":"2024-10-28T05:57:30.878010Z","iopub.status.idle":"2024-10-28T05:57:30.910650Z","shell.execute_reply.started":"2024-10-28T05:57:30.877984Z","shell.execute_reply":"2024-10-28T05:57:30.909949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\ndevice","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:30.911871Z","iopub.execute_input":"2024-10-28T05:57:30.912228Z","iopub.status.idle":"2024-10-28T05:57:31.033717Z","shell.execute_reply.started":"2024-10-28T05:57:30.912192Z","shell.execute_reply":"2024-10-28T05:57:31.032671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"set_seed(719)","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:31.036909Z","iopub.execute_input":"2024-10-28T05:57:31.037217Z","iopub.status.idle":"2024-10-28T05:57:31.071362Z","shell.execute_reply.started":"2024-10-28T05:57:31.037188Z","shell.execute_reply":"2024-10-28T05:57:31.070659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset and Dataloader","metadata":{}},{"cell_type":"code","source":"class ImageDataset(Dataset):\n    def __init__(self, monet_dir, photo_dir, size=(256, 256), normalize=True):\n        super().__init__()\n        self.monet_dir = monet_dir\n        self.photo_dir = photo_dir\n        self.monet_idx = dict()\n        self.photo_idx = dict()\n        if normalize:\n            self.transform = transforms.Compose([\n                transforms.Resize(size),\n                transforms.ToTensor(),\n                transforms.Normalize((0.5), (0.5))                                \n            ])\n        else:\n            self.transform = transforms.Compose([\n                transforms.Resize(size),\n                transforms.ToTensor()                               \n            ])\n        for i, fl in enumerate(os.listdir(self.monet_dir)):\n            self.monet_idx[i] = fl\n        for i, fl in enumerate(os.listdir(self.photo_dir)):\n            self.photo_idx[i] = fl\n\n    def __getitem__(self, idx):\n        try:\n            rand_idx = int(np.random.uniform(0, len(self.photo_idx.keys())))\n            photo_path = os.path.join(self.photo_dir, self.photo_idx[rand_idx])\n            monet_path = os.path.join(self.monet_dir, self.monet_idx[idx])\n            photo_img = Image.open(photo_path).convert(\"L\")\n            photo_img = self.transform(photo_img)\n            monet_img = Image.open(monet_path).convert(\"L\")\n            monet_img = self.transform(monet_img)\n            return photo_img, monet_img\n        except Exception as e:\n            print(f\"エラーが発生しました: {e}\")\n            print(rand_idx)\n    def __len__(self):\n        return min(len(self.monet_idx.keys()), len(self.photo_idx.keys()))","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:31.074638Z","iopub.execute_input":"2024-10-28T05:57:31.074994Z","iopub.status.idle":"2024-10-28T05:57:31.118651Z","shell.execute_reply.started":"2024-10-28T05:57:31.074961Z","shell.execute_reply":"2024-10-28T05:57:31.117672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_ds = ImageDataset(\"/kaggle/input/lumbar-coordinate-pretraining-dataset/data/processed_lsd_jpgs\", '/kaggle/input/lumbar-coordinate-pretraining-dataset/data/processed_tseg_jpgs/')","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:31.119938Z","iopub.execute_input":"2024-10-28T05:57:31.120266Z","iopub.status.idle":"2024-10-28T05:57:31.545142Z","shell.execute_reply.started":"2024-10-28T05:57:31.120238Z","shell.execute_reply":"2024-10-28T05:57:31.544507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_dl = DataLoader(img_ds, batch_size=1, pin_memory=True)","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:31.546369Z","iopub.execute_input":"2024-10-28T05:57:31.546772Z","iopub.status.idle":"2024-10-28T05:57:31.579548Z","shell.execute_reply.started":"2024-10-28T05:57:31.546734Z","shell.execute_reply":"2024-10-28T05:57:31.578626Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"photo_img, monet_img = next(iter(img_dl))","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:31.581184Z","iopub.execute_input":"2024-10-28T05:57:31.581482Z","iopub.status.idle":"2024-10-28T05:57:35.919170Z","shell.execute_reply.started":"2024-10-28T05:57:31.581454Z","shell.execute_reply":"2024-10-28T05:57:35.918279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def unnorm(img, mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]):\n    for t, m, s in zip(img, mean, std):\n        t.mul_(s).add_(s)\n        \n    return img","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:35.920759Z","iopub.execute_input":"2024-10-28T05:57:35.921065Z","iopub.status.idle":"2024-10-28T05:57:35.954758Z","shell.execute_reply.started":"2024-10-28T05:57:35.921035Z","shell.execute_reply":"2024-10-28T05:57:35.954035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f = plt.figure(figsize=(8, 8))\n\nf.add_subplot(1, 2, 1)\nplt.title('CT')\nphoto_img = unnorm(photo_img)\nplt.imshow(photo_img[0,0])\n\nf.add_subplot(1, 2, 2)\nplt.title('Sagital')\nmonet_img = unnorm(monet_img)\nplt.imshow(monet_img[0,0])","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:35.956004Z","iopub.execute_input":"2024-10-28T05:57:35.956468Z","iopub.status.idle":"2024-10-28T05:57:36.313170Z","shell.execute_reply.started":"2024-10-28T05:57:35.956437Z","shell.execute_reply":"2024-10-28T05:57:36.312215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Save and Load","metadata":{}},{"cell_type":"code","source":"def load_checkpoint(ckpt_path, map_location=None):\n    ckpt = torch.load(ckpt_path, map_location=map_location)\n    print(' [*] Loading checkpoint from %s succeed!' % ckpt_path)\n    return ckpt","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:36.314466Z","iopub.execute_input":"2024-10-28T05:57:36.314836Z","iopub.status.idle":"2024-10-28T05:57:36.350263Z","shell.execute_reply.started":"2024-10-28T05:57:36.314796Z","shell.execute_reply":"2024-10-28T05:57:36.349570Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_checkpoint(state, save_path):\n    torch.save(state, save_path)","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:36.351396Z","iopub.execute_input":"2024-10-28T05:57:36.351726Z","iopub.status.idle":"2024-10-28T05:57:36.384042Z","shell.execute_reply.started":"2024-10-28T05:57:36.351696Z","shell.execute_reply":"2024-10-28T05:57:36.383265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"def Upsample(in_ch, out_ch, use_dropout=True, dropout_ratio=0.5):\n    if use_dropout:\n        return nn.Sequential(\n            nn.ConvTranspose2d(in_ch, out_ch, 3, stride=2, padding=1, output_padding=1),\n            nn.InstanceNorm2d(out_ch),\n            nn.Dropout(dropout_ratio),\n            nn.GELU()\n        )\n    else:\n        return nn.Sequential(\n            nn.ConvTranspose2d(in_ch, out_ch, 3, stride=2, padding=1, output_padding=1),\n            nn.InstanceNorm2d(out_ch),\n            nn.GELU()\n        )","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:36.385443Z","iopub.execute_input":"2024-10-28T05:57:36.385771Z","iopub.status.idle":"2024-10-28T05:57:36.420337Z","shell.execute_reply.started":"2024-10-28T05:57:36.385719Z","shell.execute_reply":"2024-10-28T05:57:36.419509Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def Convlayer(in_ch, out_ch, kernel_size=3, stride=2, use_leaky=True, use_inst_norm=True, use_pad=True):\n    if use_pad:\n        conv = nn.Conv2d(in_ch, out_ch, kernel_size, stride, 1, bias=True)\n    else:\n        conv = nn.Conv2d(in_ch, out_ch, kernel_size, stride, 0, bias=True)\n\n    if use_leaky:\n        actv = nn.LeakyReLU(negative_slope=0.2, inplace=True)\n    else:\n        actv = nn.GELU()\n\n    if use_inst_norm:\n        norm = nn.InstanceNorm2d(out_ch)\n    else:\n        norm = nn.BatchNorm2d(out_ch)\n\n    return nn.Sequential(\n        conv,\n        norm,\n        actv\n    )","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:36.421551Z","iopub.execute_input":"2024-10-28T05:57:36.421860Z","iopub.status.idle":"2024-10-28T05:57:36.456158Z","shell.execute_reply.started":"2024-10-28T05:57:36.421833Z","shell.execute_reply":"2024-10-28T05:57:36.455369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Resblock(nn.Module):\n    def __init__(self, in_features, use_dropout=True, dropout_ratio=0.5):\n        super().__init__()\n        layers = list()\n        layers.append(nn.ReflectionPad2d(1))\n        layers.append(Convlayer(in_features, in_features, 3, 1, False, use_pad=False))\n        layers.append(nn.Dropout(dropout_ratio))\n        layers.append(nn.ReflectionPad2d(1))\n        layers.append(nn.Conv2d(in_features, in_features, 3, 1, padding=0, bias=True))\n        layers.append(nn.InstanceNorm2d(in_features))\n        self.res = nn.Sequential(*layers)\n\n    def forward(self, x):\n        return x + self.res(x)","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:36.457299Z","iopub.execute_input":"2024-10-28T05:57:36.457681Z","iopub.status.idle":"2024-10-28T05:57:36.492866Z","shell.execute_reply.started":"2024-10-28T05:57:36.457652Z","shell.execute_reply":"2024-10-28T05:57:36.492134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Generator(nn.Module):\n    def __init__(self, in_ch, out_ch, num_res_blocks=6):\n        super().__init__()\n        model = list()\n        model.append(nn.ReflectionPad2d(3))\n        model.append(Convlayer(in_ch, 64, 7, 1, False, True, False))\n        model.append(Convlayer(64, 128, 3, 2, False))\n        model.append(Convlayer(128, 256, 3, 2, False))\n        for _ in range(num_res_blocks):\n            model.append(Resblock(256))\n        model.append(Upsample(256, 128))\n        model.append(Upsample(128, 64))\n        model.append(nn.ReflectionPad2d(3))\n        model.append(nn.Conv2d(64, out_ch, kernel_size=7, padding=0))\n        model.append(nn.Tanh())\n\n        self.gen = nn.Sequential(*model)\n\n    def forward(self, x):\n        return self.gen(x)","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:36.494049Z","iopub.execute_input":"2024-10-28T05:57:36.494342Z","iopub.status.idle":"2024-10-28T05:57:36.531708Z","shell.execute_reply.started":"2024-10-28T05:57:36.494304Z","shell.execute_reply":"2024-10-28T05:57:36.530998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Discriminator(nn.Module):\n    def __init__(self, in_ch, num_layers=4):\n        super().__init__()\n        model = list()\n        model.append(nn.Conv2d(in_ch, 64, 4, stride=2, padding=1))\n        model.append(nn.LeakyReLU(0.2, inplace=True))\n        for i in range(1, num_layers):\n            in_chs = 64 * 2**(i-1)\n            out_chs = in_chs * 2\n            if i == num_layers -1:\n                model.append(Convlayer(in_chs, out_chs, 4, 1))\n            else:\n                model.append(Convlayer(in_chs, out_chs, 4, 2))\n        model.append(nn.Conv2d(512, 1, kernel_size=4, stride=1, padding=1))\n        self.disc = nn.Sequential(*model)\n\n    def forward(self, x):\n        return self.disc(x)","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:36.533037Z","iopub.execute_input":"2024-10-28T05:57:36.533289Z","iopub.status.idle":"2024-10-28T05:57:36.570310Z","shell.execute_reply.started":"2024-10-28T05:57:36.533265Z","shell.execute_reply":"2024-10-28T05:57:36.569669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def init_weights(net, init_type='normal', gain=0.02):\n    def init_func(m):\n        classname = m.__class__.__name__\n        if hasattr(m, 'weight') and (classname.find('Conv') != -1 or classname.find('Linear') != -1):\n            init.normal_(m.weight.data, 0.0, gain)\n            if hasattr(m, 'bias') and m.bias is not None:\n                init.constant_(m.bias.data, 0.0)\n        elif classname.find('BatchNorm2d') != -1:\n            init.normal_(m.weight.data, 1.0, gain)\n            init.constant_(m.bias.data, 0.0)\n    net.apply(init_func)","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:36.571385Z","iopub.execute_input":"2024-10-28T05:57:36.571668Z","iopub.status.idle":"2024-10-28T05:57:36.607184Z","shell.execute_reply.started":"2024-10-28T05:57:36.571641Z","shell.execute_reply":"2024-10-28T05:57:36.606423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Some additional classes and functions","metadata":{}},{"cell_type":"code","source":"def update_req_grad(models, requires_grad=True):\n    for model in models:\n        for param in model.parameters():\n            param.requires_grad = requires_grad","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:36.608138Z","iopub.execute_input":"2024-10-28T05:57:36.608398Z","iopub.status.idle":"2024-10-28T05:57:36.639404Z","shell.execute_reply.started":"2024-10-28T05:57:36.608372Z","shell.execute_reply":"2024-10-28T05:57:36.638715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://arxiv.org/pdf/1612.07828.pdf\n# Save 50 generated fake imgs and sample through them\n# to feed discriminators to avoid large oscillations \n# from iterations to iterations.\nclass sample_fake(object):\n    def __init__(self, max_imgs=10):\n        self.max_imgs = max_imgs\n        self.cur_img = 0\n        self.imgs = list()\n\n    def __call__(self, imgs):\n        ret = list()\n        for img in imgs:\n            if self.cur_img < self.max_imgs:\n                self.imgs.append(img)\n                ret.append(img)\n                self.cur_img += 1\n            else:\n                if np.random.ranf() > 0.5:\n                    idx = np.random.randint(0, self.max_imgs)\n                    ret.append(self.imgs[idx])\n                    self.imgs[idx] = img\n                else:\n                    ret.append(img)\n        return ret","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:36.640381Z","iopub.execute_input":"2024-10-28T05:57:36.640671Z","iopub.status.idle":"2024-10-28T05:57:36.676546Z","shell.execute_reply.started":"2024-10-28T05:57:36.640644Z","shell.execute_reply":"2024-10-28T05:57:36.675767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class lr_sched():\n    def __init__(self, decay_epochs=5, total_epochs=10):\n        self.decay_epochs = decay_epochs\n        self.total_epochs = total_epochs\n\n    def step(self, epoch_num):\n        if epoch_num <= self.decay_epochs:\n            return 1.0\n        else:\n            fract = (epoch_num - self.decay_epochs)  / (self.total_epochs - self.decay_epochs)\n            return 1.0 - fract","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:36.677813Z","iopub.execute_input":"2024-10-28T05:57:36.678179Z","iopub.status.idle":"2024-10-28T05:57:36.710304Z","shell.execute_reply.started":"2024-10-28T05:57:36.678140Z","shell.execute_reply":"2024-10-28T05:57:36.709633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AvgStats(object):\n    def __init__(self):\n        self.reset()\n        \n    def reset(self):\n        self.losses =[]\n        self.its = []\n        \n    def append(self, loss, it):\n        self.losses.append(loss)\n        self.its.append(it)","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:36.711506Z","iopub.execute_input":"2024-10-28T05:57:36.711909Z","iopub.status.idle":"2024-10-28T05:57:36.744375Z","shell.execute_reply.started":"2024-10-28T05:57:36.711859Z","shell.execute_reply":"2024-10-28T05:57:36.743594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# GAN Class","metadata":{}},{"cell_type":"code","source":"class CycleGAN(object):\n    def __init__(self, in_ch, out_ch, epochs, device, start_lr=2e-4, lmbda=10, idt_coef=0.5, decay_epoch=0):\n        self.epochs = epochs\n        self.decay_epoch = decay_epoch if decay_epoch > 0 else int(self.epochs/2)\n        self.lmbda = lmbda\n        self.idt_coef = idt_coef\n        self.device = device\n        self.gen_mtp = Generator(in_ch, out_ch)\n        self.gen_ptm = Generator(in_ch, out_ch)\n        self.desc_m = Discriminator(in_ch)\n        self.desc_p = Discriminator(in_ch)\n        self.init_models()\n        self.mse_loss = nn.MSELoss()\n        self.l1_loss = nn.L1Loss()\n        self.adam_gen = torch.optim.Adam(itertools.chain(self.gen_mtp.parameters(), self.gen_ptm.parameters()),\n                                         lr = start_lr, betas=(0.5, 0.999))\n        self.adam_desc = torch.optim.Adam(itertools.chain(self.desc_m.parameters(), self.desc_p.parameters()),\n                                          lr=start_lr, betas=(0.5, 0.999))\n        self.sample_monet = sample_fake()\n        self.sample_photo = sample_fake()\n        gen_lr = lr_sched(self.decay_epoch, self.epochs)\n        desc_lr = lr_sched(self.decay_epoch, self.epochs)\n        self.gen_lr_sched = torch.optim.lr_scheduler.LambdaLR(self.adam_gen, gen_lr.step)\n        self.desc_lr_sched = torch.optim.lr_scheduler.LambdaLR(self.adam_desc, desc_lr.step)\n        self.gen_stats = AvgStats()\n        self.desc_stats = AvgStats()\n        \n    def init_models(self):\n        init_weights(self.gen_mtp)\n        init_weights(self.gen_ptm)\n        init_weights(self.desc_m)\n        init_weights(self.desc_p)\n        self.gen_mtp = self.gen_mtp.to(self.device)\n        self.gen_ptm = self.gen_ptm.to(self.device)\n        self.desc_m = self.desc_m.to(self.device)\n        self.desc_p = self.desc_p.to(self.device)\n        \n    def train(self, photo_dl):\n        for epoch in range(self.epochs):\n            start_time = time.time()\n            avg_gen_loss = 0.0\n            avg_desc_loss = 0.0\n            t = tqdm(photo_dl, leave=False, total=photo_dl.__len__())\n            for i, (photo_real, monet_real) in enumerate(t):\n                photo_img, monet_img = photo_real.to(device), monet_real.to(device)\n                update_req_grad([self.desc_m, self.desc_p], False)\n                self.adam_gen.zero_grad()\n\n                # Forward pass through generator\n                fake_photo = self.gen_mtp(monet_img)\n                fake_monet = self.gen_ptm(photo_img)\n\n                cycl_monet = self.gen_ptm(fake_photo)\n                cycl_photo = self.gen_mtp(fake_monet)\n\n                id_monet = self.gen_ptm(monet_img)\n                id_photo = self.gen_mtp(photo_img)\n\n                # generator losses - identity, Adversarial, cycle consistency\n                idt_loss_monet = self.l1_loss(id_monet, monet_img) * self.lmbda * self.idt_coef\n                idt_loss_photo = self.l1_loss(id_photo, photo_img) * self.lmbda * self.idt_coef\n\n                cycle_loss_monet = self.l1_loss(cycl_monet, monet_img) * self.lmbda\n                cycle_loss_photo = self.l1_loss(cycl_photo, photo_img) * self.lmbda\n\n                monet_desc = self.desc_m(fake_monet)\n                photo_desc = self.desc_p(fake_photo)\n\n                real = torch.ones(monet_desc.size()).to(self.device)\n\n                adv_loss_monet = self.mse_loss(monet_desc, real)\n                adv_loss_photo = self.mse_loss(photo_desc, real)\n\n                # total generator loss\n                total_gen_loss = cycle_loss_monet + adv_loss_monet\\\n                              + cycle_loss_photo + adv_loss_photo\\\n                              + idt_loss_monet + idt_loss_photo\n                \n                avg_gen_loss += total_gen_loss.item()\n\n                # backward pass\n                total_gen_loss.backward()\n                self.adam_gen.step()\n\n                # Forward pass through Descriminator\n                update_req_grad([self.desc_m, self.desc_p], True)\n                self.adam_desc.zero_grad()\n\n                fake_monet = self.sample_monet([fake_monet.cpu().data.numpy()])[0]\n                fake_photo = self.sample_photo([fake_photo.cpu().data.numpy()])[0]\n                fake_monet = torch.tensor(fake_monet).to(self.device)\n                fake_photo = torch.tensor(fake_photo).to(self.device)\n\n                monet_desc_real = self.desc_m(monet_img)\n                monet_desc_fake = self.desc_m(fake_monet)\n                photo_desc_real = self.desc_p(photo_img)\n                photo_desc_fake = self.desc_p(fake_photo)\n\n                real = torch.ones(monet_desc_real.size()).to(self.device)\n                fake = torch.zeros(monet_desc_fake.size()).to(self.device)\n\n                # Descriminator losses\n                # --------------------\n                monet_desc_real_loss = self.mse_loss(monet_desc_real, real)\n                monet_desc_fake_loss = self.mse_loss(monet_desc_fake, fake)\n                photo_desc_real_loss = self.mse_loss(photo_desc_real, real)\n                photo_desc_fake_loss = self.mse_loss(photo_desc_fake, fake)\n\n                monet_desc_loss = (monet_desc_real_loss + monet_desc_fake_loss) / 2\n                photo_desc_loss = (photo_desc_real_loss + photo_desc_fake_loss) / 2\n                total_desc_loss = monet_desc_loss + photo_desc_loss\n                avg_desc_loss += total_desc_loss.item()\n\n                # Backward\n                monet_desc_loss.backward()\n                photo_desc_loss.backward()\n                self.adam_desc.step()\n                \n                t.set_postfix(gen_loss=total_gen_loss.item(), desc_loss=total_desc_loss.item())\n\n            save_dict = {\n                'epoch': epoch+1,\n                'gen_mtp': gan.gen_mtp.state_dict(),\n                'gen_ptm': gan.gen_ptm.state_dict(),\n                'desc_m': gan.desc_m.state_dict(),\n                'desc_p': gan.desc_p.state_dict(),\n                'optimizer_gen': gan.adam_gen.state_dict(),\n                'optimizer_desc': gan.adam_desc.state_dict()\n            }\n            save_checkpoint(save_dict, 'current.ckpt')\n            \n            avg_gen_loss /= photo_dl.__len__()\n            avg_desc_loss /= photo_dl.__len__()\n            time_req = time.time() - start_time\n            \n            self.gen_stats.append(avg_gen_loss, time_req)\n            self.desc_stats.append(avg_desc_loss, time_req)\n            \n            print(\"Epoch: (%d) | Generator Loss:%f | Discriminator Loss:%f\" % \n                                                (epoch+1, avg_gen_loss, avg_desc_loss))\n      \n            self.gen_lr_sched.step()\n            self.desc_lr_sched.step()","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:36.745691Z","iopub.execute_input":"2024-10-28T05:57:36.746026Z","iopub.status.idle":"2024-10-28T05:57:36.816755Z","shell.execute_reply.started":"2024-10-28T05:57:36.745998Z","shell.execute_reply":"2024-10-28T05:57:36.816056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gan = CycleGAN(1, 1, 10, device)","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:36.818165Z","iopub.execute_input":"2024-10-28T05:57:36.818717Z","iopub.status.idle":"2024-10-28T05:57:37.256518Z","shell.execute_reply.started":"2024-10-28T05:57:36.818677Z","shell.execute_reply":"2024-10-28T05:57:37.255813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Save before train\nsave_dict = {\n    'epoch': 0,\n    'gen_mtp': gan.gen_mtp.state_dict(),\n    'gen_ptm': gan.gen_ptm.state_dict(),\n    'desc_m': gan.desc_m.state_dict(),\n    'desc_p': gan.desc_p.state_dict(),\n    'optimizer_gen': gan.adam_gen.state_dict(),\n    'optimizer_desc': gan.adam_desc.state_dict()\n}","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:37.257825Z","iopub.execute_input":"2024-10-28T05:57:37.258095Z","iopub.status.idle":"2024-10-28T05:57:37.293101Z","shell.execute_reply.started":"2024-10-28T05:57:37.258069Z","shell.execute_reply":"2024-10-28T05:57:37.292372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_checkpoint(save_dict, 'init.ckpt')","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:37.294304Z","iopub.execute_input":"2024-10-28T05:57:37.294656Z","iopub.status.idle":"2024-10-28T05:57:37.417131Z","shell.execute_reply.started":"2024-10-28T05:57:37.294609Z","shell.execute_reply":"2024-10-28T05:57:37.416166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gan.train(img_dl)","metadata":{"execution":{"iopub.status.busy":"2024-10-28T05:57:37.418763Z","iopub.execute_input":"2024-10-28T05:57:37.419168Z","iopub.status.idle":"2024-10-28T06:49:24.856825Z","shell.execute_reply.started":"2024-10-28T05:57:37.419124Z","shell.execute_reply":"2024-10-28T06:49:24.855917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.xlabel(\"Epochs\")\nplt.ylabel(\"Losses\")\nplt.plot(gan.gen_stats.losses, 'r', label='Generator Loss')\nplt.plot(gan.desc_stats.losses, 'b', label='Descriminator Loss')\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-28T06:49:24.858002Z","iopub.execute_input":"2024-10-28T06:49:24.858318Z","iopub.status.idle":"2024-10-28T06:49:25.056133Z","shell.execute_reply.started":"2024-10-28T06:49:24.858288Z","shell.execute_reply":"2024-10-28T06:49:25.055130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_, ax = plt.subplots(5, 2, figsize=(12, 12))\nfor i in range(5):\n    photo_img, _ = next(iter(img_dl))\n    pred_monet = gan.gen_ptm(photo_img.to(device)).cpu().detach()\n    photo_img = unnorm(photo_img)\n    pred_monet = unnorm(pred_monet)\n    \n    ax[i, 0].imshow(photo_img[0,0])\n    ax[i, 1].imshow(pred_monet[0,0])\n    ax[i, 0].set_title(\"Input Photo\")\n    ax[i, 1].set_title(\"Monet-esque Photo\")\n    ax[i, 0].axis(\"off\")\n    ax[i, 1].axis(\"off\")\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Run Generator over all images","metadata":{}},{"cell_type":"code","source":"class PhotoDataset(Dataset):\n    def __init__(self, photo_dir, size=(256, 256), normalize=True):\n        super().__init__()\n        self.photo_dir = photo_dir\n        self.photo_idx = dict()\n        if normalize:\n            self.transform = transforms.Compose([\n                transforms.Resize(size),\n                transforms.ToTensor(),\n                transforms.Normalize((0.5), (0.5))                                \n            ])\n        else:\n            self.transform = transforms.Compose([\n                transforms.Resize(size),\n                transforms.ToTensor()                               \n            ])\n        for i, fl in enumerate(os.listdir(self.photo_dir)):\n            self.photo_idx[i] = fl\n\n    def __getitem__(self, idx):\n        photo_path = os.path.join(self.photo_dir, self.photo_idx[idx])\n        photo_img = Image.open(photo_path).convert(\"L\")\n        photo_img = self.transform(photo_img)\n        return photo_img\n\n    def __len__(self):\n        return len(self.photo_idx.keys())","metadata":{"execution":{"iopub.status.busy":"2024-10-28T06:49:25.706897Z","iopub.execute_input":"2024-10-28T06:49:25.707169Z","iopub.status.idle":"2024-10-28T06:49:25.749932Z","shell.execute_reply.started":"2024-10-28T06:49:25.707143Z","shell.execute_reply":"2024-10-28T06:49:25.749043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ph_ds = PhotoDataset('/kaggle/input/lumbar-coordinate-pretraining-dataset/data/processed_tseg_jpgs/')","metadata":{"execution":{"iopub.status.busy":"2024-10-28T06:49:25.754883Z","iopub.execute_input":"2024-10-28T06:49:25.755182Z","iopub.status.idle":"2024-10-28T06:49:25.789405Z","shell.execute_reply.started":"2024-10-28T06:49:25.755154Z","shell.execute_reply":"2024-10-28T06:49:25.788616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ph_dl = DataLoader(ph_ds, batch_size=1, pin_memory=True)","metadata":{"execution":{"iopub.status.busy":"2024-10-28T06:49:25.791054Z","iopub.execute_input":"2024-10-28T06:49:25.791437Z","iopub.status.idle":"2024-10-28T06:49:25.827251Z","shell.execute_reply.started":"2024-10-28T06:49:25.791396Z","shell.execute_reply":"2024-10-28T06:49:25.826385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir ../images","metadata":{"execution":{"iopub.status.busy":"2024-10-28T06:49:25.828577Z","iopub.execute_input":"2024-10-28T06:49:25.828928Z","iopub.status.idle":"2024-10-28T06:49:26.912872Z","shell.execute_reply.started":"2024-10-28T06:49:25.828892Z","shell.execute_reply":"2024-10-28T06:49:26.911561Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trans = transforms.ToPILImage()","metadata":{"execution":{"iopub.status.busy":"2024-10-28T06:49:26.914490Z","iopub.execute_input":"2024-10-28T06:49:26.914848Z","iopub.status.idle":"2024-10-28T06:49:26.952311Z","shell.execute_reply.started":"2024-10-28T06:49:26.914812Z","shell.execute_reply":"2024-10-28T06:49:26.951258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"t = tqdm(ph_dl, leave=False, total=ph_dl.__len__())\nfor i, photo in enumerate(t):\n    with torch.no_grad():\n        pred_monet = gan.gen_ptm(photo.to(device)).cpu().detach()\n    pred_monet = unnorm(pred_monet)\n    img = trans(pred_monet[0]).convert(\"RGB\")\n    img.save(\"../images/\" + str(i+1) + \".jpg\")","metadata":{"execution":{"iopub.status.busy":"2024-10-28T06:49:26.953766Z","iopub.execute_input":"2024-10-28T06:49:26.954160Z","iopub.status.idle":"2024-10-28T06:49:40.849837Z","shell.execute_reply.started":"2024-10-28T06:49:26.954115Z","shell.execute_reply":"2024-10-28T06:49:40.848835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"shutil.make_archive(\"/kaggle/working/images\", 'zip', \"/kaggle/images\")","metadata":{"execution":{"iopub.status.busy":"2024-10-28T06:49:40.851368Z","iopub.execute_input":"2024-10-28T06:49:40.851963Z","iopub.status.idle":"2024-10-28T06:49:41.122612Z","shell.execute_reply.started":"2024-10-28T06:49:40.851920Z","shell.execute_reply":"2024-10-28T06:49:41.121501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}