{"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":"markdown","source":"# Herbarium 2022","metadata":{}},{"cell_type":"code","source":"# PyTorch/Lightning\nimport torch\nimport torchvision\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision.transforms as T\nimport pytorch_lightning as pl\nimport torchmetrics.functional as metrics\nfrom torch.utils.data import Dataset, DataLoader\n\n# Pandas/Numpy/etc\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\n\n# Other\nimport os\nimport json\nimport random\nfrom PIL import Image\nfrom pathlib import Path\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:03.911653Z","iopub.execute_input":"2022-03-12T18:57:03.912349Z","iopub.status.idle":"2022-03-12T18:57:07.574966Z","shell.execute_reply.started":"2022-03-12T18:57:03.912202Z","shell.execute_reply":"2022-03-12T18:57:07.573939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.__version__, torchvision.__version__, pl.__version__","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:07.580671Z","iopub.execute_input":"2022-03-12T18:57:07.581076Z","iopub.status.idle":"2022-03-12T18:57:07.596148Z","shell.execute_reply.started":"2022-03-12T18:57:07.581034Z","shell.execute_reply":"2022-03-12T18:57:07.595351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls ../input/herbarium-2022-fgvc9","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:07.600465Z","iopub.execute_input":"2022-03-12T18:57:07.602359Z","iopub.status.idle":"2022-03-12T18:57:08.285431Z","shell.execute_reply.started":"2022-03-12T18:57:07.602323Z","shell.execute_reply":"2022-03-12T18:57:08.284629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRAIN_JSON = '../input/herbarium-2022-fgvc9/train_metadata.json'\nTEST_JSON = '../input/herbarium-2022-fgvc9/test_metadata.json'\nTRAIN_IMGS = '../input/herbarium-2022-fgvc9/train_images/'\nTEST_IMGS = '../input/herbarium-2022-fgvc9/test_images/'","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:08.287749Z","iopub.execute_input":"2022-03-12T18:57:08.288032Z","iopub.status.idle":"2022-03-12T18:57:08.292414Z","shell.execute_reply.started":"2022-03-12T18:57:08.287996Z","shell.execute_reply":"2022-03-12T18:57:08.291711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils","metadata":{}},{"cell_type":"code","source":"def load_annotations(metadata):\n    '''\n    Args:\n        metadata (dict): JSON with annotations.\n    Returns:\n        dataframe (DataFrame): Dataframe with annotations. \n    '''\n    metadata_list = []\n    categories = {category['category_id']: (category['family'], category['genus'], category['species']) \n                  for category in metadata['categories']}\n    for img, anns in tqdm(zip(metadata['images'], metadata['annotations'])):\n        category_id = anns['category_id']\n        family, genus, species = categories[category_id]\n        row = {\n            'file_name': img['file_name'],\n            'img_id': img['image_id'],\n            'category_id': category_id,\n            'family': family,\n            'genus': genus,\n            'species': species\n        }\n        metadata_list.append(row)\n    return pd.DataFrame.from_dict(metadata_list)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:08.295455Z","iopub.execute_input":"2022-03-12T18:57:08.29606Z","iopub.status.idle":"2022-03-12T18:57:08.304404Z","shell.execute_reply.started":"2022-03-12T18:57:08.296021Z","shell.execute_reply":"2022-03-12T18:57:08.303761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EDA","metadata":{}},{"cell_type":"code","source":"with open(TRAIN_JSON) as f:\n    train_metadata = json.load(f)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:08.305762Z","iopub.execute_input":"2022-03-12T18:57:08.306042Z","iopub.status.idle":"2022-03-12T18:57:21.813333Z","shell.execute_reply.started":"2022-03-12T18:57:08.306008Z","shell.execute_reply":"2022-03-12T18:57:21.81261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# All metadata possible keys\nprint(f'metadata keys: {[*train_metadata]}\\n')\n\nfor key in train_metadata.keys():\n    print(f'{key}: \\n\\t{[*train_metadata[key][0]]} \\n\\tcount: {len(train_metadata[key])}')","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:21.814842Z","iopub.execute_input":"2022-03-12T18:57:21.815101Z","iopub.status.idle":"2022-03-12T18:57:21.823838Z","shell.execute_reply.started":"2022-03-12T18:57:21.815053Z","shell.execute_reply":"2022-03-12T18:57:21.82286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Example annotation\ntrain_metadata['annotations'][0]","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:21.825648Z","iopub.execute_input":"2022-03-12T18:57:21.825988Z","iopub.status.idle":"2022-03-12T18:57:21.833148Z","shell.execute_reply.started":"2022-03-12T18:57:21.825946Z","shell.execute_reply":"2022-03-12T18:57:21.832394Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Example category\ntrain_metadata['categories'][0]","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:21.834948Z","iopub.execute_input":"2022-03-12T18:57:21.835263Z","iopub.status.idle":"2022-03-12T18:57:21.842811Z","shell.execute_reply.started":"2022-03-12T18:57:21.83523Z","shell.execute_reply":"2022-03-12T18:57:21.841999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Loading all required train metadata into one dataframe \ntrain_df = load_annotations(train_metadata)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:21.844294Z","iopub.execute_input":"2022-03-12T18:57:21.844593Z","iopub.status.idle":"2022-03-12T18:57:24.495314Z","shell.execute_reply.started":"2022-03-12T18:57:21.844558Z","shell.execute_reply":"2022-03-12T18:57:24.494569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.sample(5)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:24.496726Z","iopub.execute_input":"2022-03-12T18:57:24.496977Z","iopub.status.idle":"2022-03-12T18:57:24.538514Z","shell.execute_reply.started":"2022-03-12T18:57:24.496942Z","shell.execute_reply":"2022-03-12T18:57:24.537722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Top K Distribution","metadata":{}},{"cell_type":"code","source":"columns = ['family', 'genus', 'species']\n\ndef plot_top_K_barh(metadata, column, K=10):\n    ax = metadata[column] \\\n           .value_counts() \\\n           .head(K) \\\n           .plot(title=f'Top {K} {column}', kind='barh')\n    for container in ax.containers:\n        ax.bar_label(container)\n    return ax","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:24.540003Z","iopub.execute_input":"2022-03-12T18:57:24.54025Z","iopub.status.idle":"2022-03-12T18:57:24.545243Z","shell.execute_reply.started":"2022-03-12T18:57:24.540217Z","shell.execute_reply":"2022-03-12T18:57:24.544536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure(figsize=(30, 8))\nfig.suptitle(f'Train data', fontsize=22)\nfor idx, column in enumerate(columns):\n    fig.add_subplot(1, 3, idx+1)\n    plot_top_K_barh(train_df, column)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:24.546662Z","iopub.execute_input":"2022-03-12T18:57:24.547223Z","iopub.status.idle":"2022-03-12T18:57:25.575682Z","shell.execute_reply.started":"2022-03-12T18:57:24.547181Z","shell.execute_reply":"2022-03-12T18:57:25.57501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Saving metadata as CSV","metadata":{}},{"cell_type":"code","source":"train_df.to_csv('train_metadata.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:25.579187Z","iopub.execute_input":"2022-03-12T18:57:25.579877Z","iopub.status.idle":"2022-03-12T18:57:28.90374Z","shell.execute_reply.started":"2022-03-12T18:57:25.579836Z","shell.execute_reply":"2022-03-12T18:57:28.903001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Image Visualization","metadata":{}},{"cell_type":"code","source":"def plot_images(metadata, img_dir, by=None, name=None):\n    '''\n    Args:\n        metadata (DataFrame): DataFrame with annotations.\n        img_dir (str): Path to the image directory.\n        by (str): Sample field (or randomly): [family, genus, species, None].\n        name (str): Name of the example, in cases of non-random sampling. \n    '''\n    \n    # Expected values: family, genus, species.\n    if by and name is not None:\n        metadata = metadata[metadata[by] == name]\n    \n    sample = metadata.sample(16)\n    filenames = sample['file_name'].to_list()\n    family = sample['family'].to_list()\n    genus = sample['genus'].to_list()\n    species = sample['species'].to_list()\n    \n    fig, axes = plt.subplots(4, 4, figsize=(12, 16))\n    title = f'{by} ({name})' if by is not None else \"random\"\n    fig.suptitle(f'Select by {title}\\n', fontsize=22)\n    for idx, ax in enumerate(axes.flatten()):\n        img_path = Path(img_dir).joinpath(filenames[idx])\n        img = np.array(Image.open(img_path))\n        ax.imshow(img)\n        ax.title.set_text(f'family: {family[idx]}\\n genus:' \\\n                          f'{genus[idx]}\\n species: {species[idx]}')\n        ax.set_axis_off()\n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:28.904995Z","iopub.execute_input":"2022-03-12T18:57:28.90524Z","iopub.status.idle":"2022-03-12T18:57:28.913562Z","shell.execute_reply.started":"2022-03-12T18:57:28.905209Z","shell.execute_reply":"2022-03-12T18:57:28.912876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_images(train_df, img_dir=TRAIN_IMGS)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:28.914799Z","iopub.execute_input":"2022-03-12T18:57:28.915242Z","iopub.status.idle":"2022-03-12T18:57:31.735768Z","shell.execute_reply.started":"2022-03-12T18:57:28.915206Z","shell.execute_reply":"2022-03-12T18:57:31.733455Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_images(train_df, img_dir=TRAIN_IMGS, by='family', name='Asteraceae')","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:31.736741Z","iopub.execute_input":"2022-03-12T18:57:31.736994Z","iopub.status.idle":"2022-03-12T18:57:34.188758Z","shell.execute_reply.started":"2022-03-12T18:57:31.736959Z","shell.execute_reply":"2022-03-12T18:57:34.185812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_images(train_df, img_dir=TRAIN_IMGS, by='genus', name='Carex')","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:34.190197Z","iopub.execute_input":"2022-03-12T18:57:34.190643Z","iopub.status.idle":"2022-03-12T18:57:36.742365Z","shell.execute_reply.started":"2022-03-12T18:57:34.190607Z","shell.execute_reply":"2022-03-12T18:57:36.741716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_images(train_df, img_dir=TRAIN_IMGS, by='species', name='californica')","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:36.743541Z","iopub.execute_input":"2022-03-12T18:57:36.743861Z","iopub.status.idle":"2022-03-12T18:57:39.094316Z","shell.execute_reply.started":"2022-03-12T18:57:36.743832Z","shell.execute_reply":"2022-03-12T18:57:39.093667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## PyTorch Dataset","metadata":{}},{"cell_type":"code","source":"class HerbariumDataset(Dataset):\n    def __init__(self, img_dir, metadata_csv, transform):\n        self.img_dir = img_dir\n        self.metadata = pd.read_csv(metadata_csv)\n        self.transform = transform\n    \n    def __getitem__(self, idx):\n        filename = self.metadata['file_name'][idx]\n        label = self.metadata['category_id'][idx]\n        \n        img_path = Path(self.img_dir).joinpath(filename)\n        img = Image.open(img_path)\n        img = self.transform(img)\n        \n        return img, label\n    \n    def __len__(self):\n        return len(self.metadata)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:39.095539Z","iopub.execute_input":"2022-03-12T18:57:39.095886Z","iopub.status.idle":"2022-03-12T18:57:39.103888Z","shell.execute_reply.started":"2022-03-12T18:57:39.095854Z","shell.execute_reply":"2022-03-12T18:57:39.103191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Init all necessary transforms\nRESIZE_H, RESIZE_W = 360, 360\ntrain_transform = T.Compose([\n    T.ToTensor(),\n    T.Resize((RESIZE_H, RESIZE_W)),\n    T.RandomHorizontalFlip(p=0.5),\n    T.RandomVerticalFlip(p=0.5)\n])\n\ntest_transform = T.Compose([\n    T.ToTensor(),\n    T.Resize((RESIZE_H, RESIZE_W))\n])\n\ntrain_dataset = HerbariumDataset(TRAIN_IMGS, 'train_metadata.csv', \n                                 transform=train_transform)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:39.105336Z","iopub.execute_input":"2022-03-12T18:57:39.10612Z","iopub.status.idle":"2022-03-12T18:57:40.348848Z","shell.execute_reply.started":"2022-03-12T18:57:39.106083Z","shell.execute_reply":"2022-03-12T18:57:40.347714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axes = plt.subplots(3, 3, figsize=(10, 10))\nfor i, ax in enumerate(axes.flatten()):\n    img, label = random.choice(train_dataset)\n    ax.imshow(img.permute(1, 2, 0))\n    ax.set_axis_off()\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:40.353469Z","iopub.execute_input":"2022-03-12T18:57:40.353733Z","iopub.status.idle":"2022-03-12T18:57:41.625804Z","shell.execute_reply.started":"2022-03-12T18:57:40.353698Z","shell.execute_reply":"2022-03-12T18:57:41.625211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Proof of Concept","metadata":{}},{"cell_type":"markdown","source":"Training on 10 *category_id* classes on a custom Convolutional Neural Network","metadata":{}},{"cell_type":"markdown","source":"### Dataset subsample for 10 classes (by count of examples)","metadata":{}},{"cell_type":"code","source":"labels = train_df['category_id'] \\\n          .value_counts()[:10] \\\n          .index \\\n          .to_list()\n\nlabels_dict = {label: idx for idx, label in enumerate(labels)}\n\nsubsample_df = train_df[train_df['category_id'].isin(labels)] \\\n                .reset_index()\n\nsubsample_df['category_id'] = subsample_df['category_id'] \\\n                               .map(labels_dict)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:41.627136Z","iopub.execute_input":"2022-03-12T18:57:41.627657Z","iopub.status.idle":"2022-03-12T18:57:41.659438Z","shell.execute_reply.started":"2022-03-12T18:57:41.627621Z","shell.execute_reply":"2022-03-12T18:57:41.658798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Encoded labels\nprint(labels_dict)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:41.660764Z","iopub.execute_input":"2022-03-12T18:57:41.661019Z","iopub.status.idle":"2022-03-12T18:57:41.664966Z","shell.execute_reply.started":"2022-03-12T18:57:41.660984Z","shell.execute_reply":"2022-03-12T18:57:41.664322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"subsample_df","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:41.666272Z","iopub.execute_input":"2022-03-12T18:57:41.66674Z","iopub.status.idle":"2022-03-12T18:57:41.685267Z","shell.execute_reply.started":"2022-03-12T18:57:41.666703Z","shell.execute_reply":"2022-03-12T18:57:41.684426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"subsample_df['category_id'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:41.686415Z","iopub.execute_input":"2022-03-12T18:57:41.686777Z","iopub.status.idle":"2022-03-12T18:57:41.697429Z","shell.execute_reply.started":"2022-03-12T18:57:41.68674Z","shell.execute_reply":"2022-03-12T18:57:41.696644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Train/Test Split","metadata":{}},{"cell_type":"code","source":"train_subsample_df = subsample_df.sample(frac=0.8, random_state=200)\ntest_subsample_df = subsample_df.drop(train_subsample_df.index)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:41.698478Z","iopub.execute_input":"2022-03-12T18:57:41.698671Z","iopub.status.idle":"2022-03-12T18:57:41.707607Z","shell.execute_reply.started":"2022-03-12T18:57:41.698649Z","shell.execute_reply":"2022-03-12T18:57:41.706842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ax = train_subsample_df['category_id'] \\\n      .value_counts() \\\n      .plot(kind='barh')\nfor container in ax.containers:\n    ax.bar_label(container)\nplt.title('Train dataset')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:41.708782Z","iopub.execute_input":"2022-03-12T18:57:41.709053Z","iopub.status.idle":"2022-03-12T18:57:41.930901Z","shell.execute_reply.started":"2022-03-12T18:57:41.709019Z","shell.execute_reply":"2022-03-12T18:57:41.93024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ax = test_subsample_df['category_id'] \\\n      .value_counts() \\\n      .plot(kind='barh')\nfor container in ax.containers:\n    ax.bar_label(container)\nplt.title('Test dataset')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:41.931943Z","iopub.execute_input":"2022-03-12T18:57:41.932338Z","iopub.status.idle":"2022-03-12T18:57:42.15256Z","shell.execute_reply.started":"2022-03-12T18:57:41.932301Z","shell.execute_reply":"2022-03-12T18:57:42.151869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_subsample_df.to_csv('train_subsample_metadata.csv', index=False)\ntest_subsample_df.to_csv('test_subsample_metadata.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:42.153645Z","iopub.execute_input":"2022-03-12T18:57:42.15432Z","iopub.status.idle":"2022-03-12T18:57:42.165337Z","shell.execute_reply.started":"2022-03-12T18:57:42.15428Z","shell.execute_reply":"2022-03-12T18:57:42.164624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### PyTorch Lightning Custom CNN","metadata":{}},{"cell_type":"code","source":"LR = 0.01\nEPOCHS = 60\nCLASSES = 10\nNUM_WORKERS = os.cpu_count()\nAVAIL_GPUS = torch.cuda.device_count()\nBATCH_SIZE = 32 if AVAIL_GPUS else 16","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:42.166859Z","iopub.execute_input":"2022-03-12T18:57:42.167127Z","iopub.status.idle":"2022-03-12T18:57:42.210728Z","shell.execute_reply.started":"2022-03-12T18:57:42.167086Z","shell.execute_reply":"2022-03-12T18:57:42.209922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LitCNN(pl.LightningModule):\n    '''\n    Custom CNN with PyTorch Lightning.\n    '''\n    def __init__(self):\n        super().__init__()\n        self.conv1 = self._conv_module(3, 16)\n        self.conv2 = self._conv_module(16, 32)\n        self.conv3 = self._conv_module(32, 64)\n        self.conv4 = self._conv_module(64, 128)\n        self.flatten = nn.Flatten()\n        self.drop = nn.Dropout(p=0.2)\n        self.fc1 = nn.Linear(20*20*128, 256)\n        self.fc2 = nn.Linear(256, 128)\n        self.fc3 = nn.Linear(128, CLASSES)\n        self.relu = nn.ReLU()\n        self.loss_fn = nn.CrossEntropyLoss()\n        \n    def _conv_module(self, in_shape, out_shape):\n        return nn.Sequential(\n            nn.Conv2d(in_shape, out_shape, kernel_size=3, stride=1),\n            nn.BatchNorm2d(out_shape),\n            nn.LeakyReLU(),\n            nn.MaxPool2d(kernel_size=2, stride=2)\n        )\n    \n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.conv2(x)\n        x = self.conv3(x)\n        x = self.conv4(x)\n        x = self.flatten(x)\n        x = self.drop(x)\n        x = self.relu(self.fc1(x))\n        x = self.relu(self.fc2(x))\n        x = self.fc3(x)\n        return x\n        \n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        logits = self(x)\n        loss = self.loss_fn(logits, y)\n        self.log('loss', loss, on_step=True, on_epoch=True, prog_bar=True, logger=True)\n        return loss\n    \n    def test_step(self, batch, batch_idx):\n        x, y = batch\n        logits = self(x)\n        loss = self.loss_fn(logits, y)\n        accuracy = metrics.accuracy(logits, y)\n        metrics_dict = {'loss': loss, 'accuracy': accuracy}\n        self.log_dict(metrics_dict, on_epoch=True, prog_bar=True)\n        return metrics_dict\n    \n    def training_epoch_end(self, outputs):\n        avg_loss = torch.tensor([out['loss'] for out in outputs]).mean()\n        self.log('train_loss', avg_loss, logger=True, prog_bar=True)\n    \n    def test_epoch_end(self, outputs):\n        avg_loss = torch.tensor([out['loss'] for out in outputs]).mean()\n        avg_acc = torch.tensor([out['accuracy'] for out in outputs]).mean()\n        self.log('test_loss', avg_loss, on_epoch=True)\n        self.log('test_accuracy', avg_acc, on_epoch=True)\n    \n    def configure_optimizers(self):\n        optimizer = optim.Adam(self.parameters(), lr=LR)\n        scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.1)\n        return {'optimizer': optimizer, 'lr_scheduler': scheduler}\n    \n    def train_dataloader(self):\n        train_dataset = HerbariumDataset(TRAIN_IMGS, 'train_subsample_metadata.csv', \n                                         transform=train_transform)\n        return DataLoader(train_dataset, batch_size=BATCH_SIZE, \n                          shuffle=True, num_workers=NUM_WORKERS)\n        \n    def test_dataloader(self):\n        test_dataset = HerbariumDataset(TRAIN_IMGS, 'test_subsample_metadata.csv', \n                                        transform=test_transform)\n        return DataLoader(test_dataset, batch_size=BATCH_SIZE, \n                          num_workers=NUM_WORKERS)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:57:42.214037Z","iopub.execute_input":"2022-03-12T18:57:42.214348Z","iopub.status.idle":"2022-03-12T18:57:42.23924Z","shell.execute_reply.started":"2022-03-12T18:57:42.21432Z","shell.execute_reply":"2022-03-12T18:57:42.238524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = LitCNN()\nmodel","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:58:56.177004Z","iopub.execute_input":"2022-03-12T18:58:56.177275Z","iopub.status.idle":"2022-03-12T18:58:56.307121Z","shell.execute_reply.started":"2022-03-12T18:58:56.177244Z","shell.execute_reply":"2022-03-12T18:58:56.306261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = pl.Trainer(log_every_n_steps=10, \n                     gpus=AVAIL_GPUS, \n                     max_epochs=EPOCHS)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:58:58.724643Z","iopub.execute_input":"2022-03-12T18:58:58.724893Z","iopub.status.idle":"2022-03-12T18:58:58.734529Z","shell.execute_reply.started":"2022-03-12T18:58:58.724866Z","shell.execute_reply":"2022-03-12T18:58:58.733846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.fit(model)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T18:59:00.548983Z","iopub.execute_input":"2022-03-12T18:59:00.549486Z","iopub.status.idle":"2022-03-12T19:10:35.123311Z","shell.execute_reply.started":"2022-03-12T18:59:00.549439Z","shell.execute_reply":"2022-03-12T19:10:35.122535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.test(model, verbose=False)","metadata":{"execution":{"iopub.status.busy":"2022-03-12T19:10:52.939143Z","iopub.execute_input":"2022-03-12T19:10:52.940086Z","iopub.status.idle":"2022-03-12T19:10:56.542961Z","shell.execute_reply.started":"2022-03-12T19:10:52.940036Z","shell.execute_reply":"2022-03-12T19:10:56.542251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EfficientNet?","metadata":{}},{"cell_type":"markdown","source":"# Summary?","metadata":{}}]}