{"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":"markdown","source":"## Data","metadata":{}},{"cell_type":"code","source":"import os\nimport os.path as osp\nfrom PIL import Image\nfrom torch.utils.data import Dataset\nimport torchvision.transforms as T\n\n\nclass Monet(Dataset):\n    def __init__(self, image_dir: str, transform=T.ToTensor()):\n        self.image_dir = image_dir\n        self.file_names = sorted(os.listdir(image_dir))\n        self.transform = transform\n\n    def __getitem__(self, idx):\n        file_name = self.file_names[idx]\n        image_file = osp.join(self.image_dir, file_name)\n\n        img = Image.open(image_file)\n        img = self.transform(img)\n\n        return img\n    \n    def __len__(self):\n        return len(self.file_names)\n\ndataset = Monet('../input/gan-getting-started/monet_jpg')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-26T06:59:03.916189Z","iopub.execute_input":"2022-07-26T06:59:03.916568Z","iopub.status.idle":"2022-07-26T06:59:06.160469Z","shell.execute_reply.started":"2022-07-26T06:59:03.916501Z","shell.execute_reply":"2022-07-26T06:59:06.159456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nplt.imshow(dataset[0].permute(1,2,0))\nplt.axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-26T06:59:06.162622Z","iopub.execute_input":"2022-07-26T06:59:06.163181Z","iopub.status.idle":"2022-07-26T06:59:06.386778Z","shell.execute_reply.started":"2022-07-26T06:59:06.163144Z","shell.execute_reply":"2022-07-26T06:59:06.385963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\nclass Generator(nn.Module):\n    def __init__(self, nz=100):\n        super().__init__()\n        self.gen = nn.Sequential(\n            # input is Z, going into a convolution\n            nn.ConvTranspose2d(nz, 1024, 4, 1, 0, bias=False),\n            nn.BatchNorm2d(1024),\n            nn.ReLU(True),\n            # state size. (1024) x 4 x 4\n            nn.ConvTranspose2d(1024, 512, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(True),\n            # state size. (512) x 8 x 8\n            nn.ConvTranspose2d(512, 256, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(True),\n            # state size. (256) x 16 x 16\n            nn.ConvTranspose2d(256, 128, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(True),\n            # state size. (128) x 32 x 32\n            nn.ConvTranspose2d(128, 64, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(True),\n            # state size. (64) x 64 x 64\n            nn.ConvTranspose2d(64, 32, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(32),\n            nn.ReLU(True),\n            # state size. (32) x 128 x 128\n            nn.ConvTranspose2d(32, 3, 4, 2, 1, bias=False),\n            nn.Tanh()\n            # state size. (3) x 256 x 256\n        )\n\n    def forward(self, input):\n        return self.gen(input)\n\n\nclass Discriminator(nn.Module):\n    def __init__(self, nc=3, ndf=32):\n        super().__init__()\n        self.disc = nn.Sequential(\n            # input is (nc) x 256 x 256\n            nn.Conv2d(nc, ndf, 4, 2, 1, bias=False),\n            nn.LeakyReLU(0.2, inplace=True),\n            # state size. (ndf) x 128 x 128\n            nn.Conv2d(ndf, ndf * 2, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(ndf * 2),\n            nn.LeakyReLU(0.2, inplace=True),\n            # state size. (ndf*2) x 64 x 64\n            nn.Conv2d(ndf * 2, ndf * 4, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(ndf * 4),\n            nn.LeakyReLU(0.2, inplace=True),\n            # state size. (ndf*4) x 32 x 32\n            nn.Conv2d(ndf * 4, ndf * 8, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(ndf * 8),\n            nn.LeakyReLU(0.2, inplace=True),\n            # state size. (ndf*8) x 16 x 16\n            nn.Conv2d(ndf * 8, ndf * 16, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(ndf * 16),\n            nn.LeakyReLU(0.2, inplace=True),\n            # state size. (ndf*16) x 8 x 8\n            nn.Conv2d(ndf * 16, ndf * 32, 4, 2, 1, bias=False),\n            nn.BatchNorm2d(ndf * 32),\n            nn.LeakyReLU(0.2, inplace=True),\n            # state size. (ndf*16) x 4 x 4\n            nn.Conv2d(ndf * 32, 1, 4, 1, 0, bias=False),\n            nn.Sigmoid()\n        )\n\n    def forward(self, input):\n        return self.disc(input)","metadata":{"execution":{"iopub.status.busy":"2022-07-26T06:59:27.586705Z","iopub.execute_input":"2022-07-26T06:59:27.587087Z","iopub.status.idle":"2022-07-26T06:59:27.604584Z","shell.execute_reply.started":"2022-07-26T06:59:27.587054Z","shell.execute_reply":"2022-07-26T06:59:27.603816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\ntorch.manual_seed(1)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False\n\nbatch_size = 64\nnz = 100\nlr = 0.0002\nbeta1 = 0.5\n\nnetG = Generator(nz)\nnetG = netG.to(device)\nnetD = Discriminator()\nnetD = netD.to(device)\n\ncriterion = nn.BCELoss()\n\nfixed_noise = torch.randn(batch_size, nz, 1, 1, device=device)\n\nreal_label = 1.\nfake_label = 0.\n\noptD = torch.optim.Adam(netD.parameters(), lr=lr, betas=(beta1, 0.999))\noptG = torch.optim.Adam(netG.parameters(), lr=lr, betas=(beta1, 0.999))","metadata":{"execution":{"iopub.status.busy":"2022-07-26T06:59:41.643557Z","iopub.execute_input":"2022-07-26T06:59:41.644121Z","iopub.status.idle":"2022-07-26T06:59:44.694806Z","shell.execute_reply.started":"2022-07-26T06:59:41.644074Z","shell.execute_reply":"2022-07-26T06:59:44.693864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader\nimport torchvision.utils as vutils\n\nimg_list = []\nG_losses = []\nD_losses = []\niters = 0\n\nnum_epochs = 300\n\ndataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True)\n\nprint('Starting Training Loop...')\n# For each epoch\nfor epoch in range(num_epochs):\n    for i, data in enumerate(dataloader):\n        ############################\n        # (1) Update D network: maximize log(D(x)) + log(1 - D(G(z)))\n        ###########################\n        ## Train with all-real batch\n        netD.zero_grad()\n        # Format batch\n        data = data.to(device)\n        label = torch.full((data.shape[0],), real_label, dtype=torch.float, device=device)\n        # Forward pass real batch through D\n        output = netD(data).view(-1)\n        # Calculate loss on all-real batch\n        errD_real = criterion(output, label)\n        # Calculate gradients for D in backward pass\n        errD_real.backward()\n        D_x = output.mean().item()\n\n        ## Train with all-fake batch\n        # Generate batch of latent vectors\n        noise = torch.randn(data.shape[0], nz, 1, 1, device=device)\n        # Generate fake image batch with G\n        fake = netG(noise)\n        label.fill_(fake_label)\n        # Classify all fake batch with D\n        output = netD(fake.detach()).view(-1)\n        # Calculate D's loss on the all-fake batch\n        errD_fake = criterion(output, label)\n        # Calculate the gradients for this batch, accumulated (summed) with previous gradients\n        errD_fake.backward()\n        D_G_z1 = output.mean().item()\n        # Compute error of D as sum over the fake and the real batches\n        errD = errD_real + errD_fake\n        # Update D\n        optD.step()\n\n        ############################\n        # (2) Update G network: maximize log(D(G(z)))\n        ###########################\n        netG.zero_grad()\n        label.fill_(real_label)  # fake labels are real for generator cost\n        # Since we just updated D, perform another forward pass of all-fake batch through D\n        output = netD(fake).view(-1)\n        # Calculate G's loss based on this output\n        errG = criterion(output, label)\n        # Calculate gradients for G\n        errG.backward()\n        D_G_z2 = output.mean().item()\n        # Update G\n        optG.step()\n\n        # Output training stats\n        print('[%d/%d][%d/%d]\\tLoss_D: %.4f\\tLoss_G: %.4f\\tD(x): %.4f\\tD(G(z)): %.4f / %.4f'\n                  % (epoch, num_epochs, i, len(dataloader),\n                     errD.item(), errG.item(), D_x, D_G_z1, D_G_z2))\n\n        # Save Losses for plotting later\n        G_losses.append(errG.item())\n        D_losses.append(errD.item())\n\n        # Check how the generator is doing by saving G's output on fixed_noise\n        if (iters % 500 == 0) or ((epoch == num_epochs-1) and (i == len(dataloader)-1)):\n            with torch.no_grad():\n                fake = netG(fixed_noise).detach().cpu()\n            img_list.append(vutils.make_grid(fake, padding=2, normalize=True))\n\n        iters += 1","metadata":{"execution":{"iopub.status.busy":"2022-06-14T09:24:05.342965Z","iopub.execute_input":"2022-06-14T09:24:05.343533Z","iopub.status.idle":"2022-06-14T09:40:09.688381Z","shell.execute_reply.started":"2022-06-14T09:24:05.343494Z","shell.execute_reply":"2022-06-14T09:40:09.687547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10,5))\nplt.title(\"Generator and Discriminator Loss During Training\")\nplt.plot(G_losses,label=\"G\")\nplt.plot(D_losses,label=\"D\")\nplt.xlabel(\"Iterations\")\nplt.ylabel(\"Loss\")\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-06-14T09:40:09.689692Z","iopub.execute_input":"2022-06-14T09:40:09.690036Z","iopub.status.idle":"2022-06-14T09:40:09.889856Z","shell.execute_reply.started":"2022-06-14T09:40:09.690008Z","shell.execute_reply":"2022-06-14T09:40:09.889113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import HTML\nimport numpy as np\nimport matplotlib.animation as animation\n\nfig = plt.figure(figsize=(8,8))\nplt.axis(\"off\")\nims = [[plt.imshow(np.transpose(i,(1,2,0)), animated=True)] for i in img_list]\nani = animation.ArtistAnimation(fig, ims, interval=1000, repeat_delay=1000, blit=True)\n\nHTML(ani.to_jshtml())","metadata":{"execution":{"iopub.status.busy":"2022-06-14T09:40:09.891004Z","iopub.execute_input":"2022-06-14T09:40:09.893617Z","iopub.status.idle":"2022-06-14T09:40:13.705029Z","shell.execute_reply.started":"2022-06-14T09:40:09.893572Z","shell.execute_reply":"2022-06-14T09:40:13.704241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Competition Output","metadata":{}},{"cell_type":"code","source":"import os\nimport shutil\nfrom PIL import Image\n\nos.mkdir('../images')\n\nfor i in range(7000):\n    noise = torch.randn(1, nz, 1, 1, device=device)\n    im = netG(noise).detach().squeeze(0).permute(1,2,0).cpu().numpy()\n    im = (im * 255).astype(np.uint8)\n    im = Image.fromarray(im)\n    im.save('../images/' + str(i) + '.jpg')\n\nshutil.make_archive('/kaggle/working/images', 'zip', '/kaggle/images')","metadata":{},"execution_count":null,"outputs":[]}]}