{"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":"# Import","metadata":{}},{"cell_type":"code","source":"# import pytorch and use fasterRCNN \nimport cv2\nimport time\nimport random\nimport numpy as np\nimport pandas as pd\nfrom matplotlib import pyplot as plt\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\nimport torch\nimport torchvision\nfrom torch.utils.data import DataLoader\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection import FasterRCNN","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:28:33.152203Z","iopub.execute_input":"2022-11-15T22:28:33.152854Z","iopub.status.idle":"2022-11-15T22:28:38.034474Z","shell.execute_reply.started":"2022-11-15T22:28:33.152761Z","shell.execute_reply":"2022-11-15T22:28:38.033466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# the path and some values \nDEVICE = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\nBASE_DIR = \"../input/tensorflow-great-barrier-reef/train_images/\"\nNUM_EPOCHS = 8","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:29:17.496842Z","iopub.execute_input":"2022-11-15T22:29:17.497283Z","iopub.status.idle":"2022-11-15T22:29:17.565197Z","shell.execute_reply.started":"2022-11-15T22:29:17.497247Z","shell.execute_reply":"2022-11-15T22:29:17.564235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Spliting Data","metadata":{}},{"cell_type":"code","source":"#read the csv file\ndf = pd.read_csv(\"/kaggle/input/tensorflow-great-barrier-reef/train.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:28:44.771161Z","iopub.execute_input":"2022-11-15T22:28:44.772355Z","iopub.status.idle":"2022-11-15T22:28:44.822215Z","shell.execute_reply.started":"2022-11-15T22:28:44.772308Z","shell.execute_reply":"2022-11-15T22:28:44.821317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# non-empty\ndf_train = df.copy().loc[df[\"annotations\"].astype(str) != \"[]\"]\ndf_train['annotations'] = df_train['annotations'].apply(eval)\n# Add image path \ndef image_path(r):\n    video_id = r['video_id']\n    video_frame = r['video_frame']\n    return  \"video_\" + str(video_id) + \"/\" + str(video_frame) + \".jpg\"\ndf_train['image_path'] = df_train.apply(lambda x: image_path(x), axis=1)\n# split the image with annotation into train and test\nimage_ids = df_train['image_id'].unique()\nfrom sklearn.model_selection import train_test_split\ntrain_ids, test_ids = train_test_split(\n    image_ids,  test_size=0.2, random_state=42)\n\ntrain_df = df_train[df_train['image_id'].isin(train_ids)]\ntest_df = df_train[df_train['image_id'].isin(test_ids)]\ntest_df.shape, train_df.shape","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:28:46.995942Z","iopub.execute_input":"2022-11-15T22:28:46.996321Z","iopub.status.idle":"2022-11-15T22:28:47.220998Z","shell.execute_reply.started":"2022-11-15T22:28:46.996289Z","shell.execute_reply":"2022-11-15T22:28:47.220123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## empty\ndf_empty=df.copy().loc[df[\"annotations\"].astype(str) == \"[]\"]\ndf_empty['annotations'] = df_empty['annotations'].apply(eval)\n#Add image path \ndef image_path(r):\n    video_id = r['video_id']\n    video_frame = r['video_frame']\n    return  \"video_\" + str(video_id) + \"/\" + str(video_frame) + \".jpg\"\ndf_empty['image_path'] = df_empty.apply(lambda x: image_path(x), axis=1)\ndf_empty=df_empty.sample(frac=0.1)\n# split the image without annotation into train and test\nimage_empty = df_empty['image_id'].unique()\nfrom sklearn.model_selection import train_test_split\ntrain_empty, test_empty = train_test_split(\n    image_empty,  test_size=0.2, random_state=42)\ntrain_empty = df_empty[df_empty['image_id'].isin(train_empty)]\ntest_empty = df_empty[df_empty['image_id'].isin(test_empty)]\ntest_empty.shape, train_empty.shape","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:28:49.914435Z","iopub.execute_input":"2022-11-15T22:28:49.914820Z","iopub.status.idle":"2022-11-15T22:28:50.261165Z","shell.execute_reply.started":"2022-11-15T22:28:49.914786Z","shell.execute_reply":"2022-11-15T22:28:50.260299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# combine empty and none empty dataform\ntrain_df = pd.concat([train_df,train_empty])\ntest_df = pd.concat([test_df,test_empty])\ndf_train.shape, test_df.shape","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:28:52.823690Z","iopub.execute_input":"2022-11-15T22:28:52.824076Z","iopub.status.idle":"2022-11-15T22:28:52.837691Z","shell.execute_reply.started":"2022-11-15T22:28:52.824041Z","shell.execute_reply":"2022-11-15T22:28:52.836305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# split to train and valid dataform\ntrain_d, val_d = train_test_split(train_df, test_size=0.2, random_state=42)\ntrain_df=train_d\nval_df = val_d","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:28:55.332796Z","iopub.execute_input":"2022-11-15T22:28:55.333847Z","iopub.status.idle":"2022-11-15T22:28:55.343030Z","shell.execute_reply.started":"2022-11-15T22:28:55.333799Z","shell.execute_reply":"2022-11-15T22:28:55.341876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_df.shape,train_df.shape","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:28:57.985298Z","iopub.execute_input":"2022-11-15T22:28:57.986191Z","iopub.status.idle":"2022-11-15T22:28:57.992953Z","shell.execute_reply.started":"2022-11-15T22:28:57.986145Z","shell.execute_reply":"2022-11-15T22:28:57.991843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"#Define a map style Datasets\nclass StarfishDataset:\n\n    def __init__(self, df, transforms=None):\n        self.df = df\n        self.transforms = transforms\n    # read the image,faster-rcnn model expects input to be in range [0-1]\n    def get_image(self, row):\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        # Picture normalization\n        image /= 255.0\n        return image\n    \n    # Get the bounding boxes\n    def get_boxes(self, row):\n        boxes = pd.DataFrame(row['annotations'], columns=['x', 'y', 'width', 'height']).astype(float).values \n        boxes[:, 2] = boxes[:, 0] + boxes[:, 2]\n        boxes[:, 3] = boxes[:, 1] + boxes[:, 3]\n        # check if the bboxes are are valid，image pixels are 1280*720\n        boxes[:, 2] = np.clip(boxes[:, 2], 0, 1280)\n        boxes[:, 3] = np.clip(boxes[:, 3], 0, 720)\n        return boxes\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        # Create a target dictionary\n        target = {\n            'boxes': torch.as_tensor(boxes, dtype=torch.float32),\n            'area': torch.as_tensor(area, dtype=torch.float32),\n            'image_id': torch.tensor([i]),\n            'labels': torch.ones((n_boxes,), dtype=torch.int64),\n            'iscrowd': torch.zeros((n_boxes,), dtype=torch.int64)            \n        }\n\n        if self.transforms:\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-11-15T22:29:00.216273Z","iopub.execute_input":"2022-11-15T22:29:00.217406Z","iopub.status.idle":"2022-11-15T22:29:00.231517Z","shell.execute_reply.started":"2022-11-15T22:29:00.217351Z","shell.execute_reply":"2022-11-15T22:29:00.230389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# here we experiment image augmentation using simple transformations \n# 50% probability of horizontal flip and 50% probability of vertical flip\n# we resize all of pictures as 720*1280\ndef get_train_transform():\n    return A.Compose([\n        A.Flip(0.5),\n        A.HorizontalFlip(p=0.5),   \n        A.VerticalFlip(p=0.5),\n        A.Resize(height=720, width=1280, p=1.0),\n        ToTensorV2(p=1.0)\n    ], bbox_params={'format': 'pascal_voc', 'label_fields': ['labels']})\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-11-15T22:29:04.486081Z","iopub.execute_input":"2022-11-15T22:29:04.486781Z","iopub.status.idle":"2022-11-15T22:29:04.493614Z","shell.execute_reply.started":"2022-11-15T22:29:04.486745Z","shell.execute_reply":"2022-11-15T22:29:04.492615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_train = StarfishDataset(train_df, get_train_transform())\nds_val = StarfishDataset(val_df, get_valid_transform())","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:29:08.479076Z","iopub.execute_input":"2022-11-15T22:29:08.479744Z","iopub.status.idle":"2022-11-15T22:29:08.484486Z","shell.execute_reply.started":"2022-11-15T22:29:08.479710Z","shell.execute_reply":"2022-11-15T22:29:08.483434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Check the sample","metadata":{}},{"cell_type":"code","source":"image, targets = ds_train[2560]","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:29:22.263024Z","iopub.execute_input":"2022-11-15T22:29:22.263729Z","iopub.status.idle":"2022-11-15T22:29:22.373931Z","shell.execute_reply.started":"2022-11-15T22:29:22.263690Z","shell.execute_reply":"2022-11-15T22:29:22.372920Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# In image,show the starfish in the boxes \nboxes = 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-11-15T22:29:24.388641Z","iopub.execute_input":"2022-11-15T22:29:24.389004Z","iopub.status.idle":"2022-11-15T22:29:25.068504Z","shell.execute_reply.started":"2022-11-15T22:29:24.388972Z","shell.execute_reply":"2022-11-15T22:29:25.067262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create model and load data to it","metadata":{}},{"cell_type":"code","source":"# Create PyTorch DataLoader\n#load the data for the model\ndef collate_fn(batch):\n    return tuple(zip(*batch))\n#Due to the train speed, we adjust batch_size=8 and num_worker=2\ndl_train = DataLoader(ds_train, batch_size=8, shuffle=False, num_workers=2, collate_fn=collate_fn)\ndl_val = DataLoader(ds_val, batch_size=8, shuffle=False, num_workers=2, collate_fn=collate_fn)","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:29:32.103814Z","iopub.execute_input":"2022-11-15T22:29:32.104218Z","iopub.status.idle":"2022-11-15T22:29:32.111933Z","shell.execute_reply.started":"2022-11-15T22:29:32.104164Z","shell.execute_reply":"2022-11-15T22:29:32.109808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model():\n    # load a resnet50 model and pre-train on COCO\n    model = torchvision.models.detection.fasterrcnn_resnet50_fpn(pretrained=True)\n    # 1 class (starfish) + background\n    num_classes = 2  \n    # get number of input features for the classifier\n    in_features = model.roi_heads.box_predictor.cls_score.in_features\n    # replace the pre-trained head with a new one\n    model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes)\n    model.to(DEVICE)\n    return model\n\nmodel = get_model()","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:29:35.171390Z","iopub.execute_input":"2022-11-15T22:29:35.172272Z","iopub.status.idle":"2022-11-15T22:29:40.617966Z","shell.execute_reply.started":"2022-11-15T22:29:35.172229Z","shell.execute_reply":"2022-11-15T22:29:40.616897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"# start to train and use torch.SGD to optimize\nparams = [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)\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}   [{elapsed:.0f} secs]\")","metadata":{"execution":{"iopub.status.busy":"2022-11-15T10:56:37.859219Z","iopub.execute_input":"2022-11-15T10:56:37.859515Z","iopub.status.idle":"2022-11-15T12:34:43.933548Z","shell.execute_reply.started":"2022-11-15T10:56:37.859487Z","shell.execute_reply":"2022-11-15T12:34:43.931313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# save the model after train\ntorch.save(model, \"my_model.pth\") ","metadata":{"execution":{"iopub.status.busy":"2022-11-15T12:35:42.684981Z","iopub.execute_input":"2022-11-15T12:35:42.685350Z","iopub.status.idle":"2022-11-15T12:35:43.027584Z","shell.execute_reply.started":"2022-11-15T12:35:42.685317Z","shell.execute_reply":"2022-11-15T12:35:43.026483Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict","metadata":{}},{"cell_type":"code","source":"import random as rand\n\ndef showbbox(model, img):     \n    model.eval()\n    with torch.no_grad():\n        prediction = model([img.to(device)])\n    print(prediction)\n    \n    # we set the threshold to 0.5\n    threshold=0.5\n    count=0\n    d=[]\n    \n    #tranfrom C,H,W to H,W,C\n    img = img.permute(1,2,0) \n    #rotate 0-255\n    img = (img * 255).byte().data.cpu() \n    img = np.array(img)  \n    \n    # draw the boxes\n    for i in range(prediction[0]['boxes'].cpu().shape[0]):\n        count+=1\n        if prediction[0]['scores'][i]<threshold:\n            count-=1\n            continue\n        xmin = round(prediction[0]['boxes'][i][0].item())\n        ymin = round(prediction[0]['boxes'][i][1].item())\n        xmax = round(prediction[0]['boxes'][i][2].item())\n        ymax = round(prediction[0]['boxes'][i][3].item())\n        d.append([xmin,ymin,xmax,ymax])\n        label = prediction[0]['labels'][i].item()\n        \n        # put the label on the boxes \n        if label == 1:\n            cv2.rectangle(img, (xmin, ymin), (xmax, ymax), (int(rand.random() * 255),int(rand.random() * 255),int(rand.random() * 255)), thickness=2)\n            cv2.putText(img, 'Starfish', (xmin, ymin), cv2.FONT_HERSHEY_SIMPLEX, 0.7, (int(rand.random() * 255),int(rand.random() * 255),int(rand.random() * 255)),\n                               thickness=2)\n    \n    plt.figure(figsize=(20,15))\n    plt.imshow(img)\n    print(f'The number of starfish is equal to {count}.')\n    return d","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:29:43.208365Z","iopub.execute_input":"2022-11-15T22:29:43.209341Z","iopub.status.idle":"2022-11-15T22:29:43.239397Z","shell.execute_reply.started":"2022-11-15T22:29:43.209267Z","shell.execute_reply":"2022-11-15T22:29:43.237878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = torch.load(r'/kaggle/input/my-model-rcnn/my_model.pth')\ndevice = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\nmodel.to(device)\n\nimg, _ = ds_val[700] \na=showbbox(model, img)","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:29:46.184945Z","iopub.execute_input":"2022-11-15T22:29:46.185327Z","iopub.status.idle":"2022-11-15T22:29:55.399582Z","shell.execute_reply.started":"2022-11-15T22:29:46.185293Z","shell.execute_reply":"2022-11-15T22:29:55.397377Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Test","metadata":{}},{"cell_type":"code","source":"test_set = StarfishDataset(test_df, get_train_transform())","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:30:02.002395Z","iopub.execute_input":"2022-11-15T22:30:02.002770Z","iopub.status.idle":"2022-11-15T22:30:02.008256Z","shell.execute_reply.started":"2022-11-15T22:30:02.002738Z","shell.execute_reply":"2022-11-15T22:30:02.007071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def showbbox_2(model, img):     \n    model.eval()\n    with torch.no_grad():\n        prediction = model([img.to(device)])\n    \n    # we set the threshold to 0.5\n    threshold=0.5\n    count=0\n    d=[]\n    \n    #tranfrom C,H,W to H,W,C\n    img = img.permute(1,2,0) \n    #rotate 0-255\n    img = (img * 255).byte().data.cpu() \n    img = np.array(img)  \n    \n    # draw the boxes\n    for i in range(prediction[0]['boxes'].cpu().shape[0]):\n        count+=1\n        if prediction[0]['scores'][i]<threshold:\n            count-=1\n            continue\n        xmin = round(prediction[0]['boxes'][i][0].item())\n        ymin = round(prediction[0]['boxes'][i][1].item())\n        xmax = round(prediction[0]['boxes'][i][2].item())\n        ymax = round(prediction[0]['boxes'][i][3].item())\n        d.append([xmin,ymin,xmax,ymax])\n        label = prediction[0]['labels'][i].item()\n        \n        # put the label on the boxes \n        if label == 1:\n            cv2.rectangle(img, (xmin, ymin), (xmax, ymax), (int(rand.random() * 255),int(rand.random() * 255),int(rand.random() * 255)), thickness=2)\n            cv2.putText(img, 'Starfish', (xmin, ymin), cv2.FONT_HERSHEY_SIMPLEX, 0.7, (int(rand.random() * 255),int(rand.random() * 255),int(rand.random() * 255)),\n                               thickness=2)\n    return d","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:30:05.553062Z","iopub.execute_input":"2022-11-15T22:30:05.553778Z","iopub.status.idle":"2022-11-15T22:30:05.566922Z","shell.execute_reply.started":"2022-11-15T22:30:05.553741Z","shell.execute_reply":"2022-11-15T22:30:05.565527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Iou function\ndef iou(a,b):\n    ximin=max(a[0],b[0])\n    ximax=min(a[2],b[2])\n    yimin=max(a[1],b[1])\n    yimax=min(a[3],b[3])\n    intersection=(ximax-ximin)*(yimax-yimin)\n    \n    AeraA = (a[2]-a[0])*(a[3]-a[1])\n    AeraB = (b[2]-b[0])*(b[3]-b[1])\n    union = AeraA+AeraB-intersection\n    \n    return intersection/union","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:30:08.448839Z","iopub.execute_input":"2022-11-15T22:30:08.449318Z","iopub.status.idle":"2022-11-15T22:30:08.458166Z","shell.execute_reply.started":"2022-11-15T22:30:08.449277Z","shell.execute_reply":"2022-11-15T22:30:08.457130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# f2 evaluation\nf2list=[]\nfor IOU in range(30,85,5):\n    FP,TP,FN=0,0,0\n    # for every image\n    for image_index in range(len(test_set)):\n        img,_ = test_set[image_index] # answer\n        a = showbbox_2(model,img) # prediction\n        unmatched=[1 for i in range(len(_[\"boxes\"]))]\n        if unmatched == []:\n            if len(a)==0:\n                TP+=1\n                continue\n            else: \n                FP+=1\n                continue\n        # for every box in the prediction\n        for i in range(len(a)):\n            iouMax=0\n            matchedBoxIndex=0\n            # get iouMax of one box\n            for j in range(len(_[\"boxes\"])):\n                if iouMax < iou(_[\"boxes\"][j],a[i]):\n                    iouMax=iou(_[\"boxes\"][j],a[i])\n                    matchedBoxIndex=j\n            # compare iouMax with IOU\n            if iouMax > IOU*0.01:\n                TP+=1 \n            else :\n                FP+=1\n            unmatched[matchedBoxIndex]=0\n        FN += sum(unmatched)\n    f2=5*TP/(5*TP+4*FN+FP)\n    print(f2)\n    f2list.append(f2)","metadata":{"execution":{"iopub.status.busy":"2022-11-15T22:46:30.501444Z","iopub.execute_input":"2022-11-15T22:46:30.501803Z","iopub.status.idle":"2022-11-15T23:16:57.865541Z","shell.execute_reply.started":"2022-11-15T22:46:30.501771Z","shell.execute_reply":"2022-11-15T23:16:57.864599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mean = sum(f2list)/11\nprint(mean)","metadata":{"execution":{"iopub.status.busy":"2022-11-15T23:17:01.586236Z","iopub.execute_input":"2022-11-15T23:17:01.586602Z","iopub.status.idle":"2022-11-15T23:17:01.592119Z","shell.execute_reply.started":"2022-11-15T23:17:01.586566Z","shell.execute_reply":"2022-11-15T23:17:01.591049Z"},"trusted":true},"execution_count":null,"outputs":[]}]}