{"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":"%%capture\n!pip install deepspeed\n!pip install --upgrade wandb","metadata":{"execution":{"iopub.status.busy":"2021-07-30T10:28:05.410533Z","iopub.execute_input":"2021-07-30T10:28:05.410977Z","iopub.status.idle":"2021-07-30T10:28:33.465918Z","shell.execute_reply.started":"2021-07-30T10:28:05.410890Z","shell.execute_reply":"2021-07-30T10:28:33.464881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\nimport multiprocessing as mp\nfrom pathlib import Path\nfrom typing import Any, Callable, List, Tuple\n\nfrom deepspeed.ops.adam import FusedAdam\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport pytorch_lightning as pl\nfrom pytorch_lightning.loggers import WandbLogger\nfrom sklearn.model_selection import train_test_split\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import io, models, transforms\nimport torchvision.transforms.functional as TF\nfrom tqdm.auto import tqdm\n\n# Wandb login:\nfrom kaggle_secrets import UserSecretsClient\nimport wandb\nuser_secrets = UserSecretsClient()\nsecret_value = user_secrets.get_secret(\"wandb_api_key\")\nwandb.login(key=secret_value)\n\ndevice = torch.device(\"cuda\") if torch.cuda.is_available() else torch.device(\"cpu\")\n\n%matplotlib inline\nprint(torch.__version__, pl.__version__, wandb.__version__)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-07-30T10:30:15.091390Z","iopub.execute_input":"2021-07-30T10:30:15.091752Z","iopub.status.idle":"2021-07-30T10:30:18.715933Z","shell.execute_reply.started":"2021-07-30T10:30:15.091709Z","shell.execute_reply":"2021-07-30T10:30:18.714881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT_PATH = Path(\"/kaggle/input/cassava-leaf-disease-classification/\")\nIMAGE_SIZE = (128, 128)\nBATCH_SIZE = 64\nLR = 1e-3\nEPOCHS = 3","metadata":{"execution":{"iopub.status.busy":"2021-07-30T10:30:22.134646Z","iopub.execute_input":"2021-07-30T10:30:22.135048Z","iopub.status.idle":"2021-07-30T10:30:22.140033Z","shell.execute_reply.started":"2021-07-30T10:30:22.135012Z","shell.execute_reply":"2021-07-30T10:30:22.138988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(ROOT_PATH / \"train.csv\")\nwith open(ROOT_PATH / \"label_num_to_disease_map.json\", \"r\") as f:\n    label_map = json.load(f)\nlabel_map = {int(k): v for k, v in label_map.items()}\ntrain_df, valid_df = train_test_split(df, stratify=df[\"label\"].values)\n    \nplt.figure(figsize=(12, 5))\nprint(df[\"label\"].map(label_map).value_counts())\ndf.sample(5)","metadata":{"execution":{"iopub.status.busy":"2021-07-30T10:30:24.230355Z","iopub.execute_input":"2021-07-30T10:30:24.230703Z","iopub.status.idle":"2021-07-30T10:30:24.321312Z","shell.execute_reply.started":"2021-07-30T10:30:24.230649Z","shell.execute_reply":"2021-07-30T10:30:24.320314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Data(Dataset):\n    def __init__(self, df: pd.DataFrame, transforms=None):\n        self.files = [ROOT_PATH / \"train_images\" / file for file in df[\"image_id\"].values]\n        self.y = df[\"label\"].values.tolist()\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.y)\n    \n    def __getitem__(self, i):\n        img = Image.open(self.files[i])\n        label = self.y[i]\n        if self.transforms is not None:\n            img = self.transforms(img)\n            \n        return img, label","metadata":{"execution":{"iopub.status.busy":"2021-07-30T10:30:27.093553Z","iopub.execute_input":"2021-07-30T10:30:27.093953Z","iopub.status.idle":"2021-07-30T10:30:27.101617Z","shell.execute_reply.started":"2021-07-30T10:30:27.093919Z","shell.execute_reply":"2021-07-30T10:30:27.100700Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_tfms = transforms.Compose(\n    [\n        transforms.Resize(IMAGE_SIZE),\n        transforms.RandomHorizontalFlip(),\n        transforms.RandomRotation(10),\n        transforms.ToTensor(),\n        transforms.Normalize([0.4766, 0.4527, 0.3926], [0.2275, 0.2224, 0.2210])\n    ]\n)\n\nvalid_tfms = transforms.Compose(\n    [\n        transforms.Resize(IMAGE_SIZE),\n        transforms.ToTensor(),\n        transforms.Normalize([0.4766, 0.4527, 0.3926], [0.2275, 0.2224, 0.2210])\n    ]\n)\n\ntrain_ds = Data(train_df, train_tfms)\nvalid_ds = Data(valid_df, valid_tfms)\n\ntrain_dl = DataLoader(\n    train_ds,\n    BATCH_SIZE, \n    shuffle=True, \n    drop_last=True, \n    num_workers=4,\n    pin_memory=True,\n)\n\nvalid_dl = DataLoader(\n    valid_ds, \n    BATCH_SIZE*2, \n    shuffle=False, \n    drop_last=False, \n    num_workers=4,\n    pin_memory=True,\n)","metadata":{"execution":{"iopub.status.busy":"2021-07-30T10:30:28.924040Z","iopub.execute_input":"2021-07-30T10:30:28.924397Z","iopub.status.idle":"2021-07-30T10:30:29.273354Z","shell.execute_reply.started":"2021-07-30T10:30:28.924363Z","shell.execute_reply":"2021-07-30T10:30:29.272504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x, y = next(iter(train_dl))\nx.shape, y.shape","metadata":{"execution":{"iopub.status.busy":"2021-07-30T10:30:35.239225Z","iopub.execute_input":"2021-07-30T10:30:35.239564Z","iopub.status.idle":"2021-07-30T10:30:45.236049Z","shell.execute_reply.started":"2021-07-30T10:30:35.239532Z","shell.execute_reply":"2021-07-30T10:30:45.235140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self, num_classes: int):\n        super().__init__()\n        self.base = models.resnet34(pretrained=True)\n        self.linear1 = nn.Linear(self.base.fc.in_features, self.base.fc.in_features // 2)\n        self.linear2 = nn.Linear(self.base.fc.in_features // 2, num_classes)\n        self.norm1 = nn.BatchNorm1d(self.base.fc.in_features)\n        self.norm2 = nn.BatchNorm1d(self.base.fc.in_features // 2)\n        self.dropout1 = nn.Dropout(p=0.5)\n        self.dropout2 = nn.Dropout(p=0.5)\n        self.base.fc = nn.Identity()\n        \n        for p in self.base.parameters():\n            p.requires_grad = False\n        \n    def forward(self, x):\n        out1 = self.dropout1(self.norm1(F.leaky_relu(self.base(x))))\n        out2 = self.dropout2(self.norm2(F.leaky_relu(self.linear1(out1))))\n        out3 = self.linear2(out2)\n        \n        return out3","metadata":{"execution":{"iopub.status.busy":"2021-07-30T10:30:49.625197Z","iopub.execute_input":"2021-07-30T10:30:49.625563Z","iopub.status.idle":"2021-07-30T10:30:49.636048Z","shell.execute_reply.started":"2021-07-30T10:30:49.625529Z","shell.execute_reply":"2021-07-30T10:30:49.634805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LightningModel(pl.LightningModule):\n    def __init__(self, model: nn.Module, loss_fn: Callable, lr: float):\n        super().__init__()\n        self.model = model\n        self.loss_fn = loss_fn\n        self.lr = lr\n        \n    def common_step(self, batch):\n        x, y = batch\n        logits = self.model(x)\n        loss = self.loss_fn(logits, y)\n        accuracy = (logits.argmax(-1) == y).float().mean()\n\n        return loss, accuracy\n        \n    def training_step(self, batch: Tuple[torch.FloatTensor, torch.LongTensor], *args: List[Any]):\n        loss, accuracy = self.common_step(batch)\n        self.log(\"training_loss\", loss, on_step=True, on_epoch=True)\n        self.log(\"training_accuracy\", accuracy, on_step=True, on_epoch=True)\n        \n        return loss\n        \n    def on_epoch_end(self, *args):\n        if self.current_epoch == 0:\n            for p in self.model.base.parameters():\n                p.requires_grad = True\n        \n    def validation_step(self, batch: Tuple[torch.FloatTensor, torch.LongTensor], *args: List[Any]):\n        loss, accuracy = self.common_step(batch)\n        self.log(\"validation_loss\", loss, on_step=False, on_epoch=True)\n        self.log(\"validation_accuracy\", accuracy, on_step=False, on_epoch=True)\n        \n    def configure_optimizers(self):\n        return FusedAdam(self.model.parameters(), lr=self.lr)","metadata":{"execution":{"iopub.status.busy":"2021-07-30T10:46:36.331741Z","iopub.execute_input":"2021-07-30T10:46:36.332130Z","iopub.status.idle":"2021-07-30T10:46:36.344988Z","shell.execute_reply.started":"2021-07-30T10:46:36.332095Z","shell.execute_reply":"2021-07-30T10:46:36.343537Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"label_counts = train_df[\"label\"].value_counts().sort_index()\nclass_weights = max(label_counts) / label_counts.values\nlabel_counts, class_weights","metadata":{"execution":{"iopub.status.busy":"2021-07-30T10:46:40.231406Z","iopub.execute_input":"2021-07-30T10:46:40.231767Z","iopub.status.idle":"2021-07-30T10:46:40.250611Z","shell.execute_reply.started":"2021-07-30T10:46:40.231732Z","shell.execute_reply":"2021-07-30T10:46:40.249780Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir /kaggle/working/logs\nmodel = Model(df[\"label\"].nunique())\nloss_fn = nn.CrossEntropyLoss(weight=torch.FloatTensor(class_weights))\nlightning_model = LightningModel(model, loss_fn, LR)\n\nlogger = WandbLogger(\"cassava-1\", \"/kaggle/working/logs/\", project=\"Kaggle-Cassava\")\ntrainer = pl.Trainer(\n    max_epochs=EPOCHS,\n    gpus=torch.cuda.device_count(),\n    gradient_clip_val=1.0,\n    logger=logger,\n    precision=16,\n)\ntrainer.fit(lightning_model, train_dl, valid_dl)","metadata":{"execution":{"iopub.status.busy":"2021-07-30T10:46:54.722482Z","iopub.execute_input":"2021-07-30T10:46:54.722847Z","iopub.status.idle":"2021-07-30T10:51:41.848939Z","shell.execute_reply.started":"2021-07-30T10:46:54.722812Z","shell.execute_reply":"2021-07-30T10:51:41.847929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TestData(Dataset):\n    def __init__(self, transforms=None):\n        self.files = [path for path in (ROOT_PATH / \"test_images\").glob(\"*.jpg\")]\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.files)\n    \n    def __getitem__(self, i):\n        img = Image.open(self.files[i])\n        if self.transforms is not None:\n            img = self.transforms(img)\n            \n        return img\n    \ntest_ds = TestData(valid_tfms)\ntest_dl = DataLoader(\n    test_ds, \n    BATCH_SIZE*2, \n    shuffle=False, \n    drop_last=False, \n    num_workers=4,\n    pin_memory=True,\n)\n\nmodel = model.eval().to(device)\ny_preds = []\nwith torch.no_grad():\n    for x in tqdm(test_dl):\n        y_preds.append(model(x.to(device)).argmax(dim=-1).cpu())\n\ny_preds = torch.cat(y_preds).cpu().numpy()\nfile_names = [test_ds.files[i].name for i in range(len(test_ds))]\npd.DataFrame({\"image_id\": file_names, \"label\": y_preds}).to_csv(\"./submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2021-07-30T10:54:15.555444Z","iopub.execute_input":"2021-07-30T10:54:15.555812Z","iopub.status.idle":"2021-07-30T10:54:15.971785Z","shell.execute_reply.started":"2021-07-30T10:54:15.555775Z","shell.execute_reply":"2021-07-30T10:54:15.970769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y","metadata":{"execution":{"iopub.status.busy":"2021-07-30T09:58:46.309097Z","iopub.execute_input":"2021-07-30T09:58:46.309775Z","iopub.status.idle":"2021-07-30T09:58:46.384172Z","shell.execute_reply.started":"2021-07-30T09:58:46.309684Z","shell.execute_reply":"2021-07-30T09:58:46.382830Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}