{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"}],"dockerImageVersionId":30587,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport torch\nfrom PIL import Image\nfrom torch.utils.data import Dataset\nfrom typing import Optional, List\nfrom tqdm import tqdm\nimport rasterio\nfrom functools import lru_cache\nimport numpy as np\nimport pandas as pd\n\nimport warnings\nwarnings.filterwarnings(\"ignore\", category=rasterio.errors.NotGeoreferencedWarning)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-11-22T15:50:19.822355Z","iopub.execute_input":"2023-11-22T15:50:19.823068Z","iopub.status.idle":"2023-11-22T15:50:19.830952Z","shell.execute_reply.started":"2023-11-22T15:50:19.823022Z","shell.execute_reply":"2023-11-22T15:50:19.829755Z"},"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_decode(mask_rle: str, img_shape: tuple = None) -> np.ndarray:\n    seq = mask_rle.split()\n    starts = np.array(list(map(int, seq[0::2])))\n    lengths = np.array(list(map(int, seq[1::2])))\n    assert len(starts) == len(lengths)\n    ends = starts + lengths\n    img = np.zeros((np.product(img_shape),), dtype=np.uint8)\n    for begin, end in zip(starts, ends):\n        img[begin:end] = 1\n    # https://stackoverflow.com/a/46574906/4521646\n    img.shape = img_shape\n    return img\n\n@lru_cache(maxsize=64)\ndef ropen(img_fpath):\n    return rasterio.open(img_fpath)","metadata":{"execution":{"iopub.status.busy":"2023-11-22T15:50:19.833263Z","iopub.execute_input":"2023-11-22T15:50:19.833977Z","iopub.status.idle":"2023-11-22T15:50:19.846533Z","shell.execute_reply.started":"2023-11-22T15:50:19.833935Z","shell.execute_reply":"2023-11-22T15:50:19.844995Z"},"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# In-Memory Tiling\n\nOne technique I've observed folks using is to resize the input image down to a constant scale. This introduces distortions and issues with the number of pixels-on-target disappearing (e.g., a 10 pixel target downsampled by 1/10 becomes a 1 pixel target, while a 5 pixel target becomes a 0 pixel target).\n\nRather than resize the imagery, I've adopted a technique that's common in geospatial processing- tiling. Tiling preserves the original target sizes and avoids distortions by simply applying a sliding window to the data. The window can be square or unsquare, avoiding aspect ratio issues. The window also has a stride associated with it, allowing the network multiple opportunities to \"detect\" a target if it misses. Critically, by tiling, this increases the percieved target scale as observed by the model (e.g., a 10 pixel target looks smaller in a 1000x1000 image than a 100x100 image). Finally, tiling the inflates the number of samples available from training and increases sample diversity","metadata":{}},{"cell_type":"code","source":"class SenNetHOATiledDataset(Dataset):\n\n    def __init__(\n            self,\n            df_data: pd.DataFrame,\n            path_img_dir: str,\n            transforms=None,\n            tile_size: Optional[List[int]] = None,\n            overlap_pct: float = 0.2,\n            empty_tile_pct: float = 0.0,\n            sample_limit: int = 30000*10,\n            cache_dir: str = None\n    ):\n        \"\"\"\n        Generates an in-memory tiled dataset using a sliding window approach (common in geospatial processing communities)\n\n        Args:\n            df_data (pd.DataFrame): Training dataframe.\n            path_img_dir (str): Path to root directory containing data.\n            transforms (torchvision.transforms.Transform): Albumentation transforms.\n            tile_size (Optional[List[int]]): List describing the tile size. If None, tiling is not performed and the\n                full scene is utilized.\n            overlap_pct (float): Percentage of tile size to use as overlap between tiles. This effectively\n                controls the \"stride\" of the window as it's slid across the full scene.\n            empty_tile_pct (float): Percentage of final dataset that should be empty tiles. Default is 0.0 (no empty tiles)\n            sample_limit (int): Upper-bound on the number of samples to return.\n            cache_dir (str): Cache dir to read/write to. If None, no cache directory is utilized.\n        \"\"\"\n        self.path_img_dir = path_img_dir\n        self.transforms = transforms\n        self.data = df_data\n        self.samples = []\n        self.tile_size = np.array(tile_size) if tile_size is not None else None\n        self.overlap_pct = overlap_pct\n        self.empty_tile_pct = empty_tile_pct\n        self.sample_limit = sample_limit\n        self.cache_dir = cache_dir\n\n        for _, row in self.data.iterrows():\n            p_img = os.path.join(self.path_img_dir, row[\"dataset\"], \"images\", f'{row[\"slice\"]}.tif')\n\n            with rasterio.open(p_img) as reader:\n                width, height = reader.width, reader.height\n                img = reader.read()\n                px_max, px_min = img.max(), img.min()\n\n            if not os.path.isfile(p_img):\n                continue\n            self.samples.append((p_img, row['rle'], [], [], 1, [px_min, px_max], [width, height]))\n\n        if self.tile_size is not None:\n            empty = 0\n            nonempty = 0\n\n            empty_tiles = []\n            populated_tiles = []\n\n            for file_path, rle, _, _, _, px_stats, img_dims in tqdm(self.samples, total=len(self.samples), desc='Generating tiles'):\n                width, height = img_dims\n\n                mask = rle_decode(rle, img_shape=[height, width])\n\n                min_overlap = float(overlap_pct) * 0.01\n                max_stride = self.tile_size * (1.0 - min_overlap)\n                num_patches = np.ceil(np.array([height, width]) / max_stride).astype(np.int64)\n                # compute the cutoff points for the x-y dimensions\n                starts = [np.int64(np.linspace(0, width - self.tile_size[1], num_patches[1])),\n                          np.int64(np.linspace(0, height - self.tile_size[0], num_patches[0]))]\n                stops = [starts[0] + self.tile_size[0], starts[1] + self.tile_size[1]]\n                for y1, y2 in zip(starts[1], stops[1]):\n                    for x1, x2 in zip(starts[0], stops[0]):\n                        this_region = mask[y1:y2, x1:x2]\n                        is_empty = np.all(this_region == 0)\n\n                        if self.empty_tile_pct == 0.0:\n                            is_empty = is_empty or (this_region.sum() < (0.05 * self.tile_size[0]))\n\n                        if is_empty:\n                            empty += 1\n                            empty_tiles.append((file_path, rle, [x1, y1, x2 - x1, y2 - y1], [height, width], 0, px_stats, img_dims))\n\n                        else:\n                            nonempty += 1\n                            populated_tiles.append((file_path, rle, [x1, y1, x2 - x1, y2 - y1], [height, width], 1, px_stats, img_dims))\n\n            num_empty_tiles_to_sample = int(self.sample_limit * self.empty_tile_pct)\n            num_pos_tiles_to_sample = int(self.sample_limit * (1 - self.empty_tile_pct))\n\n            empty_idxs_to_sample = np.random.choice(len(empty_tiles), min(num_empty_tiles_to_sample, len(empty_tiles)), replace=False)\n            pos_idxs_to_sample = np.random.choice(len(populated_tiles), min(num_pos_tiles_to_sample, len(populated_tiles)), replace=False)\n\n            neg_samples = list(map(empty_tiles.__getitem__, empty_idxs_to_sample))\n            pos_samples = list(map(populated_tiles.__getitem__, pos_idxs_to_sample))\n\n            new_samples = pos_samples + neg_samples\n\n            self.samples = new_samples\n            if self.empty_tile_pct == 0.0:\n                print(f'Dropped {empty} empty tiles.')\n            print(f'Dataset contains {len(neg_samples)} empty and {len(pos_samples)} non-empty tiles.')\n\n    def __getitem__(self, idx: int) -> tuple:\n        \"\"\"\n        Returns a sample of data\n\n        \"\"\"\n        # Grab the sample from the sample list\n        img_fpath, rle, bbox, original_img_size, target, px_stats, img_dims = self.samples[idx]\n\n        # If the user points us to a cache directory, generate the file paths to the cached imagery\n        cache_file_img = None\n        cache_file_mask = None\n        if self.cache_dir is not None:\n            cache_file_img = os.path.join(self.cache_dir, img_fpath.split('/')[-3], os.path.basename(img_fpath).split('.')[\n                0] + f'_{bbox[0]}_{bbox[1]}_{bbox[2]}_{bbox[3]}.png')\n            cache_file_mask = os.path.join(self.cache_dir, img_fpath.split('/')[-3],\n                                       os.path.basename(img_fpath).split('.')[\n                                           0] + f'_{bbox[0]}_{bbox[1]}_{bbox[2]}_{bbox[3]}_mask.png')\n\n        # If the cached image exists, load it. Otherwise, read the window from the full scene\n        if cache_file_img and os.path.exists(cache_file_img):\n            img = np.array(Image.open(cache_file_img))\n        else:\n            img = ropen(img_fpath).read(1, window=rasterio.windows.Window(*bbox) if len(bbox) > 0 else None)\n\n            if len(original_img_size) == 0:\n                original_img_size = img.shape\n\n            if img.ndim == 3:\n                img = np.mean(img, axis=2)\n\n            # If we read the window from the full scene, compress the dynamic range to UINT8\n            # TODO: Save full scene statistics and compress relative to those, rather than tile level statistics\n            img = (img - px_stats[0]) / (px_stats[1] - px_stats[0])\n            img *= 255.0\n            img = img.astype(np.uint8)\n\n        # If the cached mask exists, load it. Otherwise, decode the rle and grab the window\n        if cache_file_mask and os.path.exists(cache_file_mask):\n            mask = np.array(Image.open(cache_file_mask))\n        else:\n            mask = rle_decode(rle, img_shape=original_img_size)\n\n            if len(bbox) > 0:\n                x, y, w, h = bbox\n                mask = mask[y:y + h, x:x + w]\n\n        # If the user requested use of a cache dir and the images don't exist on-disk, write them.\n        if self.cache_dir is not None:\n            os.makedirs(os.path.dirname(cache_file_img), exist_ok=True)\n            if not os.path.exists(cache_file_img):\n                im = Image.fromarray(img)\n                im.save(cache_file_img)\n\n            if not os.path.exists(cache_file_mask):\n                im = Image.fromarray(mask)\n                im.save(cache_file_mask)\n\n        # Transform the image and mask using Albumentations\n        if self.transforms:\n            transformed = self.transforms(image=img, mask=mask)\n            img = transformed['image']\n            mask = transformed['mask']\n\n        # The augmentations may scale the mask to the range 0-1, with 1's becoming ~0.0039 (1/255).\n        # Make sure we overwrite this in any case\n        mask[mask > 0] = 1\n\n        # Return the sample\n        return img, mask, img_fpath, bbox, original_img_size, target\n\n    def __len__(self) -> int:\n        return len(self.samples)","metadata":{"execution":{"iopub.status.busy":"2023-11-22T16:18:53.210780Z","iopub.execute_input":"2023-11-22T16:18:53.211187Z","iopub.status.idle":"2023-11-22T16:18:53.244370Z","shell.execute_reply.started":"2023-11-22T16:18:53.211153Z","shell.execute_reply":"2023-11-22T16:18:53.243186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Construct a dataset and dataloader with some Augmentations\n\nThe in-memory tiling process isn't free- it takes some time at startup. The time required to tile the imagery scales inversely proportional with tile size (i.e., smaller tiles = longer wait time).","metadata":{}},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n\nDATASET_FOLDER = '/kaggle/input/blood-vessel-segmentation'\nIMG_PATH = DATASET_FOLDER + '/train'\nTO_VISUALIZE = 'kidney_1_dense'\n\ndf_train = pd.read_csv(os.path.join(DATASET_FOLDER, \"train_rles.csv\"))\ndf_train[[\"dataset\", \"slice\"]] = df_train['id'].str.rsplit(pat='_', n=1, expand=True)\n\nvalid_transforms = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.ToRGB(),\n    ToTensorV2()\n])\n\ndataset = SenNetHOATiledDataset(\n                            df_train.loc[df_train.dataset == TO_VISUALIZE],\n                            path_img_dir=IMG_PATH,\n                            tile_size=[800, 800],\n                            empty_tile_pct=0.0,\n                            transforms=valid_transforms)\n\nloader = torch.utils.data.DataLoader(\n    dataset,\n    batch_size=8,\n    num_workers=4,\n    shuffle=False,\n    pin_memory=True\n)","metadata":{"execution":{"iopub.status.busy":"2023-11-22T15:56:04.590365Z","iopub.execute_input":"2023-11-22T15:56:04.590755Z","iopub.status.idle":"2023-11-22T15:57:04.691798Z","shell.execute_reply.started":"2023-11-22T15:56:04.590726Z","shell.execute_reply":"2023-11-22T15:57:04.690820Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize the tiles\n\nNow that the data is tiled, we can visualize the resulting pixel data.","metadata":{"execution":{"iopub.status.busy":"2023-11-22T16:09:37.644254Z","iopub.execute_input":"2023-11-22T16:09:37.644898Z","iopub.status.idle":"2023-11-22T16:09:37.652203Z","shell.execute_reply.started":"2023-11-22T16:09:37.644853Z","shell.execute_reply":"2023-11-22T16:09:37.650558Z"}}},{"cell_type":"code","source":"from skimage import color\nimport matplotlib.pyplot as plt\n\nfor i in range(10, 13):\n    fig, axarr = plt.subplots(ncols=3, figsize=(12, 6))\n    img, mask, _, _, _, _ = dataset[i]\n    img = img.cpu().numpy().transpose(1, 2, 0)\n    mask = mask.cpu().numpy()\n    axarr[0].imshow(img, cmap=\"gray\")\n    axarr[1].imshow(color.label2rgb(mask, img, bg_label=0, bg_color=(1.,1.,1.), alpha=0.25))\n    axarr[2].imshow(mask, vmin=0, interpolation='antialiased', interpolation_stage='rgba')\n    \n    for i in range(3):\n        axarr[i].set_axis_off()\n    \n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-11-22T16:08:52.715247Z","iopub.execute_input":"2023-11-22T16:08:52.716065Z","iopub.status.idle":"2023-11-22T16:08:55.853325Z","shell.execute_reply.started":"2023-11-22T16:08:52.716024Z","shell.execute_reply":"2023-11-22T16:08:55.852111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Compute dataset statistics\n\nSomething a significant number of people overlook is computing the normalization statistics of their input data. I am guilty of simply using the imagenet/coco/pretrained model statistics, regardless of my underlying dataset statistics. However, I've executed some (probably poorly configured) experiments in the past that have shown that, depending on the pretrained model source and experiment configuration, using the statistics for your dataset can improve stability of convergence.","metadata":{}},{"cell_type":"code","source":"psum = torch.tensor([0.0, 0.0, 0.0])\npsum_sq = torch.tensor([0.0, 0.0, 0.0])\nshape = None\n\n# loop through images\nfor inputs in tqdm(loader):\n    img, _, _, _, _, _ = inputs\n    if shape is None:\n        shape = torch.tensor(img.shape[2:]).prod()\n    img = img.float() / 255.0\n    psum += img.sum(axis=[0, 2, 3])\n    psum_sq += (img ** 2).sum(axis=[0, 2, 3])\n    \ncount = (len(dataset) * shape).item()\n    \n# mean and STD\ntotal_mean = (psum / count) / 255\ntotal_var = ((psum_sq / count) / 255) - (total_mean ** 2)\ntotal_std = torch.sqrt(total_var)\n    \n# output\nprint('Training data stats:')\nprint(total_mean)\nprint(total_std)","metadata":{"execution":{"iopub.status.busy":"2023-11-22T16:11:52.845190Z","iopub.execute_input":"2023-11-22T16:11:52.846075Z","iopub.status.idle":"2023-11-22T16:15:25.368884Z","shell.execute_reply.started":"2023-11-22T16:11:52.846036Z","shell.execute_reply":"2023-11-22T16:15:25.367446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Write the tiles to disk\n\nWe can write the tiles to disk be specifying the `cache_dir` argument to the dataset. It seems counter intuitive, but hear me out- once you settle on a tile size, this will speed up training by reducing file I/O. But wouldn't reading more files from disk mean more file I/O? Yes and No. The size of the data being read from disk matters just as much as the location. When randomly sampling small windows (e.g., window size of 256x256 at train time), it becomes very expensive to perform random access from full scenes (i.e., load the full scene just to grab one small 256x256 window from it, before moving to a new scene). This could be mitigated with an intelligent sampler that samples tiles randomly from a single image before moving on, but the easiest path forward is to write these tiles to disk.","metadata":{}},{"cell_type":"code","source":"df_to_tile = df_train.iloc[490:500] # I'm subsampling the dataframe- writing to disk can be slow.\nprint(df_to_tile)\ndataset = SenNetHOATiledDataset(\n                            df_to_tile,\n                            path_img_dir=IMG_PATH,\n                            tile_size=[800, 800],\n                            empty_tile_pct=0.0,\n                            transforms=valid_transforms,\n                            cache_dir='800x800')\n\nloader = torch.utils.data.DataLoader(\n    dataset,\n    batch_size=5,\n    num_workers=4,\n    shuffle=False,\n    pin_memory=True\n)\n\nfor x in tqdm(loader, desc='tiling dataset to disk'):\n        pass","metadata":{"execution":{"iopub.status.busy":"2023-11-22T16:18:57.397655Z","iopub.execute_input":"2023-11-22T16:18:57.398074Z","iopub.status.idle":"2023-11-22T16:18:59.995620Z","shell.execute_reply.started":"2023-11-22T16:18:57.398041Z","shell.execute_reply":"2023-11-22T16:18:59.994209Z"},"trusted":true},"execution_count":null,"outputs":[]}]}