{"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"},{"sourceType":"kernelVersion","sourceId":152123399}],"dockerImageVersionId":30587,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport math\nimport glob\nimport tqdm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision\nimport torchvision.transforms as T\nimport torchvision.transforms.functional as TF\nimport cv2\nimport numpy as np","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-01T03:10:17.481068Z","iopub.execute_input":"2023-12-01T03:10:17.481861Z","iopub.status.idle":"2023-12-01T03:10:20.667887Z","shell.execute_reply.started":"2023-12-01T03:10:17.481793Z","shell.execute_reply":"2023-12-01T03:10:20.666580Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_volume(dataset, labeled=True, slice_range=None):\n    ''' Load slices into a volume. Keeps the memory requirement\n        as low as possible by using uint8 and uint16 in CPU memory.\n    '''\n    if labeled:\n        path = os.path.join(dataset, \"labels\", \"*.tif\")\n    else:\n        path = os.path.join(dataset, \"images\", \"*.tif\")\n        \n    dataset = sorted(glob.glob(path))\n    volume = None\n    target = None\n    keys = []\n    offset = 0 if slice_range is None else slice_range[0]\n    depth = len(dataset) if slice_range is None else slice_range[1]-slice_range[0]\n    \n    for z, path in enumerate(tqdm.tqdm(dataset)):\n        if slice_range is not None:\n            if z < slice_range[0]: continue\n            if z >= slice_range[1]: continue\n        \n        parts = path.split(os.path.sep)\n        key = parts[-3] + \"_\" + parts[-1].split(\".\")[0]\n        keys.append(key)\n                \n        if labeled:\n            label = cv2.imread(path, cv2.IMREAD_ANYDEPTH)\n            label = np.array(label,dtype=np.uint8)\n            if target is None:\n                target = np.zeros((1,depth, *label.shape[-2:]), dtype=np.uint8)\n            target[:,z-offset] = label\n        \n        path = path.replace(\"/labels/\",\"/images/\")\n        path = path.replace(\"/kidney_3_dense/\",\"/kidney_3_sparse/\")\n        image = cv2.imread(path, cv2.IMREAD_ANYDEPTH)\n        image = np.array(image,dtype=np.uint16)\n        \n        if volume is None:\n            volume = np.zeros((1,depth, *image.shape[-2:]), dtype=np.uint16)\n        volume[:,z-offset] = image\n    \n    return volume, target, keys","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-01T03:10:20.672617Z","iopub.execute_input":"2023-12-01T03:10:20.673660Z","iopub.status.idle":"2023-12-01T03:10:20.689409Z","shell.execute_reply.started":"2023-12-01T03:10:20.673605Z","shell.execute_reply":"2023-12-01T03:10:20.688063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RandomVolumetricDataset(torch.utils.data.Dataset):\n    ''' Dataset for segmentation of a sparse class. Keeps\n        track of positive samples and favors samples that\n        contain a positive sample.\n        WARNING: do not use in a distributed setting.\n    '''\n    def __init__(self, datasets, shape=(256,256,256), length=1000, transform=None):\n        self.volumes = []\n        self.targets = []\n        self.length = length\n        self.shape = shape\n        self.transform = transform\n        self.nonzero = []\n        \n        for dataset in datasets:\n            print(\"loading volume\", dataset)\n            volume, target, _ = load_volume(dataset)\n            self.volumes.append(volume)\n            self.targets.append(target)\n            self.nonzero.append(np.argwhere(target > 0))\n        \n    def __len__(self):\n        return self.length\n\n    def __getitem__(self, idx):\n        vidx = torch.randint(len(self.volumes), (1,)).item()\n        volume = self.volumes[vidx]\n        target = self.targets[vidx]\n        nonzero = self.nonzero[vidx]\n        random = torch.rand(1)\n        \n        if random > 0.9:\n            # Load a random subvolume\n            z,y,x = torch.randint(volume.shape[-3]-self.shape[-3], (1,)).item(), \\\n                    torch.randint(volume.shape[-2]-self.shape[-2], (1,)).item(), \\\n                    torch.randint(volume.shape[-1]-self.shape[-1], (1,)).item()\n        else:\n            # Load a subvolume containing a random sample\n            idx = torch.randint(len(nonzero), (1,)).item()\n            c,z,y,x = nonzero[idx]\n            \n            z += torch.randint(self.shape[-3],(1,)).sub(self.shape[-3]//2).item()\n            y += torch.randint(self.shape[-2],(1,)).sub(self.shape[-2]//2).item()\n            x += torch.randint(self.shape[-1],(1,)).sub(self.shape[-1]//2).item()\n            \n            z = min(max(0,z+self.shape[-3]//2), volume.shape[-3]-self.shape[-3])\n            y = min(max(0,y+self.shape[-2]//2), volume.shape[-2]-self.shape[-2])\n            x = min(max(0,x+self.shape[-3]//2), volume.shape[-1]-self.shape[-1])\n            \n        volume = volume[:,z:z+self.shape[-3], y:y+self.shape[-2], x:x+self.shape[-1]]\n        target = target[:,z:z+self.shape[-3], y:y+self.shape[-2], x:x+self.shape[-1]]\n\n        volume = torch.from_numpy((volume/65536).astype(np.float32))\n        target = torch.from_numpy(target > 0).float()\n        if self.transform is not None:\n            rng = torch.get_rng_state()\n            volume = self.transform(volume)\n            torch.set_rng_state(rng)\n            target = self.transform(target)\n        \n        return volume, target\n                ","metadata":{"execution":{"iopub.status.busy":"2023-12-01T03:10:20.691277Z","iopub.execute_input":"2023-12-01T03:10:20.692086Z","iopub.status.idle":"2023-12-01T03:10:20.715570Z","shell.execute_reply.started":"2023-12-01T03:10:20.692046Z","shell.execute_reply":"2023-12-01T03:10:20.714209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# The augmentations\n\nclass RandomRotationNd(nn.Module):\n    ''' This augmentation first permutes the dimensions as an initial rotation\n        to select the rotation axis, then rotates around the (fixed) z axis. \n        The result is zoomed in to remove empty space and finally permuted \n        once more to move randomize the rotation axis.\n    '''\n    def __init__(self, dims):\n        super().__init__()\n        self.dims = dims\n\n    def forward(self, x):\n        angle = torch.rand(1).item() * 360\n        keep = torch.arange(x.dim() - self.dims)\n        perm = -torch.randperm(self.dims)-1\n        x = x.clone().permute(*[k.item() for k in keep], *[p.item() for p in perm])\n        rad = math.pi * angle / 180\n        scale = abs(math.sin(rad)) + abs(math.cos(rad))\n        for i in range(0, x.shape[-3],8): # presumptuous\n            v = x[...,i:i+8,:,:]\n            w = v.view(-1, *v.shape[-3:])\n            w = TF.rotate(w, angle)\n            v = w.view(*v.shape)\n            x[...,i:i+8,:,:] = v\n        s = x.shape\n        x = F.interpolate(x, scale_factor=scale, mode=\"bilinear\")\n        x = TF.center_crop(x, s[-2:])\n        perm = -torch.randperm(self.dims)-1\n        x = x.permute(*[k.item() for k in keep], *[p.item() for p in perm])\n        return x\n\nclass RandomRot90Nd(nn.Module):\n    def __init__(self, dims):\n        super().__init__()\n        self.dims = dims\n\n    def forward(self, x):\n        dims = -torch.randperm(self.dims)[:2]-1\n        dims = [d.item() for d in dims]\n        rot = torch.randint(4, (1,)).item()\n        return x.rot90(rot, dims)\n\nclass RandomPermuteNd(nn.Module):\n    def __init__(self, dims):\n        super().__init__()\n        self.dims = dims\n\n    def forward(self, x):\n        perm = -torch.randperm(self.dims)-1\n        keep = torch.arange(x.dim() - self.dims)\n        return x.permute(*[k.item() for k in keep], *[p.item() for p in perm])\n\nclass RandomFlipNd(nn.Module):\n    def  __init__(self, dims, p=0.5):\n        super().__init__()\n        self.dims = dims\n        self.p = p\n        \n    def forward(self, x):\n        for i in range(self.dims):\n            if torch.rand(1) < self.p:\n                x = x.flip(-i-1)\n        return x\n\nclass ToDevice(nn.Module):\n    ''' Sometimes it helps to move the tensor to the gpu before augmentations like\n        rotation. Note however that you need to set num_workers to 0 in the dataloader\n    '''\n    def __init__(self, device=None):\n        super().__init__()\n        self.device = device or (\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n    def forward(self, x):\n        return x.to(self.device)\n ","metadata":{"execution":{"iopub.status.busy":"2023-12-01T03:10:20.720407Z","iopub.execute_input":"2023-12-01T03:10:20.720803Z","iopub.status.idle":"2023-12-01T03:10:20.743020Z","shell.execute_reply.started":"2023-12-01T03:10:20.720770Z","shell.execute_reply":"2023-12-01T03:10:20.741745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Constructing the dataset\n\ntransform = T.Compose((ToDevice(), RandomRotationNd(3), RandomFlipNd(3)))\nds = RandomVolumetricDataset([\n    \"/kaggle/input/blood-vessel-segmentation/train/kidney_1_dense\",\n    \"/kaggle/input/blood-vessel-segmentation/train/kidney_3_dense\"\n], length=1000, transform=transform)","metadata":{"execution":{"iopub.status.busy":"2023-12-01T03:10:20.744640Z","iopub.execute_input":"2023-12-01T03:10:20.745148Z","iopub.status.idle":"2023-12-01T03:12:09.283774Z","shell.execute_reply.started":"2023-12-01T03:10:20.745103Z","shell.execute_reply":"2023-12-01T03:12:09.281993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import inline\n\nprint(\"Each sample returned from the dataset is random and augmented\")\ntorch.manual_seed(125)\nvolume, target = ds[0] # ^ irregardless of idx, which is why it \n                       # doesn't work in distributed settings \n\nvolume = volume.sub(volume.mean()).div(volume.std().add(1e-5))\ninline.plot(torch.stack((inline.disp(volume[0]), inline.disp(target[0]))), width=10)\n\nprint(\"For show: more augmentation of the same subvolume\")\n# Showing the same subvolume, with random rotations\nrot = RandomRotationNd(3)\nrng = torch.get_rng_state()\nvolumes = torch.stack([rot(volume)[0] for _ in range(8)])\ntorch.set_rng_state(rng)\ntargets = torch.stack([rot(target)[0] for _ in range(8)])\ninline.plot(volumes.mul(0.288).add(0.5)[:,[100]])\ninline.plot(volumes)\ninline.plot(targets)","metadata":{"execution":{"iopub.status.busy":"2023-12-01T03:14:44.642981Z","iopub.execute_input":"2023-12-01T03:14:44.643421Z","iopub.status.idle":"2023-12-01T03:14:56.240224Z","shell.execute_reply.started":"2023-12-01T03:14:44.643390Z","shell.execute_reply":"2023-12-01T03:14:56.239038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}