{"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":"import pandas as pd\nimport numpy as np\nimport cv2\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport albumentations as A\nimport pytorch_lightning as pl\nimport matplotlib.pyplot as plt\nimport torchvision\nimport torchmetrics\n\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom albumentations.core.composition import Compose, OneOf\nfrom albumentations.augmentations.transforms import CLAHE, GaussNoise, ISONoise\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\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.utils import shuffle\n\nfrom iterstrat.ml_stratifiers import MultilabelStratifiedKFold","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-06-15T20:58:56.533875Z","iopub.execute_input":"2021-06-15T20:58:56.534199Z","iopub.status.idle":"2021-06-15T20:59:01.085204Z","shell.execute_reply.started":"2021-06-15T20:58:56.534169Z","shell.execute_reply":"2021-06-15T20:59:01.084058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install timm","metadata":{"execution":{"iopub.status.busy":"2021-06-15T20:58:20.578241Z","iopub.execute_input":"2021-06-15T20:58:20.57855Z","iopub.status.idle":"2021-06-15T20:58:27.760009Z","shell.execute_reply.started":"2021-06-15T20:58:20.578521Z","shell.execute_reply":"2021-06-15T20:58:27.759099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install iterative-stratification","metadata":{"execution":{"iopub.status.busy":"2021-06-15T20:58:48.997737Z","iopub.execute_input":"2021-06-15T20:58:48.998096Z","iopub.status.idle":"2021-06-15T20:58:54.835929Z","shell.execute_reply.started":"2021-06-15T20:58:48.998063Z","shell.execute_reply":"2021-06-15T20:58:54.834889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    seed = 42\n    model_name = 'resnet50'\n    pretrained = True\n    img_size = 384\n    num_classes = 12\n    lr = 5e-4\n    min_lr = 1e-6\n    t_max = 20\n    num_epochs = 15\n    batch_size = 8\n    accum = 1\n    precision = 16\n    n_fold = 5\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2021-06-15T21:10:31.349316Z","iopub.execute_input":"2021-06-15T21:10:31.349726Z","iopub.status.idle":"2021-06-15T21:10:31.358268Z","shell.execute_reply.started":"2021-06-15T21:10:31.34969Z","shell.execute_reply":"2021-06-15T21:10:31.357444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    seed = 42\n    model_name = 'tf_efficientnet_b5_ns'\n    pretrained = True\n    img_size = 384\n    num_classes = 6\n    lr = 1e-4\n    max_lr = 1e-3\n    pct_start = 0.3\n    div_factor = 1.0e+3\n    final_div_factor = 1.0e+3\n    num_epochs = 20\n    batch_size = 16\n    accum = 1\n    precision = 16\n    n_fold = 5\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2021-06-15T21:10:35.916066Z","iopub.execute_input":"2021-06-15T21:10:35.916378Z","iopub.status.idle":"2021-06-15T21:10:35.92184Z","shell.execute_reply.started":"2021-06-15T21:10:35.916351Z","shell.execute_reply":"2021-06-15T21:10:35.920853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PATH = \"../input/plant-pathology-2021-fgvc8/\"\n\nTRAIN_DIR = \"../input/resized-plant2021/img_sz_384/\"\nTEST_DIR = PATH + 'test_images/'","metadata":{"execution":{"iopub.status.busy":"2021-06-15T21:10:50.55398Z","iopub.execute_input":"2021-06-15T21:10:50.554315Z","iopub.status.idle":"2021-06-15T21:10:50.558531Z","shell.execute_reply.started":"2021-06-15T21:10:50.554285Z","shell.execute_reply":"2021-06-15T21:10:50.557551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_mult = pd.read_csv(PATH + \"train.csv\")\n\nfrom collections import defaultdict\n\n\ndct = defaultdict(list)\n\nfor i, label in enumerate(df_mult.labels):\n    for category in label.split():\n        dct[category].append(i)\n\ndct = {key: np.array(val) for key, val in dct.items()}\n\nnew_df = pd.DataFrame(np.zeros((df_mult.shape[0], len(dct.keys())), dtype=np.int8), columns=dct.keys())\n\nfor key, val in dct.items():\n    new_df.loc[val, key] = 1\n\nnew_df.head()","metadata":{"execution":{"iopub.status.busy":"2021-06-15T21:10:50.952339Z","iopub.execute_input":"2021-06-15T21:10:50.952639Z","iopub.status.idle":"2021-06-15T21:10:51.006824Z","shell.execute_reply.started":"2021-06-15T21:10:50.952609Z","shell.execute_reply":"2021-06-15T21:10:51.005977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_mult = pd.concat([df_mult, new_df], axis=1)\ndf_mult.to_csv('better_train.csv', index = False)\ndf_mult.head()\n","metadata":{"execution":{"iopub.status.busy":"2021-06-15T21:10:52.160068Z","iopub.execute_input":"2021-06-15T21:10:52.160402Z","iopub.status.idle":"2021-06-15T21:10:52.238429Z","shell.execute_reply.started":"2021-06-15T21:10:52.160371Z","shell.execute_reply":"2021-06-15T21:10:52.237727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"msss = MultilabelStratifiedKFold(n_splits=CFG.n_fold, shuffle=True, random_state=CFG.seed)\n\nfor train_idx, valid_idx in msss.split(df_mult['image'], df_mult.loc[:, list(df_mult.columns[2:].values)]):\n    df_train = df_mult.iloc[train_idx]\n    df_valid = df_mult.iloc[valid_idx]\n\nprint(f\"train size: {len(df_train)}\")\nprint(f\"valid size: {len(df_valid)}\")","metadata":{"execution":{"iopub.status.busy":"2021-06-15T21:10:52.394572Z","iopub.execute_input":"2021-06-15T21:10:52.394923Z","iopub.status.idle":"2021-06-15T21:10:53.020782Z","shell.execute_reply.started":"2021-06-15T21:10:52.394894Z","shell.execute_reply":"2021-06-15T21:10:53.019868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PlantDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.image_id = df['image'].values\n        self.labels = df.iloc[:, 2:].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 = TRAIN_DIR + image_id\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":"2021-06-15T21:10:56.087851Z","iopub.execute_input":"2021-06-15T21:10:56.088363Z","iopub.status.idle":"2021-06-15T21:10:56.103525Z","shell.execute_reply.started":"2021-06-15T21:10:56.08832Z","shell.execute_reply":"2021-06-15T21:10:56.098735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2021-06-15T21:10:57.435344Z","iopub.execute_input":"2021-06-15T21:10:57.435685Z","iopub.status.idle":"2021-06-15T21:10:57.445467Z","shell.execute_reply.started":"2021-06-15T21:10:57.435654Z","shell.execute_reply":"2021-06-15T21:10:57.444606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = PlantDataset(df_train, get_transform('train'))\nvalid_dataset = PlantDataset(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=4)\nvalid_loader = DataLoader(valid_dataset, batch_size=CFG.batch_size, shuffle=False, pin_memory=True, num_workers=4)","metadata":{"execution":{"iopub.status.busy":"2021-06-15T21:11:07.08511Z","iopub.execute_input":"2021-06-15T21:11:07.085484Z","iopub.status.idle":"2021-06-15T21:11:07.093242Z","shell.execute_reply.started":"2021-06-15T21:11:07.085455Z","shell.execute_reply":"2021-06-15T21:11:07.091421Z"},"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":"2021-06-15T21:11:10.168562Z","iopub.execute_input":"2021-06-15T21:11:10.168934Z","iopub.status.idle":"2021-06-15T21:11:10.174455Z","shell.execute_reply.started":"2021-06-15T21:11:10.168904Z","shell.execute_reply":"2021-06-15T21:11:10.173527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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":"2021-06-15T21:11:14.407169Z","iopub.execute_input":"2021-06-15T21:11:14.407492Z","iopub.status.idle":"2021-06-15T21:11:14.413848Z","shell.execute_reply.started":"2021-06-15T21:11:14.407462Z","shell.execute_reply":"2021-06-15T21:11:14.412704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LitCassava(pl.LightningModule):\n    def __init__(self, model):\n        super(LitCassava, self).__init__()\n        self.model = model\n        self.metric = torchmetrics.F1(CFG.num_classes, average='weighted')\n        self.criterion = nn.BCEWithLogitsLoss()\n        self.sigmoid = nn.Sigmoid()\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']\n        output = self.model(image)\n        loss = self.criterion(output, target)\n        score = self.metric(self.sigmoid(output), target.clone().detach().to(torch.int32))\n        logs = {'train_loss': loss, 'train_f1': 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']\n        output = self.model(image)\n        loss = self.criterion(output, target)\n        score = self.metric(self.sigmoid(output), target.clone().detach().to(torch.int32))\n        logs = {'valid_loss': loss, 'valid_f1': 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":"2021-06-15T21:11:15.417347Z","iopub.execute_input":"2021-06-15T21:11:15.417697Z","iopub.status.idle":"2021-06-15T21:11:15.430096Z","shell.execute_reply.started":"2021-06-15T21:11:15.417664Z","shell.execute_reply":"2021-06-15T21:11:15.429209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CustomEffNet(model_name=CFG.model_name, pretrained=CFG.pretrained)\nlit_model = LitCassava(model.model)","metadata":{"execution":{"iopub.status.busy":"2021-06-15T21:11:16.664551Z","iopub.execute_input":"2021-06-15T21:11:16.664907Z","iopub.status.idle":"2021-06-15T21:11:19.15167Z","shell.execute_reply.started":"2021-06-15T21:11:16.664877Z","shell.execute_reply":"2021-06-15T21:11:19.150883Z"},"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_f1',\n                                      save_top_k=1,\n                                      save_last=True,\n                                      save_weights_only=True,\n                                      filename='{epoch:02d}-{valid_loss:.4f}-{valid_f1:.4f}',\n                                      verbose=False,\n                                      mode='max')\n\ntrainer = Trainer(\n    max_epochs=CFG.num_epochs,\n    gpus=1,\n    accumulate_grad_batches=CFG.accum,\n    precision=CFG.precision,\n    checkpoint_callback=checkpoint_callback,\n    logger=logger,\n    weights_summary='top',\n    amp_backend='native',\n)","metadata":{"execution":{"iopub.status.busy":"2021-06-15T21:11:22.355353Z","iopub.execute_input":"2021-06-15T21:11:22.355696Z","iopub.status.idle":"2021-06-15T21:11:22.365415Z","shell.execute_reply.started":"2021-06-15T21:11:22.355665Z","shell.execute_reply":"2021-06-15T21:11:22.364707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.fit(lit_model, train_dataloader=train_loader, val_dataloaders=valid_loader)","metadata":{"execution":{"iopub.status.busy":"2021-06-15T22:16:21.949749Z","iopub.execute_input":"2021-06-15T22:16:21.950111Z","iopub.status.idle":"2021-06-16T00:54:18.368984Z","shell.execute_reply.started":"2021-06-15T22:16:21.950077Z","shell.execute_reply":"2021-06-16T00:54:18.367908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metrics = pd.read_csv(f'{trainer.logger.log_dir}/metrics.csv')\n\ntrain_acc = metrics['train_f1'].dropna().reset_index(drop=True)\nvalid_acc = metrics['valid_f1'].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/f1')\nplt.plot(valid_acc, color=\"b\", marker=\"x\", label='valid/f1')\nplt.ylabel('F1', fontsize=24)\nplt.xlabel('Epoch', fontsize=24)\nplt.legend(loc='lower right', fontsize=18)\nplt.savefig(f'{trainer.logger.log_dir}/f1.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')\n","metadata":{"execution":{"iopub.status.busy":"2021-06-16T00:56:09.821929Z","iopub.execute_input":"2021-06-16T00:56:09.82228Z","iopub.status.idle":"2021-06-16T00:56:10.544601Z","shell.execute_reply.started":"2021-06-16T00:56:09.822247Z","shell.execute_reply":"2021-06-16T00:56:10.543743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}