{"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":"# Pytorch starter - FasterRCNN Train\nIn this notebook I enabled the GPU and the Internet access (needed for the pre-trained weights). We can not use Internet during inference, so I'll create another notebook for commiting. Stay tuned!\n\nYou can find the [inference notebook here](https://www.kaggle.com/pestipeti/pytorch-starter-fasterrcnn-inference)\n\n- FasterRCNN from torchvision\n- Use Resnet50 backbone\n- Albumentation enabled (simple flip for now)\n","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport cv2\nimport os\nimport re\nimport pydicom\nimport warnings\n\nfrom PIL import Image\n\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\n\nimport torch\nimport torchvision\n\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection import FasterRCNN\nfrom torchvision.models.detection.rpn import AnchorGenerator\n\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.utils.data.sampler import SequentialSampler\n\nfrom matplotlib import pyplot as plt\n\n\nwarnings.filterwarnings(\"ignore\")\n\n\n\nDIR_INPUT = '/kaggle/input/vinbigdata-chest-xray-abnormalities-detection'\nDIR_TRAIN = f'{DIR_INPUT}/train'\nDIR_TEST = f'{DIR_INPUT}/test'","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","execution":{"iopub.status.busy":"2021-12-07T09:52:34.728226Z","iopub.execute_input":"2021-12-07T09:52:34.728594Z","iopub.status.idle":"2021-12-07T09:52:35.63181Z","shell.execute_reply.started":"2021-12-07T09:52:34.728541Z","shell.execute_reply":"2021-12-07T09:52:35.630967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(f'{DIR_INPUT}/train.csv')\ndf = df[((df.class_id==4) | (df.class_id==6) | ( df.class_id==7)| (df.class_id==14)) & (df.rad_id=='R9')].reset_index(drop = True)\ndf.replace(np.nan, 0, inplace=True)\ndf.loc[df[\"class_id\"] == 14, ['x_max', 'y_max']] = 1.0\ndf.loc[df[\"class_id\"] == 14, [\"class_id\"]] = 0\ndf.loc[df[\"class_id\"] == 4, [\"class_id\"]] = 1\ndf.loc[df[\"class_id\"] == 6, [\"class_id\"]] = 1\ndf.loc[df[\"class_id\"] == 7, [\"class_id\"]] = 1\ndf.loc[df[\"class_id\"]==1,[\"class_name\"]] = \"Lung Opacity\"","metadata":{"execution":{"iopub.status.busy":"2021-12-07T09:52:37.055408Z","iopub.execute_input":"2021-12-07T09:52:37.055809Z","iopub.status.idle":"2021-12-07T09:52:37.182987Z","shell.execute_reply.started":"2021-12-07T09:52:37.05575Z","shell.execute_reply":"2021-12-07T09:52:37.182218Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.tail()","metadata":{"execution":{"iopub.status.busy":"2021-12-07T09:52:38.736001Z","iopub.execute_input":"2021-12-07T09:52:38.736371Z","iopub.status.idle":"2021-12-07T09:52:38.758443Z","shell.execute_reply.started":"2021-12-07T09:52:38.736318Z","shell.execute_reply":"2021-12-07T09:52:38.75767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2021-12-07T09:52:40.365856Z","iopub.execute_input":"2021-12-07T09:52:40.366238Z","iopub.status.idle":"2021-12-07T09:52:40.372007Z","shell.execute_reply.started":"2021-12-07T09:52:40.366178Z","shell.execute_reply":"2021-12-07T09:52:40.371059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_df = df[df['image_id'].isin(valid_ids)]\ntrain_df = df[df['image_id'].isin(train_ids)]","metadata":{"execution":{"iopub.status.busy":"2021-12-07T09:52:41.770263Z","iopub.execute_input":"2021-12-07T09:52:41.770627Z","iopub.status.idle":"2021-12-07T09:52:41.780332Z","shell.execute_reply.started":"2021-12-07T09:52:41.770573Z","shell.execute_reply":"2021-12-07T09:52:41.779444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_df.head()","metadata":{"execution":{"iopub.status.busy":"2021-12-07T07:58:32.513663Z","iopub.execute_input":"2021-12-07T07:58:32.51419Z","iopub.status.idle":"2021-12-07T07:58:32.529251Z","shell.execute_reply.started":"2021-12-07T07:58:32.513973Z","shell.execute_reply":"2021-12-07T07:58:32.52812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VinBigDataset(Dataset):\n    \n    def __init__(self, dataframe, image_dir, transforms=None):\n        super().__init__()\n        \n        self.image_ids = dataframe[\"image_id\"].unique()\n        self.df = dataframe\n        self.image_dir = image_dir\n        self.transforms = transforms\n        \n    def __getitem__(self, index):\n        \n        image_id = self.image_ids[index]\n        records = self.df[(self.df['image_id'] == image_id)]\n        records = records.reset_index(drop=True)\n\n        dicom = pydicom.dcmread(f\"{self.image_dir}/{image_id}.dicom\")\n        \n        image = dicom.pixel_array\n        \n        if \"PhotometricInterpretation\" in dicom:\n            if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n                image = np.amax(image) - image\n        \n        intercept = dicom.RescaleIntercept if \"RescaleIntercept\" in dicom else 0.0\n        slope = dicom.RescaleSlope if \"RescaleSlope\" in dicom else 1.0\n        \n        if slope != 1:\n            image = slope * image.astype(np.float64)\n            image = image.astype(np.int16)\n            \n        image += np.int16(intercept)        \n        \n        image = np.stack([image, image, image])\n        image = image.astype('float32')\n        image = image - image.min()\n        image = image / image.max()\n        image = image * 255.0\n        image = image.transpose(1,2,0)\n       \n        if records.loc[0, \"class_id\"] == 0:\n            records = records.loc[[0], :]\n        \n        boxes = records[['x_min', 'y_min', 'x_max', 'y_max']].values\n        area = (boxes[:, 3] - boxes[:, 1]) * (boxes[:, 2] - boxes[:, 0])\n        area = torch.as_tensor(area, dtype=torch.float32)\n        labels = torch.tensor(records[\"class_id\"].values, dtype=torch.int64)\n\n        # suppose all instances are not crowd\n        iscrowd = torch.zeros((records.shape[0],), dtype=torch.int64)\n\n        target = {}\n        target['boxes'] = boxes\n        target['labels'] = labels\n        target['image_id'] = torch.tensor([index])\n        target['area'] = area\n        target['iscrowd'] = iscrowd\n\n        if self.transforms:\n            sample = {\n                'image': image,\n                'bboxes': target['boxes'],\n                'labels': labels\n            }\n            sample = self.transforms(**sample)\n            image = sample['image']\n            \n            target['boxes'] = torch.tensor(sample['bboxes']).type(torch.float32)\n\n        if target[\"boxes\"].shape[0] == 0:\n            # Albumentation cuts the target (class 14, 1x1px in the corner)\n            target[\"boxes\"] = torch.from_numpy(np.array([[0.0, 0.0, 1.0, 1.0]])).type(torch.float32)\n            target[\"area\"] = torch.tensor([1.0], dtype=torch.float32)\n            target[\"labels\"] = torch.tensor([0], dtype=torch.int64)\n            \n        return image, target\n    \n    def __len__(self):\n        return self.image_ids.shape[0]","metadata":{"execution":{"iopub.status.busy":"2021-12-07T09:52:46.701513Z","iopub.execute_input":"2021-12-07T09:52:46.701997Z","iopub.status.idle":"2021-12-07T09:52:46.724442Z","shell.execute_reply.started":"2021-12-07T09:52:46.701933Z","shell.execute_reply":"2021-12-07T09:52:46.723294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Albumentations\ndef get_train_transform():\n    return A.Compose([\n        A.Resize (512, 512,always_apply=True, p=1),\n        A.ShiftScaleRotate(scale_limit=0.1, rotate_limit=45, p=0.25),\n        A.LongestMaxSize(max_size=800, p=1.0),\n        A.RandomBrightnessContrast(p=0.25),\n        # FasterRCNN will normalize.\n        A.Normalize(mean=(0, 0, 0), std=(1, 1, 1), max_pixel_value=255.0, 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        A.Resize (512, 512,always_apply=True, p=1),\n        A.Normalize(mean=(0, 0, 0), std=(1, 1, 1), max_pixel_value=255.0, p=1.0),\n        ToTensorV2(p=1.0)\n    ], bbox_params={'format': 'pascal_voc', 'label_fields': ['labels']})\n","metadata":{"execution":{"iopub.status.busy":"2021-12-07T09:52:49.991812Z","iopub.execute_input":"2021-12-07T09:52:49.992198Z","iopub.status.idle":"2021-12-07T09:52:50.002496Z","shell.execute_reply.started":"2021-12-07T09:52:49.992122Z","shell.execute_reply":"2021-12-07T09:52:50.001464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# METRIC","metadata":{}},{"cell_type":"code","source":"'''\nhttps://www.kaggle.com/pestipeti/competition-metric-details-script\n'''\n\n\ndef calculate_iou(gt, pr, form='pascal_voc') -> float:\n    \"\"\"Calculates the Intersection over Union.\n\n    Args:\n        gt: (np.ndarray[Union[int, float]]) coordinates of the ground-truth box\n        pr: (np.ndarray[Union[int, float]]) coordinates of the prdected box\n        form: (str) gt/pred coordinates format\n            - pascal_voc: [xmin, ymin, xmax, ymax]\n            - coco: [xmin, ymin, w, h]\n    Returns:\n        (float) Intersection over union (0.0 <= iou <= 1.0)\n    \"\"\"\n    if form == 'coco':\n        gt = gt.copy()\n        pr = pr.copy()\n\n        gt[2] = gt[0] + gt[2]\n        gt[3] = gt[1] + gt[3]\n        pr[2] = pr[0] + pr[2]\n        pr[3] = pr[1] + pr[3]\n\n    # Calculate overlap area\n    dx = min(gt[2], pr[2]) - max(gt[0], pr[0]) + 1\n    \n    if dx < 0:\n        return 0.0\n    \n    dy = min(gt[3], pr[3]) - max(gt[1], pr[1]) + 1\n\n    if dy < 0:\n        return 0.0\n\n    overlap_area = dx * dy\n\n    # Calculate union area\n    union_area = (\n            (gt[2] - gt[0] + 1) * (gt[3] - gt[1] + 1) +\n            (pr[2] - pr[0] + 1) * (pr[3] - pr[1] + 1) -\n            overlap_area\n    )\n\n    return overlap_area / union_area\n\n\ndef find_best_match(gts, pred, pred_idx, threshold = 0.5, form = 'pascal_voc', ious=None) -> int:\n    \"\"\"Returns the index of the 'best match' between the\n    ground-truth boxes and the prediction. The 'best match'\n    is the highest IoU. (0.0 IoUs are ignored).\n\n    Args:\n        gts: (List[List[Union[int, float]]]) Coordinates of the available ground-truth boxes\n        pred: (List[Union[int, float]]) Coordinates of the predicted box\n        pred_idx: (int) Index of the current predicted box\n        threshold: (float) Threshold\n        form: (str) Format of the coordinates\n        ious: (np.ndarray) len(gts) x len(preds) matrix for storing calculated ious.\n\n    Return:\n        (int) Index of the best match GT box (-1 if no match above threshold)\n    \"\"\"\n    best_match_iou = -np.inf\n    best_match_idx = -1\n\n    for gt_idx in range(len(gts)):\n        \n        if gts[gt_idx][0] < 0:\n            # Already matched GT-box\n            continue\n        \n        iou = -1 if ious is None else ious[gt_idx][pred_idx]\n\n        if iou < 0:\n            iou = calculate_iou(gts[gt_idx], pred, form=form)\n            \n            if ious is not None:\n                ious[gt_idx][pred_idx] = iou\n\n        if iou < threshold:\n            continue\n\n        if iou > best_match_iou:\n            best_match_iou = iou\n            best_match_idx = gt_idx\n\n    return best_match_idx\n\n\ndef calculate_precision(gts, preds, threshold = 0.5, form = 'coco', ious=None) -> float:\n    \"\"\"Calculates precision for GT - prediction pairs at one threshold.\n\n    Args:\n        gts: (List[List[Union[int, float]]]) Coordinates of the available ground-truth boxes\n        preds: (List[List[Union[int, float]]]) Coordinates of the predicted boxes,\n               sorted by confidence value (descending)\n        threshold: (float) Threshold\n        form: (str) Format of the coordinates\n        ious: (np.ndarray) len(gts) x len(preds) matrix for storing calculated ious.\n\n    Return:\n        (float) Precision\n    \"\"\"\n    n = len(preds)\n    tp = 0\n    fp = 0\n    \n    # for pred_idx, pred in enumerate(preds_sorted):\n    for pred_idx in range(n):\n\n        best_match_gt_idx = find_best_match(gts, preds[pred_idx], pred_idx,\n                                            threshold=threshold, form=form, ious=ious)\n\n        if best_match_gt_idx >= 0:\n            # True positive: The predicted box matches a gt box with an IoU above the threshold.\n            tp += 1\n            # Remove the matched GT box\n            gts[best_match_gt_idx] = -1\n\n        else:\n            # No match\n            # False positive: indicates a predicted box had no associated gt box.\n            fp += 1\n\n    # False negative: indicates a gt box had no associated predicted box.\n    fn = (gts.sum(axis=1) > 0).sum()\n\n    return tp / (tp + fp + fn)\n\n\n\ndef calculate_image_precision(gts, preds, thresholds = (0.5, ), form = 'coco') -> float:\n    \"\"\"Calculates image precision.\n       The mean average precision at different intersection over union (IoU) thresholds.\n\n    Args:\n        gts: (List[List[Union[int, float]]]) Coordinates of the available ground-truth boxes\n        preds: (List[List[Union[int, float]]]) Coordinates of the predicted boxes,\n               sorted by confidence value (descending)\n        thresholds: (float) Different thresholds\n        form: (str) Format of the coordinates\n\n    Return:\n        (float) Precision\n    \"\"\"\n    n_threshold = len(thresholds)\n    image_precision = 0.0\n    \n    ious = np.ones((len(gts), len(preds))) * -1\n    # ious = None\n\n    for threshold in thresholds:\n        precision_at_threshold = calculate_precision(gts.copy(), preds, threshold=threshold,\n                                                     form=form, ious=ious)\n        image_precision += precision_at_threshold / n_threshold\n\n    return image_precision","metadata":{"execution":{"iopub.status.busy":"2021-12-07T08:59:41.595542Z","iopub.execute_input":"2021-12-07T08:59:41.595893Z","iopub.status.idle":"2021-12-07T08:59:41.621466Z","shell.execute_reply.started":"2021-12-07T08:59:41.595837Z","shell.execute_reply":"2021-12-07T08:59:41.620174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = torchvision.models.detection.fasterrcnn_resnet50_fpn(pretrained=True)","metadata":{"execution":{"iopub.status.busy":"2021-12-07T09:52:56.275841Z","iopub.execute_input":"2021-12-07T09:52:56.276362Z","iopub.status.idle":"2021-12-07T09:52:56.907279Z","shell.execute_reply.started":"2021-12-07T09:52:56.276142Z","shell.execute_reply":"2021-12-07T09:52:56.90629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_classes = 2\n\n# get number of input features for the classifier\nin_features = model.roi_heads.box_predictor.cls_score.in_features\n\n# replace the pre-trained head with a new one\nmodel.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes)","metadata":{"execution":{"iopub.status.busy":"2021-12-07T09:52:58.464369Z","iopub.execute_input":"2021-12-07T09:52:58.464726Z","iopub.status.idle":"2021-12-07T09:52:58.470333Z","shell.execute_reply.started":"2021-12-07T09:52:58.464666Z","shell.execute_reply":"2021-12-07T09:52:58.46904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def collate_fn(batch):\n    return tuple(zip(*batch))\n\ntrain_dataset = VinBigDataset(train_df, DIR_TRAIN, get_train_transform())\nvalid_dataset = VinBigDataset(valid_df, DIR_TRAIN, get_valid_transform())\n\n\n\n\ntrain_data_loader = DataLoader(\n    train_dataset,\n    batch_size=8,\n    shuffle=True,\n    num_workers=4,\n    collate_fn=collate_fn\n)\n\nvalid_data_loader = DataLoader(\n    valid_dataset,\n    batch_size=8,\n    shuffle=False,\n    num_workers=4,\n    collate_fn=collate_fn\n)","metadata":{"execution":{"iopub.status.busy":"2021-12-07T09:53:00.005715Z","iopub.execute_input":"2021-12-07T09:53:00.006118Z","iopub.status.idle":"2021-12-07T09:53:00.015633Z","shell.execute_reply.started":"2021-12-07T09:53:00.00605Z","shell.execute_reply":"2021-12-07T09:53:00.014284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')","metadata":{"execution":{"iopub.status.busy":"2021-12-07T09:53:01.733002Z","iopub.execute_input":"2021-12-07T09:53:01.733535Z","iopub.status.idle":"2021-12-07T09:53:01.767224Z","shell.execute_reply.started":"2021-12-07T09:53:01.733458Z","shell.execute_reply":"2021-12-07T09:53:01.766265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# SAMPLE","metadata":{}},{"cell_type":"code","source":"images, targets = next(iter(train_data_loader))","metadata":{"execution":{"iopub.status.busy":"2021-12-07T08:56:18.491538Z","iopub.execute_input":"2021-12-07T08:56:18.491885Z","iopub.status.idle":"2021-12-07T08:58:21.714878Z","shell.execute_reply.started":"2021-12-07T08:56:18.491831Z","shell.execute_reply":"2021-12-07T08:58:21.713889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images = list(image.to(device) for image in images)\ntargets = [{k: v.to(device) for k, v in t.items()} for t in targets]","metadata":{"execution":{"iopub.status.busy":"2021-12-07T08:06:08.455028Z","iopub.execute_input":"2021-12-07T08:06:08.455397Z","iopub.status.idle":"2021-12-07T08:06:08.504878Z","shell.execute_reply.started":"2021-12-07T08:06:08.455344Z","shell.execute_reply":"2021-12-07T08:06:08.504081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"boxes = targets[2]['boxes'].cpu().numpy().astype(np.int32)\nsample = images[2].permute(1,2,0).cpu().numpy()","metadata":{"execution":{"iopub.status.busy":"2021-12-07T08:06:14.045071Z","iopub.execute_input":"2021-12-07T08:06:14.045437Z","iopub.status.idle":"2021-12-07T08:06:14.053146Z","shell.execute_reply.started":"2021-12-07T08:06:14.045385Z","shell.execute_reply":"2021-12-07T08:06:14.052084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax = plt.subplots(1, 1, figsize=(16, 8))\n\nfor box in boxes:\n    cv2.rectangle(sample,\n                  (box[0], box[1]),\n                  (box[2], box[3]),\n                  (220, 0, 0), 3)\n    \nax.set_axis_off()\nax.imshow(sample)","metadata":{"execution":{"iopub.status.busy":"2021-12-07T08:06:16.07169Z","iopub.execute_input":"2021-12-07T08:06:16.072046Z","iopub.status.idle":"2021-12-07T08:06:16.290782Z","shell.execute_reply.started":"2021-12-07T08:06:16.071986Z","shell.execute_reply":"2021-12-07T08:06:16.289835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Averager:\n    def __init__(self):\n        self.current_total = 0.0\n        self.iterations = 0.0\n\n    def send(self, value):\n        self.current_total += value\n        self.iterations += 1\n\n    @property\n    def value(self):\n        if self.iterations == 0:\n            return 0\n        else:\n            return 1.0 * self.current_total / self.iterations\n\n    def reset(self):\n        self.current_total = 0.0\n        self.iterations = 0.0","metadata":{"execution":{"iopub.status.busy":"2021-12-07T09:53:05.717067Z","iopub.execute_input":"2021-12-07T09:53:05.717454Z","iopub.status.idle":"2021-12-07T09:53:05.724538Z","shell.execute_reply.started":"2021-12-07T09:53:05.7174Z","shell.execute_reply":"2021-12-07T09:53:05.723382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.to(device)\nparams = [p for p in model.parameters() if p.requires_grad]\noptimizer = torch.optim.SGD(params, lr=0.005, momentum=0.9, weight_decay=0.0005)\nlr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=4, gamma=0.1)\n\nnum_epochs = 30","metadata":{"execution":{"iopub.status.busy":"2021-12-07T09:53:07.758634Z","iopub.execute_input":"2021-12-07T09:53:07.758985Z","iopub.status.idle":"2021-12-07T09:53:09.358012Z","shell.execute_reply.started":"2021-12-07T09:53:07.758926Z","shell.execute_reply":"2021-12-07T09:53:09.357154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_hist = Averager()\nitr = 1\n\nfor epoch in range(num_epochs):\n    loss_hist.reset()\n    \n    for images, targets in train_data_loader:\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        loss_dict = model(images, targets)\n\n        losses = sum(loss for loss in loss_dict.values())\n        loss_value = losses.item()\n\n        loss_hist.send(loss_value)\n\n        optimizer.zero_grad()\n        losses.backward()\n        optimizer.step()\n\n        if itr % 100 == 0:\n            print(f\"Iteration #{itr} loss: {loss_hist.value}\")\n\n        itr += 1\n        \n    \n    # update the learning rate\n    if lr_scheduler is not None:\n        lr_scheduler.step()\n\n    print(f\"Epoch #{epoch} loss: {loss_hist.value}\")   \n    print(\"Saving epoch's state...\")\n    torch.save(model.state_dict(), f\"model_state_epoch_{epoch}.pth\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(torch.load('./model_state_epoch_10.pth'))\nmodel.eval()\n\nx = model.to(device)","metadata":{"execution":{"iopub.status.busy":"2021-12-07T15:33:23.587301Z","iopub.execute_input":"2021-12-07T15:33:23.587734Z","iopub.status.idle":"2021-12-07T15:33:23.767181Z","shell.execute_reply.started":"2021-12-07T15:33:23.587671Z","shell.execute_reply":"2021-12-07T15:33:23.766365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images, targets = next(iter(valid_data_loader))","metadata":{"execution":{"iopub.status.busy":"2021-12-07T15:37:50.478519Z","iopub.execute_input":"2021-12-07T15:37:50.478891Z","iopub.status.idle":"2021-12-07T15:38:48.147256Z","shell.execute_reply.started":"2021-12-07T15:37:50.478833Z","shell.execute_reply":"2021-12-07T15:38:48.146291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images = list(img.to(device) for img in images)\ntargets = [{k: v.to(device) for k, v in t.items()} for t in targets]","metadata":{"execution":{"iopub.status.busy":"2021-12-07T15:39:16.467347Z","iopub.execute_input":"2021-12-07T15:39:16.467742Z","iopub.status.idle":"2021-12-07T15:39:16.496368Z","shell.execute_reply.started":"2021-12-07T15:39:16.467676Z","shell.execute_reply":"2021-12-07T15:39:16.495565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"boxes = targets[2]['boxes'].cpu().numpy().astype(np.int32)\nsample = images[2].permute(1,2,0).cpu().numpy()","metadata":{"execution":{"iopub.status.busy":"2021-12-07T15:39:47.654696Z","iopub.execute_input":"2021-12-07T15:39:47.655062Z","iopub.status.idle":"2021-12-07T15:39:47.662383Z","shell.execute_reply.started":"2021-12-07T15:39:47.655007Z","shell.execute_reply":"2021-12-07T15:39:47.661273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"outputs = model(images)\noutputs = [{k: v.to(device) for k, v in t.items()} for t in outputs]","metadata":{"execution":{"iopub.status.busy":"2021-12-07T15:39:49.370251Z","iopub.execute_input":"2021-12-07T15:39:49.370612Z","iopub.status.idle":"2021-12-07T15:39:49.739513Z","shell.execute_reply.started":"2021-12-07T15:39:49.370558Z","shell.execute_reply":"2021-12-07T15:39:49.738718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(outputs)","metadata":{"execution":{"iopub.status.busy":"2021-12-07T15:50:01.256422Z","iopub.execute_input":"2021-12-07T15:50:01.256847Z","iopub.status.idle":"2021-12-07T15:50:01.369051Z","shell.execute_reply.started":"2021-12-07T15:50:01.25679Z","shell.execute_reply":"2021-12-07T15:50:01.367928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax = plt.subplots(1, 1, figsize=(16, 8))\n\nfor box in boxes:\n    cv2.rectangle(sample,\n                  (box[0], box[1]),\n                  (box[2], box[3]),\n                  (220, 0, 0), 3)\n    \nax.set_axis_off()\nax.imshow(sample)","metadata":{"execution":{"iopub.status.busy":"2021-12-07T15:39:51.656005Z","iopub.execute_input":"2021-12-07T15:39:51.656392Z","iopub.status.idle":"2021-12-07T15:39:51.847982Z","shell.execute_reply.started":"2021-12-07T15:39:51.656337Z","shell.execute_reply":"2021-12-07T15:39:51.847117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"boxes = targets[3]['boxes'].cpu().numpy().astype(np.int32)\nsample = images[3].permute(1,2,0).cpu().numpy()\n\noutputs = model(images)\noutputs = [{k: v.to(device) for k, v in t.items()} for t in outputs]\n\nfig, ax = plt.subplots(1, 1, figsize=(16, 8))\n\nfor box in boxes:\n    cv2.rectangle(sample,\n                  (box[0], box[1]),\n                  (box[2], box[3]),\n                  (220, 0, 0), 3)\n    \nax.set_axis_off()\nax.imshow(sample)","metadata":{"execution":{"iopub.status.busy":"2021-12-07T15:41:27.865762Z","iopub.execute_input":"2021-12-07T15:41:27.866185Z","iopub.status.idle":"2021-12-07T15:41:28.450151Z","shell.execute_reply.started":"2021-12-07T15:41:27.866106Z","shell.execute_reply":"2021-12-07T15:41:28.449252Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"boxes = targets[4]['boxes'].cpu().numpy().astype(np.int32)\nsample = images[4].permute(1,2,0).cpu().numpy()\n\noutputs = model(images)\noutputs = [{k: v.to(device) for k, v in t.items()} for t in outputs]\n\nfig, ax = plt.subplots(1, 1, figsize=(16, 8))\n\nfor box in boxes:\n    cv2.rectangle(sample,\n                  (box[0], box[1]),\n                  (box[2], box[3]),\n                  (220, 0, 0), 3)\n    \nax.set_axis_off()\nax.imshow(sample)","metadata":{"execution":{"iopub.status.busy":"2021-12-07T15:41:56.505891Z","iopub.execute_input":"2021-12-07T15:41:56.506297Z","iopub.status.idle":"2021-12-07T15:41:57.095623Z","shell.execute_reply.started":"2021-12-07T15:41:56.506219Z","shell.execute_reply":"2021-12-07T15:41:57.094255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run_training()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}