{"cells":[{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"import json\nimport os\nimport tqdm\nimport time\nimport numpy as np\nimport pandas as pd\nimport PIL.Image as Image\nimport matplotlib.pyplot as plt\n\nfrom sklearn.model_selection import train_test_split\nimport torch\nimport torchvision.transforms as transforms\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn as nn\nimport torchvision.models as models\n\n\nmap_json_path = '/kaggle/input/cassava-leaf-disease-classification/label_num_to_disease_map.json'\ntrain_data_path = '/kaggle/input/cassava-leaf-disease-classification/train_images/'\ntest_data_path = '/kaggle/input/cassava-leaf-disease-classification/test_images/'\ntrain_csv_path = '/kaggle/input/cassava-leaf-disease-classification/train.csv'\n\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\nmodel_pretrained = False\nbatch_size = 8\nimg_resize = (100, 100)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class get_dataset(Dataset):\n    def __init__(self, data_path, csv_label_path, train, train_size, transforms=None):\n        self.train = train\n        self.data_path = data_path\n        self.transforms = transforms\n        \n        self.label_csv = pd.read_csv(csv_label_path)\n        self.label_dict = self.label_csv.set_index('image_id')['label'].to_dict()\n       \n        images_name_list = os.listdir(data_path)[:20]\n        train_image, test_image = train_test_split(images_name_list, train_size=train_size, random_state=0)\n        self.image_list = train_image if self.train else test_image\n    \n    def __getitem__(self, index):\n        image_name = self.image_list[index]\n        label = self.label_dict[image_name]\n        image = Image.open(os.path.join(self.data_path, image_name))\n        \n        if self.transforms: image = self.transforms(image)\n        return image, label\n    \n    def __len__(self):\n        return len(self.image_list)\n    \n    \nclass get_test_dataset(Dataset):\n    def __init__(self, data_path, transforms=None):\n        self.data_path = data_path\n        self.transforms = transforms\n        self.image_list = os.listdir(data_path)\n    \n    def __getitem__(self, index):\n        image_name = self.image_list[index]\n        image = Image.open(os.path.join(self.data_path, image_name))\n        if self.transforms: image = self.transforms(image)\n        return image, image_name\n    \n    def __len__(self):\n        return len(self.image_list)   \n    \n\nmytransforms = transforms.Compose([\n    transforms.Resize(img_resize),\n    transforms.RandomVerticalFlip(),\n    transforms.RandomHorizontalFlip(),\n    transforms.ToTensor(),\n    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)),\n    transforms.RandomErasing(),\n])\n\n\ntrain_dataset = get_dataset(train_data_path, train_csv_path, True, 0.9, mytransforms)\nvalidation_dataset = get_dataset(train_data_path, train_csv_path, False, 0.9, mytransforms)\ntrain_dataloader = DataLoader(train_dataset, batch_size=batch_size)\nvalidation_dataloader = DataLoader(validation_dataset, batch_size=batch_size)\n\ntest_dataset = get_test_dataset(test_data_path, mytransforms)\ntest_dataloader = DataLoader(test_dataset, batch_size=batch_size)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"with open(map_json_path, 'r') as f:\n    map_json_data = json.load(f)\n\nmap_json_data","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"plt.figure(figsize=(20, 6))\nfor i in range(10):\n    plt.subplot(2, 5, i+1)\n    \n    image = train_dataset[i][0]\n    image = image*0.5 + 0.5\n    image = transforms.ToPILImage()(image)\n    \n    plt.title(map_json_data[str(train_dataset[i][1])])\n    plt.imshow(image)\n    plt.xticks([])\n    plt.yticks([]) ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = models.vgg16_bn(pretrained=model_pretrained)\nsequential = list(model.classifier[:3])\nsequential.append(nn.Linear(4096, 5))\nmodel.classifier = nn.Sequential(*sequential)\nmodel.to(device)\n    \noptimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9, weight_decay=1e-5)\ncriterion = nn.CrossEntropyLoss()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def train(cur_epoch, dataloader, compute_grid=True):\n    tq_description = 'epoch %d'%cur_epoch\n    tqbar = tqdm.tqdm(enumerate(dataloader), total=len(dataloader))\n\n    total_loss = 0\n    preds_list = []\n    labels_list = []\n\n    for i, item in tqbar:\n        tqbar.set_description(tq_description)\n        images, labels = item\n        images = images.to(device)\n        labels = labels.to(device)\n            \n        model_out = model(images)\n        loss = criterion(model_out, labels)\n        _, preds = torch.max(model_out, 1)\n        \n        if compute_grid:\n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n\n        total_loss += loss\n        preds_list += preds.tolist()\n        labels_list += labels.tolist()\n        \n    return preds_list, labels_list, total_loss\n\n\ndef generate_submission_csv():\n    tq_description = 'generate csv'\n    tqbar = tqdm.tqdm(enumerate(test_dataloader), total=len(test_dataloader))\n    model.load_state_dict(torch.load('model.pkl'))\n    \n    names_list = []\n    preds_list = []\n    for i, item in tqbar:\n        tqbar.set_description(tq_description)\n        images, names = item\n        \n        images = images.to(device)\n        model_out = model(images)\n        _, preds = torch.max(model_out, 1)\n        \n        names_list += list(names)\n        preds_list += preds.tolist()\n        \n    submission = pd.DataFrame({ 'image_id': names_list, 'label': preds_list })\n    submission.to_csv(\"submission.csv\", index=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def compute_recall(preds_list, labels_list, class_num):\n    recall_arr = np.zeros(class_num)\n    \n    for i in range(class_num):\n        i_labels_list = labels_list == i\n        i_preds_list = preds_list == i \n        total_i_class_num = np.sum(i_labels_list)\n        preds_i_class_num = np.sum(i_preds_list * i_labels_list)\n        recall_arr[i] = preds_i_class_num / total_i_class_num if total_i_class_num != 0 else 0\n    return recall_arr\n    \n    \ndef compute_accuracy(preds_list, labels_list):\n    return np.sum(preds_list == labels_list) / len(labels_list)\n\n    \ndef do_train(epoch):       \n    train_loss_list = []\n    train_accuracy_list = []\n    train_recall_list = []\n    \n    val_loss_list = []\n    val_accuracy_list = []\n    val_recall_list = []\n    \n    best_accuracy = [-1 ,-1] #(epoch, value)\n    train_image_num = len(train_dataset)\n    val_image_num = len(test_dataset)\n    \n    print('info:')\n    print('train image number: ', train_image_num)\n    print('validation image number:', val_image_num)\n    print('train on: %s'%device)\n    print('train epoch: %d'%epoch)\n    \n    for i in range(epoch):\n        preds_list, labels_list, total_loss = train(i, train_dataloader, True)\n        accuracy = compute_accuracy(preds_list, labels_list)\n        recall = compute_recall(preds_list, labels_list, 5)\n        train_loss_list.append(total_loss)\n        train_accuracy_list.append(accuracy)\n        train_recall_list.append(recall)\n        print('train loss: %f'%total_loss)\n        print('train accuracy: %f'%accuracy)\n        print('train recall:', recall)\n        \n        preds_list, labels_list, total_loss = train(i, validation_dataloader, False)\n        accuracy = compute_accuracy(preds_list, labels_list)\n        recall = compute_recall(preds_list, labels_list, 5)\n        val_loss_list.append(total_loss)\n        val_accuracy_list.append(accuracy)\n        val_recall_list.append(recall)\n        print('test loss: %f'%total_loss)\n        print('test accuracy: %f'%accuracy)\n        print('test recall:', recall)\n        \n        if best_accuracy[1] < accuracy:\n            best_accuracy[0] = epoch\n            best_accuracy[1] = accuracy\n            torch.save(model.state_dict(), 'model.pkl')\n            \n    # plot data\n    plt.figure()\n    plt.plot(train_loss_list)\n    plt.plot(val_loss_list)\n    plt.title('loss')\n    plt.legend(labels=['train','validation'])\n    \n    plt.figure()\n    plt.plot(train_accuracy_list)\n    plt.plot(val_accuracy_list)\n    plt.title('accuracy')\n    plt.legend(labels=['train','validation'])\n    \n    plt.figure()\n    plt.plot(train_recall_list)\n    plt.plot(val_recall_list)\n    plt.title('recall')\n    plt.legend(labels=['train','validation'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"do_train(20)\ngenerate_submission_csv()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## 改进方法\n1. 一张图像分成多个小图像，取模型预测平均值\n2. 查看患病分布，给予每个图像分类不同的权重"}],"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}