{"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":"# 🐠 Reef - Pytorch Starter - FasterRCNN Train\n\n## A self-contained, simple, pure pytorch 🔥 Faster R-CNN implementation with `LB=0.413`\n\n![](https://storage.googleapis.com/kaggle-competitions/kaggle/31703/logos/header.png)\n\n#### FasterR-CNN is one of the SOTA models for Object detection.\n\n### In this notebook we present a simple solution using a pure pytorch Faster R-CNN with pretrained weights, and finetuning it for few epochs.\n\nIt is an adapted version of [this notebook](https://www.kaggle.com/pestipeti/pytorch-starter-fasterrcnn-train) mentioned in [this comment](https://www.kaggle.com/c/tensorflow-great-barrier-reef/discussion/290016).\n\n## You can find the [inference notebook here](https://www.kaggle.com/julian3833/coral-reef-pytorch-fasterrcnn-infer-0-xxx).\n\n## Details: \n- FasterRCNN from torchvision\n- Use Resnet50 backbone\n\n**Update**: Added simple train/validation split in this version, using the \"subsequence\" split of this notebook: [🐠 Reef - CV strategy: subsequences!](https://www.kaggle.com/julian3833/reef-cv-strategy-subsequences)\n\nStill dropping all the images with no objects, as the model doesn't support them out-of-the-box. The other starters are removing the empty images as well, so it might be a general condition of Object Detection. I'm quite noob in the field to be honest.\n\n# Please, _DO_ upvote if you find this useful!!\n\n\n&nbsp;\n&nbsp;\n&nbsp;\n\n#### Changelog\n\n| Version | Description| Dataset| Best LB |\n| --- | ----| --- | --- |\n| [**V8**](https://www.kaggle.com/julian3833/reef-starter-torch-fasterrcnn-train-lb-0-293?scriptVersionId=80517118)  | 2 epochs - Save last epoch | [coral-reef-pytorch-starter-fasterrcnn-weights](https://www.kaggle.com/julian3833/coral-reef-pytorch-starter-fasterrcnn-weights)| `0.293`|\n| [**V16**](https://www.kaggle.com/julian3833/reef-starter-torch-fasterrcnn-train-lb-0-293?scriptVersionId=80601095) | 4 epochs - Save all epochs | [reef-starter-torch-fasterrcnn-4e](https://www.kaggle.com/julian3833/reef-starter-torch-fasterrcnn-4e)| `0.361` |\n| [**V17**](https://www.kaggle.com/julian3833/reef-starter-torch-fasterrcnn-train-lb-0-369?scriptVersionId=80604402) | Add **validation**. 95-5 split. 8 epochs, keeping track of validation loss. | [reef-starter-torch-fasterrcnn-8e](https://www.kaggle.com/julian3833/reef-starter-torch-fasterrcnn-8e)| `0.369` |\n| [**V19**](https://www.kaggle.com/julian3833/reef-starter-torch-fasterrcnn-train-lb-0-369?scriptVersionId=80610403) | 12 epochs, lower LR | [reef-starter-torch-fasterrcnn-12e](https://www.kaggle.com/julian3833/reef-starter-torch-fasterrcnn-12e)| `0.413` |\n| [**V24**](https://www.kaggle.com/julian3833/reef-starter-torch-fasterrcnn-train-lb-0-369) | V19 with 90-10 train-validation split. Tidy up code. Add Flip. Correct problem with augmentations. | -- | `??` |\n\n\n","metadata":{}},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"# Very few imports. This is a pure torch solution!\nimport cv2\nimport time\n\nimport numpy as np\nimport pandas as pd\nfrom matplotlib import pyplot as plt\n\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\n\nimport torch\nimport torchvision\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection import FasterRCNN","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","execution":{"iopub.status.busy":"2022-04-18T16:24:29.250003Z","iopub.execute_input":"2022-04-18T16:24:29.250362Z","iopub.status.idle":"2022-04-18T16:24:31.698511Z","shell.execute_reply.started":"2022-04-18T16:24:29.250309Z","shell.execute_reply":"2022-04-18T16:24:31.697625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Constants","metadata":{}},{"cell_type":"code","source":"DEVICE = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\nBASE_DIR = \"../input/tensorflow-great-barrier-reef/train_images/\"\n\nNUM_EPOCHS = 2","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:02.187128Z","iopub.execute_input":"2022-04-18T16:28:02.187495Z","iopub.status.idle":"2022-04-18T16:28:02.192863Z","shell.execute_reply.started":"2022-04-18T16:28:02.187433Z","shell.execute_reply":"2022-04-18T16:28:02.191622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load `df`\n\n### See: [🐠 Reef - CV strategy: subsequences!](https://www.kaggle.com/julian3833/reef-cv-strategy-subsequences)","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(\"../input/reef-cv-strategy-subsequences-dataframes/train-validation-split/train-0.1.csv\")\n\n# Turn annotations from strings into lists of dictionaries\ndf['annotations'] = df['annotations'].apply(eval)\n\n# Create the image path for the row\ndf['image_path'] = \"video_\" + df['video_id'].astype(str) + \"/\" + df['video_frame'].astype(str) + \".jpg\"\n\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:10.48475Z","iopub.execute_input":"2022-04-18T16:28:10.485209Z","iopub.status.idle":"2022-04-18T16:28:10.934651Z","shell.execute_reply.started":"2022-04-18T16:28:10.485155Z","shell.execute_reply":"2022-04-18T16:28:10.933737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train, df_val = df[df['is_train']], df[~df['is_train']]","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:13.295507Z","iopub.execute_input":"2022-04-18T16:28:13.295939Z","iopub.status.idle":"2022-04-18T16:28:13.307287Z","shell.execute_reply.started":"2022-04-18T16:28:13.295872Z","shell.execute_reply":"2022-04-18T16:28:13.30638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# The model doesn't support images with no annotations\n# It raises an error that suggest that it just doesn't support them:\n# V    alueError: No ground-truth boxes available for one of the images during training\n# I'm dropping those images for now\n# https://discuss.pytorch.org/t/fasterrcnn-images-with-no-objects-present-cause-an-error/117974/3\ndf_train = df_train[df_train.annotations.str.len() > 0 ].reset_index(drop=True)\ndf_val = df_val[df_val.annotations.str.len() > 0 ].reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:17.711311Z","iopub.execute_input":"2022-04-18T16:28:17.711646Z","iopub.status.idle":"2022-04-18T16:28:17.738533Z","shell.execute_reply.started":"2022-04-18T16:28:17.711594Z","shell.execute_reply":"2022-04-18T16:28:17.737437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.shape[0], df_val.shape[0]","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:20.577322Z","iopub.execute_input":"2022-04-18T16:28:20.577767Z","iopub.status.idle":"2022-04-18T16:28:20.586622Z","shell.execute_reply.started":"2022-04-18T16:28:20.577716Z","shell.execute_reply":"2022-04-18T16:28:20.585553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset class","metadata":{}},{"cell_type":"code","source":"class ReefDataset:\n\n    def __init__(self, df, transforms=None):\n        self.df = df\n        self.transforms = transforms\n\n    def can_augment(self, boxes):\n        \"\"\" Check if bounding boxes are OK to augment\n        \n        \n        For example: image_id 1-490 has a bounding box that is partially outside of the image\n        It breaks albumentation\n        Here we check the margins are within the image to make sure the augmentation can be applied\n        \"\"\"\n        \n        box_outside_image = ((boxes[:, 0] < 0).any() or (boxes[:, 1] < 0).any() \n                             or (boxes[:, 2] > 1280).any() or (boxes[:, 3] > 720).any())\n        return not box_outside_image\n\n    def get_boxes(self, row):\n        \"\"\"Returns the bboxes for a given row as a 3D matrix with format [x_min, y_min, x_max, y_max]\"\"\"\n        \n        boxes = pd.DataFrame(row['annotations'], columns=['x', 'y', 'width', 'height']).astype(float).values\n        \n        # Change from [x_min, y_min, w, h] to [x_min, y_min, x_max, y_max]\n        boxes[:, 2] = boxes[:, 0] + boxes[:, 2]\n        boxes[:, 3] = boxes[:, 1] + boxes[:, 3]\n        return boxes\n    \n    def get_image(self, row):\n        \"\"\"Gets the image for a given row\"\"\"\n        \n        image = cv2.imread(f'{BASE_DIR}/{row[\"image_path\"]}', cv2.IMREAD_COLOR)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB).astype(np.float32) \n        image /= 255.0     \n                           \n        return image\n    \n    def __getitem__(self, i):\n\n        row = self.df.iloc[i]\n        image = self.get_image(row)\n        boxes = self.get_boxes(row)\n        \n        n_boxes = boxes.shape[0]\n        \n        # Calculate the area\n        area = (boxes[:, 3] - boxes[:, 1]) * (boxes[:, 2] - boxes[:, 0])\n        \n        \n        target = {\n            'boxes': torch.as_tensor(boxes, dtype=torch.float32),\n            'area': torch.as_tensor(area, dtype=torch.float32),\n            \n            'image_id': torch.tensor([i]),\n            \n            # There is only one class\n            'labels': torch.ones((n_boxes,), dtype=torch.int64),\n            \n            # Suppose all instances are not crowd\n            'iscrowd': torch.zeros((n_boxes,), dtype=torch.int64)            \n        }\n\n        if self.transforms and self.can_augment(boxes):\n            sample = {\n                'image': image,\n                'bboxes': target['boxes'],\n                'labels': target['labels']\n            }\n            sample = self.transforms(**sample)\n            image = sample['image']\n            \n            if n_boxes > 0:\n                target['boxes'] = torch.stack(tuple(map(torch.tensor, zip(*sample['bboxes'])))).permute(1, 0)\n        else:\n            image = ToTensorV2(p=1.0)(image=image)['image']\n\n        return image, target\n\n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:23.376496Z","iopub.execute_input":"2022-04-18T16:28:23.376885Z","iopub.status.idle":"2022-04-18T16:28:23.399655Z","shell.execute_reply.started":"2022-04-18T16:28:23.376829Z","shell.execute_reply":"2022-04-18T16:28:23.398506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Original Augmentation code\n# class HE_HSV(A.ImageOnlyTransform):\n#     def __init__(self, p: float = 0.5, always_apply=False):\n#         super().__init__(always_apply, p)\n        \n#     def apply(self, image,**params):\n#         img_hsv = cv2.cvtColor(image, cv2.COLOR_RGB2HSV)\n\n#         # Histogram equalisation on the V-channel\n#         img_hsv[:, :, 2] = cv2.equalizeHist(img_hsv[:, :, 2])\n\n#         # convert image back from HSV to RGB\n#         image_hsv = cv2.cvtColor(img_hsv, cv2.COLOR_HSV2RGB)\n\n#         return image_hsv\n    \n# class RecoverHE(A.ImageOnlyTransform):\n#     def __init__(self, p: float = 0.5, always_apply=False):\n#         super().__init__(always_apply, p)\n        \n#     def apply(self, sceneRadiance,**params):\n#         for i in range(3):\n#             sceneRadiance[:, :, i] =  cv2.equalizeHist(sceneRadiance[:, :, i])\n#         return sceneRadiance\n\n# class CLAHE_HSV(A.ImageOnlyTransform):\n    \n#     def __init__(self, p: float = 0.5, always_apply=False):\n#         super().__init__(always_apply, p)\n        \n#     def apply(self, img, **params):\n#         hsv_img = cv2.cvtColor(img, cv2.COLOR_BGR2HSV)\n\n#         h, s, v = hsv_img[:,:,0], hsv_img[:,:,1], hsv_img[:,:,2]\n#         clahe = cv2.createCLAHE(clipLimit = 15.0, tileGridSize = (20,20))\n#         v = clahe.apply(v)\n\n#         hsv_img = np.dstack((h,s,v))\n\n#         rgb = cv2.cvtColor(hsv_img, cv2.COLOR_HSV2RGB)\n\n#         return rgb\n\n# class RecoverCLAHE(A.ImageOnlyTransform):\n    \n#     def __init__(self, p: float = 0.5, always_apply=False):\n#         super().__init__(always_apply, p)\n        \n#     def apply(self, sceneRadiance, **params):\n#         clahe = cv2.createCLAHE(clipLimit=7, tileGridSize=(14, 14))\n#         for i in range(3):\n#             sceneRadiance[:, :, i] = clahe.apply((sceneRadiance[:, :, i]))\n\n#         return sceneRadiance\n\n# class Gamma_enhancement(A.ImageOnlyTransform):\n    \n#     def __init__(self, p: float = 0.5, always_apply=False):\n#         super().__init__(always_apply, p)\n#         self.gamma = 1/0.6\n#         self.R = 255.0\n        \n#     def apply(self, image, **params):\n#         return (self.R * np.power(image.astype(np.uint32)/self.R, self.gamma)).astype(np.uint8)\n\n# class RecoverGC(A.ImageOnlyTransform):\n    \n#     def __init__(self, p: float = 0.5, always_apply=False):\n#         super().__init__(always_apply, p)\n        \n#     def apply(self, sceneRadiance, **params):\n#         sceneRadiance = sceneRadiance/255.0\n#         # clahe = cv2.createCLAHE(clipLimit=2, tileGridSize=(2, 2))\n#         for i in range(3):\n#             sceneRadiance[:, :, i] =  np.power(sceneRadiance[:, :, i] / float(np.max(sceneRadiance[:, :, i])), 3.2)\n#         sceneRadiance = np.clip(sceneRadiance*255, 0, 255)\n#         sceneRadiance = np.uint8(sceneRadiance)\n#         return sceneRadiance\n\n# class RecoverICM(A.ImageOnlyTransform):\n    \n#     def __init__(self, p: float = 0.5, always_apply=False):\n#         super().__init__(always_apply, p)\n        \n#     def apply(self, image, **params):\n#         img_stre = stretching(iamge)\n#         sceneRadiance = sceneRadianceRGB(img_stre)\n#         sceneRadiance = HSVStretching(sceneRadiance)\n#         sceneRadiance = sceneRadianceRGB(sceneRadiance)\n\n#         return sceneRadiance","metadata":{"execution":{"iopub.status.busy":"2022-04-05T04:44:13.767816Z","iopub.execute_input":"2022-04-05T04:44:13.76817Z","iopub.status.idle":"2022-04-05T04:44:13.778823Z","shell.execute_reply.started":"2022-04-05T04:44:13.768126Z","shell.execute_reply":"2022-04-05T04:44:13.777776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Original\ndef get_train_transform():\n    return A.Compose([\n        A.Flip(0.5),\n        ToTensorV2(p=1.0)\n    ], bbox_params={'format': 'pascal_voc', 'label_fields': ['labels']})\n\n# def get_train_transform():\n#     return A.Compose([\n#                 A.OneOf([\n#                     HE_HSV(0.75)\n# #                     ,CLAHE_HSV(0.75)\n# #                     ,Gamma_enhancement(0.75)\n#                 ]), ToTensorV2(p=1.0),\n# #                 A.ShiftScaleRotate(scale_limit = 0, rotate_limit=30, p=0.3, border_mode=0)\n#             ], bbox_params={'format': 'pascal_voc', 'label_fields': ['labels']})\n# #     return A.Compose([\n# #                 A.OneOf([\n# #                     HE_HSV(0.75),\n# #                     Gamma_enhancement(0.75)\n# #                 ]), ToTensorV2(p=1.0),\n# # #                 A.ShiftScaleRotate(scale_limit = 0, rotate_limit=30, p=0.3, border_mode=0)\n# #             ], bbox_params={'format': 'pascal_voc', 'label_fields': ['labels']})\n\n\ndef get_valid_transform():\n    return A.Compose([\n        ToTensorV2(p=1.0)\n    ], bbox_params={'format': 'pascal_voc', 'label_fields': ['labels']})","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:34.16745Z","iopub.execute_input":"2022-04-18T16:28:34.167937Z","iopub.status.idle":"2022-04-18T16:28:34.175311Z","shell.execute_reply.started":"2022-04-18T16:28:34.167856Z","shell.execute_reply":"2022-04-18T16:28:34.174127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define datasets\nds_train = ReefDataset(df_train, get_train_transform())\nds_val = ReefDataset(df_val, get_valid_transform())","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:38.600501Z","iopub.execute_input":"2022-04-18T16:28:38.600864Z","iopub.status.idle":"2022-04-18T16:28:38.606417Z","shell.execute_reply.started":"2022-04-18T16:28:38.600807Z","shell.execute_reply":"2022-04-18T16:28:38.60542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(ds_train)","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:41.061537Z","iopub.execute_input":"2022-04-18T16:28:41.061916Z","iopub.status.idle":"2022-04-18T16:28:41.066615Z","shell.execute_reply.started":"2022-04-18T16:28:41.06186Z","shell.execute_reply":"2022-04-18T16:28:41.065589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Check one sample","metadata":{}},{"cell_type":"code","source":"# Let's get an interesting one ;)\ndf_train[df_train.annotations.str.len() > 12].head()","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:43.237739Z","iopub.execute_input":"2022-04-18T16:28:43.23812Z","iopub.status.idle":"2022-04-18T16:28:43.285864Z","shell.execute_reply.started":"2022-04-18T16:28:43.238068Z","shell.execute_reply":"2022-04-18T16:28:43.284805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image, targets = ds_train[2200]\nimage","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:45.644573Z","iopub.execute_input":"2022-04-18T16:28:45.64496Z","iopub.status.idle":"2022-04-18T16:28:45.795868Z","shell.execute_reply.started":"2022-04-18T16:28:45.644905Z","shell.execute_reply":"2022-04-18T16:28:45.794813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"targets","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:48.463262Z","iopub.execute_input":"2022-04-18T16:28:48.463615Z","iopub.status.idle":"2022-04-18T16:28:48.475135Z","shell.execute_reply.started":"2022-04-18T16:28:48.463559Z","shell.execute_reply":"2022-04-18T16:28:48.474437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"boxes = targets['boxes'].cpu().numpy().astype(np.int32)\nimg = image.permute(1,2,0).cpu().numpy()\nfig, ax = plt.subplots(1, 1, figsize=(16, 8))\n\nfor box in boxes:\n    cv2.rectangle(img,\n                  (box[0], box[1]),\n                  (box[2], box[3]),\n                  (220, 0, 0), 3)\n    \nax.set_axis_off()\nax.imshow(img);","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:51.103231Z","iopub.execute_input":"2022-04-18T16:28:51.103574Z","iopub.status.idle":"2022-04-18T16:28:51.646036Z","shell.execute_reply.started":"2022-04-18T16:28:51.103524Z","shell.execute_reply":"2022-04-18T16:28:51.644872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## DataLoaders","metadata":{}},{"cell_type":"code","source":"def collate_fn(batch):\n    return tuple(zip(*batch))\n\ndl_train = DataLoader(ds_train, batch_size=8, shuffle=False, num_workers=4, collate_fn=collate_fn)\ndl_val = DataLoader(ds_val, batch_size=8, shuffle=False, num_workers=4, collate_fn=collate_fn)","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:56.048409Z","iopub.execute_input":"2022-04-18T16:28:56.048837Z","iopub.status.idle":"2022-04-18T16:28:56.055426Z","shell.execute_reply.started":"2022-04-18T16:28:56.048782Z","shell.execute_reply":"2022-04-18T16:28:56.054419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create the model","metadata":{}},{"cell_type":"code","source":"def get_model():\n    # load a model; pre-trained on COCO\n    model = torchvision.models.detection.fasterrcnn_resnet50_fpn(pretrained=True)\n\n    num_classes = 2  # 1 class (starfish) + background\n\n    # get number of input features for the classifier\n    in_features = model.roi_heads.box_predictor.cls_score.in_features\n\n    # replace the pre-trained head with a new one\n    model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes)\n\n    model.to(DEVICE)\n    return model\n\nmodel = get_model()","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:28:59.678312Z","iopub.execute_input":"2022-04-18T16:28:59.678683Z","iopub.status.idle":"2022-04-18T16:29:13.038107Z","shell.execute_reply.started":"2022-04-18T16:28:59.678633Z","shell.execute_reply":"2022-04-18T16:29:13.037069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"params = [p for p in model.parameters() if p.requires_grad]\noptimizer = torch.optim.SGD(params, lr=0.0025, momentum=0.9, weight_decay=0.0005)\n# lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.1)\nlr_scheduler = None\n\nn_batches, n_batches_val = len(dl_train), len(dl_val)\nvalidation_losses = []\n\n\nfor epoch in range(NUM_EPOCHS):\n    time_start = time.time()\n    loss_accum = 0\n    \n    for batch_idx, (images, targets) in enumerate(dl_train, 1):\n        \n        images = list(image.to(DEVICE) for image in images)\n        targets = [{k: v.to(DEVICE) for k, v in t.items()} for t in targets]\n\n        # Predict\n        loss_dict = model(images, targets)\n        losses = sum(loss for loss in loss_dict.values())\n        loss_value = losses.item()\n\n        loss_accum += loss_value\n\n        # Back-prop\n        optimizer.zero_grad()\n        losses.backward()\n        optimizer.step()\n\n    \n    # update the learning rate\n    if lr_scheduler is not None:\n        lr_scheduler.step()\n\n    # Validation \n    val_loss_accum = 0\n        \n    with torch.no_grad():\n        for batch_idx, (images, targets) in enumerate(dl_val, 1):\n            images = list(image.to(DEVICE) for image in images)\n            targets = [{k: v.to(DEVICE) for k, v in t.items()} for t in targets]\n            \n            val_loss_dict = model(images, targets)\n            val_batch_loss = sum(loss for loss in val_loss_dict.values())\n            val_loss_accum += val_batch_loss.item()\n    \n    # Logging\n    val_loss = val_loss_accum / n_batches_val\n    train_loss = loss_accum / n_batches\n    validation_losses.append(val_loss)\n    \n    # Save model\n    chk_name = f'fasterrcnn_resnet50_fpn-e{epoch}.bin'\n    torch.save(model.state_dict(), chk_name)\n    \n    \n    elapsed = time.time() - time_start\n    \n    print(f\"[Epoch {epoch+1:2d} / {NUM_EPOCHS:2d}] Train loss: {train_loss:.3f}. Val loss: {val_loss:.3f} --> {chk_name}  [{elapsed:.0f} secs]\")   ","metadata":{"execution":{"iopub.status.busy":"2022-04-18T16:29:13.122174Z","iopub.execute_input":"2022-04-18T16:29:13.122532Z","iopub.status.idle":"2022-04-18T16:31:11.467371Z","shell.execute_reply.started":"2022-04-18T16:29:13.12248Z","shell.execute_reply":"2022-04-18T16:31:11.465516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"validation_losses","metadata":{"execution":{"iopub.status.busy":"2022-04-05T04:44:17.342703Z","iopub.status.idle":"2022-04-05T04:44:17.345152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.argmin(validation_losses)","metadata":{"execution":{"iopub.status.busy":"2022-04-05T04:44:17.350147Z","iopub.status.idle":"2022-04-05T04:44:17.356064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Check result","metadata":{}},{"cell_type":"code","source":"idx = 0\n\nimages, targets = next(iter(dl_val))\nimages = list(img.to(DEVICE) for img in images)\ntargets = [{k: v.to(DEVICE) for k, v in t.items()} for t in targets]\n\nboxes = targets[idx]['boxes'].cpu().numpy().astype(np.int32)\nsample = images[idx].permute(1,2,0).cpu().numpy()\n\nmodel.eval()\n\noutputs = model(images)\noutputs = [{k: v.detach().cpu().numpy() for k, v in t.items()} for t in outputs]","metadata":{"execution":{"iopub.status.busy":"2022-04-05T04:44:17.36203Z","iopub.status.idle":"2022-04-05T04:44:17.364946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax = plt.subplots(1, 1, figsize=(16, 8))\n\n# Red for ground truth\nfor box in boxes:\n    cv2.rectangle(sample,\n                  (box[0], box[1]),\n                  (box[2], box[3]),\n                  (220, 0, 0), 3)\n\n    \n# Green for predictions\n# Print the first 5\nfor box in outputs[idx]['boxes'][:5]:\n    cv2.rectangle(sample,\n                  (box[0], box[1]),\n                  (box[2], box[3]),\n                  (0, 220, 0), 3)\n\nax.set_axis_off()\nax.imshow(sample);","metadata":{"execution":{"iopub.status.busy":"2022-04-05T04:44:17.366259Z","iopub.status.idle":"2022-04-05T04:44:17.377818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Please, _DO_ upvote if you found it useful!","metadata":{}}]}