{"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":"code","source":"from pathlib import Path\nimport pydicom\nimport numpy as np\nimport cv2\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:56:45.929305Z","iopub.execute_input":"2022-01-12T15:56:45.929868Z","iopub.status.idle":"2022-01-12T15:56:46.294241Z","shell.execute_reply.started":"2022-01-12T15:56:45.929757Z","shell.execute_reply":"2022-01-12T15:56:46.293517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = pd.read_csv(\"../input/rsna-heart-detection/rsna_heart_detection.csv\")\nlabels.head()","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:56:48.294014Z","iopub.execute_input":"2022-01-12T15:56:48.296523Z","iopub.status.idle":"2022-01-12T15:56:48.339677Z","shell.execute_reply.started":"2022-01-12T15:56:48.296466Z","shell.execute_reply":"2022-01-12T15:56:48.338943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels.shape","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:56:48.559780Z","iopub.execute_input":"2022-01-12T15:56:48.560186Z","iopub.status.idle":"2022-01-12T15:56:48.571599Z","shell.execute_reply.started":"2022-01-12T15:56:48.560137Z","shell.execute_reply":"2022-01-12T15:56:48.570780Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT_PATH = Path(\"../input/rsna-pneumonia-detection-challenge/stage_2_train_images/\")\nSAVE_PATH = Path(\"../Processed-Heart-Detection\")","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:56:48.802750Z","iopub.execute_input":"2022-01-12T15:56:48.803387Z","iopub.status.idle":"2022-01-12T15:56:48.807979Z","shell.execute_reply.started":"2022-01-12T15:56:48.803356Z","shell.execute_reply":"2022-01-12T15:56:48.806421Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axis = plt.subplots(2 ,2)\nc = 0\nfor i in range(2):\n    for j in range(2):\n        data = labels.iloc[c]\n        patient_id = data[\"name\"]\n        dcm_path = ROOT_PATH/str(patient_id)  # Create the path to the dcm file\n        dcm_path =dcm_path.with_suffix(\".dcm\")\n        \n        dcm = pydicom.read_file(dcm_path)\n        dcm_array = dcm.pixel_array\n        dcm_array = cv2.resize(dcm_array, (224,224))\n        \n        x = data[\"x0\"]\n        y = data[\"y0\"]\n        width = data[\"w\"]\n        height = data[\"h\"]\n        \n        axis[i][j].imshow(dcm_array, cmap=\"bone\")\n        rect = patches.Rectangle((x,y), width, height, linewidth=1, edgecolor=\"r\", facecolor=\"none\")\n        axis[i][j].add_patch(rect)\n        c+=1","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:56:49.381486Z","iopub.execute_input":"2022-01-12T15:56:49.381711Z","iopub.status.idle":"2022-01-12T15:56:49.941302Z","shell.execute_reply.started":"2022-01-12T15:56:49.381684Z","shell.execute_reply":"2022-01-12T15:56:49.940517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We use a similar preprocessing routine to the one used for the classification task.\n\nTo be able to distinguish between train and validation subjects, we store them in two lists and later save these lists.","metadata":{}},{"cell_type":"code","source":"sums = 0\nsums_squared = 0\ntrain_ids = []\nval_ids = []\n\nfor counter, patient_id in enumerate(list(labels.name)):\n    dcm_path = ROOT_PATH/str(patient_id)  # Create the path to the dcm file\n    dcm_path = dcm_path.with_suffix(\".dcm\")  # And add the .dcm suffix\n    \n    dcm = pydicom.read_file(dcm_path)  # Read the dicom file with pydicom\n    \n    # Retrieve the actual image \n    dcm_array = dcm.pixel_array\n        \n    # Resize the image to drastically improve training speed\n    # In order to reduce the space when storing the image we convert it to float16\n    # Standardize to 0-1 range\n    dcm_array = (cv2.resize(dcm_array, (224, 224)) / 255).astype(np.float16)\n          \n    # 4/5 train split, 1/5 val split\n    train_or_val = \"train\" if counter < 400 else \"val\" \n    \n    # Add to corresponding train or validation patient index list\n    if train_or_val ==\"train\":\n        train_ids.append(patient_id)\n    else:\n        val_ids.append(patient_id)\n        \n    current_save_path = SAVE_PATH/train_or_val # Define save path and create if necessary\n    current_save_path.mkdir(parents=True, exist_ok=True)\n    \n    np.save(current_save_path/patient_id, dcm_array) # Save the array in the corresponding directory\n    \n    normalizer = 224*224\n    if train_or_val == \"train\":\n        sums += np.sum(dcm_array) / normalizer\n        sums_squared += (dcm_array**2).sum() / normalizer","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:56:50.401901Z","iopub.execute_input":"2022-01-12T15:56:50.402149Z","iopub.status.idle":"2022-01-12T15:56:58.734624Z","shell.execute_reply.started":"2022-01-12T15:56:50.402121Z","shell.execute_reply":"2022-01-12T15:56:58.733869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.save(\"../Processed-Heart-Detection/train_subjects\", train_ids)\nnp.save(\"../Processed-Heart-Detection/val_subjects\", val_ids)","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:56:58.736317Z","iopub.execute_input":"2022-01-12T15:56:58.736576Z","iopub.status.idle":"2022-01-12T15:56:58.744301Z","shell.execute_reply.started":"2022-01-12T15:56:58.736541Z","shell.execute_reply":"2022-01-12T15:56:58.743638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mean = sums /len(train_ids)\nstd = np.sqrt((sums_squared / len(train_ids)) - mean**2)\nmean, std","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:56:58.745613Z","iopub.execute_input":"2022-01-12T15:56:58.745895Z","iopub.status.idle":"2022-01-12T15:56:58.857892Z","shell.execute_reply.started":"2022-01-12T15:56:58.745861Z","shell.execute_reply":"2022-01-12T15:56:58.857188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pathlib import Path\nimport torch\nimport numpy as np\nimport pandas as pd\nimport imgaug   #imgaug to set a random seed for augmentations\nfrom imgaug.augmentables.bbs import BoundingBox \n# BoundingBox from imgaug to automatically handle the coordinates when augmenting the image","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:56:58.860595Z","iopub.execute_input":"2022-01-12T15:56:58.860925Z","iopub.status.idle":"2022-01-12T15:57:01.883043Z","shell.execute_reply.started":"2022-01-12T15:56:58.860888Z","shell.execute_reply":"2022-01-12T15:57:01.882305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CardiacDataset(torch.utils.data.Dataset):\n    \n    def __init__(self, path_to_labels_csv, patients, root_path, augs):\n        \n        self.labels = pd.read_csv(path_to_labels_csv)\n        self.patients = np.load(patients)\n        self.root_path = Path(root_path)\n        self.augment = augs\n        \n    def __len__(self):\n        \"\"\"\n        Returns the length of the dataset\n        \"\"\"\n        return len(self.patients)\n    \n    def __getitem__(self, idx):\n        \"\"\"\n        Returns an image paired with bbox around the heart\n        \"\"\"\n        \n        patient = self.patients[idx]\n        \n        # Get data according to index\n        data = self.labels[self.labels[\"name\"]==patient]\n        \n        # Get entries of given patient\n        # Extract coordinates\n        x_min = data[\"x0\"].item()\n        y_min = data[\"y0\"].item()\n        x_max = x_min + data[\"w\"].item()\n        y_max = y_min + data[\"h\"].item()\n        bbox = [x_min, y_min, x_max, y_max]\n        \n        # Load file and convert to float32\n        file_path = self.root_path/patient  # Create the path to the file\n        img = np.load(f\"{file_path}.npy\").astype(np.float32)\n        \n        # Apply imgaug augmentations to image and bounding box\n        if self.augment:\n            bb = BoundingBox(x1=bbox[0], y1=bbox[1], x2=bbox[2], y2=bbox[3])\n            random_seed = torch.randint(0, 100000, (1,)).item()\n            imgaug.seed(random_seed)\n            \n            img, aug_bbox = self.augment(image=img, bounding_boxes=bb)\n            bbox = aug_bbox[0][0], aug_bbox[0][1], aug_bbox[1][0], aug_bbox[1][1]\n            \n        # Normalize the image according to the values computed in Preprocessing    \n        img = (img - 0.494) / 0.252\n        img = torch.tensor(img).unsqueeze(0)\n        bbox = torch.tensor(bbox)\n        return img, bbox","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:01.884395Z","iopub.execute_input":"2022-01-12T15:57:01.884628Z","iopub.status.idle":"2022-01-12T15:57:01.897226Z","shell.execute_reply.started":"2022-01-12T15:57:01.884596Z","shell.execute_reply":"2022-01-12T15:57:01.895580Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import imgaug.augmenters as iaa\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:01.898549Z","iopub.execute_input":"2022-01-12T15:57:01.899008Z","iopub.status.idle":"2022-01-12T15:57:01.910739Z","shell.execute_reply.started":"2022-01-12T15:57:01.898974Z","shell.execute_reply":"2022-01-12T15:57:01.910071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# First create the augmentation object (augmentation pipeline define)\nseq = iaa.Sequential([\n    iaa.GammaContrast(),\n    iaa.Affine(\n        scale=(0.8, 1.2),\n        rotate=(-10, 10),\n        translate_px=(-10, 10))\n])","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:01.913641Z","iopub.execute_input":"2022-01-12T15:57:01.913873Z","iopub.status.idle":"2022-01-12T15:57:01.921308Z","shell.execute_reply.started":"2022-01-12T15:57:01.913826Z","shell.execute_reply":"2022-01-12T15:57:01.920577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels_path = \"../input/rsna-heart-detection/rsna_heart_detection.csv\"\npatients_path = \"../Processed-Heart-Detection/train_subjects.npy\"\ntrain_root = \"../Processed-Heart-Detection/train/\"\ndataset = CardiacDataset(labels_path, patients_path, train_root, seq)","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:01.923020Z","iopub.execute_input":"2022-01-12T15:57:01.923629Z","iopub.status.idle":"2022-01-12T15:57:01.937196Z","shell.execute_reply.started":"2022-01-12T15:57:01.923595Z","shell.execute_reply":"2022-01-12T15:57:01.936557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset[0]","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:01.938410Z","iopub.execute_input":"2022-01-12T15:57:01.938647Z","iopub.status.idle":"2022-01-12T15:57:02.031370Z","shell.execute_reply.started":"2022-01-12T15:57:01.938609Z","shell.execute_reply":"2022-01-12T15:57:02.030719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img, bbox = dataset[48]\n\nfig, axis = plt.subplots(1, 1)\naxis.imshow(img[0], cmap=\"bone\")\nrect = patches.Rectangle((bbox[0], bbox[1]), bbox[2]-bbox[0], bbox[3]-bbox[1], edgecolor=\"r\", facecolor=\"none\")\naxis.add_patch(rect)\n\n# run this cell more than one. every run correctly sign the heart","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:02.034786Z","iopub.execute_input":"2022-01-12T15:57:02.035101Z","iopub.status.idle":"2022-01-12T15:57:02.267132Z","shell.execute_reply.started":"2022-01-12T15:57:02.035065Z","shell.execute_reply":"2022-01-12T15:57:02.266495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img, label = dataset[48]\n\nfig, axis = plt.subplots(1, 1)\naxis.imshow(img[0], cmap=\"bone\")\nspot1 = patches.Rectangle((label[0], label[1]), label[2] - label[0], label[3] - label[1], edgecolor='r', facecolor='none')\naxis.add_patch(spot1)\n\naxis.set_title(\"X-RAY with BBOX around heart\")\nprint(label)","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:02.268120Z","iopub.execute_input":"2022-01-12T15:57:02.268465Z","iopub.status.idle":"2022-01-12T15:57:02.499289Z","shell.execute_reply.started":"2022-01-12T15:57:02.268429Z","shell.execute_reply":"2022-01-12T15:57:02.498668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import module we'll need to import our custom module\nfrom shutil import copyfile\n\n# copy our file into the working directory (make sure it has .py suffix)\ncopyfile(src = \"../input/dataset/dataset_.py\", dst = \"../working/dataset_.py\") ","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:02.500587Z","iopub.execute_input":"2022-01-12T15:57:02.501313Z","iopub.status.idle":"2022-01-12T15:57:02.511146Z","shell.execute_reply.started":"2022-01-12T15:57:02.501278Z","shell.execute_reply":"2022-01-12T15:57:02.510189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torchvision\nimport pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import ModelCheckpoint\nfrom pytorch_lightning.loggers import TensorBoardLogger\nimport numpy as np\nimport cv2\nimport imgaug.augmenters as iaa\nfrom dataset_ import CardiacDataset # IMPORT THE DATASET","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:02.512196Z","iopub.execute_input":"2022-01-12T15:57:02.512372Z","iopub.status.idle":"2022-01-12T15:57:04.217348Z","shell.execute_reply.started":"2022-01-12T15:57:02.512350Z","shell.execute_reply":"2022-01-12T15:57:04.216520Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_root_path = \"../Processed-Heart-Detection/train\"\ntrain_subjects = \"../Processed-Heart-Detection/train_subjects.npy\"\nval_root_path = \"../Processed-Heart-Detection/val\"\nval_subjects = \"../Processed-Heart-Detection/val_subjects.npy\"\n\ntrain_transformers = iaa.Sequential([\n    iaa.GammaContrast(),\n    iaa.Affine(\n        scale=(0.8, 1.2),\n        rotate=(-10, 10),\n        translate_px=(-10, 10))\n])","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:04.219204Z","iopub.execute_input":"2022-01-12T15:57:04.219458Z","iopub.status.idle":"2022-01-12T15:57:04.226171Z","shell.execute_reply.started":"2022-01-12T15:57:04.219423Z","shell.execute_reply":"2022-01-12T15:57:04.225583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = CardiacDataset(\"../input/rsna-heart-detection/rsna_heart_detection.csv\", train_subjects, train_root_path, train_transformers)\nval_dataset = CardiacDataset(\"../input/rsna-heart-detection/rsna_heart_detection.csv\", val_subjects, val_root_path, None)","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:04.227370Z","iopub.execute_input":"2022-01-12T15:57:04.227801Z","iopub.status.idle":"2022-01-12T15:57:04.304915Z","shell.execute_reply.started":"2022-01-12T15:57:04.227763Z","shell.execute_reply":"2022-01-12T15:57:04.304223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 8\nnum_workers = 2\n\ntrain_loader = torch.utils.data.DataLoader(train_dataset, batch_size=batch_size, num_workers=num_workers, shuffle=True)\nval_loader = torch.utils.data.DataLoader(val_dataset, batch_size=batch_size, num_workers=num_workers, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:04.306258Z","iopub.execute_input":"2022-01-12T15:57:04.306738Z","iopub.status.idle":"2022-01-12T15:57:04.312472Z","shell.execute_reply.started":"2022-01-12T15:57:04.306702Z","shell.execute_reply":"2022-01-12T15:57:04.311739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torchvision.models.resnet18()","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:04.313976Z","iopub.execute_input":"2022-01-12T15:57:04.314442Z","iopub.status.idle":"2022-01-12T15:57:04.516379Z","shell.execute_reply.started":"2022-01-12T15:57:04.314407Z","shell.execute_reply":"2022-01-12T15:57:04.515695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model Creation\n\n4 outputs: Instead of predicting a binary label we need to estimate the location of the heart (xmin, ymin, xmax, ymax).\n\nLoss function: Mean Squared Error as we are dealing with continuous values.","metadata":{}},{"cell_type":"code","source":"class CardiacDetectionModel(pl.LightningModule):\n    \n    def __init__(self):\n        super().__init__()\n        \n        self.model = torchvision.models.resnet18(pretrained=True)  #transfer learning\n        self.model.conv1 = torch.nn.Conv2d(1, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n        self.model.fc = torch.nn.Linear(in_features=512, out_features=4)\n        \n        self.optimizer = torch.optim.Adam(self.model.parameters(), lr=1e-4)\n        self.loss_fn = torch.nn.MSELoss()\n        \n    def forward(self, data):\n        return self.model(data)\n    \n    def training_step(self, batch, batch_idx):\n        x_ray, label = batch\n        label = label.float()  # Convert label to float (just needed for loss computation)\n        pred = self(x_ray)\n        loss = self.loss_fn(pred, label)  # Compute the loss\n        \n        self.log(\"Train Loss\", loss)\n        \n        if batch_idx % 50 ==0:\n            self.log_images(x_ray.cpu(), pred.cpu(), label.cpu(), \"Train\")\n        return loss\n    \n    def validation_step(self, batch, batch_idx):\n        x_ray, label = batch\n        label = label.float()\n        pred = self(x_ray)\n        loss = self.loss_fn(pred, label)\n        \n        self.log(\"Val Loss\", loss)\n        \n        if batch_idx % 50 ==0:\n            self.log_images(x_ray.cpu(), pred.cpu(), label.cpu(), \"Val\")\n        return loss\n    \n    def log_images(self, x_ray, pred, label, name):\n        results = []\n        \n        # Here we create a grid consisting of 4 predictions\n        for i in range(4):\n            coords_labels = label[i]\n            coords_pred = pred[i]\n            \n            img = ((x_ray[i] * 0.252) + 0.494).numpy()[0]\n            \n            # Extract the coordinates from the label\n            x0, y0 = coords_labels[0].int().item(), coords_labels[1].int().item()\n            x1, y1 = coords_labels[2].int().item(), coords_labels[3].int().item()\n            img = cv2.rectangle(img, (x0, y0), (x1, y1), (0, 0, 0), 2)\n            \n            # Extract the coordinates from the prediction\n            x0, y0 = coords_pred[0].int().item(), coords_pred[1].int().item()\n            x1, y1 = coords_pred[2].int().item(), coords_pred[3].int().item()\n            img = cv2.rectangle(img, (x0, y0), (x1, y1), (0, 0, 0), 2)\n            \n            results.append(torch.tensor(img).unsqueeze(0))\n        \n        grid = torchvision.utils.make_grid(results, 2)\n        self.logger.experiment.add_image(name, grid, self.global_step)\n        \n    def configure_optimizers(self):\n        return [self.optimizer]","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:04.518873Z","iopub.execute_input":"2022-01-12T15:57:04.519128Z","iopub.status.idle":"2022-01-12T15:57:04.536909Z","shell.execute_reply.started":"2022-01-12T15:57:04.519093Z","shell.execute_reply":"2022-01-12T15:57:04.536270Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create the model object\nmodel = CardiacDetectionModel()","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:04.538081Z","iopub.execute_input":"2022-01-12T15:57:04.538401Z","iopub.status.idle":"2022-01-12T15:57:07.240683Z","shell.execute_reply.started":"2022-01-12T15:57:04.538365Z","shell.execute_reply":"2022-01-12T15:57:07.239887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint_callback = ModelCheckpoint(\n    monitor=\"Val Loss\",\n    save_top_k=10,\n    mode=\"min\"    \n)","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:07.242013Z","iopub.execute_input":"2022-01-12T15:57:07.242336Z","iopub.status.idle":"2022-01-12T15:57:07.247004Z","shell.execute_reply.started":"2022-01-12T15:57:07.242300Z","shell.execute_reply":"2022-01-12T15:57:07.246198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create the trainer\n# Change the gpus parameter to the number of available gpus in your computer. Use 0 for CPU training\n#trainer = pl.Trainer(logger=TensorBoardLogger(\"./logs\"), log_every_n_steps=1, callbacks=checkpoint_callback, max_epochs=10)\n\ngpus = 1 #TODO\ntrainer = pl.Trainer(gpus=gpus, logger=TensorBoardLogger(\"./logs\"), log_every_n_steps=1,\n                     callbacks=checkpoint_callback,\n                     max_epochs=100)","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:07.250140Z","iopub.execute_input":"2022-01-12T15:57:07.250760Z","iopub.status.idle":"2022-01-12T15:57:07.313943Z","shell.execute_reply.started":"2022-01-12T15:57:07.250724Z","shell.execute_reply":"2022-01-12T15:57:07.313210Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.fit(model, train_loader, val_loader)","metadata":{"execution":{"iopub.status.busy":"2022-01-12T15:57:07.317287Z","iopub.execute_input":"2022-01-12T15:57:07.317519Z","iopub.status.idle":"2022-01-12T16:03:45.597973Z","shell.execute_reply.started":"2022-01-12T15:57:07.317489Z","shell.execute_reply":"2022-01-12T16:03:45.597055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport matplotlib.patches as patches","metadata":{"execution":{"iopub.status.busy":"2022-01-12T16:03:50.293643Z","iopub.execute_input":"2022-01-12T16:03:50.293926Z","iopub.status.idle":"2022-01-12T16:03:50.299720Z","shell.execute_reply.started":"2022-01-12T16:03:50.293892Z","shell.execute_reply":"2022-01-12T16:03:50.297131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.load_from_checkpoint(\"logs/default/version_0/checkpoints/epoch=99-step=4999.ckpt\")\nmodel.eval();\nmodel.to(device);","metadata":{"execution":{"iopub.status.busy":"2022-01-12T16:04:02.702648Z","iopub.execute_input":"2022-01-12T16:04:02.703451Z","iopub.status.idle":"2022-01-12T16:04:03.078059Z","shell.execute_reply.started":"2022-01-12T16:04:02.703410Z","shell.execute_reply":"2022-01-12T16:04:03.077305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = []\nlabels =[]\n\nwith torch.no_grad():\n    for data, label in val_dataset:\n        data = data.to(device).float().unsqueeze(0)\n        pred = model(data)[0].cpu()\n        preds.append(pred)\n        labels.append(label)\n        \npreds=torch.stack(preds)\nlabels=torch.stack(labels)","metadata":{"execution":{"iopub.status.busy":"2022-01-12T16:04:04.265115Z","iopub.execute_input":"2022-01-12T16:04:04.265773Z","iopub.status.idle":"2022-01-12T16:04:04.749609Z","shell.execute_reply.started":"2022-01-12T16:04:04.265732Z","shell.execute_reply":"2022-01-12T16:04:04.748898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds","metadata":{"execution":{"iopub.status.busy":"2022-01-12T16:04:20.938161Z","iopub.execute_input":"2022-01-12T16:04:20.938418Z","iopub.status.idle":"2022-01-12T16:04:20.951234Z","shell.execute_reply.started":"2022-01-12T16:04:20.938389Z","shell.execute_reply":"2022-01-12T16:04:20.950507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels","metadata":{"execution":{"iopub.status.busy":"2022-01-12T16:22:18.286660Z","iopub.execute_input":"2022-01-12T16:22:18.286926Z","iopub.status.idle":"2022-01-12T16:22:18.302127Z","shell.execute_reply.started":"2022-01-12T16:22:18.286897Z","shell.execute_reply":"2022-01-12T16:22:18.301302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels.shape","metadata":{"execution":{"iopub.status.busy":"2022-01-12T16:22:38.070370Z","iopub.execute_input":"2022-01-12T16:22:38.070622Z","iopub.status.idle":"2022-01-12T16:22:38.075931Z","shell.execute_reply.started":"2022-01-12T16:22:38.070594Z","shell.execute_reply":"2022-01-12T16:22:38.075253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"abs(preds-labels).mean(0) # --> 5/224","metadata":{"execution":{"iopub.status.busy":"2022-01-12T16:14:11.088773Z","iopub.execute_input":"2022-01-12T16:14:11.089096Z","iopub.status.idle":"2022-01-12T16:14:11.095960Z","shell.execute_reply.started":"2022-01-12T16:14:11.089062Z","shell.execute_reply":"2022-01-12T16:14:11.095292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"abs(preds-labels).mean(0).sum() / 4 / 224 * 100","metadata":{"execution":{"iopub.status.busy":"2022-01-12T16:10:59.194039Z","iopub.execute_input":"2022-01-12T16:10:59.194656Z","iopub.status.idle":"2022-01-12T16:10:59.202788Z","shell.execute_reply.started":"2022-01-12T16:10:59.194618Z","shell.execute_reply":"2022-01-12T16:10:59.201887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IDX = 50  # Feel free to inspect all validation samples by changing the index\nimg, label = val_dataset[IDX]\ncurrent_pred = preds[IDX]\n\nfig, axis = plt.subplots(1, 1)\naxis.imshow(img[0], cmap=\"bone\")\nheart = patches.Rectangle((current_pred[0], current_pred[1]), current_pred[2]-current_pred[0],\n                          current_pred[3]-current_pred[1], linewidth=1, edgecolor='r', facecolor='none')\naxis.add_patch(heart)\n\nprint(label)","metadata":{"execution":{"iopub.status.busy":"2022-01-12T16:54:33.740245Z","iopub.execute_input":"2022-01-12T16:54:33.740508Z","iopub.status.idle":"2022-01-12T16:54:33.983244Z","shell.execute_reply.started":"2022-01-12T16:54:33.740479Z","shell.execute_reply":"2022-01-12T16:54:33.982588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds[50]","metadata":{"execution":{"iopub.status.busy":"2022-01-12T16:06:01.379788Z","iopub.execute_input":"2022-01-12T16:06:01.380564Z","iopub.status.idle":"2022-01-12T16:06:01.390524Z","shell.execute_reply.started":"2022-01-12T16:06:01.380527Z","shell.execute_reply":"2022-01-12T16:06:01.389812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}