{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\n\nimport matplotlib\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport missingno as msno\n\nimport torchvision\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport os\nimport time\nimport shutil\n\nimport itertools\n\nfrom torchvision import transforms\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\n\nfrom tqdm.notebook import tqdm\n\nfrom PIL import Image","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-03-24T07:30:47.513383Z","iopub.execute_input":"2022-03-24T07:30:47.513914Z","iopub.status.idle":"2022-03-24T07:30:47.520437Z","shell.execute_reply.started":"2022-03-24T07:30:47.513877Z","shell.execute_reply":"2022-03-24T07:30:47.519782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preprocessing","metadata":{}},{"cell_type":"code","source":"class trainDataset(Dataset):\n    def __init__(self, data_dir,mode = 'train',transforms=None):\n            A_dir = os.path.join(data_dir, 'monet_jpg')\n            B_dir = os.path.join(data_dir, 'photo_jpg')\n            \n            if mode == 'train':\n                self.A = [os.path.join(A_dir, name) for name in sorted(os.listdir(A_dir))[:300]]\n                self.B = [os.path.join(B_dir, name) for name in sorted(os.listdir(B_dir))[:300]]\n            elif mode == 'test':\n                self.A = [os.path.join(B_dir, name) for name in sorted(os.listdir(B_dir))]\n                self.B = [os.path.join(B_dir, name) for name in sorted(os.listdir(B_dir))]\n\n            self.transforms = transforms\n\n    def __len__(self):\n        return len(self.A)\n\n    def __getitem__(self, index):\n        A = self.A[index]\n        B = self.B[index]\n\n        A = Image.open(A)\n        B = Image.open(B)\n\n        if self.transforms is not None:\n            A = self.transforms(A)\n            B = self.transforms(B)\n\n        return B,A","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:30:48.71441Z","iopub.execute_input":"2022-03-24T07:30:48.714947Z","iopub.status.idle":"2022-03-24T07:30:48.726197Z","shell.execute_reply.started":"2022-03-24T07:30:48.714908Z","shell.execute_reply":"2022-03-24T07:30:48.725398Z"},"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":"2022-03-24T07:30:52.347114Z","iopub.execute_input":"2022-03-24T07:30:52.34769Z","iopub.status.idle":"2022-03-24T07:30:52.353479Z","shell.execute_reply.started":"2022-03-24T07:30:52.347651Z","shell.execute_reply":"2022-03-24T07:30:52.35262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform_train = transforms.Compose([\n    transforms.Resize((256,256)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.ToTensor(),\n    transforms.Normalize((0.5,0.5,0.5), (0.5, 0.5, 0.5))\n])\ntransform_test = transforms.Compose([\n    transforms.Resize((256,256)),\n    transforms.ToTensor(),\n    transforms.Normalize((0.5,0.5,0.5), (0.5, 0.5, 0.5))\n])\n","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:30:50.300463Z","iopub.execute_input":"2022-03-24T07:30:50.300736Z","iopub.status.idle":"2022-03-24T07:30:50.307271Z","shell.execute_reply.started":"2022-03-24T07:30:50.300686Z","shell.execute_reply":"2022-03-24T07:30:50.306213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainLoader = DataLoader(\n    trainDataset('../input/gan-getting-started','train',transform_train),\n    batch_size = 1,\n    shuffle = True,\n    pin_memory = True\n)\ntestLoader = DataLoader(\n    trainDataset('../input/gan-getting-started','test',transform_test),\n    batch_size = 1,\n    shuffle = True,\n    pin_memory = True\n)","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:30:50.626917Z","iopub.execute_input":"2022-03-24T07:30:50.627175Z","iopub.status.idle":"2022-03-24T07:30:50.673008Z","shell.execute_reply.started":"2022-03-24T07:30:50.627149Z","shell.execute_reply":"2022-03-24T07:30:50.672312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Some additional classes and functions","metadata":{}},{"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)\n","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:30:52.711059Z","iopub.execute_input":"2022-03-24T07:30:52.711835Z","iopub.status.idle":"2022-03-24T07:30:52.716961Z","shell.execute_reply.started":"2022-03-24T07:30:52.71179Z","shell.execute_reply":"2022-03-24T07:30:52.716166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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":"2022-03-24T07:30:51.98909Z","iopub.execute_input":"2022-03-24T07:30:51.989284Z","iopub.status.idle":"2022-03-24T07:30:51.997247Z","shell.execute_reply.started":"2022-03-24T07:30:51.989262Z","shell.execute_reply":"2022-03-24T07:30:51.996491Z"},"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":"2022-03-24T07:30:53.083973Z","iopub.execute_input":"2022-03-24T07:30:53.084437Z","iopub.status.idle":"2022-03-24T07:30:53.089578Z","shell.execute_reply.started":"2022-03-24T07:30:53.084403Z","shell.execute_reply":"2022-03-24T07:30:53.088759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def update_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":"2022-03-24T07:30:49.925979Z","iopub.execute_input":"2022-03-24T07:30:49.926686Z","iopub.status.idle":"2022-03-24T07:30:49.931364Z","shell.execute_reply.started":"2022-03-24T07:30:49.926643Z","shell.execute_reply":"2022-03-24T07:30:49.930178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class ResidualBlock(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(ResidualBlock,self).__init__()\n        self.seq = nn.Sequential(\n        nn.ReflectionPad2d(1),\n        nn.Conv2d(in_channels = in_channels,out_channels = out_channels ,kernel_size = (3,3),stride = 1),\n        nn.LeakyReLU(negative_slope=0.2, inplace=True),\n        nn.InstanceNorm2d(out_channels),\n        nn.Dropout(0.5),\n        nn.ReflectionPad2d(1),\n        nn.Conv2d(in_channels = in_channels,out_channels = out_channels,kernel_size = (3,3),stride = 1),\n        nn.InstanceNorm2d(out_channels)\n        )\n        \n        \n        \n    def forward(self,X):\n        return X+self.seq(X)","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:30:51.929395Z","iopub.execute_input":"2022-03-24T07:30:51.929667Z","iopub.status.idle":"2022-03-24T07:30:51.936608Z","shell.execute_reply.started":"2022-03-24T07:30:51.92964Z","shell.execute_reply":"2022-03-24T07:30:51.935764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Discriminator(nn.Module):\n    def __init__(self, in_channels):\n        super(Discriminator, self).__init__()\n       \n        self.seq = nn.Sequential(\n            \n        nn.Conv2d(in_channels = in_channels,out_channels = 64,kernel_size = (4,4),stride = 2),\n        nn.LeakyReLU(negative_slope=0.2, inplace=True),\n        nn.InstanceNorm2d(num_features = 64),\n        nn.Conv2d(in_channels = 64,out_channels = 128,kernel_size = (4,4),stride = 2),\n        nn.InstanceNorm2d(num_features = 128),\n        nn.LeakyReLU(negative_slope=0.2, inplace=True),\n        nn.Conv2d(in_channels = 128,out_channels = 256,kernel_size = (4,4),stride = 2),\n        nn.LeakyReLU(negative_slope=0.2, inplace=True),\n        nn.InstanceNorm2d(num_features = 256),\n        nn.Conv2d(in_channels = 256,out_channels = 512,kernel_size = (4,4),stride = 2),\n        nn.InstanceNorm2d(num_features = 512),\n        nn.Conv2d(in_channels = 512,out_channels = 1,kernel_size = (4,4),stride = 1)\n            \n        )\n    \n    \n    def forward(self,X):\n        output = self.seq(X)\n        return output","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:30:51.938851Z","iopub.execute_input":"2022-03-24T07:30:51.939111Z","iopub.status.idle":"2022-03-24T07:30:51.950605Z","shell.execute_reply.started":"2022-03-24T07:30:51.939077Z","shell.execute_reply":"2022-03-24T07:30:51.94957Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Generator(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(Generator,self).__init__()\n        \n        self.seq = nn.Sequential(\n        nn.ReflectionPad2d(3),\n        nn.Conv2d(in_channels = in_channels,out_channels = 64,kernel_size = (7,7),stride = 1,padding = (0,0)),\n        nn.InstanceNorm2d(64),\n        nn.Conv2d(in_channels = 64,out_channels = 128,kernel_size = (3,3),stride = 2,padding = (1,1)),\n        nn.InstanceNorm2d(128),\n        nn.Conv2d(in_channels =128,out_channels = 256,kernel_size = (3,3),stride = 2,padding = (1,1)),\n        nn.InstanceNorm2d(256),\n        ResidualBlock(in_channels = 256,out_channels = 256),\n        ResidualBlock(in_channels = 256,out_channels = 256),\n        ResidualBlock(in_channels = 256,out_channels = 256),\n        ResidualBlock(in_channels = 256,out_channels = 256),\n        ResidualBlock(in_channels = 256,out_channels = 256),\n        ResidualBlock(in_channels = 256,out_channels = 256),\n        ResidualBlock(in_channels = 256,out_channels = 256),\n        ResidualBlock(in_channels = 256,out_channels = 256),\n        ResidualBlock(in_channels = 256,out_channels = 256),\n        nn.ConvTranspose2d(in_channels = 256,out_channels = 128,kernel_size = (3,3),stride = 2,padding=1, output_padding=1),\n        nn.InstanceNorm2d(128),\n        nn.Dropout(0.5),\n        nn.GELU(),\n        nn.ConvTranspose2d(in_channels = 128,out_channels = 64,kernel_size = (3,3),stride = 2,padding=1, output_padding=1),\n        nn.InstanceNorm2d(64),\n        nn.Dropout(0.5),\n        nn.GELU(),\n        nn.ReflectionPad2d(3),\n        nn.Conv2d(in_channels = 64,out_channels = 3,kernel_size = (7,7),stride = 1,padding = (0,0)),\n        nn.Tanh()\n        )\n    \n    def forward(self,X):\n        output = self.seq(X)\n        return output","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:30:51.95221Z","iopub.execute_input":"2022-03-24T07:30:51.952513Z","iopub.status.idle":"2022-03-24T07:30:51.969403Z","shell.execute_reply.started":"2022-03-24T07:30:51.952479Z","shell.execute_reply":"2022-03-24T07:30:51.96839Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# GAN Architecture","metadata":{}},{"cell_type":"code","source":"class CycleGan(object):\n    def __init__(self,in_channels,out_channels,epochs,device,decay_epoch,lmbda,idt_coef):\n        self.epochs = epochs\n        self.decay_epoch = decay_epoch\n        self.device = device\n        self.GeneratorAB = Generator(in_channels = in_channels,out_channels = out_channels).to(device) \n        self.GeneratorBA = Generator(in_channels = in_channels,out_channels = out_channels).to(device) \n        self.DiscriminatorA = Discriminator(in_channels = in_channels).to(device)\n        self.DiscriminatorB = Discriminator(in_channels = in_channels).to(device)\n\n        self.AdamGenerator = torch.optim.Adam(itertools.chain(self.GeneratorBA.parameters(),\n                                                              self.GeneratorAB.parameters()),\n                                              lr = 2e-4,betas=(0.5, 0.999))\n        self.AdamDiscriminator = torch.optim.Adam(itertools.chain(self.DiscriminatorB.parameters(), \n                                                                  self.DiscriminatorA.parameters()),\n                                              lr = 2e-4,betas=(0.5, 0.999))\n        \n\n        self.mse_loss = nn.MSELoss()\n        self.l1_loss = nn.L1Loss()\n        self.sample_B = sample_fake()\n        self.sample_A = sample_fake()  \n        \n        self.lmbda = lmbda\n        self.idt_coef = idt_coef\n        \n        self.Generator_stats = AvgStats()\n        self.Discriminator_stats = AvgStats()\n        \n        Generator_lr = lr_sched(self.decay_epoch, self.epochs)\n        Discriminator_lr = lr_sched(self.decay_epoch, self.epochs)\n        self.Generator_lr_sched = torch.optim.lr_scheduler.LambdaLR(self.AdamGenerator, Generator_lr.step)\n        self.Discriminator_lr_sched = torch.optim.lr_scheduler.LambdaLR(self.AdamDiscriminator, Discriminator_lr.step)\n\n    def train(self,trainLoader):\n        for epoch in range(self.epochs):\n            start_time = time.time()\n            avg_gen_loss = 0.0\n            avg_disc_loss = 0.0\n            t = tqdm(trainLoader,leave = False,total = trainLoader.__len__())\n            for i,(A,B) in enumerate(t):\n\n\n                A_real,B_real = A.to(device),B.to(device)\n\n                update_grad([self.DiscriminatorA,self.DiscriminatorB],False)\n                self.AdamGenerator.zero_grad()\n\n                fake_B = self.GeneratorAB(A_real)\n                fake_A = self.GeneratorBA(fake_B)\n\n\n                cycle_A = self.GeneratorBA(fake_B) \n                cycle_B = self.GeneratorAB(fake_A) \n\n\n                id_B = self.GeneratorAB(B_real) \n                id_A = self.GeneratorBA(A_real) \n                \n                \n                loss_id_B =  self.l1_loss(cycle_B,id_B) * self.lmbda * self.idt_coef\n                loss_id_A =  self.l1_loss(cycle_A,id_A) * self.lmbda * self.idt_coef\n                \n                \n                loss_cycle_B =  self.l1_loss(cycle_B,B_real) * self.lmbda\n                loss_cycle_A =  self.l1_loss(cycle_A,A_real) * self.lmbda\n\n\n                disc_A = self.DiscriminatorA(fake_A)\n                disc_B = self.DiscriminatorB(fake_B)\n                \n                real = torch.ones(disc_A.size()).to(device)\n\n                \n                loss_adversial_B = self.mse_loss(disc_B,real)\n                loss_adversial_A = self.mse_loss(disc_A,real)\n\n                generator_loss =    loss_id_B+loss_id_A + \\\n                                    loss_cycle_B + loss_cycle_A +\\\n                                    loss_adversial_B+loss_adversial_A\n                \n                avg_gen_loss += generator_loss.item()\n\n                generator_loss.backward()\n                self.AdamGenerator.step()\n\n                update_grad([self.DiscriminatorA,self.DiscriminatorB],True)\n\n                self.AdamDiscriminator.zero_grad()\n\n\n                fake_B = self.sample_B([fake_B.cpu().data.numpy()])[0]\n                fake_A = self.sample_A([fake_A.cpu().data.numpy()])[0]\n                fake_B = torch.tensor(fake_B).to(self.device)\n                fake_A = torch.tensor(fake_A).to(self.device)\n\n                disc_A_real = self.DiscriminatorA(A_real)\n                disc_B_real = self.DiscriminatorB(B_real)\n                disc_A_fake = self.DiscriminatorA(fake_A)\n                disc_B_fake = self.DiscriminatorB(fake_B)\n\n                real = torch.ones(disc_A_real.size()).to(device)\n                fake = torch.zeros(disc_A_fake.size()).to(device)\n\n\n                B_desc_real_loss = self.mse_loss(disc_B_real, real)\n                B_desc_fake_loss = self.mse_loss(disc_B_fake, fake)\n                A_desc_real_loss = self.mse_loss(disc_A_real, real)\n                A_desc_fake_loss = self.mse_loss(disc_A_fake, fake)\n\n                B_desc_loss = (B_desc_real_loss + B_desc_fake_loss) / 2\n                A_desc_loss = (A_desc_real_loss + A_desc_fake_loss) / 2\n                disc_loss = B_desc_loss + A_desc_loss\n                avg_disc_loss += disc_loss.item()\n\n                B_desc_loss.backward()\n                A_desc_loss.backward()\n                self.AdamDiscriminator.step()\n\n                t.set_postfix(gen_loss=generator_loss.item(), disc_loss=disc_loss.item())\n    \n\n            save_dict = {\n                    'epoch': epoch+1,\n                    'GeneratorAB': gan.GeneratorAB.state_dict(),\n                    'GeneratorBA': gan.GeneratorBA.state_dict(),\n                    'disc_m': gan.DiscriminatorA.state_dict(),\n                    'disc_p': gan.DiscriminatorB.state_dict(),\n                    'optimizer_gan': gan.AdamGenerator.state_dict(),\n                    'optimizer_disc': gan.AdamDiscriminator.state_dict()\n                }\n            save_checkpoint(gan, './current.ckpt')\n\n            avg_gen_loss /= trainLoader.__len__()\n            avg_disc_loss /= trainLoader.__len__()\n            time_req = time.time() - start_time\n\n            self.Generator_stats.append(avg_gen_loss, time_req)\n            self.Discriminator_stats.append(avg_disc_loss, time_req)\n\n            print(\"Epoch: (%d) | Generator Loss:%f | Discriminator Loss:%f\" % \n                                                    (epoch+1, avg_gen_loss, avg_disc_loss))\n            self.Generator_lr_sched.step()\n            self.Discriminator_lr_sched.step()","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:30:54.130495Z","iopub.execute_input":"2022-03-24T07:30:54.130799Z","iopub.status.idle":"2022-03-24T07:30:54.159419Z","shell.execute_reply.started":"2022-03-24T07:30:54.130765Z","shell.execute_reply":"2022-03-24T07:30:54.158035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Save and load model","metadata":{}},{"cell_type":"code","source":"def save_checkpoint(state, save_path):\n    torch.save(state, save_path)","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:30:53.457004Z","iopub.execute_input":"2022-03-24T07:30:53.457455Z","iopub.status.idle":"2022-03-24T07:30:53.462963Z","shell.execute_reply.started":"2022-03-24T07:30:53.457419Z","shell.execute_reply":"2022-03-24T07:30:53.462135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2022-03-24T07:32:22.953808Z","iopub.execute_input":"2022-03-24T07:32:22.955942Z","iopub.status.idle":"2022-03-24T07:32:22.96129Z","shell.execute_reply.started":"2022-03-24T07:32:22.955905Z","shell.execute_reply":"2022-03-24T07:32:22.960517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model evaluation","metadata":{}},{"cell_type":"code","source":"device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\ngan = CycleGan(3, 3, 200, device = device,decay_epoch = 100,lmbda = 10,idt_coef = 0.5)","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:30:54.456072Z","iopub.execute_input":"2022-03-24T07:30:54.456322Z","iopub.status.idle":"2022-03-24T07:30:54.708006Z","shell.execute_reply.started":"2022-03-24T07:30:54.456295Z","shell.execute_reply":"2022-03-24T07:30:54.707256Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_dict = {\n    'epoch': 0,\n    'gen_mtp': gan.GeneratorAB.state_dict(),\n    'gen_ptm': gan.GeneratorBA.state_dict(),\n    'desc_m': gan.DiscriminatorA.state_dict(),\n    'desc_p': gan.DiscriminatorB.state_dict(),\n    'optimizer_gen': gan.AdamGenerator.state_dict(),\n    'optimizer_desc': gan.AdamDiscriminator.state_dict()\n}\nsave_checkpoint(save_dict, './init.ckpt')","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:30:54.809129Z","iopub.execute_input":"2022-03-24T07:30:54.809633Z","iopub.status.idle":"2022-03-24T07:30:55.064722Z","shell.execute_reply.started":"2022-03-24T07:30:54.809585Z","shell.execute_reply":"2022-03-24T07:30:55.063676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gan.train(trainLoader)","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:30:55.182782Z","iopub.execute_input":"2022-03-24T07:30:55.183292Z","iopub.status.idle":"2022-03-24T07:32:22.948809Z","shell.execute_reply.started":"2022-03-24T07:30:55.183257Z","shell.execute_reply":"2022-03-24T07:32:22.948005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"load_checkpoint('./current.ckpt')","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:32:22.962793Z","iopub.execute_input":"2022-03-24T07:32:22.96306Z","iopub.status.idle":"2022-03-24T07:32:23.784079Z","shell.execute_reply.started":"2022-03-24T07:32:22.963018Z","shell.execute_reply":"2022-03-24T07:32:23.783402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.xlabel(\"Epochs\")\nplt.ylabel(\"Losses\")\nplt.plot(gan.Generator_stats.losses, 'r', label='Generator Loss')\nplt.plot(gan.Discriminator_stats.losses, 'b', label='Discriminator Loss')\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:32:23.785935Z","iopub.execute_input":"2022-03-24T07:32:23.786266Z","iopub.status.idle":"2022-03-24T07:32:23.967266Z","shell.execute_reply.started":"2022-03-24T07:32:23.786226Z","shell.execute_reply":"2022-03-24T07:32:23.966571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Run Generator over all images ","metadata":{}},{"cell_type":"code","source":"!mkdir ../images","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:32:23.968634Z","iopub.execute_input":"2022-03-24T07:32:23.968913Z","iopub.status.idle":"2022-03-24T07:32:24.711135Z","shell.execute_reply.started":"2022-03-24T07:32:23.968877Z","shell.execute_reply":"2022-03-24T07:32:24.710267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"t = tqdm(testLoader, leave=False, total=testLoader.__len__())\nfor i, photo in enumerate(t):\n    with torch.no_grad():\n        pred_ = gan.GeneratorBA(photo[0].to(device)).cpu().detach()\n    pred_ = unnorm(pred_)\n    img = transforms.ToPILImage()(pred_[0]).convert(\"RGB\")\n    img.save(\"../images/\" + str(i+1) + \".jpg\")","metadata":{"execution":{"iopub.status.busy":"2022-03-24T07:32:41.254557Z","iopub.execute_input":"2022-03-24T07:32:41.255027Z","iopub.status.idle":"2022-03-24T07:32:49.05622Z","shell.execute_reply.started":"2022-03-24T07:32:41.254989Z","shell.execute_reply":"2022-03-24T07:32:49.05499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"shutil.make_archive(\"/kaggle/working/images\", 'zip', \"/kaggle/images\")","metadata":{"execution":{"iopub.status.busy":"2022-03-23T20:02:55.103366Z","iopub.execute_input":"2022-03-23T20:02:55.103947Z","iopub.status.idle":"2022-03-23T20:02:57.459549Z","shell.execute_reply.started":"2022-03-23T20:02:55.103908Z","shell.execute_reply":"2022-03-23T20:02:57.45862Z"},"trusted":true},"execution_count":null,"outputs":[]}]}