{"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":"# Imports","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append('../input/d/kozodoi/timm-pytorch-image-models/pytorch-image-models-master')\nimport timm","metadata":{"papermill":{"duration":0.015936,"end_time":"2021-03-21T08:24:58.959372","exception":false,"start_time":"2021-03-21T08:24:58.943436","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pytorch_lightning as pl\nimport torch\nimport pandas as pd\nimport torch.nn as nn\n\nfrom PIL import Image\nfrom sklearn.model_selection import KFold\nfrom torchvision import transforms as tsfm\nfrom torch.utils.data import Dataset, DataLoader\nfrom pytorch_lightning import Trainer, seed_everything\nfrom pytorch_lightning.loggers import CSVLogger\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping\nfrom pytorch_lightning.metrics import Metric\nfrom typing import List, Dict\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport matplotlib.pyplot as plt\nfrom scipy.optimize import minimize\n\nfrom glob import glob","metadata":{"papermill":{"duration":5.910497,"end_time":"2021-03-21T08:25:04.879091","exception":false,"start_time":"2021-03-21T08:24:58.968594","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"class CFG:\n    # data config\n    root_dir_origin = \"../input/plant-pathology-2021-fgvc8\"\n    root_dir_resized = \"../input/resized-plantpathology2021fgvc8-train-data/resized_plant-pathology-2021-fgvc8_train_data\"\n    \n    train_csv_path = os.path.join(root_dir_origin, 'train.csv')\n    folds_csv_path = \"../input/pp2021-kfold-tfrecords-0/folds.csv\"\n#     folds_csv_path = \"../input/pp2021-dataset-gnueih/6folds_pp2021.csv\"\n\n    train_imgs_dir = os.path.join(root_dir_resized, 'resized_train_images_360_512')\n    test_imgs_dir = os.path.join(root_dir_origin, 'test_images')\n    \n    num_classes = 5\n    labels = np.array(['powdery_mildew',\n                     'scab',\n                     'complex',\n                     'frog_eye_leaf_spot',\n                     'rust',])\n    \n    # model config\n    model_name = 'tf_efficientnet_b4_ns'\n    \n    model_paths = glob('../input/pp2021-models/ef4_ns_v18_5fold5/*')\n#     model_paths = [ '../input/final-pp2021-training/ckpt/tf_efficientnet_b4_ns_kag_final_v18/ftf_efficientnet_b4_ns_epoch=02-valid_f1=0.9091.ckpt',\n#                     '../input/final-pp2021-training/ckpt/tf_efficientnet_b4_ns_kag_final_v18/ftf_efficientnet_b4_ns_epoch=02-valid_f1=0.9053.ckpt',\n#                     '../input/final-pp2021-training/ckpt/tf_efficientnet_b4_ns_kag_final_v18/ftf_efficientnet_b4_ns_epoch=02-valid_f1=0.9160.ckpt',\n#                     '../input/final-pp2021-training/ckpt/tf_efficientnet_b4_ns_kag_final_v18/ftf_efficientnet_b4_ns_epoch=02-valid_f1=0.9081.ckpt',\n#                     '../input/final-pp2021-training/ckpt/tf_efficientnet_b4_ns_kag_final_v18/ftf_efficientnet_b4_ns_epoch=02-valid_f1=0.9086.ckpt',]\n    # training hyper-parameters\n    seed = 42\n    batch_size = 32\n    n_fold = 5\n    num_workers = 4\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"papermill":{"duration":0.881316,"end_time":"2021-03-21T08:25:05.811987","exception":false,"start_time":"2021-03-21T08:25:04.930671","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# CFG.model_paths.sort()\nCFG.model_paths","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{"papermill":{"duration":0.008651,"end_time":"2021-03-21T08:25:05.831717","exception":false,"start_time":"2021-03-21T08:25:05.823066","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Hàm predict một model\ndef predict_one_model(model, dataloader, device, tta=1, valid=False):\n    model.eval()\n    model.to(device)\n    outputs = []\n    with torch.no_grad():\n        for b in dataloader:\n            if valid:\n                imgs = b[0]\n            else:\n                imgs = b\n            imgs = imgs.to(device)\n            y_pred = model(imgs).detach()\n            outputs.append(y_pred)\n    return torch.cat(outputs, dim=0)\n\n# Hàm predict nhiều model, sau đó lấy trung bình các predictions \ndef predict_multi_model(models, dataloader, device, tta=1, valid=False):\n    preds = None\n    for model in models:\n        pred = predict_one_model(model, dataloader, CFG.device, tta=tta, valid=valid)\n        if preds is None:\n            preds = torch.sigmoid(pred)\n        else:\n            preds += torch.sigmoid(pred)\n    return preds / len(models)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ImageDataset(Dataset):\n    \"\"\" Leaf Disease Dataset \"\"\"\n    def __init__(self,\n                image_names,\n                labels,\n                image_dir, \n                transforms):        \n        self.image_names = image_names\n        self.image_dir = image_dir\n        self.transforms = transforms                \n        self.labels = labels\n\n    def __len__(self) -> int:\n        return len(self.image_names)\n    \n    def get_orig_img(self, idx: int):\n        return Image.open(os.path.join(self.image_dir, self.image_names[idx]))\n    \n    def __getitem__(self, idx: int):\n        image = np.array(self.get_orig_img(idx))        \n        transformed_image = self.transforms(image=image)['image']\n        if self.labels is not None:\n            target = self.labels[idx]\n            return transformed_image, target\n        return transformed_image","metadata":{"_kg_hide-input":true,"papermill":{"duration":0.029166,"end_time":"2021-03-21T08:25:05.871596","exception":false,"start_time":"2021-03-21T08:25:05.84243","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_transform = A.Compose([\n    A.Resize(height=360, width=512, p=1.0),\n    A.Normalize(),\n    ToTensorV2(),\n])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_val_loader(valid_df):\n    valid_dataset = ImageDataset(image_names=valid_df.image.values, \n                                labels=valid_df[CFG.labels].values, \n                                image_dir=CFG.train_imgs_dir, \n                                transforms=valid_transform)\n    valid_loader = DataLoader(\n                    valid_dataset,\n                    batch_size=CFG.batch_size,\n                    num_workers=CFG.num_workers,\n                    shuffle=False,\n                    pin_memory=True)\n    return valid_loader","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nDefine F1 score metric\n\"\"\"\nclass F1Score(Metric):\n    def __init__(self, threshold: float = 0.5, dist_sync_on_step=False):\n        super().__init__(dist_sync_on_step=dist_sync_on_step)\n        self.threshold = threshold\n        self.add_state(\"tp\", default=torch.tensor(0), dist_reduce_fx=\"sum\")\n        self.add_state(\"fp\", default=torch.tensor(0), dist_reduce_fx=\"sum\")\n        self.add_state(\"fn\", default=torch.tensor(0), dist_reduce_fx=\"sum\")\n\n    def update(self, preds: torch.Tensor, target: torch.Tensor, sigmoid=True):\n        assert preds.shape == target.shape\n        with torch.no_grad():\n            if sigmoid: preds = torch.sigmoid(preds)\n            preds = (preds > self.threshold).type(torch.long)\n\n            target_healthy = 1 - torch.clip(target.sum(dim=-1, keepdim=True), 0, 1)\n            pred_healthy = 1 - torch.clip(preds.sum(dim=-1, keepdim=True), 0, 1)\n            preds = torch.cat([preds, pred_healthy], -1)\n            target = torch.cat([target, target_healthy], -1)\n\n            tp = (preds*target).sum()\n            fp = preds.sum() - tp\n            fn = ((1 - preds)*target).sum()\n        \n        self.tp += tp.item()\n        self.fp += fp.item()\n        self.fn += fn.item()\n\n    def compute(self):\n        f1 = 2.0 * self.tp / (2.0 * self.tp + self.fn + self.fp)\n        return f1","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import multilabel_confusion_matrix\nimport seaborn as sns\n\ndef plot_confusion_matrix(\n    y_test, \n    y_pred_proba, \n    threshold, \n    label_names=CFG.labels\n)-> None:\n    \"\"\"\n    \"\"\"\n    y_pred = np.where(y_pred_proba > threshold, 1, 0)\n    c_matrices = multilabel_confusion_matrix(y_test, y_pred)\n    \n    cmap = plt.get_cmap('Blues')\n    fig, axes = plt.subplots(nrows=2, ncols=3, figsize=(15, 8))\n\n    for cm, label, ax in zip(c_matrices, label_names, axes.flatten()):\n        sns.heatmap(cm, annot=True, fmt='g', ax=ax, cmap=cmap);\n\n        ax.set_xlabel('Predicted labels');\n        ax.set_ylabel('True labels'); \n        ax.set_title(f'{label}');\n\n    plt.tight_layout()    \n    plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model\n\nLoad các model có path trong `CFG.model_paths`","metadata":{}},{"cell_type":"code","source":"models = []\nfor path in CFG.model_paths:\n    print(path)\n    pretrain = torch.load(path, map_location=CFG.device)\n    state_dict = {k[6:]:v for k,v in pretrain['state_dict'].items() if 'model' in k}\n    model = timm.create_model(CFG.model_name, pretrained=False, num_classes=CFG.num_classes)\n    model.load_state_dict(state_dict)\n    models.append(model)\nprint(len(models))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Threshold configuration\nMột số phương pháp tìm threshold cho các class","metadata":{}},{"cell_type":"code","source":"def one_hot_encoded_df(dataset_df):\n    # copy dataframe\n    unique_labels = dataset_df.labels.unique()\n    new_column_names = list(set(' '.join(unique_labels).split()))\n    # initialize columns with zero\n    dataset_df[new_column_names] = 0        \n    # one-hot-encoding using the column names\n    for labels in unique_labels:                \n        label_indices = dataset_df[dataset_df['labels'] == labels].index\n        splited_labels = labels.split()\n        dataset_df.loc[label_indices, splited_labels] = 1\n    return dataset_df\n\nfolds_df = pd.read_csv(CFG.folds_csv_path)\ndf = one_hot_encoded_df(pd.read_csv(CFG.train_csv_path))\ndf = folds_df.merge(df, on='image')\ndf.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Predict trên validation dataset","metadata":{}},{"cell_type":"code","source":"# pred_df = df.drop([*CFG.labels, 'healthy', 'labels'], axis=1)\n# pred_df[CFG.labels] = 0\n# pred_df.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Out of fold predictions","metadata":{}},{"cell_type":"code","source":"# for fold_num in range(0, 5):\n# #     fold_num = 5\n#     print(fold_num)\n#     valid_df = df[df.fold == fold_num].reset_index()\n#     valid_loader = get_val_loader(valid_df)\n#     valid_pred = predict_one_model(models[fold_num], valid_loader, CFG.device, valid=True)\n#     pred_df.loc[pred_df.fold == fold_num, CFG.labels] = torch.sigmoid(valid_pred).cpu().numpy()    ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Sử dụng scipy optimizer tìm threshold","metadata":{}},{"cell_type":"code","source":"# def my_metric(thresholds):\n#     torch_metric = F1Score()\n#     torch_metric.threshold = torch.tensor(thresholds)\n#     return 1 - torch_metric(torch.tensor(pred_df[pred_df.fold == 4][CFG.labels].values), \n#                       torch.tensor(df[df.fold == 4][CFG.labels].values, \n#                                    dtype=torch.long), \n#                                    False).cpu().numpy()\n# opt_thresh = minimize(my_metric, np.array([0.22 for i in range(CFG.num_classes)]),method='POWELL', bounds=[(0.2, 0.7) for i in range(CFG.num_classes)])\n# print(opt_thresh)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Plot Confusion matrix","metadata":{}},{"cell_type":"code","source":"# y_true, y_pred_proba = df[CFG.labels].values, pred_df[CFG.labels].values\n# plot_confusion_matrix(y_true, y_pred_proba, threshold=opt_thresh.x)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Một cách khác tìm threshold","metadata":{}},{"cell_type":"code","source":"# import tensorflow_addons as tfa\n# y_true, y_pred = df[CFG.labels].values, pred_df[CFG.labels].values\n# thresholds = np.arange(.01, 1., .01)\n# scores = []\n\n# for threshold in thresholds:\n#     m = tfa.metrics.F1Score(\n#         num_classes=5, \n#         average=None, \n#         threshold=threshold)\n#     m.update_state(y_true, y_pred)\n#     scores.append(m.result().numpy())\n    \n# pdf = pd.DataFrame(columns=CFG.labels, data=scores, index=pd.Index(thresholds, name='threshold'))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# pdf.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# thresholds2 = []#lưu threshold mà class có giá trị lớn nhất\n# scores = []#lưu giá trị f1 lớn nhất của classs\n\n# for x in CFG.labels:\n#     thresholds2.append(pdf[x].idxmax())#tìm threshold mà x có giá trị lớn nhất\n#     scores.append(pdf[x].max())#tìm scores có giá tị lớn nhất\n#     print(f'{x}: {pdf.loc[.5, x]:.4f} >>> {pdf.loc[thresholds2[-1], x]:.4f} ({thresholds2[-1]:.2f})')#lấy threshold 0.5 để so sánh giữa khách quan nhất với lớn nhất\n# # print(df.loc[0.39])\n# # print(df.complex.sort_values())\n# # df.loc[thresholds[-1]]\n# print(f'\\nmean score: {pdf.loc[.5].mean():.4f} >>> {np.mean(scores):.4f}')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot_confusion_matrix(y_true, y_pred_proba, threshold=np.array(thresholds2))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Một cách khác nữa","metadata":{}},{"cell_type":"code","source":"# thresholds = np.linspace(0.2, 0.7, 31)\n# scores = [1 - my_metric(np.ones(5)*t) for t in thresholds]\n\n# threshold_best_index = np.argmax(scores) \n# score_best = scores[threshold_best_index]\n# threshold_best = thresholds[threshold_best_index]\n\n# plt.plot(thresholds, scores)\n# plt.plot(threshold_best, score_best, \"xr\", label=\"Best threshold\")\n# plt.xlabel(\"Threshold\")\n# plt.ylabel(\"IoU\")\n# plt.title(\"Threshold vs IoU ({}, {})\".format(threshold_best, score_best))\n# plt.legend()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot_confusion_matrix(y_true, y_pred_proba, threshold=np.ones((5,))*threshold_best)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot_confusion_matrix(y_true, y_pred_proba, threshold=np.ones((5,))*0.4333)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Pred test","metadata":{}},{"cell_type":"code","source":"\ntest_dataset = ImageDataset(image_names=os.listdir(CFG.test_imgs_dir), \n                            labels=None, \n                            image_dir=CFG.test_imgs_dir, \n                            transforms=valid_transform,)\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=CFG.batch_size,\n    num_workers=CFG.num_workers,\n    shuffle=False,\n    pin_memory=True,\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"figure, axes = plt.subplots(1, 3, figsize=[20, 10])\n\nfor i, ax in enumerate(axes):\n    image = test_dataset.get_orig_img(i)\n    ax.imshow(image)\n    ax.axis('off')\n    \nplt.show()","metadata":{"_kg_hide-input":true,"papermill":{"duration":4.163951,"end_time":"2021-03-21T08:25:12.535026","exception":false,"start_time":"2021-03-21T08:25:08.371075","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# predicts = predict(model, test_loader, CFG.device)\n# predicts = torch.sigmoid(predicts)\nlogits = predict_multi_model(models, test_loader, CFG.device, valid=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Create submission.csv","metadata":{"papermill":{"duration":0.034282,"end_time":"2021-03-21T08:25:29.874649","exception":false,"start_time":"2021-03-21T08:25:29.840367","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# thresholds = torch.round(torch.FloatTensor([0.49692652, 0.47886734, 0.45990339, 0.50891365, 0.50]) * 10000) / 10000\n# thresholds = torch.round(torch.FloatTensor([0.47, 0.47886734, 0.45990339, 0.50891365, 0.50]) * 10000) / 10000\n\n# print(thresholds)\n# tensor([0.5398, 0.5112, 0.5599, 0.5087, 0.5776])\n# powdery_mildew: 0.9641 >>> 0.9670 (0.54)\n# scab: 0.9118 >>> 0.9131 (0.53)\n# complex: 0.7404 >>> 0.7459 (0.56)\n# frog_eye_leaf_spot: 0.8815 >>> 0.8834 (0.56)\n# rust: 0.9297 >>> 0.9350 (0.58)\n\n# mean score: 0.8855 >>> 0.8889","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for i in range(len(predicts)):\n# #     predicts[i] = predicts[i] > thresholds\n# predicts = logits > thresholds.to(CFG.device)\npredicts = logits > 0.4333\npredicts = predicts.type(torch.bool).cpu().numpy()\nlabels = []\n\nfor i in range(len(predicts)):\n    labels.append(' '.join(CFG.labels[predicts[i]]))\n    \nlabels = ['healthy' if ('healthy' in x or x == '') else x for x in labels]\n    \nsdf = pd.DataFrame({\n    'image': test_dataset.image_names,\n    'labels': labels})\n\nsdf.to_csv('submission.csv', index=False)\ndisplay(sdf.head())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}