{"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":"# 🕵Explore the provdided data","metadata":{}},{"cell_type":"code","source":"!ls -l /kaggle/input/happy-whale-and-dolphin\n\nPATH_DATASET = \"/kaggle/input/happy-whale-and-dolphin\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-03-03T17:33:20.30709Z","iopub.execute_input":"2022-03-03T17:33:20.307578Z","iopub.status.idle":"2022-03-03T17:33:21.005715Z","shell.execute_reply.started":"2022-03-03T17:33:20.30749Z","shell.execute_reply":"2022-03-03T17:33:21.004768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Browsing the metadata","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport seaborn as sn\nimport matplotlib.pyplot as plt\n\nsn.set()\n\ndf_train = pd.read_csv(os.path.join(PATH_DATASET, \"train.csv\"))\ndisplay(df_train.head())\nprint(f\"Dataset size: {len(df_train)}\")\nprint(f\"Unique ids: {len(df_train['individual_id'].unique())}\")","metadata":{"execution":{"iopub.status.busy":"2022-03-03T17:33:21.008573Z","iopub.execute_input":"2022-03-03T17:33:21.009997Z","iopub.status.idle":"2022-03-03T17:33:21.976483Z","shell.execute_reply.started":"2022-03-03T17:33:21.009945Z","shell.execute_reply":"2022-03-03T17:33:21.975343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Lets see how many speaced we have in the database...","metadata":{}},{"cell_type":"code","source":"counts_imgs = df_train[\"species\"].value_counts()\ncounts_inds = df_train.drop_duplicates(\"individual_id\")[\"species\"].value_counts()\n\nax = pd.concat({\"per Images\": counts_imgs, \"per Individuals\": counts_inds}, axis=1).plot.barh(grid=True, figsize=(7, 10))\nax.set_xscale('log')","metadata":{"execution":{"iopub.status.busy":"2022-03-03T17:33:21.977611Z","iopub.execute_input":"2022-03-03T17:33:21.978124Z","iopub.status.idle":"2022-03-03T17:33:23.249268Z","shell.execute_reply.started":"2022-03-03T17:33:21.978078Z","shell.execute_reply":"2022-03-03T17:33:23.248607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"And compare they with unique individuals... \n\n**Note:** that the counts are in log scale","metadata":{}},{"cell_type":"code","source":"import numpy as np\nfrom pprint import pprint\n\nspecies_individuals = {}\nfor name, dfg in df_train.groupby(\"species\"):\n    species_individuals[name] = dfg[\"individual_id\"].value_counts()\n\nsi_max = max(list(map(len, species_individuals.values())))\nsi = {n: [0] * si_max for n in species_individuals}\nfor n, counts in species_individuals.items():\n    si[n][:len(counts)] = list(np.log(counts))\nsi = pd.DataFrame(si)","metadata":{"execution":{"iopub.status.busy":"2022-03-03T17:33:23.25121Z","iopub.execute_input":"2022-03-03T17:33:23.251578Z","iopub.status.idle":"2022-03-03T17:33:23.32274Z","shell.execute_reply.started":"2022-03-03T17:33:23.251537Z","shell.execute_reply":"2022-03-03T17:33:23.322068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import seaborn as sn\n\nfig = plt.figure(figsize=(10, 8))\nax = sn.heatmap(si[:500].T, cmap=\"BuGn\", ax=fig.gca())","metadata":{"execution":{"iopub.status.busy":"2022-03-03T17:33:23.324753Z","iopub.execute_input":"2022-03-03T17:33:23.325146Z","iopub.status.idle":"2022-03-03T17:33:24.245294Z","shell.execute_reply.started":"2022-03-03T17:33:23.325112Z","shell.execute_reply":"2022-03-03T17:33:24.244622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"And see the top individulas","metadata":{}},{"cell_type":"code","source":"ax = df_train[\"individual_id\"].value_counts(ascending=True)[-50:].plot.barh(figsize=(3, 8), grid=True)  # ascending=True","metadata":{"execution":{"iopub.status.busy":"2022-03-03T17:33:24.246277Z","iopub.execute_input":"2022-03-03T17:33:24.246532Z","iopub.status.idle":"2022-03-03T17:33:25.277314Z","shell.execute_reply.started":"2022-03-03T17:33:24.246496Z","shell.execute_reply":"2022-03-03T17:33:25.276652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Browse some images","metadata":{}},{"cell_type":"code","source":"nb_species = len(df_train[\"species\"].unique())\nfig, axarr = plt.subplots(ncols=5, nrows=nb_species, figsize=(12, nb_species * 2))\n\nfor i, (name, dfg) in enumerate(df_train.groupby(\"species\")):\n    axarr[i, 0].set_title(name)\n    for j, (_, row) in enumerate(dfg[:5].iterrows()):\n        im_path = os.path.join(PATH_DATASET, \"train_images\", row[\"image\"])\n        img = plt.imread(im_path)\n        axarr[i, j].imshow(img)\n        axarr[i, j].set_axis_off()","metadata":{"execution":{"iopub.status.busy":"2022-03-03T17:33:25.278688Z","iopub.execute_input":"2022-03-03T17:33:25.279166Z","iopub.status.idle":"2022-03-03T17:34:37.725745Z","shell.execute_reply.started":"2022-03-03T17:33:25.279122Z","shell.execute_reply":"2022-03-03T17:34:37.723877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Baseline: species classification with Lightning⚡Flash\n\nFollow the example: https://lightning-flash.readthedocs.io/en/stable/reference/image_classification.html","metadata":{}},{"cell_type":"code","source":"!pip install -q effdet \"icevision[all]\" 'lightning-flash[image]'\n# !pip install -q \"pytorch-lightning==1.4.*\"\n!pip uninstall -y wandb","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-03-03T17:34:37.727025Z","iopub.execute_input":"2022-03-03T17:34:37.727296Z","iopub.status.idle":"2022-03-03T17:35:53.813586Z","shell.execute_reply.started":"2022-03-03T17:34:37.727255Z","shell.execute_reply":"2022-03-03T17:35:53.812778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip download -q effdet \"icevision[all]\" 'lightning-flash[image]' --dest frozen_packages --prefer-binary\n!rm frozen_packages/torch-*\n!ls -l frozen_packages","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-03-03T17:35:53.815519Z","iopub.execute_input":"2022-03-03T17:35:53.815828Z","iopub.status.idle":"2022-03-03T17:37:45.793425Z","shell.execute_reply.started":"2022-03-03T17:35:53.815775Z","shell.execute_reply":"2022-03-03T17:37:45.79262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\nimport flash\nfrom flash.core.data.utils import download_data\nfrom flash.image import ImageClassificationData, ImageClassifier","metadata":{"execution":{"iopub.status.busy":"2022-03-03T17:37:45.797765Z","iopub.execute_input":"2022-03-03T17:37:45.797994Z","iopub.status.idle":"2022-03-03T17:37:57.012727Z","shell.execute_reply.started":"2022-03-03T17:37:45.797966Z","shell.execute_reply":"2022-03-03T17:37:57.011995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 1. Create the DataModule 🗄️","metadata":{}},{"cell_type":"code","source":"datamodule = ImageClassificationData.from_data_frame(\n    input_field=\"image\",\n    target_fields=\"species\",\n    # for simplicity take just half of the data\n    train_data_frame=df_train[:len(df_train) // 2],\n    train_images_root=os.path.join(PATH_DATASET, \"train_images\"),\n    batch_size=64,\n    transform_kwargs={\"image_size\": (300, 300)},\n    val_split=0.1,\n    num_workers=2,\n)","metadata":{"execution":{"iopub.status.busy":"2022-03-03T17:37:57.014268Z","iopub.execute_input":"2022-03-03T17:37:57.014518Z","iopub.status.idle":"2022-03-03T17:38:14.897631Z","shell.execute_reply.started":"2022-03-03T17:37:57.014482Z","shell.execute_reply":"2022-03-03T17:38:14.896926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Build the task ⚙️","metadata":{}},{"cell_type":"code","source":"from torchmetrics import F1\n\nmodel = ImageClassifier(\n    backbone=\"efficientnet_b3\",\n    labels=datamodule.labels,\n    metrics=F1(),\n    pretrained=True,\n    optimizer=\"AdamW\",\n    learning_rate=0.005,\n)","metadata":{"execution":{"iopub.status.busy":"2022-03-03T17:38:14.898914Z","iopub.execute_input":"2022-03-03T17:38:14.899142Z","iopub.status.idle":"2022-03-03T17:38:17.600362Z","shell.execute_reply.started":"2022-03-03T17:38:14.89911Z","shell.execute_reply":"2022-03-03T17:38:17.599635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3. Finetune the model 🛠️","metadata":{}},{"cell_type":"code","source":"from pytorch_lightning.loggers import CSVLogger\n# from pytorch_lightning.callbacks import StochasticWeightAveraging\n\n# Trainer Args\nGPUS = int(torch.cuda.is_available())  # Set to 1 if GPU is enabled for notebook\n\n# swa = StochasticWeightAveraging(swa_epoch_start=0.6)\nlogger = CSVLogger(save_dir='logs/')\n\ntrainer = flash.Trainer(\n    max_epochs=5,\n    # gradient_clip_val=0.01,\n    gpus=GPUS,\n    precision=16 if GPUS else 32,\n    logger=logger,\n)","metadata":{"execution":{"iopub.status.busy":"2022-03-03T17:38:17.601756Z","iopub.execute_input":"2022-03-03T17:38:17.602019Z","iopub.status.idle":"2022-03-03T17:38:17.615144Z","shell.execute_reply.started":"2022-03-03T17:38:17.601984Z","shell.execute_reply":"2022-03-03T17:38:17.614482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.finetune(model, datamodule=datamodule, strategy=(\"freeze_unfreeze\", 2))\n# trainer.finetune(model, datamodule=datamodule, strategy=\"no_freeze\")\n\ntrainer.save_checkpoint(\"image_classification_model.pt\")","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-03-03T17:38:17.616483Z","iopub.execute_input":"2022-03-03T17:38:17.616839Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metrics = pd.read_csv(f'{trainer.logger.log_dir}/metrics.csv')\ndel metrics[\"step\"]\nmetrics.set_index(\"epoch\", inplace=True)\ndisplay(metrics.dropna(axis=1, how=\"all\").head())\ng = sn.relplot(data=metrics, kind=\"line\")\nplt.gcf().set_size_inches(12, 4)\nplt.grid()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}