{"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":"### **Library Imports**","metadata":{"papermill":{"duration":0.015435,"end_time":"2022-03-22T13:13:45.640278","exception":false,"start_time":"2022-03-22T13:13:45.624843","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nimport re\nimport cv2\nimport torch\nimport pickle\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom time import time\nfrom torch import nn, optim\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader as DL\nfrom torch.nn.utils import weight_norm as WN\nfrom torchvision import models, transforms\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.preprocessing import StandardScaler, LabelEncoder","metadata":{"papermill":{"duration":2.851592,"end_time":"2022-03-22T13:13:48.508281","exception":false,"start_time":"2022-03-22T13:13:45.656689","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-03-23T11:47:34.072981Z","iopub.execute_input":"2022-03-23T11:47:34.073937Z","iopub.status.idle":"2022-03-23T11:47:36.805032Z","shell.execute_reply.started":"2022-03-23T11:47:34.073799Z","shell.execute_reply":"2022-03-23T11:47:36.804281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Utilities and Constants**","metadata":{"papermill":{"duration":0.015854,"end_time":"2022-03-22T13:13:48.540567","exception":false,"start_time":"2022-03-22T13:13:48.524713","status":"completed"},"tags":[]}},{"cell_type":"code","source":"le = LabelEncoder()\n\nSEED = 42\nSIZE = 256\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nTRANSFORM_FINAL = transforms.Compose([transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])])\nTRANSFORM = transforms.Compose([transforms.ToTensor(), transforms.Normalize([0.43032, 0.49672, 0.31341], [0.20665, 0.21170, 0.18763])])","metadata":{"papermill":{"duration":0.079387,"end_time":"2022-03-22T13:13:48.635448","exception":false,"start_time":"2022-03-22T13:13:48.556061","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-03-23T11:48:30.966353Z","iopub.execute_input":"2022-03-23T11:48:30.966617Z","iopub.status.idle":"2022-03-23T11:48:30.973546Z","shell.execute_reply.started":"2022-03-23T11:48:30.966587Z","shell.execute_reply":"2022-03-23T11:48:30.972573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def breaker(num: int = 50, char: str = \"*\") -> None:\n    print(\"\\n\" + num*char + \"\\n\")\n\n\n# def show_graphs(L: list, A: list) -> None:\n#     TL, VL, TA, VA = [], [], [], []\n#     for i in range(len(L)):\n#         TL.append(L[i][\"train\"])\n#         VL.append(L[i][\"valid\"])\n#         TA.append(A[i][\"train\"])\n#         VA.append(A[i][\"valid\"])\n#     x_Axis = np.arange(1, len(TL) + 1)\n#     plt.figure(figsize=(8, 6))\n#     plt.subplot(1, 2, 1)\n#     plt.plot(x_Axis, TL, \"r\", label=\"Train\")\n#     plt.plot(x_Axis, VL, \"b\", label=\"Valid\")\n#     plt.legend()\n#     plt.grid()\n#     plt.title(\"Loss Graph\")\n#     plt.subplot(1, 2, 2)\n#     plt.plot(x_Axis, TA, \"r\", label=\"Train\")\n#     plt.plot(x_Axis, VA, \"b\", label=\"Valid\")\n#     plt.legend()\n#     plt.grid()\n#     plt.title(\"Accuracy Graph\")\n#     plt.show()\n\n    \n# def load_image(path: str, size: int, make_rgb: bool=False) -> np.ndarray:\n#     image = cv2.imread(path, cv2.IMREAD_COLOR).astype(\"uint8\")\n#     if make_rgb:\n#         image = cv2.cvtColor(src=image, code=cv2.COLOR_BGR2RGB)\n#     image = cv2.resize(src=image, dsize=(size, size), interpolation=cv2.INTER_AREA)\n#     return image\n\n\ndef load_image(path: str, size: int) -> np.ndarray:\n    image = cv2.imread(path, cv2.IMREAD_COLOR)\n    image = cv2.cvtColor(src=image, code=cv2.COLOR_BGR2RGB)\n    image = cv2.resize(src=image, dsize=(size, size), interpolation=cv2.INTER_AREA)\n    return image\n\n\ndef show_image(image: np.ndarray, cmap: str = \"gnuplot2\") -> None:\n    plt.figure()\n    plt.imshow(image, cmap=cmap)\n    plt.axis(\"off\")\n    plt.show()","metadata":{"papermill":{"duration":0.031342,"end_time":"2022-03-22T13:13:48.682682","exception":false,"start_time":"2022-03-22T13:13:48.65134","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-03-23T11:48:51.361979Z","iopub.execute_input":"2022-03-23T11:48:51.36243Z","iopub.status.idle":"2022-03-23T11:48:51.371471Z","shell.execute_reply.started":"2022-03-23T11:48:51.362397Z","shell.execute_reply":"2022-03-23T11:48:51.37057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Dataset Template**","metadata":{"papermill":{"duration":0.013648,"end_time":"2022-03-22T13:13:48.710196","exception":false,"start_time":"2022-03-22T13:13:48.696548","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class DS(Dataset):\n    def __init__(self, base_path: str, filenames: np.ndarray, labels: np.ndarray = None, mode: str = \"train\", transform = None):\n\n        assert re.match(r\"^train$\", mode, re.IGNORECASE) or re.match(r\"^valid$\", mode, re.IGNORECASE) or re.match(r\"^test$\", mode, re.IGNORECASE), \"Invalid Mode\"\n        \n        self.mode = mode\n        self.base_path = base_path\n        self.filenames = filenames\n        self.transform = transform\n\n        if re.match(r\"^train$\", mode, re.IGNORECASE) or re.match(r\"^valid$\", mode, re.IGNORECASE):\n            self.labels = labels\n\n    def __len__(self):\n        return self.filenames.shape[0]\n\n    def __getitem__(self, idx):\n        image = load_image(os.path.join(self.base_path, self.filenames[idx]), SIZE)\n        if re.match(r\"^train$\", self.mode, re.IGNORECASE) or re.match(r\"^valid$\", self.mode, re.IGNORECASE):\n            return self.transform(image), torch.LongTensor(self.labels[idx])\n        else:\n            return self.transform(image)","metadata":{"papermill":{"duration":0.025781,"end_time":"2022-03-22T13:13:48.750167","exception":false,"start_time":"2022-03-22T13:13:48.724386","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-03-23T11:48:53.528927Z","iopub.execute_input":"2022-03-23T11:48:53.529593Z","iopub.status.idle":"2022-03-23T11:48:53.54027Z","shell.execute_reply.started":"2022-03-23T11:48:53.529551Z","shell.execute_reply":"2022-03-23T11:48:53.53929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Model**","metadata":{"papermill":{"duration":0.013675,"end_time":"2022-03-22T13:13:48.777652","exception":false,"start_time":"2022-03-22T13:13:48.763977","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class CNN(nn.Module):\n    def __init__(self, mode: str, model_name: str):\n        \n        super(CNN, self).__init__()\n        \n        self.model_name = model_name\n        self.mode = mode\n        \n        if re.match(r\"^vgg$\", self.model_name, re.IGNORECASE):\n            if re.match(r\"^full$\", self.mode, re.IGNORECASE):\n                self.model = models.vgg16_bn(pretrained=False, progress=True)\n                self.model.classifier[-1] = nn.Linear(in_features=self.model.classifier[-1].in_features, out_features=10)\n            elif re.match(r\"^semi$\", self.mode, re.IGNORECASE) or re.match(r\"^final$\", self.mode, re.IGNORECASE):\n                self.model = models.vgg16_bn(pretrained=True, progress=True)\n                self.freeze()\n                self.model.classifier[-1] = nn.Linear(in_features=self.model.classifier[-1].in_features, out_features=10)\n            \n        elif re.match(r\"^resnet$\", self.model_name, re.IGNORECASE):\n            if re.match(r\"^full$\", self.mode, re.IGNORECASE):\n                self.model = models.resnet50(pretrained=False, progress=True)\n                self.model.fc = nn.Linear(in_features=self.model.fc.in_features, out_features=10)\n            elif re.match(r\"^semi$\", self.mode, re.IGNORECASE) or re.match(r\"^final$\", self.mode, re.IGNORECASE):\n                self.model = models.resnet50(pretrained=True, progress=True)\n                self.freeze()\n                self.model.fc = nn.Linear(in_features=self.model.fc.in_features, out_features=10)\n        \n        elif re.match(r\"^densenet$\", self.model_name, re.IGNORECASE):\n            if re.match(r\"^full$\", self.mode, re.IGNORECASE):\n                self.model = models.densenet169(pretrained=False, progress=True)\n                self.model.classifier = nn.Linear(in_features=self.model.classifier.in_features, out_features=10)\n            elif re.match(r\"^semi$\", self.mode, re.IGNORECASE) or re.match(r\"^final$\", self.mode, re.IGNORECASE):\n                self.model = models.densenet169(pretrained=True, progress=True)\n                self.freeze()\n                self.model.classifier = nn.Linear(in_features=self.model.classifier.in_features, out_features=10)\n        \n        elif re.match(r\"^mobilenet$\", self.model_name, re.IGNORECASE):\n            if re.match(r\"^full$\", self.mode, re.IGNORECASE):\n                self.model = models.mobilenet_v3_small(pretrained=False, progress=True)\n                self.model.classifier[-1] = nn.Linear(in_features=self.model.classifier[-1].in_features, out_features=10)\n            elif re.match(r\"^semi$\", self.mode, re.IGNORECASE) or re.match(r\"^final$\", self.mode, re.IGNORECASE):\n                self.model = models.mobilenet_v3_small(pretrained=True, progress=True)\n                self.freeze()\n                self.model.classifier[-1] = nn.Linear(in_features=self.model.classifier[-1].in_features, out_features=10)\n\n    def freeze(self):\n        for params in self.parameters():\n            params.requires_grad = False\n\n        if re.match(r\"^vgg$\", self.model_name, re.IGNORECASE):\n            if re.match(r\"^semi$\", self.mode, re.IGNORECASE):\n                for names, params in self.named_parameters():\n                    if re.match(r\".*features.3[4-9].*\", names, re.IGNORECASE) or re.match(r\".*features.4[0-9].*\", names, re.IGNORECASE) or re.match(r\".*classifier.*\", names, re.IGNORECASE):\n                        params.requires_grad = True\n        \n        elif re.match(r\"^resnet$\", self.model_name, re.IGNORECASE):\n            if re.match(r\"^semi$\", self.mode, re.IGNORECASE):\n                for names, params in self.named_parameters():\n                    if re.match(r\".*layer4.*\", names, re.IGNORECASE):\n                        params.requires_grad = True\n        \n        elif re.match(r\"^densenet$\", self.model_name, re.IGNORECASE):\n            if re.match(r\"^semi$\", self.mode, re.IGNORECASE):\n                for names, params in self.named_parameters():\n                    if re.match(r\".*denseblock4.*\", names, re.IGNORECASE) or re.match(r\".*norm5.*\", names, re.IGNORECASE):\n                        params.requires_grad = True\n        \n        elif re.match(r\"^mobilenet$\", self.model_name, re.IGNORECASE):\n            if re.match(r\"^semi$\", self.mode, re.IGNORECASE):\n                for names, params in self.named_parameters():\n                    if re.match(r\".*features.9.*\", names, re.IGNORECASE) or re.match(r\".*features.1[0-2].*\", names, re.IGNORECASE) or re.match(r\".*classifier.*\", names, re.IGNORECASE):\n                        params.requires_grad = True\n    \n    def get_optimizer(self, lr=1e-3, wd=0.0):\n        params = [p for p in self.parameters() if p.requires_grad]\n        return optim.Adam(params, lr=lr, weight_decay=wd)\n\n    def get_plateau_scheduler(self, optimizer=None, patience=5, eps=1e-8):\n        return optim.lr_scheduler.ReduceLROnPlateau(optimizer=optimizer, patience=patience, eps=eps)\n    \n    def forward(self, x):\n        return nn.LogSoftmax(dim=1)(self.model(x))","metadata":{"papermill":{"duration":0.147219,"end_time":"2022-03-22T13:13:48.938868","exception":false,"start_time":"2022-03-22T13:13:48.791649","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-03-23T11:48:55.696107Z","iopub.execute_input":"2022-03-23T11:48:55.696918Z","iopub.status.idle":"2022-03-23T11:48:55.850318Z","shell.execute_reply.started":"2022-03-23T11:48:55.696859Z","shell.execute_reply":"2022-03-23T11:48:55.849274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Fit and Predict Helpers**","metadata":{"papermill":{"duration":0.013756,"end_time":"2022-03-22T13:13:48.966983","exception":false,"start_time":"2022-03-22T13:13:48.953227","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# def fit(model=None, optimizer=None, scheduler=None, epochs=None, early_stopping_patience=None, dataloaders=None, verbose=False) -> tuple:\n    \n#     def get_accuracy(y_pred, y_true):\n#         y_pred = torch.argmax(y_pred, dim=1)\n#         return torch.count_nonzero(y_pred == y_true).item() / len(y_pred)\n    \n#     if verbose:\n#         breaker()\n#         print(f\"Training ...\")\n#         breaker()\n\n#     bestLoss, bestAccs = {\"train\" : np.inf, \"valid\" : np.inf}, {\"train\" : 0.0, \"valid\" : 0.0}\n#     Losses, Accuracies = [], []\n#     name = f\"state.pt\"\n\n#     start_time = time()\n#     for e in range(epochs):\n#         e_st = time()\n#         epochLoss, epochAccs = {\"train\" : 0.0, \"valid\" : 0.0}, {\"train\" : 0.0, \"valid\" : 0.0}\n\n#         for phase in [\"train\", \"valid\"]:\n#             if phase == \"train\":\n#                 model.train()\n#             else:\n#                 model.eval()\n            \n#             lossPerPass, accsPerPass = [], []\n\n#             for X, y in dataloaders[phase]:\n#                 X, y = X.to(DEVICE), y.to(DEVICE).view(-1)\n\n#                 optimizer.zero_grad()\n#                 with torch.set_grad_enabled(phase == \"train\"):\n#                     output = model(X)\n#                     loss = torch.nn.NLLLoss()(output, y)\n#                     if phase == \"train\":\n#                         loss.backward()\n#                         optimizer.step()\n#                 lossPerPass.append(loss.item())\n#                 accsPerPass.append(get_accuracy(output, y))\n#             epochLoss[phase] = np.mean(np.array(lossPerPass))\n#             epochAccs[phase] = np.mean(np.array(accsPerPass))\n#         Losses.append(epochLoss)\n#         Accuracies.append(epochAccs)\n\n#         if early_stopping_patience:\n#             if epochLoss[\"valid\"] < bestLoss[\"valid\"]:\n#                 bestLoss = epochLoss\n#                 BLE = e + 1\n#                 torch.save({\"model_state_dict\": model.state_dict(),\n#                             \"optim_state_dict\": optimizer.state_dict()},\n#                            os.path.join(SAVE_PATH, name))\n#                 early_stopping_step = 0\n#             else:\n#                 early_stopping_step += 1\n#                 if early_stopping_step > early_stopping_patience:\n#                     print(\"\\nEarly Stopping at Epoch {}\".format(e + 1))\n#                     break\n        \n#         if epochLoss[\"valid\"] < bestLoss[\"valid\"]:\n#             bestLoss = epochLoss\n#             BLE = e + 1\n#             torch.save({\"model_state_dict\" : model.state_dict(),\n#                         \"optim_state_dict\" : optimizer.state_dict()},\n#                         os.path.join(SAVE_PATH, name))\n        \n#         if epochAccs[\"valid\"] > bestAccs[\"valid\"]:\n#             bestAccs = epochAccs\n#             BAE = e + 1\n        \n#         if scheduler:\n#             scheduler.step(epochLoss[\"valid\"])\n        \n#         if verbose:\n#             print(\"Epoch: {} | Train Loss: {:.5f} | Valid Loss: {:.5f} |\\\n# Train Accs: {:.5f} | Valid Accs: {:.5f} | Time: {:.2f} seconds\".format(e+1, \n#                                                                        epochLoss[\"train\"], epochLoss[\"valid\"], \n#                                                                        epochAccs[\"train\"], epochAccs[\"valid\"], \n#                                                                        time()-e_st))\n\n#     if verbose:                                           \n#         breaker()\n#         print(f\"Best Validation Loss at Epoch {BLE}\")\n#         breaker()\n#         print(f\"Best Validation Accs at Epoch {BAE}\")\n#         breaker()\n#         print(\"Time Taken [{} Epochs] : {:.2f} minutes\".format(len(Losses), (time()-start_time)/60))\n#         breaker()\n#         print(\"Training Completed\")\n#         breaker()\n\n#     return Losses, Accuracies, BLE, BAE, name\n\n\ndef predict(model=None, dataloader=None, path=None) -> np.ndarray:\n    model.load_state_dict(torch.load(path, map_location=DEVICE)[\"model_state_dict\"])\n    model.to(DEVICE)    \n    model.eval()\n    \n    y_pred = torch.zeros(1, 1).to(DEVICE)\n    \n    for X in dataloader:\n        X = X.to(DEVICE)\n        with torch.no_grad():\n            output = torch.argmax(torch.exp(model(X)), dim=1)\n        y_pred = torch.cat((y_pred, output.view(-1, 1)), dim=0)\n    \n    return y_pred[1:].detach().cpu().numpy()","metadata":{"papermill":{"duration":0.035515,"end_time":"2022-03-22T13:13:49.016232","exception":false,"start_time":"2022-03-22T13:13:48.980717","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-03-23T11:48:58.714731Z","iopub.execute_input":"2022-03-23T11:48:58.715156Z","iopub.status.idle":"2022-03-23T11:48:58.72483Z","shell.execute_reply.started":"2022-03-23T11:48:58.715118Z","shell.execute_reply":"2022-03-23T11:48:58.72395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Params**","metadata":{"papermill":{"duration":0.013584,"end_time":"2022-03-22T13:13:49.043603","exception":false,"start_time":"2022-03-22T13:13:49.030019","status":"completed"},"tags":[]}},{"cell_type":"code","source":"DEBUG: bool = False\n    \nif DEBUG:\n    n_splits = 3\n    epochs = 2\n    mode = \"semi\"\n    model_name = \"mobilenet\"\nelse:\n    n_splits = 25\n    epochs = 250\n    mode = \"semi\"\n    model_name = \"resnet\"\n\nbatch_size = 128\nlr = 1e-6\nwd = 1e-5\nearly_stopping = 5\npatience = None\neps = None","metadata":{"papermill":{"duration":0.020739,"end_time":"2022-03-22T13:13:49.078282","exception":false,"start_time":"2022-03-22T13:13:49.057543","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-03-23T11:49:03.072496Z","iopub.execute_input":"2022-03-23T11:49:03.073359Z","iopub.status.idle":"2022-03-23T11:49:03.079842Z","shell.execute_reply.started":"2022-03-23T11:49:03.073308Z","shell.execute_reply":"2022-03-23T11:49:03.07892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Load & Preprocess Data**","metadata":{"papermill":{"duration":0.01369,"end_time":"2022-03-22T13:13:49.105734","exception":false,"start_time":"2022-03-22T13:13:49.092044","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# start_time = time()\n\n# train_df = pd.read_csv(\"../input/cassava-leaf-disease-classification/train.csv\")\n# train_df.head()\n\n# filenames = train_df.iloc[:, 0].copy().values\n# labels = train_df.iloc[:, -1].copy().values","metadata":{"papermill":{"duration":2.100888,"end_time":"2022-03-22T13:13:51.22045","exception":false,"start_time":"2022-03-22T13:13:49.119562","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-03-22T13:27:37.421465Z","iopub.execute_input":"2022-03-22T13:27:37.42188Z","iopub.status.idle":"2022-03-22T13:27:37.44361Z","shell.execute_reply.started":"2022-03-22T13:27:37.421841Z","shell.execute_reply":"2022-03-22T13:27:37.442864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Train**","metadata":{"papermill":{"duration":0.014217,"end_time":"2022-03-22T13:13:51.249031","exception":false,"start_time":"2022-03-22T13:13:51.234814","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# for tr_idx, va_idx in StratifiedKFold(n_splits=n_splits, random_state=SEED, shuffle=True).split(filenames, labels):\n#     break\n\n# tr_filenames, va_filenames, tr_labels, va_labels = filenames[tr_idx], filenames[va_idx], labels[tr_idx], labels[va_idx]\n\n# tr_data_setup = DS(\"../input/cassava-leaf-disease-classification/train_images\", tr_filenames, tr_labels.reshape(-1, 1), \"train\", TRANSFORM)\n# va_data_setup = DS(\"../input/cassava-leaf-disease-classification/train_images\", va_filenames, va_labels.reshape(-1, 1), \"valid\", TRANSFORM)\n\n# dataloaders = {\n#     \"train\" : DL(tr_data_setup, batch_size=batch_size, shuffle=True, generator=torch.manual_seed(SEED)),\n#     \"valid\" : DL(va_data_setup, batch_size=batch_size, shuffle=False),\n# }\n\n# torch.manual_seed(SEED)\n# model = CNN(mode=mode, model_name=model_name).to(DEVICE)\n# optimizer = model.get_optimizer(lr=lr, wd=wd)\n# scheduler = None\n# if isinstance(patience, int) and isinstance(eps, float):\n#     scheduler = get_plateau_scheduler(optimizer=optimizer, patience=patience, eps=eps)\n\n# L, A, BLE, BAE, name = fit(model=model, optimizer=optimizer, scheduler=scheduler, \n#                            epochs=epochs, early_stopping_patience=early_stopping, \n#                            dataloaders=dataloaders, verbose=True)\n\n# show_graphs(L, A)","metadata":{"papermill":{"duration":0.019925,"end_time":"2022-03-22T13:13:51.282672","exception":false,"start_time":"2022-03-22T13:13:51.262747","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-03-22T13:27:38.025693Z","iopub.execute_input":"2022-03-22T13:27:38.026087Z","iopub.status.idle":"2022-03-22T13:27:45.300248Z","shell.execute_reply.started":"2022-03-22T13:27:38.02605Z","shell.execute_reply":"2022-03-22T13:27:45.298964Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Submission**","metadata":{"papermill":{"duration":0.013618,"end_time":"2022-03-22T13:13:51.309996","exception":false,"start_time":"2022-03-22T13:13:51.296378","status":"completed"},"tags":[]}},{"cell_type":"code","source":"ss_df = pd.read_csv(\"../input/cassava-leaf-disease-classification/sample_submission.csv\")\nts_filenames = ss_df.image_id.copy().values\n\nts_data_setup = DS(\"../input/cassava-leaf-disease-classification/test_images\", ts_filenames, None, \"test\", TRANSFORM)\nts_data = DL(ts_data_setup, batch_size=batch_size, shuffle=False)\n\nmodel = CNN(mode=\"full\", model_name=model_name).to(DEVICE)\ny_pred = predict(model=model, dataloader=ts_data, path=\"../input/cldc-r256slow/saves/state.pt\")\n\nss_df[\"label\"] = y_pred.astype(\"uint8\")\nss_df.to_csv(\"submission.csv\", index=False)","metadata":{"papermill":{"duration":0.019732,"end_time":"2022-03-22T13:13:51.343386","exception":false,"start_time":"2022-03-22T13:13:51.323654","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-03-23T11:52:53.944346Z","iopub.execute_input":"2022-03-23T11:52:53.944604Z","iopub.status.idle":"2022-03-23T11:52:54.583122Z","shell.execute_reply.started":"2022-03-23T11:52:53.944575Z","shell.execute_reply":"2022-03-23T11:52:54.582382Z"},"trusted":true},"execution_count":null,"outputs":[]}]}