{"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":"gpu","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-09-13T14:26:40.326509Z","iopub.execute_input":"2024-09-13T14:26:40.326933Z","iopub.status.idle":"2024-09-13T14:26:40.360507Z","shell.execute_reply.started":"2024-09-13T14:26:40.326897Z","shell.execute_reply":"2024-09-13T14:26:40.359725Z"},"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-09-13T14:26:40.362576Z","iopub.execute_input":"2024-09-13T14:26:40.362838Z","iopub.status.idle":"2024-09-13T14:26:42.810415Z","shell.execute_reply.started":"2024-09-13T14:26:40.362813Z","shell.execute_reply":"2024-09-13T14:26:42.809588Z"},"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-09-13T14:26:42.811773Z","iopub.execute_input":"2024-09-13T14:26:42.812161Z","iopub.status.idle":"2024-09-13T14:26:42.846928Z","shell.execute_reply.started":"2024-09-13T14:26:42.812120Z","shell.execute_reply":"2024-09-13T14:26:42.846073Z"},"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-09-13T14:26:42.848172Z","iopub.execute_input":"2024-09-13T14:26:42.848474Z","iopub.status.idle":"2024-09-13T14:26:42.951514Z","shell.execute_reply.started":"2024-09-13T14:26:42.848447Z","shell.execute_reply":"2024-09-13T14:26:42.950557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"set_seed(719)","metadata":{"execution":{"iopub.status.busy":"2024-09-13T14:26:42.955441Z","iopub.execute_input":"2024-09-13T14:26:42.955751Z","iopub.status.idle":"2024-09-13T14:26:42.992189Z","shell.execute_reply.started":"2024-09-13T14:26:42.955723Z","shell.execute_reply":"2024-09-13T14:26:42.991386Z"},"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-09-13T14:28:38.633038Z","iopub.execute_input":"2024-09-13T14:28:38.633424Z","iopub.status.idle":"2024-09-13T14:28:38.680290Z","shell.execute_reply.started":"2024-09-13T14:28:38.633385Z","shell.execute_reply":"2024-09-13T14:28:38.679427Z"},"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-09-13T14:28:39.535194Z","iopub.execute_input":"2024-09-13T14:28:39.535578Z","iopub.status.idle":"2024-09-13T14:28:39.580565Z","shell.execute_reply.started":"2024-09-13T14:28:39.535541Z","shell.execute_reply":"2024-09-13T14:28:39.579849Z"},"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-09-13T14:28:42.045472Z","iopub.execute_input":"2024-09-13T14:28:42.045811Z","iopub.status.idle":"2024-09-13T14:28:42.079593Z","shell.execute_reply.started":"2024-09-13T14:28:42.045782Z","shell.execute_reply":"2024-09-13T14:28:42.078705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"photo_img, monet_img = next(iter(img_dl))","metadata":{"execution":{"iopub.status.busy":"2024-09-13T14:28:42.342937Z","iopub.execute_input":"2024-09-13T14:28:42.343288Z","iopub.status.idle":"2024-09-13T14:28:45.871891Z","shell.execute_reply.started":"2024-09-13T14:28:42.343251Z","shell.execute_reply":"2024-09-13T14:28:45.870950Z"},"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-09-13T14:28:55.845911Z","iopub.execute_input":"2024-09-13T14:28:55.846278Z","iopub.status.idle":"2024-09-13T14:28:55.881907Z","shell.execute_reply.started":"2024-09-13T14:28:55.846247Z","shell.execute_reply":"2024-09-13T14:28:55.881084Z"},"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-09-13T14:29:51.758001Z","iopub.execute_input":"2024-09-13T14:29:51.758386Z","iopub.status.idle":"2024-09-13T14:29:52.100827Z","shell.execute_reply.started":"2024-09-13T14:29:51.758347Z","shell.execute_reply":"2024-09-13T14:29:52.099832Z"},"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-09-13T14:30:05.956793Z","iopub.execute_input":"2024-09-13T14:30:05.957121Z","iopub.status.idle":"2024-09-13T14:30:05.991057Z","shell.execute_reply.started":"2024-09-13T14:30:05.957093Z","shell.execute_reply":"2024-09-13T14:30:05.990146Z"},"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-09-13T14:30:06.358867Z","iopub.execute_input":"2024-09-13T14:30:06.359203Z","iopub.status.idle":"2024-09-13T14:30:06.393516Z","shell.execute_reply.started":"2024-09-13T14:30:06.359174Z","shell.execute_reply":"2024-09-13T14:30:06.392620Z"},"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-09-13T14:30:07.483051Z","iopub.execute_input":"2024-09-13T14:30:07.483421Z","iopub.status.idle":"2024-09-13T14:30:07.521386Z","shell.execute_reply.started":"2024-09-13T14:30:07.483387Z","shell.execute_reply":"2024-09-13T14:30:07.520514Z"},"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-09-13T14:30:09.103673Z","iopub.execute_input":"2024-09-13T14:30:09.104079Z","iopub.status.idle":"2024-09-13T14:30:09.148236Z","shell.execute_reply.started":"2024-09-13T14:30:09.104039Z","shell.execute_reply":"2024-09-13T14:30:09.147291Z"},"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-09-13T14:30:09.553180Z","iopub.execute_input":"2024-09-13T14:30:09.553551Z","iopub.status.idle":"2024-09-13T14:30:09.591912Z","shell.execute_reply.started":"2024-09-13T14:30:09.553518Z","shell.execute_reply":"2024-09-13T14:30:09.590838Z"},"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-09-13T14:30:09.890768Z","iopub.execute_input":"2024-09-13T14:30:09.891096Z","iopub.status.idle":"2024-09-13T14:30:09.931563Z","shell.execute_reply.started":"2024-09-13T14:30:09.891066Z","shell.execute_reply":"2024-09-13T14:30:09.930508Z"},"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-09-13T14:30:10.341324Z","iopub.execute_input":"2024-09-13T14:30:10.341680Z","iopub.status.idle":"2024-09-13T14:30:10.381001Z","shell.execute_reply.started":"2024-09-13T14:30:10.341645Z","shell.execute_reply":"2024-09-13T14:30:10.380046Z"},"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-09-13T14:30:11.714035Z","iopub.execute_input":"2024-09-13T14:30:11.714423Z","iopub.status.idle":"2024-09-13T14:30:11.753491Z","shell.execute_reply.started":"2024-09-13T14:30:11.714386Z","shell.execute_reply":"2024-09-13T14:30:11.752529Z"},"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-09-13T14:30:14.128551Z","iopub.execute_input":"2024-09-13T14:30:14.128886Z","iopub.status.idle":"2024-09-13T14:30:14.162991Z","shell.execute_reply.started":"2024-09-13T14:30:14.128858Z","shell.execute_reply":"2024-09-13T14:30:14.162069Z"},"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=50):\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-09-13T14:30:14.931288Z","iopub.execute_input":"2024-09-13T14:30:14.931662Z","iopub.status.idle":"2024-09-13T14:30:14.970158Z","shell.execute_reply.started":"2024-09-13T14:30:14.931630Z","shell.execute_reply":"2024-09-13T14:30:14.969313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class lr_sched():\n    def __init__(self, decay_epochs=100, total_epochs=200):\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-09-13T14:30:16.934112Z","iopub.execute_input":"2024-09-13T14:30:16.934665Z","iopub.status.idle":"2024-09-13T14:30:16.970079Z","shell.execute_reply.started":"2024-09-13T14:30:16.934624Z","shell.execute_reply":"2024-09-13T14:30:16.969300Z"},"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-09-13T14:30:17.877487Z","iopub.execute_input":"2024-09-13T14:30:17.877826Z","iopub.status.idle":"2024-09-13T14:30:17.913151Z","shell.execute_reply.started":"2024-09-13T14:30:17.877796Z","shell.execute_reply":"2024-09-13T14:30:17.912298Z"},"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-09-13T14:30:26.724482Z","iopub.execute_input":"2024-09-13T14:30:26.724861Z","iopub.status.idle":"2024-09-13T14:30:26.804285Z","shell.execute_reply.started":"2024-09-13T14:30:26.724830Z","shell.execute_reply":"2024-09-13T14:30:26.803452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gan = CycleGAN(1, 1, 50, device)","metadata":{"execution":{"iopub.status.busy":"2024-09-13T14:30:53.948698Z","iopub.execute_input":"2024-09-13T14:30:53.949068Z","iopub.status.idle":"2024-09-13T14:30:54.396814Z","shell.execute_reply.started":"2024-09-13T14:30:53.949032Z","shell.execute_reply":"2024-09-13T14:30:54.395482Z"},"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-09-13T14:30:55.118721Z","iopub.execute_input":"2024-09-13T14:30:55.119093Z","iopub.status.idle":"2024-09-13T14:30:55.155460Z","shell.execute_reply.started":"2024-09-13T14:30:55.119056Z","shell.execute_reply":"2024-09-13T14:30:55.154621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_checkpoint(save_dict, 'init.ckpt')","metadata":{"execution":{"iopub.status.busy":"2024-09-13T14:30:56.083554Z","iopub.execute_input":"2024-09-13T14:30:56.083881Z","iopub.status.idle":"2024-09-13T14:30:56.211639Z","shell.execute_reply.started":"2024-09-13T14:30:56.083852Z","shell.execute_reply":"2024-09-13T14:30:56.210849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gan.train(img_dl)","metadata":{"execution":{"iopub.status.busy":"2024-09-13T14:30:57.165586Z","iopub.execute_input":"2024-09-13T14:30:57.165919Z","iopub.status.idle":"2024-09-13T14:32:28.682176Z","shell.execute_reply.started":"2024-09-13T14:30:57.165890Z","shell.execute_reply":"2024-09-13T14:32:28.680307Z"},"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-09-13T14:32:30.996168Z","iopub.execute_input":"2024-09-13T14:32:30.996526Z","iopub.status.idle":"2024-09-13T14:32:31.183307Z","shell.execute_reply.started":"2024-09-13T14:32:30.996495Z","shell.execute_reply":"2024-09-13T14:32:31.182428Z"},"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":{"execution":{"iopub.status.busy":"2024-09-13T14:32:32.159921Z","iopub.execute_input":"2024-09-13T14:32:32.160293Z","iopub.status.idle":"2024-09-13T14:32:32.834981Z","shell.execute_reply.started":"2024-09-13T14:32:32.160253Z","shell.execute_reply":"2024-09-13T14:32:32.834150Z"},"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-09-13T14:32:37.046241Z","iopub.execute_input":"2024-09-13T14:32:37.046594Z","iopub.status.idle":"2024-09-13T14:32:37.088459Z","shell.execute_reply.started":"2024-09-13T14:32:37.046564Z","shell.execute_reply":"2024-09-13T14:32:37.087510Z"},"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-09-13T14:32:43.210814Z","iopub.execute_input":"2024-09-13T14:32:43.211150Z","iopub.status.idle":"2024-09-13T14:32:43.248191Z","shell.execute_reply.started":"2024-09-13T14:32:43.211119Z","shell.execute_reply":"2024-09-13T14:32:43.247328Z"},"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-09-13T14:32:45.169242Z","iopub.execute_input":"2024-09-13T14:32:45.169588Z","iopub.status.idle":"2024-09-13T14:32:45.203719Z","shell.execute_reply.started":"2024-09-13T14:32:45.169558Z","shell.execute_reply":"2024-09-13T14:32:45.202793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir ../images","metadata":{"execution":{"iopub.status.busy":"2024-09-13T14:32:46.967919Z","iopub.execute_input":"2024-09-13T14:32:46.968283Z","iopub.status.idle":"2024-09-13T14:32:48.028169Z","shell.execute_reply.started":"2024-09-13T14:32:46.968242Z","shell.execute_reply":"2024-09-13T14:32:48.027249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trans = transforms.ToPILImage()","metadata":{"execution":{"iopub.status.busy":"2024-09-13T14:32:48.030394Z","iopub.execute_input":"2024-09-13T14:32:48.030793Z","iopub.status.idle":"2024-09-13T14:32:48.068907Z","shell.execute_reply.started":"2024-09-13T14:32:48.030750Z","shell.execute_reply":"2024-09-13T14:32:48.068095Z"},"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-09-13T14:32:48.835471Z","iopub.execute_input":"2024-09-13T14:32:48.835839Z","iopub.status.idle":"2024-09-13T14:32:54.944836Z","shell.execute_reply.started":"2024-09-13T14:32:48.835803Z","shell.execute_reply":"2024-09-13T14:32:54.943168Z"},"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-09-13T14:32:54.946287Z","iopub.status.idle":"2024-09-13T14:32:54.946931Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}