{"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"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":155317483,"sourceType":"kernelVersion"},{"sourceId":155317571,"sourceType":"kernelVersion"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport sys\nimport glob\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.distributed as dist\nimport torch.multiprocessing as mp\nimport torchvision\nimport torchvision.transforms as T\nimport torchvision.transforms.functional as TF\nimport numpy as np\nimport warnings\nimport cv2\n\n'''\nUsage:\n\npython <script.py> <data_dir>\n- run with single GPU, no DistributedDataParallel\n\npython <script.py> <data_dir> <N>\n- run with <N> GPUs, using DistributedDataparallel\n'''\nassert len(sys.argv) > 1\nDATA_DIR = sys.argv[1]\nWORLDSIZE = int(sys.argv[2]) if len(sys.argv) == 3 else 0 \nNPROCS = WORLDSIZE\nif not WORLDSIZE:\n    warnings.warn(\"running without DistributedDataParallel, use program argument > 0\")\n\nHAS_INLINE = os.path.exists(\"inline.py\")\nif HAS_INLINE: \n    import inline\n\nHAS_METERS = os.path.exists(\"meters.py\")\nif HAS_METERS:\n    from meters import Meter, Meters\n\n\n# Your kidney segmentation model\nclass SegmentKidneyStub(nn.Module):\n    ''' Segment the full kidney using 2d slices as input '''\n    def forward(self, x):\n        return x\n        \n# Your vessel segmentation model\nclass SegmentVesselStub(nn.Module):\n    ''' Segment the blood vessels using 3d chunks as input '''\n    def forward(self, x):\n        return x\n\n\ndef load_volume(dataset_path, labeled=True, slice_range=None, padding=0, padding_mode=\"zeros\"):\n    \n    if labeled:\n        path = os.path.join(dataset_path, \"labels\", \"*.tif\")\n    else:\n        path = os.path.join(dataset_path, \"images\", \"*.tif\")\n        \n    dataset = sorted(glob.glob(path))\n    volume = None\n    target = None\n    mask = 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    if HAS_METERS:\n        meters = Meters(\n            time = Meter(style=\"timer\"),\n            loading = Meter(style=\"text\", initial_value=f\"Loading {dataset_path}\"),\n            z = Meter(style=\"progress\", max_value=len(dataset)))\n            \n    for z, path in enumerate(dataset): \n        if HAS_METERS:\n            meters.z.increment()\n            meters.print(end=\"\\r\")\n            \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+2*padding, label.shape[-2]+2*padding, label.shape[-1]+2*padding), dtype=np.uint8)\n            if padding > 0:\n                target[:,z-offset+padding,padding:-padding,padding:-padding] = label\n            else:\n                target[:,z-offset] = label\n\n            mask_path = path.replace(\"/labels/\", \"/kidney-mask/\")\n            mask_path = mask_path.replace(\".tif\",\".png\")\n            if os.path.exists(mask_path):\n                mask_label = cv2.imread(mask_path, cv2.IMREAD_ANYDEPTH)\n                mask_label = np.array(mask_label, dtype=np.uint8)\n                mask_label = F.interpolate(torch.from_numpy(mask_label).div(255)[None,None], label.shape[-2:], mode=\"bilinear\")[0,0]\n                mask_label = (mask_label > 0.5).byte().numpy()\n            else:\n                mask_label = np.full_like(label, 1)\n            \n            if mask is None:\n                mask = np.zeros_like(target)\n            if padding > 0:\n                mask[:,z-offset+padding,padding:-padding,padding:-padding] = mask_label\n            else:\n                mask[:,z-offset] = mask_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+2*padding, image.shape[-2]+2*padding,image.shape[-1]+2*padding), dtype=np.uint16)\n        if padding > 0:\n            volume[:,z-offset+padding,padding:-padding,padding:-padding] = image\n        else:\n            volume[:,z-offset] = image\n            \n            \n    if padding > 0 and padding_mode == \"reflect\":\n        volume[:,:padding] = np.flip(volume[:,padding:2*padding],-3)\n        volume[:,-padding:] = np.flip(volume[:,-2*padding:-padding],-3)\n        volume[:,:,:padding] = np.flip(volume[:,:,padding:2*padding],-2)\n        volume[:,:,-padding:] = np.flip(volume[:,:,-2*padding:-padding],-2)\n        volume[:,:,:,:padding] = np.flip(volume[:,:,:,padding:2*padding],-1)\n        volume[:,:,:,-padding:] = np.flip(volume[:,:,:,-2*padding:-padding],-1)\n\n        if target is not None:\n            target[:,:padding] = np.flip(target[:,padding:2*padding],-3)\n            target[:,-padding:] = np.flip(target[:,-2*padding:-padding],-3)\n            target[:,:,:padding] = np.flip(target[:,:,padding:2*padding],-2)\n            target[:,:,-padding:] = np.flip(target[:,:,-2*padding:-padding],-2)\n            target[:,:,:,:padding] = np.flip(target[:,:,:,padding:2*padding],-1)\n            target[:,:,:,-padding:] = np.flip(target[:,:,:,-2*padding:-padding],-1)\n    \n    if padding > 0 and padding_mode == \"replicate\":\n        volume[:,:padding] = volume[:,[padding]]\n        volume[:,-padding:] = volume[:,[-padding-1]]\n        volume[:,:,:padding] = volume[:,:,[padding]]\n        volume[:,:,-padding:] = volume[:,:,[-padding-1]]\n        volume[:,:,:,:padding] = volume[:,:,:,[padding]]\n        volume[:,:,:,-padding:] = volume[:,:,:,[-padding-1]]\n\n        if target is not None:\n            target[:,:padding] = target[:,[padding]]\n            target[:,-padding:] = target[:,[-padding-1]]\n            target[:,:,:padding] = target[:,:,[padding]]\n            target[:,:,-padding:] = target[:,:,[-padding-1]]\n            target[:,:,:,:padding] = target[:,:,:,[padding]]\n            target[:,:,:,-padding:] = target[:,:,:,[-padding-1]]\n             \n    return volume, target, mask, keys\n\ndef run(device=\"cuda\",\n        data_dir=DATA_DIR,\n        image_dir=\"images\",\n        datasets = [\"train/kidney_2\"],\n        threshold=64, \n        margin=8, \n        block_size=80, \n        segment_kidney_path=\"kidney.pth\",\n        segment_vessel_path=\"vessel.pth\",\n       ):\n\n    if HAS_INLINE and not os.path.exists(image_dir):\n        os.mkdir(image_dir)\n\n    if dist.is_initialized():\n        rank = dist.get_rank()\n        localrank = rank % torch.cuda.device_count()\n        worldsize = dist.get_world_size()\n    else:\n        rank, localrank, worldsize = 0, 0, 1\n\n    if device == \"cuda\":\n        device = \"cuda:\" + str(localrank)\n\n    if localrank == 0:\n        print(\"Running inference...\")\n        print()\n        \n    is_kaggle = \"KAGGLE_URL_BASE\" in os.environ\n    \n    # Kidney mask model\n    mask_model = SegmentKidneyStub().eval().to(device)\n    #assert os.path.exists(deeplab_path)\n    if os.path.exists(segment_kidney_path):\n        sd = torch.load(segment_kidney_path, map_location=\"cpu\")\n        mask_model.load_state_dict(sd[\"model\"])\n        if dist.is_initialized():\n            mask_model = nn.parallel.DistributedDataParallel(mask_model)\n    elif localrank == 0:\n        warnings.warn(f\"unable to load kidney segmentation model at {segment_kidney_path}\")\n    \n    # Vessel segmentation model\n    model = SegmentVesselStub().eval().to(device)\n    #assert os.path.exists(agfnet_path)\n    if os.path.exists(segment_vessel_path):\n        sd = torch.load(segment_vessel_path, map_location=\"cpu\")\n        model.load_state_dict(sd[\"model\"])\n        if dist.is_initialized():\n            model = nn.parallel.DistributedDataParallel(model)\n    elif localrank == 0:\n        warnings.warn(f\"Warning: unable to load vessel segmentation model at {segment_vessel_path}\")\n        \n    # Utility functions\n    def rle_encode(img):\n        pixels = img.flatten()\n        pixels = np.concatenate([[0], pixels, [0]])\n        runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n        runs[1::2] -= runs[::2]\n        return ' '.join(str(x) for x in runs) or \"1 0\"\n\n    def all_combinations(*args):\n        assert len(args)\n        if len(args) == 1: return [[a] for a in args[0]]\n        else: return [a + b for a in all_combinations(*args[:len(args)//2]) for b in all_combinations(*args[len(args)//2:])]\n\n    # mask for crop with margins removed or ramped\n    k = torch.cat((\n        torch.zeros(margin, device=\"cpu\"),\n        torch.ones(block_size-2*margin, device=\"cpu\"),\n        torch.zeros(margin, device=\"cpu\")))\n    k = k[None,None,:] * k[None,:,None] * k[:,None,None]\n\n    # Go...\n    with torch.no_grad():\n        csv_lines = [\"id,rle\\n\"]\n            \n        for ds in datasets:\n            \n            # Load the volume\n            volume, target, _, keys = load_volume(\n                os.path.join(data_dir,ds), labeled=(not is_kaggle), \n                padding=margin, padding_mode=\"reflect\")\n\n            # Only allocate the prediction and mask arrays on rank 0\n            if rank == 0:\n                prediction = np.zeros(volume.shape[1:], dtype=np.uint8)\n                mask = np.zeros(volume.shape[1:], dtype=np.uint8)\n\n            # Predict and combine kidney masks for all XY, XZ and YZ views\n            for d in range(3):\n                dim = -d-1\n                if HAS_METERS:\n                    meters = Meters(\n                        time = Meter(style=\"timer\"),\n                        title = Meter(initial_value=ds + \" (mask:\" + [\"z\",\"y\",\"x\"][dim] + \")\", style=\"text\"),\n                        it = Meter(style=\"progressbar\", max_value=(volume.shape[dim]+worldsize-1)//worldsize),\n                    )\n\n                # iterate over all slices in the current view with step size = world size\n                for z in range(0, volume.shape[dim], worldsize):\n                    if HAS_METERS:\n                        meters.it.increment()\n                        meters.print(end=\"\\r\")\n                    \n                    slice_dim = [...,slice(0,volume.shape[-3]),slice(0,volume.shape[-2]),slice(0,volume.shape[-1])]\n                    slice_dim[dim] = min(z + rank, volume.shape[dim]-1)\n                    \n                    s = volume[tuple(slice_dim)]\n                    s = torch.from_numpy((s/65535).astype(np.float32))[None].to(device)\n                    pred = mask_model(s)[0,0]\n                    \n                    # gather predictions for all GPUs and move them to rank 0\n                    # rank 0 will add them to their mask array (85 * 3 = 255)\n                    if dist.is_initialized():\n                        if rank != 0:\n                            dist.gather(pred)\n                        else:\n                            dst = [torch.zeros_like(pred) for _ in range(worldsize)]\n                            dist.gather(pred, dst)\n                            for j,p in enumerate(dst):\n                                if z+j < volume.shape[dim]:\n                                    p = p.cpu().numpy()\n                                    slice_dim[dim] = z+j\n                                    mask[tuple(slice_dim[1:])] += (85 * p).astype(np.uint8)\n                    else:\n                        mask[tuple(slice_dim[1:])] += (85 * pred.cpu().numpy()).astype(np.uint8)\n\n            # save some images\n            if rank == 0 and not is_kaggle and HAS_INLINE:\n                inline.save(mask, filename=os.path.join(image_dir, \"infer_mask.png\"))\n            if dist.is_initialized(): dist.barrier()\n\n            # Get the coordinates for all 3d chunks in the volume\n            s, d = block_size, block_size-2*margin\n            zs = list(range(0,volume.shape[-3]-margin*2,d))\n            ys = list(range(0,volume.shape[-2]-margin*2,d))\n            xs = list(range(0,volume.shape[-1]-margin*2,d))\n            vs = all_combinations(zs,ys,xs)\n\n            if HAS_METERS:\n                meters = Meters(\n                    time = Meter(style=\"timer\"),\n                    title = Meter(initial_value=ds, style=\"text\"),\n                    it = Meter(style=\"progressbar\", max_value=(len(vs)+worldsize-1)//worldsize))\n\n            # Predict vessel segmentation for each chunk\n            for i in range(0,len(vs), worldsize):\n                if HAS_METERS:\n                    meters.it.increment()\n                    meters.print(end=\"\\r\")\n                    \n                z,y,x = vs[min(i+rank,len(vs)-1)]\n                v = (volume[:,z:z+s,y:y+s,x:x+s] / 65536).astype(np.float32)\n                v = torch.from_numpy(v).to(device)\n                orig = v.shape\n                v = F.pad(v, (0,s-v.shape[-1],0,s-v.shape[-2],0,s-v.shape[-3]))\n                pred = model(v[None])\n\n                # gather predictions from each GPU and move them to rank 0\n                # rank 0 will add them to their prediction array\n                if dist.is_initialized():\n                    loc = torch.tensor([z,y,x], device=pred.device)    \n                    if rank != 0:\n                        dist.gather(pred)\n                        dist.gather(loc)\n                    else:\n                        preds = [torch.zeros_like(pred) for _ in range(worldsize)]\n                        locs = [torch.zeros_like(loc) for _ in range(worldsize)]\n                        dist.gather(pred, preds)\n                        dist.gather(loc, locs)\n                        \n                        for r, (loc, pred) in enumerate(zip(locs, preds)):\n                            if i + r >= len(vs): continue # r = rank\n                                \n                            z,y,x = [l.item() for l in loc]\n                            m = mask[z:z+block_size, y:y+block_size, x:x+block_size]\n                            m = (m / 255).astype(np.float32)\n                        \n                            pred = pred[0,0].cpu()[...,:m.shape[-3],:m.shape[-2],:m.shape[-1]]\n                            pred = pred * k[...,:m.shape[-3],:m.shape[-2],:m.shape[-1]]\n                            pred = pred * m\n                            pred = pred[margin:block_size-margin,\n                                        margin:block_size-margin,\n                                        margin:block_size-margin]\n                            pred = pred.numpy()\n                            \n                            prediction[z+margin:z+block_size-margin,\n                                       y+margin:y+block_size-margin,\n                                       x+margin:x+block_size-margin] += (255*pred).clip(0,255).astype(np.uint8)\n                else:\n                    m = mask[z:z+block_size, y:y+block_size, x:x+block_size]\n                    m = (m / 255).astype(np.float32)\n                \n                    pred = pred[0,0].cpu()[...,:m.shape[-3],:m.shape[-2],:m.shape[-1]]\n                    pred = pred * k[...,:m.shape[-3],:m.shape[-2],:m.shape[-1]]\n                    pred = pred * m\n                    pred = pred[margin:block_size-margin,\n                                margin:block_size-margin,\n                                margin:block_size-margin]\n                    pred = pred.numpy()\n                    \n                    prediction[z+margin:z+block_size-margin,\n                               y+margin:y+block_size-margin,\n                               x+margin:x+block_size-margin] += (255*pred).clip(0,255).astype(np.uint8)\n\n            # create the RLEs and add them to the csv lines\n            if rank == 0:\n                for z, key in enumerate(keys):\n                    if margin > 0:\n                        m = (prediction[z+margin,margin:-margin,margin:-margin] > threshold)\n                    else:\n                        m = (prediction[z] > threshold)\n                    rle = rle_encode(m)\n                    csv_lines.append(key + \",\" + rle + \"\\n\")\n\n        if rank == 0:\n            print()\n            print()\n            print(\"writing submission.csv...\")\n            \n            with open(\"submission.csv\", \"w\") as stream:\n                stream.writelines(csv_lines)\n\n            # Save some images\n            if not is_kaggle and HAS_INLINE:\n                prediction = prediction[margin:-margin,margin:-margin,margin:-margin]\n                prediction = (prediction > threshold).astype(np.float32)\n                target = target[0,margin:-margin,margin:-margin,margin:-margin]\n                target = (target > 0).astype(np.float32)\n                \n                inline.save(prediction, filename=os.path.join(image_dir, \"infer_prediction.png\"))\n                inline.save(target, filename=os.path.join(image_dir, \"infer_target.png\"))\n                inline.save((target-prediction), filename=os.path.join(image_dir, \"infer_errors.png\"))\n\n        if dist.is_initialized(): dist.barrier()\n                \ndef run_distributed(rank, args):\n    print(f\"GPU {rank} checking in...\")\n    localrank = rank % torch.cuda.device_count()\n    torch.cuda.set_device(localrank)\n    \n    assert dist.is_available()\n    dist.init_process_group(\n        \"nccl\",\n        init_method=\"tcp://localhost:23456\",\n        rank=rank,\n        world_size=WORLDSIZE)\n    assert dist.is_initialized()\n    run()\n\nif __name__ == \"__main__\":\n    if WORLDSIZE < 1:\n        run()\n    else:\n        mp.spawn(run_distributed, nprocs=NPROCS, args=(None,))\n    ","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]}]}