{"cells":[{"metadata":{},"cell_type":"markdown","source":"Simple training with efficientnet and pytorch_lightening\n\nHighlights:\n- offline efficientnet\n- easy to use pytorch_lightening module\n- focal loss\n- public score 0.862"},{"metadata":{},"cell_type":"markdown","source":"# Setup"},{"metadata":{"trusted":true},"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nfrom pylab import plt\n\nimport io\nimport itertools\nfrom copy import deepcopy\nimport sys, os\nfrom pathlib import Path\nfrom PIL import Image\n\n# %matplotlib notebook\n\ntrain_path = Path('/kaggle/input/cassava-leaf-disease-classification/train_images')\ntest_path = Path('/kaggle/input/cassava-leaf-disease-classification/test_images')\nlabels_path = Path('/kaggle/input/cassava-leaf-disease-classification/train.csv')\ntrain_df = pd.read_csv(labels_path)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"! pip install efficientnet_pytorch","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# for offline usage with my offline efficientnet_pytorch dataset\n# sys.path.insert(0, '../input/efficientnet-pytorch')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import efficientnet_pytorch","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Confusion matrix code:"},{"metadata":{"trusted":true},"cell_type":"code","source":"import itertools\nimport numpy as np\nimport matplotlib.pyplot as plt\n\ndef get_confusion_matrix(data, N, normalized=True):\n    confusion_matrix = np.zeros((N, N), dtype = np.int)\n\n    for entry in data:\n        for yi, yhi in zip(entry['y'], entry['y_hat']):\n            confusion_matrix[int(yi), int(yhi)] += 1\n    if normalized:\n        confusion_matrix = confusion_matrix / confusion_matrix.sum(axis=1, keepdims=True)\n    return confusion_matrix\n\ndef plot_confusion_matrix(cm):\n    \"\"\"\n    based on https://towardsdatascience.com/exploring-confusion-matrix-evolution-on-tensorboard-e66b39f4ac12\n    Returns a matplotlib figure containing the plotted confusion matrix.\n\n    Args:\n       cm (array, shape = [n, n]): a confusion matrix of integer classes\n       class_names (array, shape = [n]): String names of the integer classes\n    \"\"\"\n    plt.close('all')\n    figure = plt.figure(figsize=(6, 6))\n    if np.max(cm) <= 1:\n        vmin, vmax = 0, 1\n    else:\n        vmin, vmax = 0, np.max(cm)\n    plt.imshow(cm, interpolation='nearest', cmap=plt.cm.Blues, vmin=vmin, vmax=vmax);\n    plt.title(\"Confusion matrix\")\n    plt.colorbar()\n\n    # Round the confusion matrix.\n    cm = np.around(cm, decimals=2)\n\n    # Use white text if squares are dark; otherwise black.\n    threshold = cm.max() / 2.\n\n    for i, j in itertools.product(range(cm.shape[0]), range(cm.shape[1])):\n        color = \"white\" if cm[i, j] > threshold else \"black\"\n        plt.text(j, i, cm[i, j], horizontalalignment=\"center\", color=color)\n\n    plt.tight_layout()\n\n    return figure","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Focal loss code:"},{"metadata":{"trusted":true},"cell_type":"code","source":"\"\"\"from https://github.com/kornia/kornia\"\"\"\nfrom typing import Optional\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\ndef one_hot(labels: torch.Tensor,\n            num_classes: int,\n            device: Optional[torch.device] = None,\n            dtype: Optional[torch.dtype] = None,\n            eps: float = 1e-6) -> torch.Tensor:\n    r\"\"\"Converts an integer label x-D tensor to a one-hot (x+1)-D tensor.\n    Args:\n        labels (torch.Tensor) : tensor with labels of shape :math:`(N, *)`,\n                                where N is batch size. Each value is an integer\n                                representing correct classification.\n        num_classes (int): number of classes in labels.\n        device (Optional[torch.device]): the desired device of returned tensor.\n         Default: if None, uses the current device for the default tensor type\n         (see torch.set_default_tensor_type()). device will be the CPU for CPU\n         tensor types and the current CUDA device for CUDA tensor types.\n        dtype (Optional[torch.dtype]): the desired data type of returned\n         tensor. Default: if None, infers data type from values.\n    Returns:\n        torch.Tensor: the labels in one hot tensor of shape :math:`(N, C, *)`,\n    Examples:\n        >>> labels = torch.LongTensor([[[0, 1], [2, 0]]])\n        >>> one_hot(labels, num_classes=3)\n        tensor([[[[1.0000e+00, 1.0000e-06],\n                  [1.0000e-06, 1.0000e+00]],\n        <BLANKLINE>\n                 [[1.0000e-06, 1.0000e+00],\n                  [1.0000e-06, 1.0000e-06]],\n        <BLANKLINE>\n                 [[1.0000e-06, 1.0000e-06],\n                  [1.0000e+00, 1.0000e-06]]]])\n    \"\"\"\n    if not isinstance(labels, torch.Tensor):\n        raise TypeError(\"Input labels type is not a torch.Tensor. Got {}\"\n                        .format(type(labels)))\n\n    if not labels.dtype == torch.int64:\n        raise ValueError(\n            \"labels must be of the same dtype torch.int64. Got: {}\" .format(\n                labels.dtype))\n\n    if num_classes < 1:\n        raise ValueError(\"The number of classes must be bigger than one.\"\n                         \" Got: {}\".format(num_classes))\n\n    shape = labels.shape\n    one_hot = torch.zeros(\n        (shape[0], num_classes) + shape[1:], device=device, dtype=dtype\n    )\n\n    return one_hot.scatter_(1, labels.unsqueeze(1), 1.0) + eps\n\ndef focal_loss(\n        input: torch.Tensor,\n        target: torch.Tensor,\n        alpha: float,\n        gamma: float = 2.0,\n        reduction: str = 'none',\n        eps: float = 1e-8) -> torch.Tensor:\n    r\"\"\"Criterion that computes Focal loss.\n    According to :cite:`lin2018focal`, the Focal loss is computed as follows:\n    .. math::\n        \\text{FL}(p_t) = -\\alpha_t (1 - p_t)^{\\gamma} \\, \\text{log}(p_t)\n    Where:\n       - :math:`p_t` is the model's estimated probability for each class.\n    Args:\n        input (torch.Tensor): logits tensor with shape :math:`(N, C, *)` where C = number of classes.\n        target (torch.Tensor): labels tensor with shape :math:`(N, *)` where each value is :math:`0 ≤ targets[i] ≤ C−1`.\n        alpha (float): Weighting factor :math:`\\alpha \\in [0, 1]`.\n        gamma (float, optional): Focusing parameter :math:`\\gamma >= 0`. Default 2.\n        reduction (str, optional): Specifies the reduction to apply to the\n         output: ‘none’ | ‘mean’ | ‘sum’. ‘none’: no reduction will be applied,\n         ‘mean’: the sum of the output will be divided by the number of elements\n         in the output, ‘sum’: the output will be summed. Default: ‘none’.\n        eps (float, optional): Scalar to enforce numerical stabiliy. Default: 1e-8.\n    Return:\n        torch.Tensor: the computed loss.\n    Example:\n        >>> N = 5  # num_classes\n        >>> input = torch.randn(1, N, 3, 5, requires_grad=True)\n        >>> target = torch.empty(1, 3, 5, dtype=torch.long).random_(N)\n        >>> output = focal_loss(input, target, alpha=0.5, gamma=2.0, reduction='mean')\n        >>> output.backward()\n    \"\"\"\n    if not isinstance(input, torch.Tensor):\n        raise TypeError(\"Input type is not a torch.Tensor. Got {}\"\n                        .format(type(input)))\n\n    if not len(input.shape) >= 2:\n        raise ValueError(\"Invalid input shape, we expect BxCx*. Got: {}\"\n                         .format(input.shape))\n\n    if input.size(0) != target.size(0):\n        raise ValueError('Expected input batch_size ({}) to match target batch_size ({}).'\n                         .format(input.size(0), target.size(0)))\n\n    n = input.size(0)\n    out_size = (n,) + input.size()[2:]\n    if target.size()[1:] != input.size()[2:]:\n        raise ValueError('Expected target size {}, got {}'.format(\n            out_size, target.size()))\n\n    if not input.device == target.device:\n        raise ValueError(\n            \"input and target must be in the same device. Got: {} and {}\" .format(\n                input.device, target.device))\n\n    # compute softmax over the classes axis\n    input_soft: torch.Tensor = F.softmax(input, dim=1) + eps\n\n    # create the labels one hot tensor\n    target_one_hot: torch.Tensor = one_hot(\n        target, num_classes=input.shape[1],\n        device=input.device, dtype=input.dtype)\n\n    # compute the actual focal loss\n    weight = torch.pow(-input_soft + 1., gamma)\n\n    focal = -alpha * weight * torch.log(input_soft)\n    loss_tmp = torch.sum(target_one_hot * focal, dim=1)\n\n    if reduction == 'none':\n        loss = loss_tmp\n    elif reduction == 'mean':\n        loss = torch.mean(loss_tmp)\n    elif reduction == 'sum':\n        loss = torch.sum(loss_tmp)\n    else:\n        raise NotImplementedError(\"Invalid reduction mode: {}\"\n                                  .format(reduction))\n    return loss\n\n\nclass FocalLoss(nn.Module):\n    r\"\"\"Criterion that computes Focal loss.\n    According to :cite:`lin2018focal`, the Focal loss is computed as follows:\n    .. math::\n        \\text{FL}(p_t) = -\\alpha_t (1 - p_t)^{\\gamma} \\, \\text{log}(p_t)\n    Where:\n       - :math:`p_t` is the model's estimated probability for each class.\n    Args:\n        alpha (float): Weighting factor :math:`\\alpha \\in [0, 1]`.\n        gamma (float, optional): Focusing parameter :math:`\\gamma >= 0`. Default 2.\n        reduction (str, optional): Specifies the reduction to apply to the\n         output: ‘none’ | ‘mean’ | ‘sum’. ‘none’: no reduction will be applied,\n         ‘mean’: the sum of the output will be divided by the number of elements\n         in the output, ‘sum’: the output will be summed. Default: ‘none’.\n        eps (float, optional): Scalar to enforce numerical stabiliy. Default: 1e-8.\n    Shape:\n        - Input: :math:`(N, C, *)` where C = number of classes.\n        - Target: :math:`(N, *)` where each value is\n          :math:`0 ≤ targets[i] ≤ C−1`.\n    Example:\n        >>> N = 5  # num_classes\n        >>> kwargs = {\"alpha\": 0.5, \"gamma\": 2.0, \"reduction\": 'mean'}\n        >>> criterion = FocalLoss(**kwargs)\n        >>> input = torch.randn(1, N, 3, 5, requires_grad=True)\n        >>> target = torch.empty(1, 3, 5, dtype=torch.long).random_(N)\n        >>> output = criterion(input, target)\n        >>> output.backward()\n    \"\"\"\n\n    def __init__(self, alpha: float, gamma: float = 2.0,\n                 reduction: str = 'none', eps: float = 1e-8) -> None:\n        super(FocalLoss, self).__init__()\n        self.alpha: float = alpha\n        self.gamma: float = gamma\n        self.reduction: str = reduction\n        self.eps: float = eps\n\n    def forward(self, input: torch.Tensor, target: torch.Tensor) -> torch.Tensor:\n        return focal_loss(input, target, self.alpha, self.gamma, self.reduction, self.eps)\n\n\ndef binary_focal_loss_with_logits(\n        input: torch.Tensor,\n        target: torch.Tensor,\n        alpha: float = .25,\n        gamma: float = 2.0,\n        reduction: str = 'none',\n        eps: float = 1e-8) -> torch.Tensor:\n    r\"\"\"Function that computes Binary Focal loss.\n    .. math::\n        \\text{FL}(p_t) = -\\alpha_t (1 - p_t)^{\\gamma} \\, \\text{log}(p_t)\n    where:\n       - :math:`p_t` is the model's estimated probability for each class.\n    Args:\n        input (torch.Tensor): input data tensor with shape :math:`(N, 1, *)`.\n        target (torch.Tensor): the target tensor with shape :math:`(N, 1, *)`.\n        alpha (float): Weighting factor for the rare class :math:`\\alpha \\in [0, 1]`. Default: 0.25.\n        gamma (float): Focusing parameter :math:`\\gamma >= 0`. Default: 2.0.\n        reduction (str, optional): Specifies the reduction to apply to the. Default: 'none'.\n        eps (float): for numerically stability when dividing. Default: 1e-8.\n    Returns:\n        torch.tensor: the computed loss.\n    Examples:\n        >>> num_classes = 1\n        >>> kwargs = {\"alpha\": 0.25, \"gamma\": 2.0, \"reduction\": 'mean'}\n        >>> logits = torch.tensor([[[[6.325]]],[[[5.26]]],[[[87.49]]]])\n        >>> labels = torch.tensor([[[1.]],[[1.]],[[0.]]])\n        >>> binary_focal_loss_with_logits(logits, labels, **kwargs)\n        tensor(4.6052)\n    \"\"\"\n\n    if not isinstance(input, torch.Tensor):\n        raise TypeError(\"Input type is not a torch.Tensor. Got {}\"\n                        .format(type(input)))\n\n    if not len(input.shape) >= 2:\n        raise ValueError(\"Invalid input shape, we expect BxCx*. Got: {}\"\n                         .format(input.shape))\n\n    if input.size(0) != target.size(0):\n        raise ValueError('Expected input batch_size ({}) to match target batch_size ({}).'\n                         .format(input.size(0), target.size(0)))\n\n    probs = torch.sigmoid(input)\n    target = target.unsqueeze(dim=1)\n    loss_tmp = -alpha * torch.pow((1. - probs), gamma) * target * torch.log(probs + eps) \\\n               - (1 - alpha) * torch.pow(probs, gamma) * (1. - target) * torch.log(1. - probs + eps)\n    loss_tmp = loss_tmp.squeeze(dim=1)\n\n    if reduction == 'none':\n        loss = loss_tmp\n    elif reduction == 'mean':\n        loss = torch.mean(loss_tmp)\n    elif reduction == 'sum':\n        loss = torch.sum(loss_tmp)\n    else:\n        raise NotImplementedError(\"Invalid reduction mode: {}\"\n                                  .format(reduction))\n    return loss\n\n\nclass BinaryFocalLossWithLogits(nn.Module):\n    r\"\"\"Criterion that computes Focal loss.\n    According to :cite:`lin2017focal`, the Focal loss is computed as follows:\n    .. math::\n        \\text{FL}(p_t) = -\\alpha_t (1 - p_t)^{\\gamma} \\, \\text{log}(p_t)\n    where:\n       - :math:`p_t` is the model's estimated probability for each class.\n    Args:\n        alpha (float): Weighting factor for the rare class :math:`\\alpha \\in [0, 1]`.\n        gamma (float): Focusing parameter :math:`\\gamma >= 0`.\n        reduction (str, optional): Specifies the reduction to apply to the\n         output: ‘none’ | ‘mean’ | ‘sum’. ‘none’: no reduction will be applied,\n         ‘mean’: the sum of the output will be divided by the number of elements\n         in the output, ‘sum’: the output will be summed. Default: ‘none’.\n    Shape:\n        - Input: :math:`(N, 1, *)`.\n        - Target: :math:`(N, 1, *)`.\n    Examples:\n        >>> N = 1  # num_classes\n        >>> kwargs = {\"alpha\": 0.25, \"gamma\": 2.0, \"reduction\": 'mean'}\n        >>> loss = BinaryFocalLossWithLogits(**kwargs)\n        >>> input = torch.randn(1, N, 3, 5, requires_grad=True)\n        >>> target = torch.empty(1, 3, 5, dtype=torch.long).random_(N)\n        >>> output = loss(input, target)\n        >>> output.backward()\n    \"\"\"\n\n    def __init__(self, alpha: float, gamma: float = 2.0,\n                 reduction: str = 'none') -> None:\n        super(BinaryFocalLossWithLogits, self).__init__()\n        self.alpha: float = alpha\n        self.gamma: float = gamma\n        self.reduction: str = reduction\n        self.eps: float = 1e-8\n\n    def forward(self, input: torch.Tensor, target: torch.Tensor) -> torch.Tensor:\n        return binary_focal_loss_with_logits(\n            input, target, self.alpha, self.gamma, self.reduction, self.eps)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch\nfrom torch.utils.data import DataLoader\nfrom torchvision import transforms\nimport pytorch_lightning as pl","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Config"},{"metadata":{"trusted":true},"cell_type":"code","source":"N_classes = 5\neffnet_version ='5'\ncrop_size = 512 #512\nnet_input_size = 256","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Data preparation"},{"metadata":{"trusted":true},"cell_type":"code","source":"train_transforms = transforms.Compose([\n#             transforms.RandomAffine(0, shear=10, scale=(0.8,1.2)),\n            transforms.RandomCrop((crop_size,crop_size)),\n            transforms.Resize((net_input_size, net_input_size)),\n#             transforms.ColorJitter(brightness=0.2, contrast=0.15, saturation=0.1),\n            transforms.RandomHorizontalFlip()])\n\nval_transforms = transforms.Compose([\n#             transforms.RandomCrop((crop_size,crop_size)),\n            transforms.Resize((net_input_size, net_input_size))])\n\ntest_transforms = val_transforms","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Dataset class"},{"metadata":{"trusted":true},"cell_type":"code","source":"class Cassava(torch.utils.data.Dataset):\n\n    def __init__(self,\n                 img_dir,\n                 df,\n                 img_transforms = None,\n                 tensor_transforms = None, #transforms applied on Tensor\n                 output_tensor = True): #if False: image is returned\n        self.df = df\n        self.img_dir = Path(img_dir)\n        self.img_transforms = img_transforms\n        self.tensor_transforms = tensor_transforms\n        self.output_tensor = output_tensor\n        \n        self.classes = sorted(self.df.label.unique())\n\n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        img_fn, label = self.df.iloc[idx]\n        \n        img = Image.open(self.img_dir / img_fn)\n        \n        if self.img_transforms is not None:\n            img = self.img_transforms(img)\n        \n        if self.output_tensor:\n            img = transforms.functional.to_tensor(img)\n            if self.tensor_transforms is not None:\n                img = self.tensor_transforms(img)\n        return img, label","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Utility functions"},{"metadata":{"trusted":true},"cell_type":"code","source":"def accuracy(y_hat, y):\n    result = torch.mean((y_hat.detach() == y.detach()).float())\n    return result","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def log_figure(logger, fig, image_label='confusion matrix', epoch=0):\n    \"\"\"https://stackoverflow.com/questions/57316491/how-to-convert-matplotlib-figure-to-pil-image-object-without-saving-image\"\"\"\n#     pil_image = Image.frombytes('RGB', fig.canvas.get_width_height(),fig.canvas.tostring_rgb())\n    buf = io.BytesIO()\n    fig.savefig(buf, format='png', dpi = 300)\n    buf.seek(0)\n    pil_img = deepcopy(Image.open(buf))\n    buf.close()\n    \n    img = transforms.functional.to_tensor(pil_img)\n    logger.experiment.add_image(image_label, img, epoch)\n","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Data setup"},{"metadata":{"trusted":true},"cell_type":"code","source":"ds = Cassava(train_path, train_df)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"N = len(train_df)\nval_size = int(0.1 * N)\nmask_val = np.zeros(N, dtype='bool')\nmask_val[np.random.choice(np.arange(N), val_size, replace=False)] = True\nds_train = Cassava(train_path, train_df.loc[~mask_val], img_transforms=train_transforms)\nds_val = Cassava(train_path, train_df.loc[mask_val], img_transforms=val_transforms)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"ds_val.img_transforms = val_transforms\nds_train.img_transforms = train_transforms","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"ds_train.output_tensor = False\nimg, label = ds_train[np.random.randint(len(ds_train))]\nplt.imshow(img)\nplt.title(label)\nds_train.output_tensor = True","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## PL Module"},{"metadata":{"trusted":true},"cell_type":"code","source":"class MyModel(pl.LightningModule):\n\n    def __init__(self, model):\n        super().__init__()\n        self.model = model\n        self.batch_size = None\n\n    def forward(self, x):\n        y = self.model(x)\n        return y\n    \n    def training_setup(self, loss_fn=None, train_dataset=None, val_dataset=None,\n                      batch_size=None, lr=1e-3):\n        if loss_fn is not None:\n            self.loss_fn = loss_fn\n        if train_dataset is not None:\n            self.train_dataset = train_dataset\n        if val_dataset is not None:\n            self.val_dataset = val_dataset\n        if batch_size is not None:\n            self.batch_size = batch_size\n        if lr is not None:\n            self.lr = lr            \n            \n    def training_step(self, batch, batch_idx):\n        # training_step defined the train loop.\n        # It is independent of forward\n        x, y = batch\n        p = self.model(x)\n        loss = self.loss_fn(p, y)\n        \n        y_hat = torch.argmax(p.detach(), dim=1)\n#         # Logging to TensorBoard by default\n# pytorch-lightning 0.10\n        self.log('train_loss', loss)\n        self.log('train_accuracy', accuracy(y_hat, y))\n        return {'loss': loss, 'y_hat': y_hat.detach().cpu().numpy(), 'y': y.detach().cpu().numpy()}\n\n    def validation_step(self, batch, batch_idx):\n        return self._shared_eval(batch, batch_idx, 'val')\n\n    def test_step(self, batch, batch_idx):\n        return self._shared_eval(batch, batch_idx, 'test')\n\n    def _shared_eval(self, batch, batch_idx, prefix):\n        x, y = batch\n        p = self.model(x)\n        loss = self.loss_fn(p, y)\n        y_hat = torch.argmax(p.detach(), dim=1)\n\n         # Logging to TensorBoard by default\n        self.log_dict({f'{prefix}_loss': loss, f'{prefix}_accuracy': accuracy(y_hat, y)})\n        return {'loss': loss, 'y_hat': y_hat.detach().cpu().numpy(), 'y': y.detach().cpu().numpy()}        \n\n    def log_confusion_matrix(self, data, prefix, normalized=True):\n        confusion_matrix = get_confusion_matrix(data, N_classes, normalized=normalized)\n        confusion_matrix_fig = plot_confusion_matrix(confusion_matrix)\n        log_figure(self.logger, confusion_matrix_fig, f'{prefix}_confusion matrix', epoch=self.current_epoch)        \n    \n    def training_epoch_end(self, training_step_outputs):\n        self.log_confusion_matrix(training_step_outputs, 'train')\n    \n    def validation_epoch_end(self, validation_step_outputs):\n        self.log_confusion_matrix(validation_step_outputs, 'val')\n        \n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=self.lr)\n        scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.1)\n        return [optimizer], [scheduler]\n    \n    def train_dataloader(self):\n        return DataLoader(self.train_dataset, batch_size=self.batch_size, shuffle=True, num_workers=0)\n    \n    def val_dataloader(self):\n        return DataLoader(self.val_dataset, batch_size=self.batch_size, num_workers=0)\n","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Model setup"},{"metadata":{"trusted":true},"cell_type":"code","source":"net = efficientnet_pytorch.EfficientNet.from_pretrained(f'efficientnet-b{effnet_version}', num_classes=N_classes)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"mymodel = MyModel(net)\nmymodel.training_setup(loss_fn=FocalLoss(alpha=0.5, reduction='mean')) #torch.nn.CrossEntropyLoss())\nmymodel.training_setup(train_dataset=ds_train, val_dataset=ds_val)\nmymodel.training_setup(batch_size=4)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# most basic trainer, uses good defaults (auto-tensorboard, checkpoints, logs, and more)\n# trainer = pl.Trainer(gpus=8) (if you have GPUs)\nearly_stopping_callback = pl.callbacks.EarlyStopping('val_loss', patience=5)\ntrainer = pl.Trainer(callbacks=[early_stopping_callback],\n                     accumulate_grad_batches=4,\n#                       gradient_clip_val=0.5,\n                     max_epochs=200,\n#                      val_check_interval=0.25, #check_val_every_n_epoch=1\n                     gpus=1, auto_select_gpus=True,\n#                      auto_scale_batch_size='binsearch'\n                    ) # needs data loader in pl_module\n# trainer = Trainer(default_root_dir='/your/path/to/save/checkpoints')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# trainer.tune(mymodel)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Training"},{"metadata":{"trusted":true},"cell_type":"code","source":"trainer.fit(mymodel)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Test"},{"metadata":{"trusted":true},"cell_type":"code","source":"test_files = list(test_path.glob('*.jpg'))\ntrain_files = list(train_path.glob('*.jpg'))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaTest(torch.utils.data.Dataset):\n    def __init__(self, list_of_files, transforms=None):\n        self.list_of_files = list_of_files\n        self.transforms = transforms\n    def __len__(self):\n        return len(self.list_of_files)\n    def __getitem__(self, idx):\n        fn = self.list_of_files[idx]\n        x = Image.open(fn)\n        if self.transforms is not None:\n            for transform in self.transforms:\n                x = transform(x)\n        return x, idx","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_dataset = CassavaTest(test_files, [test_transforms, transforms.ToTensor()]) #train_files[:22]\ntest_dataloader = torch.utils.data.DataLoader(test_dataset, batch_size=4)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_dataset[0]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def predict(model, dataloader):\n    result = {'idx': [], 'label': []}\n    for x, idxs in dataloader:\n        p = model(x)\n        y_hat = torch.argmax(p.detach(), dim=1)\n        result['idx'].extend(list(idxs.detach().cpu().long().numpy()))\n        result['label'].extend(list(y_hat.cpu().numpy()))\n    return result","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_result = predict(mymodel, test_dataloader)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_result['image_id'] = [test_dataset.list_of_files[n].name for n in test_result['idx']]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"final_result = pd.DataFrame({'image_id': test_result['image_id'], 'label': test_result['label']})","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"final_result.to_csv('submission.csv', index=False)","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}