{"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":"\n# [CoE202] **[Homework3]** CNN Classification for CIFAR10 (Pytorch)\n## **Training Script**","metadata":{"id":"K51avYKII76d"}},{"cell_type":"markdown","source":"In this section, you are going to **train** CNN classification for CIFAR10 in the Pytorch framework.\n","metadata":{"id":"YPaf-89b8W75"}},{"cell_type":"markdown","source":"# 0. Import Library\n","metadata":{"id":"AIQrCeIU1mZ7"}},{"cell_type":"code","source":"import os\nimport sys\nimport time\nimport datetime\nfrom PIL import Image, ImageEnhance\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport math\nimport torch\nimport torch.nn as nn\nimport torchvision\nfrom torchvision.utils import save_image\nimport torchvision.transforms as transforms\nfrom tqdm.auto import tqdm\nimport torch.nn.functional as F\nimport urllib\nimport glob\nimport skimage.io as skio\nfrom torch.utils.data import Dataset\nimport random\nimport pandas as pd","metadata":{"id":"IRmvDGFrcy7g","outputId":"cefeedbd-f2a2-4e95-9c2d-6cddda8101a4","execution":{"iopub.status.busy":"2021-08-09T08:56:29.588305Z","iopub.execute_input":"2021-08-09T08:56:29.588724Z","iopub.status.idle":"2021-08-09T08:56:31.606308Z","shell.execute_reply.started":"2021-08-09T08:56:29.588614Z","shell.execute_reply":"2021-08-09T08:56:31.605083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 1. Set hyper-parameters","metadata":{"id":"MQui4iIp11u_"}},{"cell_type":"markdown","source":"\nIf you change this part, write why changing the parameter will improve the performance at the end of the .ipynb file.","metadata":{"id":"8rTve68ithP_"}},{"cell_type":"code","source":"# Set several hyperparameters and settings\nbatch_size = 50 # default = 128\nn_epoch = 200 # default = 200\nlearning_rate = 0.001 # default = 0.1\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)), \n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # to tensor before normalize\n])\nvalidation_transform = transforms.Compose([\n#     transforms.Grayscale(),\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(), # to tensor @ Normalize must be at the end\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) # to tensor before normalize\n])\n\nshuffle = False\ntarget_transform = None","metadata":{"id":"9-DePpd0c1f7","execution":{"iopub.status.busy":"2021-08-09T08:56:31.60807Z","iopub.execute_input":"2021-08-09T08:56:31.608498Z","iopub.status.idle":"2021-08-09T08:56:31.621227Z","shell.execute_reply.started":"2021-08-09T08:56:31.608448Z","shell.execute_reply":"2021-08-09T08:56:31.619858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed = 123\ntorch.manual_seed(seed)\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"id":"40ROCVg3tSXi","execution":{"iopub.status.busy":"2021-08-09T08:56:31.624179Z","iopub.execute_input":"2021-08-09T08:56:31.625141Z","iopub.status.idle":"2021-08-09T08:56:31.693633Z","shell.execute_reply.started":"2021-08-09T08:56:31.625092Z","shell.execute_reply":"2021-08-09T08:56:31.692623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_image(file):\n    return Image.open(file)\ndef image_path(root, basename, extension):\n    return os.path.join(root, f'{basename}{extension}')\n\nclass CustomImageDataset(Dataset):\n    def __init__(self, img_dir, label_dir, transform=None):\n        video_files = os.listdir(img_dir)\n        self.labels = list()\n        self.imgs = list()\n        for file_name in video_files: \n            cur_csv = pd.read_csv(os.path.join(label_dir, file_name+'.csv'))\n            cur_len = len(cur_csv.index)\n            for itr in range(cur_len):\n                img_file_dir = os.path.join(os.path.join(img_dir, file_name), cur_csv.iloc[itr, 0] + '.PNG')\n                self.imgs.append(img_file_dir) # image directory\n                self.labels.append([cur_csv.iloc[itr, 0], cur_csv.iloc[itr, 1]]) # label\n        self.img_dir = img_dir\n        self.label_dir = label_dir\n        self.transform = transform\n            \n    def __getitem__(self, index):\n        with open(self.imgs[index], 'rb') as f:\n            image = load_image(f).convert('RGB')\n        label = self.labels[index][1]\n        filename = self.labels[index][0]\n        \n#         image, label = self.spm_transform(image, label)\n        \n        if self.transform is not None:\n            image = self.transform(image)\n\n        return image, label\n\n    def __len__(self):\n        return len(self.labels)\n\nimg_dir = r'../input/hsgshackathon2021/train_data/Train'\nlabel_dir = r'../input/hsgshackathon2021/train_data/Train_labels'\ndataset = CustomImageDataset(img_dir=img_dir, label_dir=label_dir,transform=transform)\nn_train = math.floor(0.9*len(dataset)) # (default) 90% of the data for training\nn_val = len(dataset) - math.floor(0.9*len(dataset)) # (default) 10% of the data for validation\n\ntrain_dataset, validation_dataset = torch.utils.data.random_split(dataset, [n_train, n_val])\n\ntrain_loader = torch.utils.data.DataLoader(train_dataset, batch_size = batch_size, shuffle=True, drop_last = True)\nvalidation_loader = torch.utils.data.DataLoader(validation_dataset, batch_size = batch_size, shuffle=False, drop_last = True)\n\nclasses = ('Hit', 'no_label')","metadata":{"id":"g-QEFJ_jOa6V","outputId":"a2a70784-24be-47dc-e49d-c0d5f6def589","execution":{"iopub.status.busy":"2021-08-09T08:56:31.695922Z","iopub.execute_input":"2021-08-09T08:56:31.696649Z","iopub.status.idle":"2021-08-09T08:56:35.672694Z","shell.execute_reply.started":"2021-08-09T08:56:31.696571Z","shell.execute_reply":"2021-08-09T08:56:35.6717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"[one_batch_data, one_batch_label] = next(iter(train_loader))\nprint(one_batch_data.size(), one_batch_label.size())","metadata":{"id":"95yCY-txbpJR","outputId":"e2317822-8374-48cf-9fac-5d641fe57752","execution":{"iopub.status.busy":"2021-08-09T08:56:35.674187Z","iopub.execute_input":"2021-08-09T08:56:35.674584Z","iopub.status.idle":"2021-08-09T08:56:38.367538Z","shell.execute_reply.started":"2021-08-09T08:56:35.674542Z","shell.execute_reply":"2021-08-09T08:56:38.364841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def vismod_img(img):\n    npimg = img.numpy()\n    npimg = np.transpose(npimg, (1, 2, 0))\n    return npimg\n\ndef show_data(data, label, classes):\n    num_row = 2\n    num_col = 4\n    fig, axes = plt.subplots(num_row, num_col)\n    for i in range(8):\n        ax = axes[i//num_col, i%num_col]\n        ax.imshow(vismod_img(data[i].cpu()))\n        ax.set_title(f'label : {classes[label[i].cpu().item()]}')\n    plt.tight_layout()\n    plt.show()\n\ndef show_inference_result(data, label, output, classes):\n    num_row = 2\n    num_col = 4\n    fig, axes = plt.subplots(num_row, num_col)\n    for i in range(8):\n        ax = axes[i//num_col, i%num_col]\n        ax.imshow(vismod_img(data[i].cpu()))\n#         print('checkclass', label[i].cpu().item())\n        ax.set_title(f'GT : {classes[label[i].cpu().item()]} \\n output : {classes[output[i].cpu().item()]}')\n    plt.tight_layout()\n    plt.show()","metadata":{"id":"540u3DeBbqrs","execution":{"iopub.status.busy":"2021-08-09T08:56:38.369181Z","iopub.execute_input":"2021-08-09T08:56:38.369773Z","iopub.status.idle":"2021-08-09T08:56:38.380447Z","shell.execute_reply.started":"2021-08-09T08:56:38.369704Z","shell.execute_reply":"2021-08-09T08:56:38.378557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# show ground truth classification (one batch)\nprint('(Training Data) ground truth example')\nshow_data(one_batch_data, one_batch_label, classes) # 10","metadata":{"id":"QP0hdlY7bsiO","outputId":"3518f5da-6347-4008-f3b1-14fa2ad334bb","execution":{"iopub.status.busy":"2021-08-09T08:56:38.38298Z","iopub.execute_input":"2021-08-09T08:56:38.383583Z","iopub.status.idle":"2021-08-09T08:56:39.387866Z","shell.execute_reply.started":"2021-08-09T08:56:38.383537Z","shell.execute_reply":"2021-08-09T08:56:39.386578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class cLinear(nn.Module):\n    # in_channels = 512\n    def __init__(self, in_channels, out_channels, dropout=0):\n        super().__init__()\n        layers = [\n            nn.Linear(in_channels, 64), \n            nn.Dropout(p=0.3), \n            nn.Linear(64, out_channels), \n        ]\n        self.fward = nn.Sequential(*layers)\n\n    def forward(self, x):\n        return self.fward(x)\n\nmodel = torchvision.models.resnet152(pretrained=True)\nfor param in model.parameters(): \n    param.requires_grad = False\n\nfc_in_ftrs = model.fc.in_features\nprint(fc_in_ftrs)\nmodel.fc = cLinear(fc_in_ftrs, 2)\nprint(model)","metadata":{"id":"iNk-O7Q-Xz7V","execution":{"iopub.status.busy":"2021-08-09T08:56:39.391262Z","iopub.execute_input":"2021-08-09T08:56:39.391789Z","iopub.status.idle":"2021-08-09T08:56:49.939645Z","shell.execute_reply.started":"2021-08-09T08:56:39.391728Z","shell.execute_reply":"2021-08-09T08:56:49.938334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def weights_init(m):\n    if isinstance(m, nn.Linear):\n        bound = 1 / math.sqrt(m.weight.size(1))\n        nn.init.uniform_(m.weight, -bound, bound)\n        nn.init.uniform_(m.bias, -bound, bound)\n\ndef thresholding(prediction):\n    _, pred_label = torch.max(prediction, 1)\n    return pred_label\n\nmy_classifier = model  # assign classifier\nmy_classifier = my_classifier.to(device)\n# my_classifier.apply(weights_init)  # applying weights_init function to linear layers\n\n# show the performance of untrained classifier\nprint('current classification')\nprediction = my_classifier(one_batch_data.to(device))  # passing forward function of classifier, return prediction\nshow_inference_result(one_batch_data, one_batch_label, thresholding(prediction), classes)","metadata":{"id":"HkhZJrd-dmMW","outputId":"fd1faa48-b063-4d7e-a95a-6a96b4d5f586","execution":{"iopub.status.busy":"2021-08-09T08:56:49.941802Z","iopub.execute_input":"2021-08-09T08:56:49.942267Z","iopub.status.idle":"2021-08-09T08:56:57.287897Z","shell.execute_reply.started":"2021-08-09T08:56:49.942223Z","shell.execute_reply":"2021-08-09T08:56:57.286519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_n_params(model):\n    pp=0\n    for p in list(model.parameters()):\n        nn=1\n        for s in list(p.size()):\n            nn = nn*s\n        pp += nn\n    return pp\nprint(get_n_params(my_classifier))","metadata":{"execution":{"iopub.status.busy":"2021-08-09T08:56:57.28983Z","iopub.execute_input":"2021-08-09T08:56:57.2903Z","iopub.status.idle":"2021-08-09T08:56:57.307176Z","shell.execute_reply.started":"2021-08-09T08:56:57.290253Z","shell.execute_reply":"2021-08-09T08:56:57.305716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = torch.optim.Adam(my_classifier.parameters(), lr=learning_rate, weight_decay=0.001)","metadata":{"id":"ixwtt0nx_Mni","execution":{"iopub.status.busy":"2021-08-09T08:56:57.308778Z","iopub.execute_input":"2021-08-09T08:56:57.309488Z","iopub.status.idle":"2021-08-09T08:56:57.324865Z","shell.execute_reply.started":"2021-08-09T08:56:57.309443Z","shell.execute_reply":"2021-08-09T08:56:57.323469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train cnn classifier\ntrain_loss_iter = np.zeros(n_epoch, dtype=float)  # Temporary numpy array to save loss for each epoch\nvalidation_loss_iter = np.zeros(n_epoch, dtype=float)\ntrain_accuracy_iter = np.zeros(n_epoch, dtype=float)  # Temporary numpy array to save accuracy for each epoch\nvalidation_accuracy_iter = np.zeros(n_epoch, dtype=float)\n\nloss_function = torch.nn.CrossEntropyLoss()\n# Assign cross-entropy loss\nbest_model_state = None\nbest_acc = 0.0\n\nfor epoch in range(n_epoch):  # We will iteratively find optimum\n    # Training\n    my_classifier.train()\n    total_loss, total_cnt, correct_cnt = 0.0, 0.0, 0.0\n    \n    for batch_idx, (image, label) in enumerate(train_loader): # current len(train_loader) = 390\n        image, label = image.to(device), label.to(device)\n        pred = my_classifier(image)  # Find output of classifier\n        \n        optimizer.zero_grad()  # Pytorch does not overwrite gradients, it 'accumulates' them so that we need to set gradient as 0 to update parameter correctly\n        loss = loss_function(pred, label)  # Calculate cross-entropy loss between prediction and label\n        loss.backward()  # Pytorch automatically back-propagate and calculate gradients\n        optimizer.step()  # From calculated gradients, change parameters of classifier with SGD algorithm\n\n        total_loss += loss.item()\n        total_cnt += len(label)\n        correct_cnt += (pred.argmax(1) == label).type(torch.float).sum().item()\n\n    \n    accuracy = correct_cnt * 1.0 / total_cnt\n    train_loss_iter[epoch] = total_loss / total_cnt\n    train_accuracy_iter[epoch] = accuracy\n        \n    \n    print(f\"[{str(epoch).zfill(len(str(n_epoch)))}/{n_epoch}] Train Loss  : {train_loss_iter[epoch]:.6f}, Acc = {100*train_accuracy_iter[epoch]:.2f}%\")\n    \n    # validation\n    my_classifier.eval()\n    total_loss, total_cnt, correct_cnt = 0.0, 0.0, 0.0\n    for batch_idx, (image, label) in enumerate(validation_loader):\n        with torch.no_grad():\n            image, label = image.to(device), label.to(device)\n            \n            pred = my_classifier(image)\n            loss = loss_function(pred, label)\n\n            total_loss += loss.item()\n            total_cnt += len(label)\n            correct_cnt += (pred.argmax(1) == label).type(torch.float).sum().item()\n    \n    accuracy = correct_cnt * 1.0 / total_cnt\n    validation_loss_iter[epoch]  = total_loss / total_cnt\n    validation_accuracy_iter[epoch] = accuracy\n    \n    if best_acc < accuracy: \n        best_acc = accuracy\n        best_model_state = my_classifier.state_dict()\n    \n    print(f\"[{str(epoch).zfill(len(str(n_epoch)))}/{n_epoch}] Validation Loss : {validation_loss_iter[epoch]:.6f}, Acc = {100*validation_accuracy_iter[epoch]:.2f}%\")","metadata":{"id":"h_z_d7EGeyyV","execution":{"iopub.status.busy":"2021-08-09T08:56:57.328481Z","iopub.execute_input":"2021-08-09T08:56:57.328933Z","iopub.status.idle":"2021-08-09T13:42:09.547324Z","shell.execute_reply.started":"2021-08-09T08:56:57.328889Z","shell.execute_reply":"2021-08-09T13:42:09.543827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#6-1. Plot loss for train, validation set","metadata":{"id":"mhyB2Uu93jU-"}},{"cell_type":"code","source":"# show loss during training\nl1, = plt.plot(range(1,n_epoch+1), train_loss_iter)\nl2, = plt.plot(range(1,n_epoch+1), validation_loss_iter)\nplt.legend(handles=(l1, l2), labels=('Train loss', 'Valid loss'))","metadata":{"id":"thWYqsfAiNhk","outputId":"543ddcb8-83ee-4c87-f63f-6ed1112b38dc","execution":{"iopub.status.busy":"2021-08-09T13:42:18.392527Z","iopub.execute_input":"2021-08-09T13:42:18.392984Z","iopub.status.idle":"2021-08-09T13:42:18.631798Z","shell.execute_reply.started":"2021-08-09T13:42:18.392951Z","shell.execute_reply":"2021-08-09T13:42:18.630559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#6-2. Plot accuracy for train, validation set","metadata":{"id":"PQVg5SQm3ul2"}},{"cell_type":"code","source":"# show accuracy during training\nl1, = plt.plot(range(1,n_epoch+1), 100*train_accuracy_iter)\nl2, = plt.plot(range(1,n_epoch+1), 100*validation_accuracy_iter)\nplt.legend(handles=(l1, l2), labels=('Train Acc', 'Valid Acc'))","metadata":{"id":"FjfbJNBBiYNl","outputId":"389cff9b-207b-4909-e5e4-7cf07b19ef8f","execution":{"iopub.status.busy":"2021-08-09T13:42:22.676192Z","iopub.execute_input":"2021-08-09T13:42:22.676552Z","iopub.status.idle":"2021-08-09T13:42:23.007975Z","shell.execute_reply.started":"2021-08-09T13:42:22.67652Z","shell.execute_reply":"2021-08-09T13:42:23.006705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#7. Visualize classification result of validation set after training","metadata":{"id":"6kwmE0cs34QO"}},{"cell_type":"code","source":"# For validation data\nprint('After training')\n[one_batch_data, one_batch_label] = next(iter(validation_loader))\nprediction = my_classifier(one_batch_data.to(device))\nshow_inference_result(one_batch_data, one_batch_label, thresholding(prediction), classes)","metadata":{"id":"SFM918cWiedY","outputId":"81e799de-bbb0-44f5-d9c6-518f6a06ebd3","execution":{"iopub.status.busy":"2021-08-09T13:42:53.629863Z","iopub.execute_input":"2021-08-09T13:42:53.630233Z","iopub.status.idle":"2021-08-09T13:42:56.021755Z","shell.execute_reply.started":"2021-08-09T13:42:53.630202Z","shell.execute_reply":"2021-08-09T13:42:56.020684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nmodel_save_name = 'checkpoint.pth'\npath = F\"{model_save_name}\" \n\n#torch.save(checkpoint, path)\ntorch.save(best_model_state, path)","metadata":{"id":"lPu1EKG4N6Sn","execution":{"iopub.status.busy":"2021-08-09T13:57:11.86904Z","iopub.execute_input":"2021-08-09T13:57:11.869387Z","iopub.status.idle":"2021-08-09T13:57:12.440309Z","shell.execute_reply.started":"2021-08-09T13:57:11.869356Z","shell.execute_reply":"2021-08-09T13:57:12.439132Z"},"trusted":true},"execution_count":null,"outputs":[]}]}