{"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":"!pip install timm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-24T08:21:48.704272Z","iopub.execute_input":"2022-07-24T08:21:48.705163Z","iopub.status.idle":"2022-07-24T08:22:01.878361Z","shell.execute_reply.started":"2022-07-24T08:21:48.705030Z","shell.execute_reply":"2022-07-24T08:22:01.876995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport pytorch_lightning as pl\nimport torchmetrics\nimport timm\n\nfrom torch.utils.data import DataLoader\n\nimport os\nimport sys\nimport glob\nimport pathlib\n\nimport re\nimport math\nimport random\nimport numpy as np\nimport pandas as pd\n\nfrom typing import Callable, Dict, Optional, Tuple\n\nfrom tqdm.notebook import tqdm\n\nimport matplotlib.pyplot as plt\nplt.style.use(\"ggplot\")\n\nimport seaborn as sns\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations.core.composition import Compose, OneOf\n\nimport cv2\n\nprint(f'Pytorch version: {torch.__version__}')\nprint(f'PyTorch Lightning version: {pl.__version__}')\nprint(f'Albumentations version: {A.__version__}')\nprint(f'Timm version: {timm.__version__}')\nprint(f'Python version: P{sys.version}')","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:37.691841Z","iopub.execute_input":"2022-07-24T08:22:37.692306Z","iopub.status.idle":"2022-07-24T08:22:47.689742Z","shell.execute_reply.started":"2022-07-24T08:22:37.692266Z","shell.execute_reply":"2022-07-24T08:22:47.687824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:47.691797Z","iopub.execute_input":"2022-07-24T08:22:47.693756Z","iopub.status.idle":"2022-07-24T08:22:47.698597Z","shell.execute_reply.started":"2022-07-24T08:22:47.693724Z","shell.execute_reply":"2022-07-24T08:22:47.696844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"class CFG:\n    model_name = 'tf_efficientnet_b3_ns'\n    pretrained = True\n    num_classes = 0\n\n    image_size = 512\n    crop_size = 0.9\n    fold = 0\n    n_splits = 5\n    \n    num_epochs = 25\n    batch_size = 32\n    \n    embedding_size = 128\n    s = 30\n    m = 0.5\n    \n    lr = 1e-4\n    max_lr = 2e-3\n    weight_decay = 1e-6\n    \n    precision = 16\n    num_workers = 4\n    seed = 42\n    \n    steps_per_epoch = 1\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:47.700279Z","iopub.execute_input":"2022-07-24T08:22:47.700721Z","iopub.status.idle":"2022-07-24T08:22:47.781316Z","shell.execute_reply.started":"2022-07-24T08:22:47.700680Z","shell.execute_reply":"2022-07-24T08:22:47.780137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CSV Preparation","metadata":{}},{"cell_type":"code","source":"TRAIN_DIR = \"../input/paddy-disease-classification/train_images/\"\n\ndef process(data):\n    path=pathlib.Path(data)\n    filepaths=list(path.glob(r\"*/*.jpg\"))\n    labels=list(map(lambda x: os.path.split(os.path.split(x)[0])[1],filepaths))\n    df1=pd.Series(filepaths,name='file_path').astype(str)\n    df2=pd.Series(labels,name='label')\n    df=pd.concat([df1,df2],axis=1)\n    \n    return df","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:47.784582Z","iopub.execute_input":"2022-07-24T08:22:47.785108Z","iopub.status.idle":"2022-07-24T08:22:47.792806Z","shell.execute_reply.started":"2022-07-24T08:22:47.784912Z","shell.execute_reply":"2022-07-24T08:22:47.791635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = process(TRAIN_DIR)\ndf['image_id'] = df['file_path'].apply(lambda image:image.split('/')[-1])\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:47.793992Z","iopub.execute_input":"2022-07-24T08:22:47.794883Z","iopub.status.idle":"2022-07-24T08:22:48.769619Z","shell.execute_reply.started":"2022-07-24T08:22:47.794841Z","shell.execute_reply":"2022-07-24T08:22:48.768448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_classes  = df['label'].nunique()\nprint(f'num_classes : {num_classes}')\n\nCFG.num_classes = num_classes","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:48.774841Z","iopub.execute_input":"2022-07-24T08:22:48.778228Z","iopub.status.idle":"2022-07-24T08:22:48.793154Z","shell.execute_reply.started":"2022-07-24T08:22:48.778183Z","shell.execute_reply":"2022-07-24T08:22:48.791664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import StratifiedKFold\n\nencoder = LabelEncoder()\nlabel2index = {l: i for (i, l) in enumerate(encoder.fit(df[\"label\"]).classes_)}\nindex2label = {x[1]: x[0] for x in label2index.items()}\n\ndf[\"label_index\"] = encoder.fit_transform(df[\"label\"])\n\nskf = StratifiedKFold(n_splits=CFG.n_splits, shuffle=True, random_state=CFG.seed)\nfor fold, (_, val_) in enumerate(skf.split(X=df, y=df.label)):\n    df.loc[val_, \"fold\"] = fold\n\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:48.798379Z","iopub.execute_input":"2022-07-24T08:22:48.799142Z","iopub.status.idle":"2022-07-24T08:22:48.877708Z","shell.execute_reply.started":"2022-07-24T08:22:48.799095Z","shell.execute_reply":"2022-07-24T08:22:48.876679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TEST_DIR = \"../input/paddy-disease-classification/test_images/\" \n\ntest_df = pd.read_csv(\"../input/paddy-disease-classification/sample_submission.csv\")\ntest_df[\"file_path\"] = test_df[\"image_id\"].apply(lambda image: TEST_DIR + image)\ntest_df['label'] = -1\ntest_df[\"label_index\"] = -1\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:48.882554Z","iopub.execute_input":"2022-07-24T08:22:48.885333Z","iopub.status.idle":"2022-07-24T08:22:48.916312Z","shell.execute_reply.started":"2022-07-24T08:22:48.885293Z","shell.execute_reply":"2022-07-24T08:22:48.914331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tmp = np.sqrt(1 / np.sqrt(df[\"label_index\"].value_counts().sort_index().values))\ndf_margins = (tmp - tmp.min()) / (tmp.max() - tmp.min()) * 0.3 + 0.05","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:48.917719Z","iopub.execute_input":"2022-07-24T08:22:48.918132Z","iopub.status.idle":"2022-07-24T08:22:48.931352Z","shell.execute_reply.started":"2022-07-24T08:22:48.918071Z","shell.execute_reply":"2022-07-24T08:22:48.930123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"class VisionDataset(torch.utils.data.Dataset):\n    def __init__(self, df_, transforms=None, test=False):\n        self.image_path = df_['file_path'].values\n        self.labels = df_[\"label_index\"].values\n        self.ids = df_['image_id']\n        self.transforms = transforms\n        self.test = test\n\n    def __len__(self):\n        return len(self.image_path)\n    \n    def resize(self, image, interp):\n        return  cv2.resize(image, (CFG.image_size, CFG.image_size), interpolation=interp)\n\n    def __getitem__(self, index: int) -> Dict[str, torch.Tensor]:\n        label = self.labels[index]\n        image_id = self.ids[index]\n        \n        image_path = self.image_path[index]\n        image = cv2.imread(image_path, 1)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        image = self.resize(image, cv2.INTER_AREA)\n        \n        if self.transforms is not None:\n            result = self.transforms(image=image)\n            image = result['image']\n        \n        if self.test:\n            return {'image':image, 'target': image_id}\n        else:\n            return {'image':image, 'target': label}","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:48.936257Z","iopub.execute_input":"2022-07-24T08:22:48.936820Z","iopub.status.idle":"2022-07-24T08:22:48.953596Z","shell.execute_reply.started":"2022-07-24T08:22:48.936785Z","shell.execute_reply":"2022-07-24T08:22:48.951805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Lightning DataModule","metadata":{}},{"cell_type":"code","source":"class LitDataModule(pl.LightningDataModule):\n    def __init__(self):\n        super().__init__()\n        self.df = df\n        self.test_df = test_df\n        self.train_transforms, self.valid_transform = self.init_transforms()\n        \n\n    def init_transforms(self):\n        crop_size = round(CFG.image_size*CFG.crop_size)\n        train_transforms = Compose([\n            A.RandomCrop(height=crop_size, width=crop_size, always_apply=True),\n            A.Affine(rotate=(-15, 15), translate_percent=(0.0, 0.25), shear=(-3, 3), p=0.7),\n            A.Cutout(max_h_size=int(crop_size * 0.4), max_w_size=int(crop_size * 0.4), num_holes=1, p=0.5),\n            A.RandomGridShuffle(grid=(2, 2), p=0.3),\n            A.GaussianBlur(blur_limit=(3, 7), p=0.05),\n            A.RandomSnow(p=0.05),\n            A.RandomRain(p=0.05),\n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225],),\n            ToTensorV2(),\n        ])\n        \n        valid_transform = Compose([\n            A.CenterCrop(height=crop_size, width=crop_size, always_apply=True),\n            A.Normalize(mean=[0.485, 0.456, 0.406],std=[0.229, 0.224, 0.225],),\n            ToTensorV2(),\n        ])\n            \n        return train_transforms, valid_transform\n        \n\n    def setup(self, stage: Optional[str] = None):\n        \n        if stage == \"fit\" or stage is None:\n            train_df = self.df[self.df.fold%CFG.n_splits != CFG.fold].reset_index(drop=True)\n            val_df = self.df[self.df.fold%CFG.n_splits == CFG.fold].reset_index(drop=True)\n        \n            self.train_dataset = VisionDataset(train_df, transforms=self.train_transforms)\n            self.valid_dataset = VisionDataset(val_df, transforms=self.valid_transform)\n            \n            self.test_dataset = VisionDataset(self.test_df, transforms=self.valid_transform, test=True)\n   \n\n    def train_dataloader(self) -> DataLoader:\n        return self.dataloader(self.train_dataset, train=True)\n\n\n    def val_dataloader(self) -> DataLoader:\n        return self.dataloader(self.valid_dataset)\n    \n    def test_dataloader(self) -> DataLoader:\n        return self.dataloader(self.test_dataset)\n    \n\n    def dataloader(self, dataset: VisionDataset, train: bool = False) -> DataLoader:\n        return DataLoader(\n            dataset,\n            batch_size=CFG.batch_size,\n            shuffle=train,\n            num_workers=CFG.num_workers,\n            pin_memory=True,\n            drop_last=train,\n        )","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:48.955555Z","iopub.execute_input":"2022-07-24T08:22:48.956336Z","iopub.status.idle":"2022-07-24T08:22:48.985446Z","shell.execute_reply.started":"2022-07-24T08:22:48.956299Z","shell.execute_reply":"2022-07-24T08:22:48.984390Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## ArcMarginProduct","metadata":{}},{"cell_type":"code","source":"class ArcMarginProduct(nn.Module):\n    r\"\"\"Implement of large margin arc distance: :\n    Args:\n        in_features: size of each input sample\n        out_features: size of each output sample\n        s: norm of input feature\n        m: margin\n        cos(theta + m)\n    \"\"\"\n\n    def __init__(\n        self,\n        in_features: int,\n        out_features: int,\n        s: float,\n        m: float,\n        easy_margin: bool,\n        ls_eps: float,\n    ):\n        super(ArcMarginProduct, self).__init__()\n        self.in_features = in_features\n        self.out_features = out_features\n        self.s = s\n        self.m = m\n        self.ls_eps = ls_eps  # label smoothing\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n\n        self.easy_margin = easy_margin\n        self.cos_m = math.cos(m)\n        self.sin_m = math.sin(m)\n        self.th = math.cos(math.pi - m)\n        self.mm = math.sin(math.pi - m) * m\n\n    def forward(self, input: torch.Tensor, label: torch.Tensor, device: str = \"cuda\") -> torch.Tensor:\n        # --------------------------- cos(theta) & phi(theta) ---------------------\n        cosine = F.linear(F.normalize(input), F.normalize(self.weight))\n        # Enable 16 bit precision\n        cosine = cosine.to(torch.float32)\n\n        sine = torch.sqrt(1.0 - torch.pow(cosine, 2))\n        phi = cosine * self.cos_m - sine * self.sin_m\n        if self.easy_margin:\n            phi = torch.where(cosine > 0, phi, cosine)\n        else:\n            phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n        # --------------------------- convert label to one-hot ---------------------\n        # one_hot = torch.zeros(cosine.size(), requires_grad=True, device='cuda')\n        one_hot = torch.zeros(cosine.size(), device=device)\n        one_hot.scatter_(1, label.view(-1, 1).long(), 1)\n        if self.ls_eps > 0:\n            one_hot = (1 - self.ls_eps) * one_hot + self.ls_eps / self.out_features\n        # -------------torch.where(out_i = {x_i if condition_i else y_i) ------------\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        output *= self.s\n\n        return output","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:48.987254Z","iopub.execute_input":"2022-07-24T08:22:48.987741Z","iopub.status.idle":"2022-07-24T08:22:49.011715Z","shell.execute_reply.started":"2022-07-24T08:22:48.987705Z","shell.execute_reply":"2022-07-24T08:22:49.010456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DenseCrossEntropy(nn.Module):\n    def __init__(self):\n        super(DenseCrossEntropy, self).__init__()\n        \n    def forward(self, x, target):\n        x = x.float()\n        target = target.float()\n        logprobs = torch.nn.functional.log_softmax(x, dim=-1)\n\n        loss = -logprobs * target\n        loss = loss.sum(-1)\n        return loss.mean()\n\nclass ArcMarginProduct_subcenter(nn.Module):\n    def __init__(self, in_features, out_features, k=1):\n        super().__init__()\n        self.weight = nn.Parameter(torch.FloatTensor(out_features*k, in_features))\n        self.reset_parameters()\n        self.k = k\n        self.out_features = out_features\n        \n    def reset_parameters(self):\n        stdv = 1. / math.sqrt(self.weight.size(1))\n        self.weight.data.uniform_(-stdv, stdv)\n        \n    def forward(self, features):\n        cosine_all = F.linear(F.normalize(features), F.normalize(self.weight))\n        cosine_all = cosine_all.view(-1, self.out_features, self.k)\n        cosine, _ = torch.max(cosine_all, dim=2)\n        return cosine   \n    \n\nclass ArcFaceLossAdaptiveMargin(nn.modules.Module):\n    def __init__(self, margins, n_classes, s=30.0):\n        super().__init__()\n        self.crit = DenseCrossEntropy()\n        self.s = s\n        self.margins = margins\n        self.out_dim = n_classes\n            \n    def forward(self, logits, labels):\n        ms = []\n        ms = self.margins[labels.cpu().numpy()]\n        cos_m = torch.from_numpy(np.cos(ms)).float().cuda()\n        sin_m = torch.from_numpy(np.sin(ms)).float().cuda()\n        th = torch.from_numpy(np.cos(math.pi - ms)).float().cuda()\n        mm = torch.from_numpy(np.sin(math.pi - ms) * ms).float().cuda()\n        labels = F.one_hot(labels, self.out_dim).float()\n        logits = logits.float()\n        cosine = logits\n        sine = torch.sqrt(1.0 - torch.pow(cosine, 2))\n        phi = cosine * cos_m.view(-1,1) - sine * sin_m.view(-1,1)\n        phi = torch.where(cosine > th.view(-1,1), phi, cosine - mm.view(-1,1))\n        output = (labels * phi) + ((1.0 - labels) * cosine)\n        output *= self.s\n        loss = self.crit(output, labels)\n        return loss","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:49.013436Z","iopub.execute_input":"2022-07-24T08:22:49.014568Z","iopub.status.idle":"2022-07-24T08:22:49.042792Z","shell.execute_reply.started":"2022-07-24T08:22:49.014523Z","shell.execute_reply":"2022-07-24T08:22:49.041585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Lightning Module","metadata":{}},{"cell_type":"code","source":"class LitModule(pl.LightningModule):\n    def __init__(self, model_name=CFG.model_name, pretrained=True):\n        super(LitModule, self).__init__()\n        \n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        in_features = self.model.classifier.in_features\n        self.model.classifier = nn.Identity()\n        self.model.global_pool = nn.Identity()\n        \n        self.multiple_dropout = [nn.Dropout(0.25) for i in range(8)]\n        self.embedding = nn.Linear(in_features * 2, CFG.embedding_size)\n        #self.fc = ArcMarginProduct(\n        #    in_features=CFG.embedding_size,\n        #    out_features=CFG.num_classes,\n        #    s=CFG.s,\n        #    m=CFG.m,\n        #    easy_margin=False,\n        #    ls_eps=0,\n        #)\n        \n        self.fc = ArcMarginProduct_subcenter(\n            in_features=CFG.embedding_size, \n            out_features=CFG.num_classes,)\n\n        self.accuracy = torchmetrics.Accuracy(num_classes=CFG.num_classes)\n    \n        #self.criterion = nn.CrossEntropyLoss()\n        self.criterion = ArcFaceLossAdaptiveMargin(\n            margins=df_margins,\n            n_classes=CFG.num_classes,\n            s=CFG.s\n        )\n        self.lr = CFG.lr\n        self.weight_decay = CFG.weight_decay\n        \n    \n    def forward(self, images: torch.Tensor) -> torch.Tensor:\n        features = self.model(images)\n        pooled_features_avg = nn.AdaptiveAvgPool2d(output_size=(1, 1))(features)\n        pooled_features_max = nn.AdaptiveMaxPool2d(output_size=(1, 1))(features)\n        pooled_features = torch.cat([pooled_features_avg, pooled_features_max], dim=1).flatten(1)\n        pooled_features_dropout = torch.zeros((pooled_features.shape), device=self.device)\n        for i in range(8):\n            pooled_features_dropout += self.multiple_dropout[i](pooled_features)\n        pooled_features_dropout /= 8\n        embeddings = self.embedding(pooled_features_dropout)\n        return embeddings\n\n    def configure_optimizers(self):\n        self.optimizer = torch.optim.Adam(self.model.parameters(), lr=self.lr, weight_decay=self.weight_decay)\n        self.scheduler = torch.optim.lr_scheduler.OneCycleLR(self.optimizer, \n                                                             epochs=CFG.num_epochs, \n                                                             steps_per_epoch=CFG.steps_per_epoch,\n                                                             max_lr=CFG.max_lr, \n                                                             pct_start=0.2, \n                                                             div_factor=1.0e+3, \n                                                             final_div_factor=1.0e+3)\n        scheduler = {'scheduler': self.scheduler, 'interval': 'step',}\n        return [self.optimizer], [scheduler]\n\n    def training_step(self, batch: Dict[str, torch.Tensor], batch_idx: int) -> torch.Tensor:\n        images = batch['image']\n        targets = batch['target'].long()\n        \n        embeddings = self(images)\n        # For ArcMarginProduct\n        #outputs = self.fc(embeddings, targets, self.device)\n        outputs = self.fc(embeddings)\n        \n        # For ArcMarginProduct\n        #loss = self.criterion(outputs)\n        loss = self.criterion(outputs, targets)\n        accuracy = self.accuracy(outputs.argmax(1), targets)\n        \n        logs = {'train_loss': loss, 'train_acc': accuracy, '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: Dict[str, torch.Tensor], batch_idx: int) -> torch.Tensor:\n        images = batch['image']\n        targets = batch['target'].long()\n        \n        embeddings = self(images)\n        # For ArcMarginProduct\n        #outputs = self.fc(embeddings, targets, self.device)\n        outputs = self.fc(embeddings)\n        \n        # For ArcMarginProduct\n        #loss = self.criterion(outputs)\n        loss = self.criterion(outputs, targets)\n        accuracy = self.accuracy(outputs.argmax(1), targets)\n        \n        logs = {'valid_loss': loss, 'valid_acc': accuracy}\n        self.log_dict(\n            logs,\n            on_step=False, on_epoch=True, prog_bar=True, logger=True\n        )\n        return loss\n    \n    @classmethod\n    def load_eval_checkpoint(cls, checkpoint_path, device):\n        module = cls.load_from_checkpoint(checkpoint_path=checkpoint_path).to(device)\n        module.eval()\n\n        return module","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:49.044623Z","iopub.execute_input":"2022-07-24T08:22:49.045608Z","iopub.status.idle":"2022-07-24T08:22:49.081861Z","shell.execute_reply.started":"2022-07-24T08:22:49.045571Z","shell.execute_reply":"2022-07-24T08:22:49.080899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{}},{"cell_type":"code","source":"def train():\n    pl.seed_everything(CFG.seed)\n    \n    datamodule = LitDataModule()\n    datamodule.setup()\n    CFG.steps_per_epoch = len(datamodule.train_dataloader())\n    \n    module = LitModule()\n    \n    logger = pl.loggers.CSVLogger(save_dir='logs/', name=CFG.model_name)\n    logger.log_hyperparams(CFG.__dict__)\n\n    checkpoint_callback = pl.callbacks.ModelCheckpoint(\n        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    )\n    \n    trainer = pl.Trainer(\n        max_epochs=CFG.num_epochs,\n        gpus=[0],\n        accumulate_grad_batches=1,\n        precision=CFG.precision,\n        callbacks=[checkpoint_callback],\n        logger=logger,\n        weights_summary='top',\n    )\n    \n    trainer.fit(module, datamodule=datamodule)\n    \n    return trainer","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:56.776848Z","iopub.execute_input":"2022-07-24T08:22:56.777427Z","iopub.status.idle":"2022-07-24T08:22:56.786112Z","shell.execute_reply.started":"2022-07-24T08:22:56.777390Z","shell.execute_reply":"2022-07-24T08:22:56.784578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = train()","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:22:57.558329Z","iopub.execute_input":"2022-07-24T08:22:57.559291Z","iopub.status.idle":"2022-07-24T08:27:41.116370Z","shell.execute_reply.started":"2022-07-24T08:22:57.559247Z","shell.execute_reply":"2022-07-24T08:27:41.114871Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train Results","metadata":{}},{"cell_type":"code","source":"metrics = pd.read_csv(f'{trainer.logger.log_dir}/metrics.csv')\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\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\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":{"execution":{"iopub.status.busy":"2022-07-24T08:27:42.990327Z","iopub.execute_input":"2022-07-24T08:27:42.990803Z","iopub.status.idle":"2022-07-24T08:27:44.618587Z","shell.execute_reply.started":"2022-07-24T08:27:42.990754Z","shell.execute_reply":"2022-07-24T08:27:44.617678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Prediction","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef infer(path):\n    pl.seed_everything(CFG.seed)\n    \n    datamodule = LitDataModule()\n    datamodule.setup()\n    test_dataloader = datamodule.test_dataloader()\n    \n    module = LitModule(model_name=CFG.model_name).load_eval_checkpoint(path, device=CFG.device)\n    \n    image_ids = []\n    prediction_list = []\n    \n    for batch in tqdm(test_dataloader):\n        images = batch[\"image\"].to(CFG.device, dtype=torch.float)\n        ids = batch['target']\n        embeddings = module(images)\n        outputs = nn.Softmax(dim=1)(CFG.s * F.linear(F.normalize(embeddings), F.normalize(module.fc.weight)))\n        preds = outputs.argmax(1).cpu().numpy()\n        for pred in preds:\n            prediction_list.append(index2label[pred])\n        image_ids.extend(ids)\n        \n    prediction_df = pd.DataFrame({'image_id': image_ids, 'label': prediction_list})\n    prediction_df.to_csv('submission.csv', index=False)\n\n    return prediction_df","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:27:46.961482Z","iopub.execute_input":"2022-07-24T08:27:46.961979Z","iopub.status.idle":"2022-07-24T08:27:46.979377Z","shell.execute_reply.started":"2022-07-24T08:27:46.961918Z","shell.execute_reply":"2022-07-24T08:27:46.978425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path = trainer.checkpoint_callback.best_model_path\n#path = trainer.checkpoint_callback.last_model_path\n\nsubmission_df = infer(path)\nsubmission_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-24T08:27:48.538044Z","iopub.execute_input":"2022-07-24T08:27:48.538555Z","iopub.status.idle":"2022-07-24T08:28:00.710706Z","shell.execute_reply.started":"2022-07-24T08:27:48.538513Z","shell.execute_reply":"2022-07-24T08:28:00.708050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}