{"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":"package_paths = [\n    '../input/pytorch-image-library/pytorch-image-models-master/pytorch-image-models-master',\n]\n\nimport sys\n\n\nfor pth in package_paths:\n    sys.path.append(pth)","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:05.302467Z","iopub.execute_input":"2022-04-08T11:26:05.303237Z","iopub.status.idle":"2022-04-08T11:26:05.307983Z","shell.execute_reply.started":"2022-04-08T11:26:05.303184Z","shell.execute_reply":"2022-04-08T11:26:05.307194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Import","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport os\nimport cv2\nimport timm\nimport torch\nimport torch.nn as nn\nimport albumentations as A\nimport pytorch_lightning as pl\nimport matplotlib.pyplot as plt\nimport torchmetrics\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom albumentations.core.composition import Compose, OneOf\nfrom albumentations.pytorch import ToTensorV2\n\nfrom pytorch_lightning import Trainer, seed_everything\nfrom pytorch_lightning import Callback\nfrom pytorch_lightning.loggers import CSVLogger\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping\nfrom sklearn.model_selection import StratifiedKFold","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-04-08T11:26:05.366520Z","iopub.execute_input":"2022-04-08T11:26:05.366711Z","iopub.status.idle":"2022-04-08T11:26:05.373077Z","shell.execute_reply.started":"2022-04-08T11:26:05.366689Z","shell.execute_reply":"2022-04-08T11:26:05.372260Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"PyTorch Lightning version: {pl.__version__}\")","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:05.422871Z","iopub.execute_input":"2022-04-08T11:26:05.423696Z","iopub.status.idle":"2022-04-08T11:26:05.428936Z","shell.execute_reply.started":"2022-04-08T11:26:05.423659Z","shell.execute_reply":"2022-04-08T11:26:05.427962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"DEBUG = False\n\nclass CFG:\n    seed = 42\n    model_name = 'tf_efficientnet_b5_ns'\n#     model_name = \"swin_base_patch4_window12_384\" # epoch:20\n    pretrained = True\n    img_size = 512\n    num_classes = 100\n    lr = 1e-4\n    max_lr = 1e-3\n#     lr = 1e-5\n#     max_lr = 1e-4\n    pct_start = 0.2\n    div_factor = 1.0e+3\n    final_div_factor = 1.0e+3\n    num_epochs = 27\n    batch_size = 16\n    accum = 1\n    precision = 16\n    n_fold = 4\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:05.480639Z","iopub.execute_input":"2022-04-08T11:26:05.480921Z","iopub.status.idle":"2022-04-08T11:26:05.487086Z","shell.execute_reply.started":"2022-04-08T11:26:05.480893Z","shell.execute_reply":"2022-04-08T11:26:05.486068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:05.538986Z","iopub.execute_input":"2022-04-08T11:26:05.539280Z","iopub.status.idle":"2022-04-08T11:26:05.545912Z","shell.execute_reply.started":"2022-04-08T11:26:05.539253Z","shell.execute_reply":"2022-04-08T11:26:05.544890Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load images that have been pre-resized by Gaurav Dutta to speed up the learning process. \nhttps://www.kaggle.com/competitions/sorghum-id-fgvc-9/discussion/313378","metadata":{}},{"cell_type":"code","source":"PATH = \"../input/sorghum-id-fgvc-9/\"\n\n# TRAIN_DIR = PATH + 'train_images/'\nTRAIN_DIR = \"../input/sorghum-cultivar-identification-512512/train/\"\n# TRAIN_DIR = \"../input/sorghum-cultivar-identification-256256/train/\"\nTEST_DIR = PATH + 'test/'","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:05.594098Z","iopub.execute_input":"2022-04-08T11:26:05.594537Z","iopub.status.idle":"2022-04-08T11:26:05.598963Z","shell.execute_reply.started":"2022-04-08T11:26:05.594508Z","shell.execute_reply":"2022-04-08T11:26:05.598265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_all = pd.read_csv(PATH + \"train_cultivar_mapping.csv\")\nprint(len(df_all))\ndf_all.dropna(inplace=True)\nprint(len(df_all))\ndf_all.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:05.657860Z","iopub.execute_input":"2022-04-08T11:26:05.658201Z","iopub.status.idle":"2022-04-08T11:26:05.711431Z","shell.execute_reply.started":"2022-04-08T11:26:05.658174Z","shell.execute_reply":"2022-04-08T11:26:05.710691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"unique_cultivars = list(df_all[\"cultivar\"].unique())\nnum_classes = len(unique_cultivars)\n\nCFG.num_classes = num_classes\nprint(num_classes)","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:05.714579Z","iopub.execute_input":"2022-04-08T11:26:05.715122Z","iopub.status.idle":"2022-04-08T11:26:05.723104Z","shell.execute_reply.started":"2022-04-08T11:26:05.715092Z","shell.execute_reply":"2022-04-08T11:26:05.722419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DIR = \"../input/sorghum-id-fgvc-9/train_images\"\n\nprint(sum(os.path.isfile(os.path.join(DIR, name)) for name in os.listdir(DIR)))","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:05.772600Z","iopub.execute_input":"2022-04-08T11:26:05.772933Z","iopub.status.idle":"2022-04-08T11:26:14.460169Z","shell.execute_reply.started":"2022-04-08T11:26:05.772903Z","shell.execute_reply":"2022-04-08T11:26:14.459414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_all[\"file_path\"] = df_all[\"image\"].apply(lambda image: TRAIN_DIR + image)\ndf_all[\"cultivar_index\"] = df_all[\"cultivar\"].map(lambda item: unique_cultivars.index(item))\ndf_all[\"is_exist\"] = df_all[\"file_path\"].apply(lambda file_path: os.path.exists(file_path))\ndf_all = df_all[df_all.is_exist==True]\ndf_all.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:14.461989Z","iopub.execute_input":"2022-04-08T11:26:14.462506Z","iopub.status.idle":"2022-04-08T11:26:23.144703Z","shell.execute_reply.started":"2022-04-08T11:26:14.462467Z","shell.execute_reply":"2022-04-08T11:26:23.143999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG == True:\n    df_all = df_all[:200]\n    CFG.num_epochs = 2","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:23.147044Z","iopub.execute_input":"2022-04-08T11:26:23.147334Z","iopub.status.idle":"2022-04-08T11:26:23.151224Z","shell.execute_reply.started":"2022-04-08T11:26:23.147305Z","shell.execute_reply":"2022-04-08T11:26:23.150182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# StratifiedKFold","metadata":{}},{"cell_type":"code","source":"skf = StratifiedKFold(n_splits=CFG.n_fold, shuffle=True, random_state=CFG.seed)\n\nfor train_idx, valid_idx in skf.split(df_all['image'], df_all[\"cultivar_index\"]):\n    df_train = df_all.iloc[train_idx]\n    df_valid = df_all.iloc[valid_idx]\n\nprint(f\"train size: {len(df_train)}\")\nprint(f\"valid size: {len(df_valid)}\")\n\nprint(df_train.cultivar.value_counts())\nprint(df_valid.cultivar.value_counts())","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:23.153592Z","iopub.execute_input":"2022-04-08T11:26:23.153982Z","iopub.status.idle":"2022-04-08T11:26:23.180962Z","shell.execute_reply.started":"2022-04-08T11:26:23.153942Z","shell.execute_reply":"2022-04-08T11:26:23.180273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define Dataset","metadata":{}},{"cell_type":"code","source":"class SorghumDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.image_path = df['file_path'].values\n        self.labels = df[\"cultivar_index\"].values\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.labels)\n\n    def __getitem__(self, idx):\n#         image_id = self.image_id[idx]\n        label = torch.tensor(self.labels[idx], dtype=torch.float32)\n        \n        image_path = self.image_path[idx]\n        image = cv2.imread(image_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        augmented = self.transform(image=image)\n        image = augmented['image']\n        return {'image':image, 'target': label}","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:23.182390Z","iopub.execute_input":"2022-04-08T11:26:23.182873Z","iopub.status.idle":"2022-04-08T11:26:23.190348Z","shell.execute_reply.started":"2022-04-08T11:26:23.182836Z","shell.execute_reply":"2022-04-08T11:26:23.189674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Augmant","metadata":{}},{"cell_type":"code","source":"def get_transform(phase: str):\n    if phase == 'train':\n        return Compose([\n            A.RandomResizedCrop(height=CFG.img_size, width=CFG.img_size),\n            A.Flip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            A.ShiftScaleRotate(p=0.5),\n            A.HueSaturationValue(p=0.5),\n            A.OneOf([\n                A.RandomBrightnessContrast(p=0.5),\n                A.RandomGamma(p=0.5),\n            ], p=0.5),\n            A.OneOf([\n                A.Blur(p=0.1),\n                A.GaussianBlur(p=0.1),\n                A.MotionBlur(p=0.1),\n            ], p=0.1),\n            A.OneOf([\n                A.GaussNoise(p=0.1),\n                A.ISONoise(p=0.1),\n                A.GridDropout(ratio=0.5, p=0.2),\n                A.CoarseDropout(max_holes=16, min_holes=8, max_height=16, max_width=16, min_height=8, min_width=8, p=0.2)\n            ], p=0.2),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])\n    else:\n        return Compose([\n            A.Resize(height=CFG.img_size, width=CFG.img_size),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:23.191682Z","iopub.execute_input":"2022-04-08T11:26:23.192169Z","iopub.status.idle":"2022-04-08T11:26:23.204143Z","shell.execute_reply.started":"2022-04-08T11:26:23.192118Z","shell.execute_reply":"2022-04-08T11:26:23.203389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = SorghumDataset(df_train, get_transform('train'))\nvalid_dataset = SorghumDataset(df_valid, get_transform('valid'))\n\ntrain_loader = DataLoader(train_dataset, batch_size=CFG.batch_size, shuffle=True, pin_memory=True, drop_last=True, num_workers=2)\nvalid_loader = DataLoader(valid_dataset, batch_size=CFG.batch_size, shuffle=False, pin_memory=True, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:23.205423Z","iopub.execute_input":"2022-04-08T11:26:23.205903Z","iopub.status.idle":"2022-04-08T11:26:23.218210Z","shell.execute_reply.started":"2022-04-08T11:26:23.205866Z","shell.execute_reply":"2022-04-08T11:26:23.217415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CFG.steps_per_epoch = len(train_loader)\nCFG.steps_per_epoch","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:23.220938Z","iopub.execute_input":"2022-04-08T11:26:23.221313Z","iopub.status.idle":"2022-04-08T11:26:23.230188Z","shell.execute_reply.started":"2022-04-08T11:26:23.221284Z","shell.execute_reply":"2022-04-08T11:26:23.229426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define Model","metadata":{}},{"cell_type":"code","source":"# EfficientNet\nclass CustomEffNet(nn.Module):\n    def __init__(self, model_name='tf_efficientnet_b0_ns', pretrained=True):\n        super().__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        in_features = self.model.get_classifier().in_features\n#         self.model.fc = nn.Linear(in_features, CFG.num_classes)\n        self.model.classifier = nn.Sequential(\n            nn.Linear(in_features, in_features),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.5),\n            nn.Linear(in_features, CFG.num_classes)\n        )\n\n    def forward(self, x):\n        x = self.model(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:23.231177Z","iopub.execute_input":"2022-04-08T11:26:23.232916Z","iopub.status.idle":"2022-04-08T11:26:23.240358Z","shell.execute_reply.started":"2022-04-08T11:26:23.232886Z","shell.execute_reply":"2022-04-08T11:26:23.239577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ViTModel(nn.Module):\n    def __init__(self, model_name='vit_tiny_patch16_224', pretrained=True):\n        super().__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        self.model.head = nn.Linear(self.model.head.in_features, CFG.num_classes)\n    def forward(self, x):\n        x = self.model(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:23.243542Z","iopub.execute_input":"2022-04-08T11:26:23.244168Z","iopub.status.idle":"2022-04-08T11:26:23.253050Z","shell.execute_reply.started":"2022-04-08T11:26:23.244121Z","shell.execute_reply":"2022-04-08T11:26:23.252298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LitSorghum(pl.LightningModule):\n    def __init__(self, model):\n        super(LitSorghum, self).__init__()\n        self.model = model\n        self.metric = torchmetrics.Accuracy(threshold=0.5, num_classes=CFG.num_classes)\n        self.criterion = nn.CrossEntropyLoss()\n        self.lr = CFG.lr\n\n    def forward(self, x, *args, **kwargs):\n        return self.model(x)\n\n    def configure_optimizers(self):\n        self.optimizer = torch.optim.Adam(self.model.parameters(), lr=self.lr)\n        self.scheduler = torch.optim.lr_scheduler.OneCycleLR(self.optimizer, \n                                                             epochs=CFG.num_epochs, steps_per_epoch=CFG.steps_per_epoch,\n                                                             max_lr=CFG.max_lr, pct_start=CFG.pct_start, \n                                                             div_factor=CFG.div_factor, final_div_factor=CFG.final_div_factor)\n        scheduler = {'scheduler': self.scheduler, 'interval': 'step',}\n\n        return [self.optimizer], [scheduler]\n\n    def training_step(self, batch, batch_idx):\n        image = batch['image']\n        target = batch['target'].long()\n        output = self.model(image)\n        loss = self.criterion(output, target)\n        score = self.metric(output.argmax(1), target)\n        logs = {'train_loss': loss, 'train_acc': score, 'lr': self.optimizer.param_groups[0]['lr']}\n        self.log_dict(\n            logs,\n            on_step=False, on_epoch=True, prog_bar=True, logger=True\n        )\n        return loss\n    \n    def validation_step(self, batch, batch_idx):\n        image = batch['image']\n        target = batch['target'].long()\n        output = self.model(image)\n        loss = self.criterion(output, target)\n        score = self.metric(output.argmax(1), target)\n        logs = {'valid_loss': loss, 'valid_acc': score}\n        self.log_dict(\n            logs,\n            on_step=False, on_epoch=True, prog_bar=True, logger=True\n        )\n        return loss","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:23.255235Z","iopub.execute_input":"2022-04-08T11:26:23.255715Z","iopub.status.idle":"2022-04-08T11:26:23.269402Z","shell.execute_reply.started":"2022-04-08T11:26:23.255679Z","shell.execute_reply":"2022-04-08T11:26:23.268571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(CFG.model_name, CFG.pretrained)\nmodel = CustomEffNet(model_name=CFG.model_name, pretrained=CFG.pretrained)\n# model = ViTModel(model_name=CFG.model_name, pretrained=CFG.pretrained)","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:23.270605Z","iopub.execute_input":"2022-04-08T11:26:23.271253Z","iopub.status.idle":"2022-04-08T11:26:24.111215Z","shell.execute_reply.started":"2022-04-08T11:26:23.271212Z","shell.execute_reply":"2022-04-08T11:26:24.110435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lit_model = LitSorghum(model.model)","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:24.112350Z","iopub.execute_input":"2022-04-08T11:26:24.112662Z","iopub.status.idle":"2022-04-08T11:26:24.118440Z","shell.execute_reply.started":"2022-04-08T11:26:24.112620Z","shell.execute_reply":"2022-04-08T11:26:24.117625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"logger = CSVLogger(save_dir='logs/', name=CFG.model_name)\nlogger.log_hyperparams(CFG.__dict__)\ncheckpoint_callback = ModelCheckpoint(monitor='valid_loss',\n                                      save_top_k=1,\n                                      save_last=True,\n                                      save_weights_only=True,\n                                      filename='{epoch:02d}-{valid_loss:.4f}-{valid_acc:.4f}',\n                                      verbose=False,\n                                      mode='min')\n\ntrainer = Trainer(\n    max_epochs=CFG.num_epochs,\n    gpus=[0],\n    accumulate_grad_batches=CFG.accum,\n    precision=CFG.precision,\n    callbacks=[checkpoint_callback], \n    logger=logger,\n    weights_summary='top',\n)","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:24.119932Z","iopub.execute_input":"2022-04-08T11:26:24.120250Z","iopub.status.idle":"2022-04-08T11:26:24.134861Z","shell.execute_reply.started":"2022-04-08T11:26:24.120215Z","shell.execute_reply":"2022-04-08T11:26:24.133587Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"trainer.fit(lit_model, train_dataloaders=train_loader, val_dataloaders=valid_loader)","metadata":{"execution":{"iopub.status.busy":"2022-04-08T11:26:24.138005Z","iopub.execute_input":"2022-04-08T11:26:24.138351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Result","metadata":{}},{"cell_type":"code","source":"metrics = pd.read_csv(f'{trainer.logger.log_dir}/metrics.csv')\n\ntrain_acc = metrics['train_acc'].dropna().reset_index(drop=True)\nvalid_acc = metrics['valid_acc'].dropna().reset_index(drop=True)\n    \nfig = plt.figure(figsize=(7, 6))\nplt.grid(True)\nplt.plot(train_acc, color=\"r\", marker=\"o\", label='train/acc')\nplt.plot(valid_acc, color=\"b\", marker=\"x\", label='valid/acc')\nplt.ylabel('Accuracy', fontsize=24)\nplt.xlabel('Epoch', fontsize=24)\nplt.legend(loc='lower right', fontsize=18)\nplt.savefig(f'{trainer.logger.log_dir}/acc.png')\n\ntrain_loss = metrics['train_loss'].dropna().reset_index(drop=True)\nvalid_loss = metrics['valid_loss'].dropna().reset_index(drop=True)\n\nfig = plt.figure(figsize=(7, 6))\nplt.grid(True)\nplt.plot(train_loss, color=\"r\", marker=\"o\", label='train/loss')\nplt.plot(valid_loss, color=\"b\", marker=\"x\", label='valid/loss')\nplt.ylabel('Loss', fontsize=24)\nplt.xlabel('Epoch', fontsize=24)\nplt.legend(loc='upper right', fontsize=18)\nplt.savefig(f'{trainer.logger.log_dir}/loss.png')\\\n\nlr = metrics['lr'].dropna().reset_index(drop=True)\n\nfig = plt.figure(figsize=(7, 6))\nplt.grid(True)\nplt.plot(lr, color=\"g\", marker=\"o\", label='learning rate')\nplt.ylabel('LR', fontsize=24)\nplt.xlabel('Epoch', fontsize=24)\nplt.legend(loc='upper right', fontsize=18)\nplt.savefig(f'{trainer.logger.log_dir}/lr.png')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"!ls f\"./logs/{CFG.model_name}/version_0/checkpoints/\"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model = CustomEffNet(model_name=CFG.model_name, pretrained=False)\ncheckpoint = f\"./logs/{CFG.model_name}/version_0/checkpoints/last.ckpt\"\nmodel.load_state_dict(torch.load(checkpoint)['state_dict'])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.read_csv(PATH + \"sample_submission.csv\")\nsub.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub[\"file_path\"] = sub[\"filename\"].apply(lambda image: TEST_DIR + image)\nsub[\"cultivar_index\"] = 0\nsub.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG == True:\n    sub = sub[:10]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = SorghumDataset(sub, get_transform('valid'))\ntest_loader = DataLoader(test_dataset, batch_size=CFG.batch_size, shuffle=False, num_workers=2)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\n\nmodel.cuda()\nmodel.eval()\n\npredictions = []\nfor batch in tqdm(test_loader):\n    image = batch['image'].cuda()\n    with torch.no_grad():\n        outputs = model(image)\n        preds = outputs.detach().cpu()\n        predictions.append(preds.argmax(1))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tmp = predictions[0]\nfor i in range(len(predictions) - 1):\n    tmp = torch.cat((tmp, predictions[i+1]))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = [unique_cultivars[pred] for pred in tmp]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.read_csv(PATH + \"sample_submission.csv\")\nsub[\"cultivar\"] = predictions\nsub.to_csv('submission.csv', index=False)\nsub.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}