{"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":"# 🐡 Starfish detection: Flash ⚡ EfficientDet\n\nYour PyTorch AI Factory - Flash enables you to easily configure and run complex AI recipes for over 15 tasks across 7 data domains","metadata":{}},{"cell_type":"markdown","source":"## Installs\n\nInstalling the packge with additinal extras for computer vision","metadata":{}},{"cell_type":"code","source":"!pip install -q fiftyone effdet \"icevision[all]\" 'git+https://github.com/PyTorchLightning/lightning-flash.git#egg=lightning-flash[image]'\n# !pip install -q \"pytorch-lightning==1.4.*\"\n!pip uninstall -y wandb","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-02-15T10:07:59.665212Z","iopub.execute_input":"2022-02-15T10:07:59.665568Z","iopub.status.idle":"2022-02-15T10:10:12.676240Z","shell.execute_reply.started":"2022-02-15T10:07:59.665479Z","shell.execute_reply":"2022-02-15T10:10:12.675446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"import ast\nfrom pathlib import Path\n\nimport torch\nimport flash\nimport fiftyone as fo\nimport numpy as np\nimport pandas as pd\nfrom flash.image import ObjectDetectionData, ObjectDetector","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-02-15T10:10:12.679000Z","iopub.execute_input":"2022-02-15T10:10:12.679500Z","iopub.status.idle":"2022-02-15T10:10:35.014518Z","shell.execute_reply.started":"2022-02-15T10:10:12.679456Z","shell.execute_reply":"2022-02-15T10:10:35.013751Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Paths","metadata":{}},{"cell_type":"code","source":"INPUT_DIR = Path(\"/kaggle/input\")\nDATA_DIR = INPUT_DIR / \"tensorflow-great-barrier-reef\"\nTRAIN_CSV_PATH = DATA_DIR / \"train.csv\"\nCOCO_DATA_DIR = Path(\"/kaggle/working/gbr-coco\")","metadata":{"execution":{"iopub.status.busy":"2022-02-15T10:10:35.015861Z","iopub.execute_input":"2022-02-15T10:10:35.017777Z","iopub.status.idle":"2022-02-15T10:10:35.022362Z","shell.execute_reply.started":"2022-02-15T10:10:35.017721Z","shell.execute_reply":"2022-02-15T10:10:35.021641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Connvert Dataset","metadata":{}},{"cell_type":"code","source":"IMAGE_DIMS = (1280, 720)\n\ndef format_bbox(annotation):\n    xmin, ymin, width, height = annotation[\"x\"], annotation[\"y\"], annotation[\"width\"], annotation[\"height\"]\n    width = IMAGE_DIMS[0] - xmin - 1 if xmin + width >= IMAGE_DIMS[0] else width\n    height = IMAGE_DIMS[1] - ymin - 1 if ymin + height >= IMAGE_DIMS[1] else height\n    return {\"xmin\": xmin, \"ymin\": ymin, \"width\": width, \"height\": height}\n\ntrain_files, train_labels, train_bboxes = [], [], []\n\ntrain_df = pd.read_csv(TRAIN_CSV_PATH)\n\nfor idx, row in train_df.iterrows():\n    image_path = DATA_DIR / \"train_images\" / f\"video_{row['video_id']}\" / f\"{row['video_frame']}.jpg\"\n    annotations = ast.literal_eval(row[\"annotations\"])\n    \n    labels, bboxes = [], []\n    \n    for annotation in annotations:\n        labels.append(\"cots\")\n        bboxes.append(format_bbox(annotation))\n    \n    # Skip images with no annotations\n    if labels != []:\n        train_files.append(image_path)\n        train_labels.append(labels)\n        train_bboxes.append(bboxes)","metadata":{"execution":{"iopub.status.busy":"2022-02-15T10:10:35.025182Z","iopub.execute_input":"2022-02-15T10:10:35.025517Z","iopub.status.idle":"2022-02-15T10:10:37.424173Z","shell.execute_reply.started":"2022-02-15T10:10:35.025482Z","shell.execute_reply":"2022-02-15T10:10:37.423454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## ⚡ Flash in 3 Steps\n\nconfigure all training args https://lightning-flash.readthedocs.io/en/stable/reference/object_detection.html","metadata":{}},{"cell_type":"markdown","source":"### Step 1. Load your data","metadata":{}},{"cell_type":"code","source":"IMAGE_SIZE = 1024\n\ndatamodule = ObjectDetectionData.from_files(\n    train_files=train_files,\n    train_targets=train_labels,\n    train_bboxes=train_bboxes,\n    val_split=0.15,\n    transform_kwargs={\"image_size\": IMAGE_SIZE},\n    batch_size=4,\n    num_workers=4,\n)","metadata":{"execution":{"iopub.status.busy":"2022-02-15T10:10:37.425471Z","iopub.execute_input":"2022-02-15T10:10:37.425744Z","iopub.status.idle":"2022-02-15T10:10:37.473559Z","shell.execute_reply.started":"2022-02-15T10:10:37.425709Z","shell.execute_reply":"2022-02-15T10:10:37.472826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Step 2: Configure your model","metadata":{}},{"cell_type":"code","source":"model = ObjectDetector(\n    head=\"efficientdet\",\n    backbone=\"d3\",\n    num_classes=datamodule.num_classes,\n    image_size=IMAGE_SIZE,\n    pretrained=True,\n    optimizer=\"AdamW\",\n    learning_rate=0.001,\n)\n# model.adapter.model.max_detection_points = 1000","metadata":{"execution":{"iopub.status.busy":"2022-02-15T10:10:37.474813Z","iopub.execute_input":"2022-02-15T10:10:37.475205Z","iopub.status.idle":"2022-02-15T10:10:38.039623Z","shell.execute_reply.started":"2022-02-15T10:10:37.475169Z","shell.execute_reply":"2022-02-15T10:10:38.038705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Step 3: Finetune","metadata":{}},{"cell_type":"code","source":"from pytorch_lightning.loggers import CSVLogger\n# from pytorch_lightning.callbacks import StochasticWeightAveraging\n\n# Trainer Args\nGPUS = torch.cuda.device_count()  # 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    # fast_dev_run=False,\n    # callbacks=[swa],\n    gradient_clip_val=0.01,\n    gpus=GPUS,\n    max_epochs=15,\n    precision=16 if GPUS else 32,\n    logger=logger,\n)\n\nimport warnings\nwarnings.filterwarnings(\"ignore\", category=DeprecationWarning) \n\ntrainer.finetune(model, datamodule=datamodule, strategy=\"no_freeze\")  # strategy=(\"freeze_unfreeze\", 5)\n\ntrainer.save_checkpoint(\"object_detection_model.pt\")","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-02-15T10:10:38.043875Z","iopub.execute_input":"2022-02-15T10:10:38.044140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training visualizations","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\nsns.set()\n\nmetrics = 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 = sns.relplot(data=metrics, kind=\"line\")\nplt.gcf().set_size_inches(15, 5)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}