{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"import os, sys, glob, gc, copy\nimport time\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision import transforms\n# from torchsummary import summary\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport cv2\nimport matplotlib.pyplot as plt\n\nfrom tqdm import tqdm","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Check GPU is available or not\nis_cuda = torch.cuda.is_available()\nif is_cuda:\n    device = torch.device(\"cuda\")\n    print(\"GPU is available\")\nelse:\n    device = torch.device(\"cpu\")\n    print(\"GPU not available, CPU used\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"data_path = '../input/cassava-leaf-disease-classification/train_images/'\ntotal_csv = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Create Dataset"},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_img(path):\n    im_bgr = cv2.imread(path)\n    im_rgb = im_bgr[:, :, ::-1]\n    #print(im_rgb)\n    return im_rgb\n\n\nclass HDD_Dataset(Dataset):\n    def __init__(self, df, data_root, device, transform=None):\n        super().__init__()\n        self.df = df.reset_index(drop=True).copy()\n        self.data_root = data_root\n        self.device = device\n        self.transform = transform\n        self.N_class = 5  # number of classes\n    \n    def __len__(self):\n        return self.df.shape[0]\n    \n    def __getitem__(self, index: int):\n        # read label and one-hot\n        label = torch.tensor(self.df.iloc[index]['label'], device=device)\n        label_onehot = nn.functional.one_hot(label, self.N_class)\n        # prepare image\n        path = \"{}/{}\".format(self.data_root, self.df.iloc[index]['image_id'])\n        img = torch.tensor(get_img(path).copy()).float().to(device)/255  # normalize to 1\n        img = img.permute(2,0,1)\n        if self.transform:\n            img = self.transform(img)\n#         img = (img-0.5)*2  # normalize images to [-1, 1] and use tanh\n        return {'image':img, 'label':label, 'label_onehot':label_onehot.float()}","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## split data to `test` and `train`"},{"metadata":{"trusted":true},"cell_type":"code","source":"def frac_train_val(N:int, val_frac:float):\n    \"\"\"N: dataset length e.g. 2100\n        val_frac: validation fraction of total e.g. 0.2 \n        return indices of train, validation\"\"\"\n    perm = np.random.permutation(N)\n    thrshld = int(N*val_frac)\n    return perm[thrshld:].tolist(), perm[:thrshld].tolist()\n\ninds_train, inds_test = frac_train_val(total_csv.shape[0], 0.1)\ntrain_csv = total_csv.loc[inds_train]\ntest_csv = total_csv.loc[inds_test]\n#create datasets\ncrop_transform = transforms.RandomCrop(512)\ntrain_dataset = HDD_Dataset(train_csv, data_path, device, crop_transform)\ntest_dataset = HDD_Dataset(test_csv, data_path, device, crop_transform)\ngc.collect()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Define `Generator` and `Discriminator`"},{"metadata":{"trusted":true},"cell_type":"code","source":"class print_layer(nn.Module):\n    def __init__(self):\n        super(print_layer, self).__init__()\n    def forward(self, x):\n        print(x.shape)\n        return x\n\nclass Discriminator(nn.Module):\n    \"\"\"Discriminator\n        `image` input shape:(N, 3, 512, 512)\n        `c` shape: (N, 5)\"\"\"\n    def __init__(self, ydim:int):\n        super(Discriminator, self).__init__()\n        self.x_conv = nn.Sequential( \n            nn.Conv2d(3, 16, kernel_size=5, stride=2, padding=2), nn.LeakyReLU(),  # (256,256)\n            nn.Conv2d(16, 32, kernel_size=5, stride=1, padding=2),\n            nn.BatchNorm2d(32), nn.LeakyReLU(),\n            nn.Conv2d(32, 32, kernel_size=5, stride=1, padding=2),\n            nn.BatchNorm2d(32), nn.LeakyReLU(),\n            nn.Conv2d(32, 64, kernel_size=5, stride=2, padding=2),  # (128,128)\n            nn.BatchNorm2d(64), nn.LeakyReLU(),\n            nn.Conv2d(64, 64, kernel_size=5, stride=1, padding=2), nn.LeakyReLU(),\n            nn.Conv2d(64, 64, kernel_size=3, stride=2, padding=1),  # (64, 64)\n            nn.BatchNorm2d(64), nn.LeakyReLU(),\n            nn.Conv2d(64, 64, kernel_size=3, stride=1, padding=1), nn.LeakyReLU(),\n            nn.Conv2d(64, 128, kernel_size=3, stride=2, padding=1),  # (32, 32)\n            nn.BatchNorm2d(128), nn.LeakyReLU(),\n            nn.Conv2d(128, 128, kernel_size=3, stride=1, padding=1), nn.LeakyReLU(),\n            nn.Conv2d(128, 128, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(128), nn.LeakyReLU(),\n            nn.Conv2d(128, 128, kernel_size=3, stride=2, padding=1),  # (16, 16)\n            nn.Conv2d(128, 128, kernel_size=3, stride=1, padding=1), nn.LeakyReLU(),\n            nn.Conv2d(128, 256, kernel_size=3, stride=2, padding=1),  # (8, 8)\n            nn.BatchNorm2d(256), nn.LeakyReLU(),\n            nn.Conv2d(256, 256, kernel_size=3, stride=1, padding=1), nn.LeakyReLU(),\n            nn.Conv2d(256, 512, kernel_size=3, stride=2, padding=1),  # (4, 4)\n            nn.BatchNorm2d(512), nn.LeakyReLU(), \n            nn.Conv2d(512, 1024, kernel_size=4, stride=1, padding=0), nn.LeakyReLU()\n        )\n        self.x_fc = nn.Sequential(\n            nn.Linear(in_features=1024, out_features=1024),\n            nn.ReLU(),\n#             nn.Dropout()  # ??????????\n        )\n        self.y_fc = nn.Sequential(\n            nn.Linear(in_features=ydim, out_features=ydim*2), nn.LeakyReLU(),\n            nn.Linear(in_features=ydim*2, out_features=ydim*2), nn.LeakyReLU()  #????????? maybe remove\n        )\n        self.final_fc = nn.Sequential(\n            nn.Linear(in_features=1024+ydim*2, out_features=128), nn.LeakyReLU(),\n            nn.Linear(in_features=128, out_features=32), nn.LeakyReLU(),\n            nn.Linear(in_features=32, out_features=8), nn.LeakyReLU(),\n            nn.Linear(in_features=8, out_features=1), nn.Sigmoid(),\n        )\n\n    def forward(self, x, y):\n        x = self.x_conv(x)\n        x = x.view(-1, 1024)\n        x = self.x_fc(x)\n        y = self.y_fc(y)\n        cat = torch.cat([x, y],dim=1)\n        is_real = self.final_fc(cat)\n        return is_real\n\n\nclass Generator(nn.Module):\n    def __init__(self, zdim:int, ydim:int):\n        super(Generator, self).__init__()\n        self.y_decode = nn.Sequential(\n            nn.Linear(in_features=ydim, out_features=zdim*2),\n            nn.ReLU()\n        )\n        self.linearMix = nn.Sequential(\n            nn.Linear(in_features=zdim*3, out_features=256),\n            nn.ReLU()\n        )\n        self.to_conv = nn.Linear(in_features=256, out_features=256*8*8)\n        self.conv_decoder = nn.Sequential(\n            nn.Conv2d(256, 256, kernel_size=3, stride=1, padding=1), # (8, 8)\n            nn.BatchNorm2d(256), nn.ReLU(),\n            nn.ConvTranspose2d(256, 256, kernel_size=3, stride=2, padding=1, output_padding=1), nn.ReLU(),\n            nn.Conv2d(256, 128, kernel_size=3, stride=1, padding=1),  # (16, 16)\n            nn.BatchNorm2d(128), nn.ReLU(),\n            nn.Conv2d(128, 128, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(128), nn.ReLU(),\n            nn.ConvTranspose2d(128, 128, kernel_size=3, stride=2, padding=1, output_padding=1), nn.ReLU(),\n            nn.Conv2d(128, 128, kernel_size=3, stride=1, padding=1),  # (32, 32)\n            nn.BatchNorm2d(128), nn.ReLU(),\n            nn.Conv2d(128, 128, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(128), nn.ReLU(),\n            nn.ConvTranspose2d(128, 64, kernel_size=3, stride=2, padding=1, output_padding=1), nn.ReLU(),\n            nn.Conv2d(64, 64, kernel_size=3, stride=1, padding=1),  # (64, 64)\n            nn.BatchNorm2d(64), nn.ReLU(),\n            nn.Conv2d(64, 64, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(64), nn.ReLU(),\n            nn.ConvTranspose2d(64, 64, kernel_size=3, stride=2, padding=1, output_padding=1), nn.ReLU(),\n            nn.Conv2d(64, 32, kernel_size=3, stride=1, padding=1),  # (128, 128)\n            nn.BatchNorm2d(32), nn.ReLU(),\n            nn.Conv2d(32, 32, kernel_size=5, stride=1, padding=2),\n            nn.BatchNorm2d(32), nn.ReLU(),\n            nn.ConvTranspose2d(32, 32, kernel_size=3, stride=2, padding=1, output_padding=1), nn.ReLU(),\n            nn.Conv2d(32, 16, kernel_size=5, stride=1, padding=2),  # (256, 256)\n            nn.BatchNorm2d(16), nn.ReLU(),\n            nn.Conv2d(16, 16, kernel_size=5, stride=1, padding=2),\n            nn.BatchNorm2d(16), nn.ReLU(),\n            nn.ConvTranspose2d(16, 16, kernel_size=3, stride=2, padding=1, output_padding=1), nn.ReLU(),\n            nn.Conv2d(16, 8, kernel_size=5, stride=1, padding=2), # (512, 512)\n            nn.BatchNorm2d(8), nn.ReLU(),\n            nn.Conv2d(8, 4, kernel_size=5, stride=1, padding=2), nn.ReLU(),\n            nn.Conv2d(4, 3, kernel_size=5, stride=1, padding=2),\n            nn.Sigmoid()\n        )\n    \n    def forward(self, z, y):\n        y = self.y_decode(y)\n#         z = copy.deepcopy(z)    #?????????????\n        z = torch.cat((y,z),dim=1)\n        z = self.linearMix(z)\n        z = self.to_conv(z)\n        z = z.view(-1, 256,8,8)\n        z = self.conv_decoder(z)\n        return z, y","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Define model and optimizers"},{"metadata":{"trusted":true},"cell_type":"code","source":"# hyper parameters\nzdim = 100\nydim = 5\n# Create `G`, `D` instance\nG = Generator(zdim, ydim)\nG = G.to(device, non_blocking=True)  # to GPU or CPU\nD = Discriminator(ydim)\nD = D.to(device, non_blocking=True)  # to GPU or CPU\n# summary(model, input_size=(3, 512, 512))\n# Define hyperparameters\nlr = torch.tensor(0.00001).to(device)\nlr_decay = torch.tensor(0.99).to(device)  # per epoch\nlr_floor = torch.tensor(0.000001).to(device)\n# Define Loss, Optimizer\ncriterion = nn.BCELoss()\noptimizer_D = torch.optim.Adam(D.parameters(), lr=lr)\noptimizer_G = torch.optim.Adam(G.parameters(), lr=lr)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Train loop"},{"metadata":{"trusted":true},"cell_type":"code","source":"# Hyper parameters\nbatch_size = 32\nn_epochs = 10\n# dataloader\ntrain_dataloader = DataLoader(train_dataset,batch_size=batch_size, shuffle=True)\ntot_batch = len(train_dataloader)\n# loss\ntrue_labels = torch.ones(batch_size).to(device)\nfake_labels = torch.zeros(batch_size).to(device)\nG_loss = []\nD_loss = []\n# Training RNN\nfor epoch in range(1, n_epochs + 1):\n    epoch_start = time.time()\n    batch_start = time.time()\n    for num_batch, data_batch in enumerate(train_dataloader):\n        # Generate noise and move it the device\n        noise = torch.randn(batch_size, zdim).to(device)\n        y_noise = F.one_hot(torch.randint(5,size=(batch_size,)), 5).float().to(device)\n        # Forward pass         \n        generated_data,_ = G(noise, y_noise) # batch_size X 784\n        true_data = data_batch['image']\n        label_onehot = data_batch['label_onehot']    \n        ####\n        optimizer_D.zero_grad()\n        # D for true data\n        out_D_orig = D(true_data, label_onehot).view(batch_size)\n        loss_D_orig = criterion(out_D_orig, true_labels)\n        # D for fake data\n        out_D_fake = D(generated_data.detach(), y_noise).view(batch_size)\n        loss_D_fake = criterion(out_D_fake, fake_labels)\n        # calculate loss and optimizing\n        loss_D = (loss_D_orig + loss_D_fake)/2   \n        loss_D.backward()\n        optimizer_D.step()\n        D_loss.append(loss_D.data.item())\n        # G for fake data\n        optimizer_G.zero_grad()\n        optimizer_D.zero_grad()  # ???????????????????\n\n        noise = torch.randn(batch_size, zdim).to(device)\n        y_noise = F.one_hot(torch.randint(5,size=(batch_size,)), 5).float().to(device)\n        generated_data,_ = G(noise, y_noise) # batch_size X 784\n        out_G_fake = D(generated_data, y_noise).view(batch_size)\n        # alculate loss and optimizing\n        loss_G = criterion(out_G_fake, true_labels)\n#         loss_G.backward()\n        loss_G.backward(torch.tensor(-1).to(device))\n        optimizer_G.step()\n        G_loss.append(loss_G.data.item())\n        \n        batch_stop = time.time()\n        print(f\"*****Epoch {epoch}, minibatch: {num_batch+1}/{tot_batch} elapsed {batch_stop-batch_start:0.1f}s, Dloss:{loss_D:0.4f}, Gloss:{loss_G:0.4f}, D(G):{out_G_fake.mean():0.4f}, D(x):{out_D_orig.mean():0.4f} *****\")\n        batch_start = time.time()\n    epoch_stop = time.time()\n    print(f\"Epoch {epoch}/{n_epochs} has finished in {epoch_stop-epoch_start:0.1f}s\")\n    lr *= lr_decay\n    if lr < lr_floor:\n        lr = lr_floor\n    gc.collect()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## save GAN"},{"metadata":{"trusted":true},"cell_type":"code","source":"torch.save(G, 'G_cassava.h5')\ntorch.save(D, 'D_cassava.h5')","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Plot loss"},{"metadata":{"trusted":true},"cell_type":"code","source":"fig, ax = plt.subplots(1,1)\nax.plot(range(1,1+len(D_loss)), D_loss)\nax.plot(range(1,1+len(G_loss)), G_loss)\nax.set(xlabel='minibatch')\nax.set_title('GAN loss')\nplt.legend(('D loss', 'G loss'))\nplt.show()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Generating images"},{"metadata":{"trusted":true},"cell_type":"code","source":"G.eval()\n\nN = 6\nz = torch.randn(N, zdim).to(device)\nclass_label = torch.randint(5,size=(N,))\ny = F.one_hot(class_label, 5).float().to(device)\nout,_ = G(z, y)\n\nG.train()\nprint('')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"fig, ax =plt.subplots(1, N)\nfor i in range(N):\n    iimg = out[i].to(torch.device('cpu')).permute(1,2,0)\n    ax[i].imshow(iimg.detach().numpy())\n    ax[i].axis('off')\n    ax[i].set_title(f'cls:{class_label[i]}')","execution_count":null,"outputs":[]},{"metadata":{"trusted":false},"cell_type":"markdown","source":"## Load model from last training"},{"metadata":{"trusted":true},"cell_type":"code","source":"model = torch.load('./path/to/CVAE_model_v1.h5',)","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}