{"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":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n        break\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-07-20T09:17:03.461535Z","iopub.execute_input":"2023-07-20T09:17:03.462580Z","iopub.status.idle":"2023-07-20T09:17:20.182620Z","shell.execute_reply.started":"2023-07-20T09:17:03.462534Z","shell.execute_reply":"2023-07-20T09:17:20.181277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport os, sys, json, cv2, random, torchvision\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom torchvision import transforms\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import f1_score, classification_report, confusion_matrix\nimport seaborn as sns\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.utils.tensorboard import SummaryWriter\nfrom PIL import Image\nimport torch.nn as nn\nfrom torch.optim.lr_scheduler import StepLR\nfrom sklearn.metrics import auc, roc_curve\nfrom numpy import interp\nfrom itertools import cycle","metadata":{"execution":{"iopub.status.busy":"2023-07-20T10:00:19.349162Z","iopub.execute_input":"2023-07-20T10:00:19.349557Z","iopub.status.idle":"2023-07-20T10:00:19.357058Z","shell.execute_reply.started":"2023-07-20T10:00:19.349526Z","shell.execute_reply":"2023-07-20T10:00:19.356051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls /kaggle/input/cassava-leaf-disease-classification/","metadata":{"execution":{"iopub.status.busy":"2023-07-20T09:19:15.001327Z","iopub.execute_input":"2023-07-20T09:19:15.001744Z","iopub.status.idle":"2023-07-20T09:19:16.008388Z","shell.execute_reply.started":"2023-07-20T09:19:15.001705Z","shell.execute_reply":"2023-07-20T09:19:16.006605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dfMerge = pd.read_csv('/kaggle/input/cassava-leaf-disease-classification/train.csv')\ndfMerge['image_id'] = dfMerge['image_id'].apply(lambda x: '/kaggle/input/cassava-leaf-disease-classification/train_images/' + x)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T09:23:09.466791Z","iopub.execute_input":"2023-07-20T09:23:09.467534Z","iopub.status.idle":"2023-07-20T09:23:09.522682Z","shell.execute_reply.started":"2023-07-20T09:23:09.467497Z","shell.execute_reply":"2023-07-20T09:23:09.521623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dfMerge","metadata":{"execution":{"iopub.status.busy":"2023-07-20T09:23:12.125961Z","iopub.execute_input":"2023-07-20T09:23:12.126335Z","iopub.status.idle":"2023-07-20T09:23:12.139238Z","shell.execute_reply.started":"2023-07-20T09:23:12.126303Z","shell.execute_reply":"2023-07-20T09:23:12.138107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"json_file = open('/kaggle/input/cassava-leaf-disease-classification/label_num_to_disease_map.json', 'r')\nclass_indict = json.load(json_file)\nclass_indict","metadata":{"execution":{"iopub.status.busy":"2023-07-20T09:30:35.258296Z","iopub.execute_input":"2023-07-20T09:30:35.258689Z","iopub.status.idle":"2023-07-20T09:30:35.271732Z","shell.execute_reply.started":"2023-07-20T09:30:35.258658Z","shell.execute_reply":"2023-07-20T09:30:35.270716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def split_df(df, plot_image=True):\n    df.columns = ['filepaths', 'labels']\n    \n    every_class_num = []\n    for idx in df['labels'].unique():\n        every_class_num.append(df[df['labels'] == idx].shape[0])\n    \n    classes = class_indict.values()\n    train_df, test_df = train_test_split(df, train_size=.8, shuffle=True, random_state=123, stratify=df['labels'])\n\n    train_image_path = train_df['filepaths'].tolist()\n    val_image_path = test_df['filepaths'].tolist()\n\n    train_image_label = train_df['labels'].tolist()\n    val_image_label = test_df['labels'].tolist()\n\n    sample_df = train_df.sample(n=50, replace=False)\n    ht, wt, count = 0, 0, 0\n    for i in range(len(sample_df)):\n        fpath = sample_df['filepaths'].iloc[i]\n        try:\n            img = cv2.imread(fpath)\n            h = img.shape[0]\n            w = img.shape[1]\n            ht += h\n            wt += w\n            count += 1\n        except:\n            pass\n\n    have = int(ht / count)\n    wave = int(wt / count)\n    aspect_ratio = have / wave\n    print('{} images were found in the dataset.\\n{} for training, {} for validation'.format(\n        sum(every_class_num), len(train_image_path), len(val_image_path)\n    ))\n    print('average image height= ', have, '  average image width= ', wave, ' aspect ratio h/w= ', aspect_ratio)\n\n    if plot_image:\n        plt.bar(range(len(classes)), every_class_num, align='center')\n        plt.xticks(range(len(classes)), classes, rotation=45)\n\n        for i, v in enumerate(every_class_num):\n            plt.text(x=i, y=v + 5, s=str(v), ha='center')\n\n        plt.xlabel('image class')\n        plt.ylabel('number of images')\n\n        plt.title('class distribution')\n        plt.show()\n\n    return train_image_path, train_image_label, val_image_path, val_image_label\n\ntrain_image_path, train_image_label, val_image_path, val_image_label = split_df(dfMerge)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T09:31:58.905108Z","iopub.execute_input":"2023-07-20T09:31:58.905503Z","iopub.status.idle":"2023-07-20T09:32:00.019199Z","shell.execute_reply.started":"2023-07-20T09:31:58.905471Z","shell.execute_reply":"2023-07-20T09:32:00.018223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(train_image_path), len(train_image_label), len(val_image_path), len(val_image_label))","metadata":{"execution":{"iopub.status.busy":"2023-07-20T09:32:41.856858Z","iopub.execute_input":"2023-07-20T09:32:41.857224Z","iopub.status.idle":"2023-07-20T09:32:41.863291Z","shell.execute_reply.started":"2023-07-20T09:32:41.857192Z","shell.execute_reply":"2023-07-20T09:32:41.862019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MyDataset(Dataset):\n    def __init__(self, image_path, image_labels, transforms=None):\n        self.image_path = image_path\n        self.image_labels = image_labels\n        self.transforms = transforms\n\n    def __len__(self):\n        return len(self.image_path)\n\n    def __getitem__(self, item):\n        image = Image.open(self.image_path[item]).convert('RGB')\n        label = self.image_labels[item]\n        if self.transforms:\n            image = self.transforms(image)\n\n        return image, label\n\n    @staticmethod\n    def collate_fn(batch):\n        images, labels = tuple(zip(*batch))\n        images = torch.stack(images, dim=0)\n        labels = torch.as_tensor(labels)\n        return images, labels","metadata":{"execution":{"iopub.status.busy":"2023-07-20T09:33:04.178790Z","iopub.execute_input":"2023-07-20T09:33:04.179395Z","iopub.status.idle":"2023-07-20T09:33:04.194739Z","shell.execute_reply.started":"2023-07-20T09:33:04.179336Z","shell.execute_reply":"2023-07-20T09:33:04.193624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = torchvision.models.swin_v2_b(weights=torchvision.models.Swin_V2_B_Weights)\nmodel.head = nn.Linear(in_features=1024, out_features=5, bias=True)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T09:36:31.815167Z","iopub.execute_input":"2023-07-20T09:36:31.815554Z","iopub.status.idle":"2023-07-20T09:36:33.869154Z","shell.execute_reply.started":"2023-07-20T09:36:31.815522Z","shell.execute_reply":"2023-07-20T09:36:33.868194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 8\n\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\nepochs = 5\nlr = 0.0003\nbest_val_accuracy = 0\nweight_decay = 0.00001\n\ndata_transform = {\n    'train': transforms.Compose([transforms.RandomResizedCrop(224), transforms.ToTensor(),\n                                 transforms.RandomHorizontalFlip(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])]),\n    'valid': transforms.Compose([transforms.Resize((224, 224)), transforms.CenterCrop(224),\n                                 transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])])\n}","metadata":{"execution":{"iopub.status.busy":"2023-07-20T09:45:53.701073Z","iopub.execute_input":"2023-07-20T09:45:53.701457Z","iopub.status.idle":"2023-07-20T09:45:53.711214Z","shell.execute_reply.started":"2023-07-20T09:45:53.701426Z","shell.execute_reply":"2023-07-20T09:45:53.709932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_step(net, optimizer, data_loader, device, epoch, scalar=None):\n    net.train()\n    loss_function = nn.CrossEntropyLoss()\n    train_acc, train_loss, sampleNum = 0, 0, 0\n    optimizer.zero_grad()\n\n    train_bar = tqdm(data_loader, file=sys.stdout)\n    for step, data in enumerate(train_bar):\n        images, labels = data\n        sampleNum += images.shape[0]  # batch\n        images, labels = images.to(device), labels.to(device)\n        optimizer.zero_grad()\n\n        if scalar is not None:\n            with torch.cuda.amp.autocast():\n                outputs = net(images)\n                loss = loss_function(outputs, labels)\n        else:\n            outputs = net(images)\n            loss = loss_function(outputs, labels)\n\n        train_acc += (torch.argmax(outputs, dim=1) == labels).sum().item()\n        train_loss += loss.item()\n        # loss.backward()\n        # optimizer.step()\n\n        if scalar is not None:\n            scalar.scale(loss).backward()\n            scalar.step(optimizer)\n            scalar.update()\n        else:\n            loss.backward()\n            optimizer.step()\n        train_bar.desc = \"[train epoch {}] loss: {:.3f}, acc: {:.3f}\".format(epoch, train_loss / (step + 1),\n                                                                             train_acc / sampleNum)\n\n    return train_loss / (step + 1), train_acc / sampleNum\n\n\n@torch.no_grad()\ndef val_step(net, data_loader, device, epoch):\n    loss_function = nn.CrossEntropyLoss()\n    net.eval()\n    val_acc = 0\n    val_loss = 0\n    sample_num = 0\n    val_bar = tqdm(data_loader, file=sys.stdout)\n    for step, data in enumerate(val_bar):\n        images, labels = data\n        sample_num += images.shape[0]\n        images, labels = images.to(device), labels.to(device)\n        outputs = net(images)\n        loss = loss_function(outputs, labels)\n        val_loss += loss.item()\n        val_acc += (torch.argmax(outputs, dim=1) == labels).sum().item()\n        val_bar.desc = \"[valid epoch {}] loss: {:.3f}, acc: {:.3f}\".format(epoch, val_loss / (step + 1),\n                                                                           val_acc / sample_num)\n\n    return val_loss / (step + 1), val_acc / sample_num","metadata":{"execution":{"iopub.status.busy":"2023-07-20T09:45:55.177613Z","iopub.execute_input":"2023-07-20T09:45:55.177986Z","iopub.status.idle":"2023-07-20T09:45:55.192885Z","shell.execute_reply.started":"2023-07-20T09:45:55.177954Z","shell.execute_reply":"2023-07-20T09:45:55.191884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def Plot_ROC(net, val_loader, save_name, device):\n    try:\n        json_file = open('/kaggle/input/cassava-leaf-disease-classification/label_num_to_disease_map.json', 'r')\n        class_indict = json.load(json_file)\n    except Exception as e:\n        print(e)\n        exit(-1)\n\n    score_list = []  # 存储预测得分\n    label_list = []  # 存储真实标签\n\n    net.load_state_dict(torch.load(save_name)['model'])\n\n    for i, data in enumerate(val_loader):\n        images, labels = data\n        images, labels = images.to(device), labels.to(device)\n        outputs = torch.softmax(net(images), dim=1)\n        score_tmp = outputs\n        score_list.extend(score_tmp.detach().cpu().numpy())\n        label_list.extend(labels.cpu().numpy())\n\n    score_array = np.array(score_list)\n    # 将label转换成onehot形式\n    label_tensor = torch.tensor(label_list)\n    label_tensor = label_tensor.reshape((label_tensor.shape[0], 1))\n    label_onehot = torch.zeros(label_tensor.shape[0], len(class_indict.keys()))\n    label_onehot.scatter_(dim=1, index=label_tensor, value=1)\n    label_onehot = np.array(label_onehot)\n\n    print(\"score_array:\", score_array.shape)  # (batchsize, classnum)\n    print(\"label_onehot:\", label_onehot.shape)  # torch.Size([batchsize, classnum])\n\n    # 调用sklearn库，计算每个类别对应的fpr和tpr\n    fpr_dict = dict()\n    tpr_dict = dict()\n    roc_auc_dict = dict()\n    for i in range(len(class_indict.keys())):\n        fpr_dict[i], tpr_dict[i], _ = roc_curve(label_onehot[:, i], score_array[:, i])\n        roc_auc_dict[i] = auc(fpr_dict[i], tpr_dict[i])\n    # micro\n    fpr_dict[\"micro\"], tpr_dict[\"micro\"], _ = roc_curve(label_onehot.ravel(), score_array.ravel())\n    roc_auc_dict[\"micro\"] = auc(fpr_dict[\"micro\"], tpr_dict[\"micro\"])\n\n    # macro\n    # First aggregate all false positive rates\n    all_fpr = np.unique(np.concatenate([fpr_dict[i] for i in range(len(class_indict.keys()))]))\n    # Then interpolate all ROC curves at this points\n    mean_tpr = np.zeros_like(all_fpr)\n\n    for i in range(len(set(label_list))):\n        mean_tpr += interp(all_fpr, fpr_dict[i], tpr_dict[i])\n\n    # Finally average it and compute AUC\n    mean_tpr /= len(class_indict.keys())\n    fpr_dict[\"macro\"] = all_fpr\n    tpr_dict[\"macro\"] = mean_tpr\n    roc_auc_dict[\"macro\"] = auc(fpr_dict[\"macro\"], tpr_dict[\"macro\"])\n\n    # 绘制所有类别平均的roc曲线\n    plt.figure(figsize=(12, 12))\n    lw = 2\n\n    plt.plot(fpr_dict[\"micro\"], tpr_dict[\"micro\"],\n             label='micro-average ROC curve (area = {0:0.2f})'\n                   ''.format(roc_auc_dict[\"micro\"]),\n             color='deeppink', linestyle=':', linewidth=4)\n\n    plt.plot(fpr_dict[\"macro\"], tpr_dict[\"macro\"],\n             label='macro-average ROC curve (area = {0:0.2f})'\n                   ''.format(roc_auc_dict[\"macro\"]),\n             color='navy', linestyle=':', linewidth=4)\n\n    colors = cycle(['aqua', 'darkorange', 'cornflowerblue'])\n    for i, color in zip(range(len(class_indict.keys())), colors):\n        plt.plot(fpr_dict[i], tpr_dict[i], color=color, lw=lw,\n                 label='ROC curve of class {0} (area = {1:0.2f})'\n                       ''.format(class_indict[str(i)], roc_auc_dict[i]))\n\n    plt.plot([0, 1], [0, 1], 'k--', lw=lw, label='Chance', color='red')\n    plt.xlim([0.0, 1.0])\n    plt.ylim([0.0, 1.05])\n    plt.xlabel('False Positive Rate')\n    plt.ylabel('True Positive Rate')\n    plt.title('Receiver operating characteristic to multi-class')\n    plt.legend(loc=\"lower right\")\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-20T09:45:55.521503Z","iopub.execute_input":"2023-07-20T09:45:55.522668Z","iopub.status.idle":"2023-07-20T09:45:55.543097Z","shell.execute_reply.started":"2023-07-20T09:45:55.522626Z","shell.execute_reply":"2023-07-20T09:45:55.542002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = MyDataset(train_image_path, train_image_label, data_transform[\"train\"])\n\nval_dataset = MyDataset(val_image_path, val_image_label, data_transform[\"valid\"])\ntrain_loader = torch.utils.data.DataLoader(train_dataset,\n                                               batch_size=batch_size,\n                                               shuffle=True,\n                                               pin_memory=True,\n                                               num_workers=0,\n                                               collate_fn=train_dataset.collate_fn)\n\nval_loader = torch.utils.data.DataLoader(val_dataset,\n                                             batch_size=batch_size,\n                                             shuffle=False,\n                                             pin_memory=True,\n                                             num_workers=0,\n                                             collate_fn=val_dataset.collate_fn)\n\nmodel = model.to(device)\noptimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)\nlr_scheduler = StepLR(optimizer, step_size=1, gamma=0.33)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T09:45:55.824709Z","iopub.execute_input":"2023-07-20T09:45:55.825067Z","iopub.status.idle":"2023-07-20T09:45:55.848417Z","shell.execute_reply.started":"2023-07-20T09:45:55.825038Z","shell.execute_reply":"2023-07-20T09:45:55.847573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(epochs):\n        # train\n    train_loss, train_acc = train_step(net=model,\n                                           optimizer=optimizer,\n                                           data_loader=train_loader,\n                                           device=device,\n                                           epoch=epoch,\n                                           scalar=None)\n\n        # validate\n    val_loss, val_acc = val_step(net=model,\n                                     data_loader=val_loader,\n                                     device=device,\n                                     epoch=epoch)\n\n    lr_scheduler.step()\n\n    save_file = {\"model\": model.state_dict(),\n                 \"optimizer\": optimizer.state_dict(),\n                 \"lr_scheduler\": lr_scheduler.state_dict(),\n                 \"epoch\": epoch\n                }\n    torch.save(save_file, \"./model.pth\")","metadata":{"execution":{"iopub.status.busy":"2023-07-20T10:00:30.542462Z","iopub.execute_input":"2023-07-20T10:00:30.542833Z","iopub.status.idle":"2023-07-20T10:31:21.925829Z","shell.execute_reply.started":"2023-07-20T10:00:30.542802Z","shell.execute_reply":"2023-07-20T10:31:21.923709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Plot_ROC(model, val_loader, './model.pth', device)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T10:31:31.728909Z","iopub.execute_input":"2023-07-20T10:31:31.729271Z","iopub.status.idle":"2023-07-20T10:32:51.134934Z","shell.execute_reply.started":"2023-07-20T10:31:31.729240Z","shell.execute_reply":"2023-07-20T10:32:51.133770Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dfTest = pd.read_csv('/kaggle/input/cassava-leaf-disease-classification/sample_submission.csv')\ndfTest","metadata":{"execution":{"iopub.status.busy":"2023-07-20T10:43:12.323564Z","iopub.execute_input":"2023-07-20T10:43:12.323984Z","iopub.status.idle":"2023-07-20T10:43:12.337125Z","shell.execute_reply.started":"2023-07-20T10:43:12.323950Z","shell.execute_reply":"2023-07-20T10:43:12.335990Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def main(model, df):\n    df.columns = ['filepath', 'labels']\n    df['filepath'] = df['filepath'].apply(lambda x: '/kaggle/input/cassava-leaf-disease-classification/test_images/' + x)\n    img_path = df.loc[0, 'filepath']\n    print('The image path is: ', img_path)\n    \n    assert os.path.exists(img_path), \"file: '{}' dose not exist.\".format(img_path)\n    img = Image.open(img_path)\n    plt.imshow(img)\n    img = data_transform[\"valid\"](img)\n    # expand batch dimension\n    img = torch.unsqueeze(img, dim=0)\n\n    # load model weights\n    weights_path = \"./model.pth\"\n    assert os.path.exists(weights_path), \"file: '{}' dose not exist.\".format(weights_path)\n\n    model.eval()\n    with torch.no_grad():\n        # predict class\n        output = torch.squeeze(model(img.to(device))).cpu()\n        predict = torch.softmax(output, dim=0)\n        predict_cla = torch.argmax(predict).numpy()\n\n    print_res = \"class: {}   prob: {:.3}\".format(class_indict[str(predict_cla)],\n                                                 predict[predict_cla].numpy())\n    print(print_res)\n    df['labels'] = class_indict[str(predict_cla)]\n    df.to_csv('submission.csv', index=False)\n    return df","metadata":{"execution":{"iopub.status.busy":"2023-07-20T10:43:12.667123Z","iopub.execute_input":"2023-07-20T10:43:12.667509Z","iopub.status.idle":"2023-07-20T10:43:12.677861Z","shell.execute_reply.started":"2023-07-20T10:43:12.667472Z","shell.execute_reply":"2023-07-20T10:43:12.676671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = main(model, dfTest)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T10:43:13.105319Z","iopub.execute_input":"2023-07-20T10:43:13.105720Z","iopub.status.idle":"2023-07-20T10:43:13.726994Z","shell.execute_reply.started":"2023-07-20T10:43:13.105690Z","shell.execute_reply":"2023-07-20T10:43:13.726053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}