{"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":"# Imports","metadata":{}},{"cell_type":"code","source":"# Very few imports. This is a pure torch solution!\nimport os\nimport shutil\nimport cv2\nimport json\nimport torch\nimport random\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom sklearn.model_selection import train_test_split\nfrom tqdm.notebook import tqdm\nfrom torch.utils.data import DataLoader","metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","execution":{"iopub.status.busy":"2022-11-17T12:06:07.80648Z","iopub.execute_input":"2022-11-17T12:06:07.806907Z","iopub.status.idle":"2022-11-17T12:06:10.510229Z","shell.execute_reply.started":"2022-11-17T12:06:07.806801Z","shell.execute_reply":"2022-11-17T12:06:10.509207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# reletive path (src/tensorflow-great-barrier-reef)\nDATASET_PATH = \"tensorflow-great-barrier-reef\"\n\n# update dataset\nUPDATE_DATASET = False\n\n# remove no label images\nREMOVE_NON_LABELED = True\n\n# training data percentage\nTRAINING_SIZE = 0.8\n\n# turn on the debug\nDEBUGGING_PREPROCESS = False\n\n# image size\nIMAGE_SIZE = (1280, 720)\n\n# if train preprocess\nTRAIN_PREPROCESS = False","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:10.51263Z","iopub.execute_input":"2022-11-17T12:06:10.513201Z","iopub.status.idle":"2022-11-17T12:06:10.520264Z","shell.execute_reply.started":"2022-11-17T12:06:10.513161Z","shell.execute_reply":"2022-11-17T12:06:10.517936Z"},"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')\nDATASET_PATH = \"../input/tensorflow-great-barrier-reef/\"\n\nNUM_EPOCHS = 10","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:10.521771Z","iopub.execute_input":"2022-11-17T12:06:10.52246Z","iopub.status.idle":"2022-11-17T12:06:10.602716Z","shell.execute_reply.started":"2022-11-17T12:06:10.522422Z","shell.execute_reply":"2022-11-17T12:06:10.601577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load data\n","metadata":{}},{"cell_type":"code","source":"# list of String(image path)\nTOTAL_images = []\n\n# list of dict(e.g. [{'x': 559, 'y': 213, 'width': 50, 'height': 32}])\nTOTAL_labels = []\n\nimage_dict = {\n    \"0\":[],\n    \"1\":[],\n    \"2\":[]\n}\nlabel_dict = {\n    \"0\":[],\n    \"1\":[],\n    \"2\":[]\n}\n\ntrain_CSV = pd.read_csv(f\"../input/tensorflow-great-barrier-reef/train.csv\")\nsize = len(train_CSV[\"video_id\"])\nfor row in range(size):\n    video_id = train_CSV[\"video_id\"][row]\n    img_id = train_CSV[\"video_frame\"][row]\n    img = f\"{DATASET_PATH}/train_images/video_{video_id}/{img_id}.jpg\"\n\n    labels = list(eval(train_CSV[\"annotations\"][row]))\n    if REMOVE_NON_LABELED and len(labels) == 0:\n        \n        continue\n\n#     TOTAL_images.append(img)\n#     TOTAL_labels.append(labels)\n    image_dict[str(video_id)].append(img)\n    label_dict[str(video_id)].append(labels)\n\nprint(f\"Total number of images collected: {len(TOTAL_images)}\")\n\nimage_train_val = image_dict[\"0\"]\nimage_train_val.extend(image_dict[\"1\"])\nlabel_train_val = label_dict[\"0\"]\nlabel_train_val.extend(label_dict[\"1\"])\n\ntrain_set,valid_set, train_label,valid_label = train_test_split(image_train_val,label_train_val,test_size=0.2)\n# train_set = image_dict[\"0\"]\n# valid_set = image_dict[\"1\"]\ntest_set = image_dict[\"2\"]\n\n# train_label = label_dict[\"0\"]\n# valid_label = label_dict[\"1\"]\ntest_label = label_dict[\"2\"]\nprint(f\"length train set is {len(train_set)} and length of valid set is {len(valid_set)}\")\n\n\nprint(f\"Traning size: {len(train_set)}\")\nprint(f\"Test size: {len(test_set)}\")\n\n# show first image and its labels\nimg_0 = cv2.cvtColor(cv2.imread(train_set[0]), cv2.COLOR_BGR2RGB)\nprint(f\"Image size: {img_0.size}\")\nprint(str(train_label[0]))\n\n# drow labels on the image\nfor dirct in train_label[0]:\n    x = dirct[\"x\"]\n    y = dirct[\"y\"]\n    width = dirct[\"width\"]\n    height = dirct[\"height\"]\n    start_point = (x, y)\n    end_point = (x+width, y+height)\n    img_0 = cv2.rectangle(img_0, start_point, end_point, (255,255,0), 2)\nplt.imshow(img_0)\nplt.axis(\"off\")\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:10.605545Z","iopub.execute_input":"2022-11-17T12:06:10.606217Z","iopub.status.idle":"2022-11-17T12:06:11.612502Z","shell.execute_reply.started":"2022-11-17T12:06:10.606181Z","shell.execute_reply":"2022-11-17T12:06:11.611439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset class","metadata":{}},{"cell_type":"code","source":"def preProcess(img_path):\n    \"\"\"Preprocess the image\n    Args:\n        img (string): image path\n\n    Returns:\n        cv2.Mat: image\n    \"\"\"\n    # read image\n    img = cv2.imread(img_path)\n\n    # image enhancement\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    img = cv2.cvtColor(hsv_img, cv2.COLOR_HSV2RGB)\n\n    # gamma\n    R = 255.0\n    img = (R * np.power(img.astype(np.uint32)/R, 1/0.6)).astype(np.uint8)\n\n    # sharpening\n    kernel = np.array([[0, -1, 0],\n                    [-1, 5,-1],\n                    [0, -1, 0]])\n    img = cv2.filter2D(src=img, ddepth=-1, kernel=kernel)\n\n    # improve contrast\n    alpha = 1.1 # Contrast control (1.0-3.0)\n    beta = 0 # Brightness control (0-100)\n    img = cv2.convertScaleAbs(img, alpha=alpha, beta=beta)\n\n    # reduce size\n    scale = 1\n    x = int(img.shape[0] * scale)\n    y = int(img.shape[1] * scale)\n    img = cv2.resize(img, dsize=(y, x), interpolation=cv2.INTER_AREA)\n\n    return img","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:11.613676Z","iopub.execute_input":"2022-11-17T12:06:11.614116Z","iopub.status.idle":"2022-11-17T12:06:11.6271Z","shell.execute_reply.started":"2022-11-17T12:06:11.614079Z","shell.execute_reply":"2022-11-17T12:06:11.62604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUGGING_PREPROCESS:\n    # random 5 images for preProcess eval\n    random_5_img = []\n    random_5_label = []\n\n    for _ in range(5):\n        idx = random.choice(range(0, len(train_set)))\n        random_5_img.append(TOTAL_images[idx])\n        random_5_label.append(TOTAL_labels[idx])\n\n    for i in range(len(random_5_img)):\n        img = preProcess(random_5_img[i])\n\n        print(str(random_5_label[i]))\n\n        # draw rectangle\n        org_img = cv2.imread(random_5_img[i])\n        org_img = cv2.cvtColor(org_img, cv2.COLOR_BGR2RGB)\n        for lable in random_5_label[i]:\n            x = lable[\"x\"]\n            y = lable[\"y\"]\n            width = lable[\"width\"]\n            height = lable[\"height\"]\n            start_point = (x-width, y-height)\n            end_point = (x+width, y+height)\n            org_img = cv2.rectangle(org_img, start_point, end_point, (255,255,0), 2)\n            img = cv2.rectangle(img, start_point, end_point, (255,255,0), 2)\n\n        plt.imshow(org_img)\n        plt.axis(\"off\")\n        plt.show()\n\n        plt.imshow(img)\n        plt.axis(\"off\")\n        plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:11.628852Z","iopub.execute_input":"2022-11-17T12:06:11.629734Z","iopub.status.idle":"2022-11-17T12:06:11.642537Z","shell.execute_reply.started":"2022-11-17T12:06:11.629696Z","shell.execute_reply":"2022-11-17T12:06:11.641394Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"org_img = cv2.imread(train_set[3])\norg_img = cv2.cvtColor(org_img, cv2.COLOR_BGR2RGB)\nplt.imshow(org_img)\nplt.axis(\"off\")\nplt.show()\nprint(train_label[3])\n","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:11.645983Z","iopub.execute_input":"2022-11-17T12:06:11.646497Z","iopub.status.idle":"2022-11-17T12:06:11.950865Z","shell.execute_reply.started":"2022-11-17T12:06:11.646467Z","shell.execute_reply":"2022-11-17T12:06:11.950025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Dataset:\n\n    def __init__(self, images,boxes,transforms=None):\n        self.images = images \n        self.boxes = boxes\n        self.transforms = transforms\n    def can_augment(self, boxes):\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, i):\n        \"\"\"Returns boxes in the form of [x_min, y_min, x_max, y_max], requires conversion\"\"\"\n        \n        boxes = np.array([list(idx.values()) for idx in self.boxes[i]]).astype(float) # Convert into list now\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, i):\n        \"\"\"Gets the image\"\"\"\n        if TRAIN_PREPROCESS:\n            image = preProcess(self.images[i])\n        else:  \n            image = cv2.imread(self.images[i])\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        Return image and item\n        '''\n        image = self.get_image(i)\n        boxes = self.get_boxes(i)\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            # There is only one class\n            'labels': torch.ones((n_boxes,), dtype=torch.int64),           \n        }\n\n\n        sample = {\n            'image':image,\n            'bboxes': target['boxes'],\n            'labels': target['labels']\n        }\n\n        if n_boxes > 0:\n            target['boxes'] = torch.stack(tuple(map(torch.tensor, zip(*sample['bboxes'])))).permute(1, 0)\n\n        image = torch.Tensor(image).permute(2,0,1)\n        return image, target\n\n    def __len__(self):\n        return len(self.images)","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:11.952424Z","iopub.execute_input":"2022-11-17T12:06:11.953147Z","iopub.status.idle":"2022-11-17T12:06:11.968853Z","shell.execute_reply.started":"2022-11-17T12:06:11.953106Z","shell.execute_reply":"2022-11-17T12:06:11.967928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_train = Dataset(train_set,train_label)\nds_val = Dataset(valid_set,valid_label)\nds_test = Dataset(test_set,test_label)","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:11.970384Z","iopub.execute_input":"2022-11-17T12:06:11.970739Z","iopub.status.idle":"2022-11-17T12:06:11.983561Z","shell.execute_reply.started":"2022-11-17T12:06:11.970704Z","shell.execute_reply":"2022-11-17T12:06:11.982604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Faster R-CNN","metadata":{}},{"cell_type":"code","source":"LEARNING_RATE = 0.002\nBATCH_SIZE = 4\nIOU_TARGET = 0.45","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:11.988699Z","iopub.execute_input":"2022-11-17T12:06:11.989282Z","iopub.status.idle":"2022-11-17T12:06:11.993862Z","shell.execute_reply.started":"2022-11-17T12:06:11.989251Z","shell.execute_reply":"2022-11-17T12:06:11.992978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision import models\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nmodel = models.detection.fasterrcnn_resnet50_fpn()\nnum_class = 2  # object and background\n\n\nin_features = model.roi_heads.box_predictor.cls_score.in_features\n\nmodel.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_class)","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:11.995255Z","iopub.execute_input":"2022-11-17T12:06:11.996285Z","iopub.status.idle":"2022-11-17T12:06:17.327497Z","shell.execute_reply.started":"2022-11-17T12:06:11.996249Z","shell.execute_reply":"2022-11-17T12:06:17.326539Z"},"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\ntrain_loader = DataLoader(ds_train, batch_size=4, shuffle=True, drop_last=False,num_workers=2, collate_fn=collate_fn)\nval_loader = DataLoader(ds_val, batch_size=4, shuffle=True, drop_last=False,num_workers=2, collate_fn=collate_fn)\ntest_loader = DataLoader(ds_test,batch_size = 4,shuffle = True,drop_last=False,num_workers=2, collate_fn=collate_fn)","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:17.328859Z","iopub.execute_input":"2022-11-17T12:06:17.32933Z","iopub.status.idle":"2022-11-17T12:06:17.33724Z","shell.execute_reply.started":"2022-11-17T12:06:17.329294Z","shell.execute_reply":"2022-11-17T12:06:17.335454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create the model","metadata":{}},{"cell_type":"code","source":"LOAD_MODEL = False","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:17.338656Z","iopub.execute_input":"2022-11-17T12:06:17.339206Z","iopub.status.idle":"2022-11-17T12:06:17.362851Z","shell.execute_reply.started":"2022-11-17T12:06:17.339179Z","shell.execute_reply":"2022-11-17T12:06:17.356745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_checkpoint(state,epoch):\n    # .../abc.pth.tar\n    base_filename = \"../working/models/\"\n    file_name = base_filename+\"model_epoch\"+str(epoch)+\".pth.tar\"\n    print(\"> Saving the State of model now\")\n    torch.save(state,file_name)","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:17.366075Z","iopub.execute_input":"2022-11-17T12:06:17.369121Z","iopub.status.idle":"2022-11-17T12:06:17.378006Z","shell.execute_reply.started":"2022-11-17T12:06:17.369081Z","shell.execute_reply":"2022-11-17T12:06:17.376927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_checkpoint(model,optimizer,checkpoint):\n    print(\"Loading State of Model now\")\n    model.load_state_dict(checkpoint['state'])\n    optimizer.load_state_dict(checkpoint['optimizer'])\n    epoch = checkpoint['epoch']\n    losses_list = checkpoint['losses_list']\n    return model,optimizer, epoch, losses_list","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:17.379747Z","iopub.execute_input":"2022-11-17T12:06:17.380494Z","iopub.status.idle":"2022-11-17T12:06:17.391094Z","shell.execute_reply.started":"2022-11-17T12:06:17.380449Z","shell.execute_reply":"2022-11-17T12:06:17.387999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_model(model,optimizer,file_name):\n    file_name = file_name # Need to change this\n    base_filename = \"../working/models/\"\n    checkout,optimizer,epoch,losses_list = load_checkpoint(model,optimizer,torch.load(base_filename+file_name))\n    return model, optimizer , epoch, losses_list\n    ","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:17.393239Z","iopub.execute_input":"2022-11-17T12:06:17.398789Z","iopub.status.idle":"2022-11-17T12:06:17.412641Z","shell.execute_reply.started":"2022-11-17T12:06:17.398746Z","shell.execute_reply":"2022-11-17T12:06:17.409937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision import models\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nmodel = models.detection.fasterrcnn_resnet50_fpn()\nnum_class = 2  # object and background\n\n# in_features = model.roi_heads.box_predictor.cls_score.in_features\n\n# model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_class)","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:17.418885Z","iopub.execute_input":"2022-11-17T12:06:17.42207Z","iopub.status.idle":"2022-11-17T12:06:18.394821Z","shell.execute_reply.started":"2022-11-17T12:06:17.422034Z","shell.execute_reply":"2022-11-17T12:06:18.393835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\ndevice = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\n\nmodel.to(device)\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)\n\nval_list = []\nbase_epoch = 0\n                             \nt_batches = len(train_loader)\nv_batches = len(val_loader)    \n\ndef F_RCNN_val(valid_loader):\n    val_loss_sum = 0\n    loss = 0\n    with torch.no_grad():\n        for b_id ,(images,targets) in enumerate(valid_loader, 1):\n            images1 = list(image.to(device) for image in images)\n            targets1 = [{key: vaule.to(device) for key, vaule in t.items()} for t in targets] \n            val_loss_dict = model(images1, targets1)\n            val_loss = sum(loss for loss in val_loss_dict.values())\n            val_loss_value = val_loss.item()\n            val_loss_sum += val_loss_value\n\n        val_loss = val_loss_sum / v_batches\n        print(f\"Validation loss is {val_loss}\")\n    return val_loss\n\n\ndef F_RCNN_train(train_loader,valid_loader,base_epoch,val_list):\n    for epoch in tqdm(range(NUM_EPOCHS)):\n        t_loss_sum = 0\n        for b_id ,(images,targets) in enumerate(train_loader, 1):\n            images1 = list(image.to(device) for image in images)\n            targets1 = [{key: vaule.to(device) for key, vaule in t.items()} for t in targets] \n\n            t_loss_dict = model(images1, targets1)\n            t_loss = sum(loss for loss in t_loss_dict.values())\n            t_loss_value = t_loss.item()\n            t_loss_sum += t_loss_value\n\n            optimizer.zero_grad()\n            t_loss.backward()\n            optimizer.step()\n        train_loss = t_loss_sum / t_batches\n        print(f\"[Epoch {epoch+base_epoch+1:2d} / {(NUM_EPOCHS+base_epoch):2d}] Train loss: {train_loss:.3f}\")\n        val_loss = F_RCNN_val(valid_loader)\n        val_list.append(val_loss)\n    return val_list\n    \n    \n","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:18.396499Z","iopub.execute_input":"2022-11-17T12:06:18.396857Z","iopub.status.idle":"2022-11-17T12:06:21.428519Z","shell.execute_reply.started":"2022-11-17T12:06:18.396819Z","shell.execute_reply":"2022-11-17T12:06:21.427582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"val_list = F_RCNN_train(train_loader,val_loader,base_epoch,val_list)","metadata":{"execution":{"iopub.status.busy":"2022-11-17T12:06:21.430108Z","iopub.execute_input":"2022-11-17T12:06:21.430488Z","iopub.status.idle":"2022-11-17T13:43:49.961235Z","shell.execute_reply.started":"2022-11-17T12:06:21.43045Z","shell.execute_reply":"2022-11-17T13:43:49.95847Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Save model\nchk_name = f'fasterrcnn_resnet50_fpn-modelworks.bin'\ntorch.save(model.state_dict(), chk_name)\n","metadata":{"execution":{"iopub.status.busy":"2022-11-17T13:43:49.963417Z","iopub.execute_input":"2022-11-17T13:43:49.963813Z","iopub.status.idle":"2022-11-17T13:43:50.333833Z","shell.execute_reply.started":"2022-11-17T13:43:49.963772Z","shell.execute_reply":"2022-11-17T13:43:50.332849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Basic Metric Evaluation","metadata":{}},{"cell_type":"code","source":"'''\nPass in the two box and compute intersection over union\n'''\ndef intersection_over_union(box1,box2):\n        \n    # in the form x,y,w,h\n    box1_x1 = box1[0]\n    box1_y1 = box1[1]\n    box1_x2 = box1[2]\n    box1_y2 = box1[3]\n    box2_x1 = box2[0]\n    box2_y1 = box2[1]\n    box2_x2 = box2[2]\n    box2_y2 = box2[3]\n\n    # Intersection key point\n    int_x1 = torch.max(box1[0],box2[0])\n    int_y1 = torch.max(box1[1],box2[1])\n    int_x2 = torch.min(box1[2],box2[2])\n    int_y2 = torch.min(box2[3],box2[3])\n    \n   \n    \n    # Intersection area\n    intersection = (int_x2-int_x1).clamp(0)*(int_y2-int_y1).clamp(0)\n    # Seperate area\n    box1_area = abs(box1_x2-box1_x1)*abs(box1_y2-box1_y1)\n    box2_area = abs(box2_x2-box2_x1)*abs(box2_y2-box2_y1)\n\n    \n    # Union\n    union = box1_area+box2_area-intersection\n    return intersection/(union+1e-6) #\n    ","metadata":{"execution":{"iopub.status.busy":"2022-11-17T13:43:50.335316Z","iopub.execute_input":"2022-11-17T13:43:50.335656Z","iopub.status.idle":"2022-11-17T13:43:50.34509Z","shell.execute_reply.started":"2022-11-17T13:43:50.335618Z","shell.execute_reply":"2022-11-17T13:43:50.343864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nreal has shape(M,4)\npred has shape(N,4)\n'''\ndef calculate_iou_pair(box1,box2,used):\n    if box2 in used:\n        return 0\n    else:\n        return intersection_over_union(box1,box2)\n\n'''\nReturn  true positive, false positive and false negative in the image's boxes\n'''\ndef group_iou(real,pred,iou_target):\n    tp = 0\n    fp = 0\n    fn = 0\n    used = np.zeros(len(real))\n    for i,box1 in enumerate(real):\n        all_iou = [calculate_iou_pair(box1,box2,used)for box2 in pred]\n        if len(all_iou) == 0:\n            continue\n\n        max_iou = max(all_iou)\n        if max_iou > iou_target:\n            used[i] = all_iou[all_iou.index(max_iou)]\n            tp+=1\n    fp = len(pred)-tp\n    fn = len(real)-tp\n    return tp,fp,fn\n    ","metadata":{"execution":{"iopub.status.busy":"2022-11-17T13:43:50.346608Z","iopub.execute_input":"2022-11-17T13:43:50.347414Z","iopub.status.idle":"2022-11-17T13:43:50.357618Z","shell.execute_reply.started":"2022-11-17T13:43:50.347377Z","shell.execute_reply":"2022-11-17T13:43:50.356026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_metrics(tp,fp,fn):\n    if tp > 0:\n        precision = tp / (tp+fp)\n        recall = tp / (tp+fn)\n        f1_score = 2 * ((precision * recall) / (precision + recall))\n        f2_score =  (5 * precision * recall) / (4 * precision + recall)\n        iou = tp / (tp + fp + fn)\n    else:\n        f2_score = precision = recall = f1_score = iou = float('NaN')\n    return {\n        \"FP\": fp,\n        \"FN\": fn,\n        \"TP\": tp,\n        \"precision\": precision,\n        \"recall\": recall,\n        \"f1_score\": f1_score,\n        \"f2_score\":f2_score,\n        \"iou\": iou\n    }\n\n\ndef get_metric_dataframe(basic_metric):\n    complete_metric = []\n    for metric in basic_metric:\n        a,full_list= get_metrics(metric)\n        complete_metric.append(full_list)\n\n\n\n    df = pd.DataFrame (complete_metric, \n                   columns = [\"index\",\"TP\",\"FP\",\"FN\",\"Precision\",\"Recall\",\"F1 Score\",\"F2 Score\",\"IOU\"])\n    return df","metadata":{"execution":{"iopub.status.busy":"2022-11-17T13:43:50.358837Z","iopub.execute_input":"2022-11-17T13:43:50.35923Z","iopub.status.idle":"2022-11-17T13:43:50.381559Z","shell.execute_reply.started":"2022-11-17T13:43:50.359197Z","shell.execute_reply":"2022-11-17T13:43:50.380303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def evaluate_metrics(iou_target,true,pred):\n    assert iou_target <= 1\n    basic_metric = [] # a list of images metrics that is in the form [index,tp,fp,fn]\n    for index in range(len(true)):\n        # for each image\n        tp,fp,fn = group_iou(true[index],pred[index],iou_target)\n        basic_metric.append([index,tp,fp,fn])\n    return get_metric_dataframe(basic_metric)","metadata":{"execution":{"iopub.status.busy":"2022-11-17T13:43:50.383406Z","iopub.execute_input":"2022-11-17T13:43:50.384388Z","iopub.status.idle":"2022-11-17T13:43:50.397595Z","shell.execute_reply.started":"2022-11-17T13:43:50.384343Z","shell.execute_reply":"2022-11-17T13:43:50.396071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\nIOU_TARGET = 0.2\nmodel.eval()\nwith torch.no_grad():\n    tp = fp = fn = 0\n    pred_label = []\n    true_label = []\n    for b_id,(images,targets) in enumerate(val_loader,1):\n        images = list(img.to(device) for img in images)\n        targets = [{k: v.to(device) for k, v in t.items()} for t in targets]\n        outputs = model(images)\n        \n        #outputs = [{k: v.detach().cpu() for k, v in t.items()} for t in outputs]\n        #print(f\"output type is {type(outputs)} and output is {outputs}\")\n        \n        for i, image in enumerate(images):\n            boxes = outputs[i]['boxes'].data\n            labels = outputs[i]['labels'].data\n            scores = outputs[i]['scores'].data\n            gt_boxes = targets[i]['boxes'].data\n            tp_temp,fp_temp,fn_temp = group_iou(gt_boxes,boxes,IOU_TARGET)\n            tp += tp_temp\n            fp += fp_temp\n            fn += fn_temp\n    print(f\"tp is {tp} fp is {fp} and fn is {fn}\")\n\n        \n        ","metadata":{"execution":{"iopub.status.busy":"2022-11-17T13:43:50.399928Z","iopub.execute_input":"2022-11-17T13:43:50.40038Z","iopub.status.idle":"2022-11-17T13:44:58.689293Z","shell.execute_reply.started":"2022-11-17T13:43:50.400336Z","shell.execute_reply":"2022-11-17T13:44:58.688089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(get_metrics(tp,fp,fn))","metadata":{"execution":{"iopub.status.busy":"2022-11-17T13:44:58.69153Z","iopub.execute_input":"2022-11-17T13:44:58.692192Z","iopub.status.idle":"2022-11-17T13:44:58.698941Z","shell.execute_reply.started":"2022-11-17T13:44:58.692148Z","shell.execute_reply":"2022-11-17T13:44:58.697981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nPlot the outcome\n'''\n# def plot_and_compare_outcome():\n# #     model.eval()\n# #     images, targets = next(iter(test_loader))\n\n#     for idx in range(BATCH_SIZE):\n#         model.eval()\n#         images, targets = next(iter(test_loader))\n#         images = list(img.to(DEVICE) for img in images)\n#         targets = [{k: v.to(DEVICE) for k, v in t.items()} for t in targets]\n        \n#         boxes = targets[idx]['boxes'].cpu().numpy().astype(np.int32)\n#         sample = images[idx].permute(1,2,0).cpu().numpy()\n\n\n#         outputs = model(images)\n#         outputs = [{k: v.detach().cpu().numpy() for k, v in t.items()} for t in outputs]\n        \n#         fig, ax = plt.subplots(1, 1, figsize=(16, 8))\n#         # Red for ground truth\n#         for box in boxes:\n#             cv2.rectangle(sample,(box[0], box[1]),\n#                           (box[2], box[3]),\n#                           (220, 0, 0), 3)\n#         pred = outputs[idx]['boxes'][:5]\n# #         print(outputs[idx])\n#         for box in pred:\n#             cv2.rectangle(sample,(box[0], box[1]),\n#                           (box[2], box[3]),\n#                           (0, 220, 0), 3)\n#         ax.set_axis_off()\n#         ax.imshow(sample);\n'''\nPlot the outcome\n'''\ndef plot_and_compare_outcome():\n        images, targets = next(iter(test_loader))\n        images = list(img.to(DEVICE) for img in images)\n        targets = [{k: v.to(DEVICE) for k, v in t.items()} for t in targets]\n        idx = 0\n        boxes = targets[idx]['boxes'].cpu().numpy().astype(np.int32)\n        sample = images[idx].permute(1,2,0).cpu().numpy()\n\n#         print(images.shape)\n        outputs = model(images)\n        outputs = [{k: v.detach().cpu().numpy() for k, v in t.items()} for t in outputs]\n        \n        fig, ax = plt.subplots(1, 1, figsize=(16, 8))\n        # Red for ground truth\n        for box in boxes:\n            cv2.rectangle(sample,(box[0], box[1]),\n                          (box[2], box[3]),\n                          (220, 0, 0), 3)\n        pred = outputs[idx]['boxes'][:5]\n#         print(outputs[idx])\n        for box in pred:\n            cv2.rectangle(sample,(box[0], box[1]),\n                          (box[2], box[3]),\n                          (0, 220, 0), 3)\n        ax.set_axis_off()\n        ax.imshow(sample);\n","metadata":{"execution":{"iopub.status.busy":"2022-11-17T14:22:42.462616Z","iopub.execute_input":"2022-11-17T14:22:42.463546Z","iopub.status.idle":"2022-11-17T14:22:42.502983Z","shell.execute_reply.started":"2022-11-17T14:22:42.463422Z","shell.execute_reply":"2022-11-17T14:22:42.501684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Showing image in batches in the test size","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nplot_and_compare_outcome()","metadata":{"execution":{"iopub.status.busy":"2022-11-17T14:22:42.505373Z","iopub.execute_input":"2022-11-17T14:22:42.505835Z","iopub.status.idle":"2022-11-17T14:22:42.581841Z","shell.execute_reply.started":"2022-11-17T14:22:42.505791Z","shell.execute_reply":"2022-11-17T14:22:42.580285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}