{"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":"# 导包","metadata":{}},{"cell_type":"code","source":"import 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":"2021-12-01T09:21:08.059473Z","iopub.execute_input":"2021-12-01T09:21:08.059911Z","iopub.status.idle":"2021-12-01T09:21:11.082092Z","shell.execute_reply.started":"2021-12-01T09:21:08.059848Z","shell.execute_reply":"2021-12-01T09:21:11.080934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 设置常量","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 = 12","metadata":{"execution":{"iopub.status.busy":"2021-12-01T09:21:11.085301Z","iopub.execute_input":"2021-12-01T09:21:11.086400Z","iopub.status.idle":"2021-12-01T09:21:11.117259Z","shell.execute_reply.started":"2021-12-01T09:21:11.086334Z","shell.execute_reply":"2021-12-01T09:21:11.115975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 读取数据\n","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":"2021-12-01T09:21:11.125159Z","iopub.execute_input":"2021-12-01T09:21:11.126336Z","iopub.status.idle":"2021-12-01T09:21:11.667130Z","shell.execute_reply.started":"2021-12-01T09:21:11.126012Z","shell.execute_reply":"2021-12-01T09:21:11.665984Z"},"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":"2021-12-01T09:21:11.672190Z","iopub.execute_input":"2021-12-01T09:21:11.672927Z","iopub.status.idle":"2021-12-01T09:21:11.684863Z","shell.execute_reply.started":"2021-12-01T09:21:11.672856Z","shell.execute_reply":"2021-12-01T09:21:11.683523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\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":"2021-12-01T09:21:11.688675Z","iopub.execute_input":"2021-12-01T09:21:11.689416Z","iopub.status.idle":"2021-12-01T09:21:11.719110Z","shell.execute_reply.started":"2021-12-01T09:21:11.689164Z","shell.execute_reply":"2021-12-01T09:21:11.718061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.shape[0], df_val.shape[0]","metadata":{"execution":{"iopub.status.busy":"2021-12-01T09:21:11.721232Z","iopub.execute_input":"2021-12-01T09:21:11.721905Z","iopub.status.idle":"2021-12-01T09:21:11.730015Z","shell.execute_reply.started":"2021-12-01T09:21:11.721840Z","shell.execute_reply":"2021-12-01T09:21:11.728716Z"},"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        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":"2021-12-01T09:21:11.731687Z","iopub.execute_input":"2021-12-01T09:21:11.732610Z","iopub.status.idle":"2021-12-01T09:21:11.762953Z","shell.execute_reply.started":"2021-12-01T09:21:11.732544Z","shell.execute_reply":"2021-12-01T09:21:11.761377Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Augmentations","metadata":{}},{"cell_type":"code","source":"def 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\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":"2021-12-01T09:21:11.765266Z","iopub.execute_input":"2021-12-01T09:21:11.766056Z","iopub.status.idle":"2021-12-01T09:21:11.781851Z","shell.execute_reply.started":"2021-12-01T09:21:11.765816Z","shell.execute_reply":"2021-12-01T09:21:11.780515Z"},"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":"2021-12-01T09:21:11.784142Z","iopub.execute_input":"2021-12-01T09:21:11.785261Z","iopub.status.idle":"2021-12-01T09:21:11.795409Z","shell.execute_reply.started":"2021-12-01T09:21:11.785195Z","shell.execute_reply":"2021-12-01T09:21:11.793925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 查验一个样例","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":"2021-12-01T09:21:11.797920Z","iopub.execute_input":"2021-12-01T09:21:11.798926Z","iopub.status.idle":"2021-12-01T09:21:11.856153Z","shell.execute_reply.started":"2021-12-01T09:21:11.798861Z","shell.execute_reply":"2021-12-01T09:21:11.854917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image, targets = ds_train[2200]\nimage","metadata":{"execution":{"iopub.status.busy":"2021-12-01T09:21:11.857978Z","iopub.execute_input":"2021-12-01T09:21:11.858820Z","iopub.status.idle":"2021-12-01T09:21:12.057487Z","shell.execute_reply.started":"2021-12-01T09:21:11.858736Z","shell.execute_reply":"2021-12-01T09:21:12.056326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"targets","metadata":{"execution":{"iopub.status.busy":"2021-12-01T09:21:12.062159Z","iopub.execute_input":"2021-12-01T09:21:12.062501Z","iopub.status.idle":"2021-12-01T09:21:12.076029Z","shell.execute_reply.started":"2021-12-01T09:21:12.062440Z","shell.execute_reply":"2021-12-01T09:21:12.074439Z"},"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":"2021-12-01T09:21:12.077692Z","iopub.execute_input":"2021-12-01T09:21:12.078431Z","iopub.status.idle":"2021-12-01T09:21:12.585585Z","shell.execute_reply.started":"2021-12-01T09:21:12.078349Z","shell.execute_reply":"2021-12-01T09:21:12.584113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 数据加载器","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":"2021-12-01T09:21:12.587018Z","iopub.execute_input":"2021-12-01T09:21:12.587405Z","iopub.status.idle":"2021-12-01T09:21:12.596003Z","shell.execute_reply.started":"2021-12-01T09:21:12.587351Z","shell.execute_reply":"2021-12-01T09:21:12.594956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 创建模型","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":"2021-12-01T09:21:12.597782Z","iopub.execute_input":"2021-12-01T09:21:12.598678Z","iopub.status.idle":"2021-12-01T09:21:33.626274Z","shell.execute_reply.started":"2021-12-01T09:21:12.598570Z","shell.execute_reply":"2021-12-01T09:21:33.625119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 训练模型","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":"2021-12-01T09:21:33.627942Z","iopub.execute_input":"2021-12-01T09:21:33.628480Z","iopub.status.idle":"2021-12-01T11:57:56.748960Z","shell.execute_reply.started":"2021-12-01T09:21:33.628410Z","shell.execute_reply":"2021-12-01T11:57:56.747782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"validation_losses","metadata":{"execution":{"iopub.status.busy":"2021-12-01T11:57:56.751193Z","iopub.execute_input":"2021-12-01T11:57:56.751598Z","iopub.status.idle":"2021-12-01T11:57:56.759732Z","shell.execute_reply.started":"2021-12-01T11:57:56.751523Z","shell.execute_reply":"2021-12-01T11:57:56.758380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.argmin(validation_losses)","metadata":{"execution":{"iopub.status.busy":"2021-12-01T11:57:56.761732Z","iopub.execute_input":"2021-12-01T11:57:56.762587Z","iopub.status.idle":"2021-12-01T11:57:56.775250Z","shell.execute_reply.started":"2021-12-01T11:57:56.762521Z","shell.execute_reply":"2021-12-01T11:57:56.773812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 检查结果","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":"2021-12-01T11:57:56.777781Z","iopub.execute_input":"2021-12-01T11:57:56.778851Z","iopub.status.idle":"2021-12-01T11:57:59.517757Z","shell.execute_reply.started":"2021-12-01T11:57:56.778761Z","shell.execute_reply":"2021-12-01T11:57:59.516616Z"},"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":"2021-12-01T11:57:59.523124Z","iopub.execute_input":"2021-12-01T11:57:59.523468Z","iopub.status.idle":"2021-12-01T11:57:59.996089Z","shell.execute_reply.started":"2021-12-01T11:57:59.523396Z","shell.execute_reply":"2021-12-01T11:57:59.995060Z"},"trusted":true},"execution_count":null,"outputs":[]}]}