{"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":"gpu","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":7179194,"sourceType":"datasetVersion","datasetId":4148557}],"dockerImageVersionId":30616,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport cv2\nimport math\n\nimport timm\nprint('timm', timm.__version__)\nfrom timm.models.resnet import *\n\nimport matplotlib.pyplot as plt\nfrom IPython.display import Image, display\n\n###########################################33\n#helper\ndef image_show_norm(name, image, min=None, max=None, type='bgr', resize=1):\n\tif max is None: max = image.max()\n\tif min is None: min = image.min()\n\tif type == 'rgb': image = np.ascontiguousarray(image[:, :, ::-1])\n\n\tH, W = image.shape[0:2]\n\tcv2.namedWindow(name, cv2.WINDOW_GUI_NORMAL)  # WINDOW_NORMAL\n\tcv2.imshow(name, (np.clip((image - min) / (max - min), 0, 1) * 255).astype(np.uint8))\n\tcv2.resizeWindow(name, round(resize * W), round(resize * H))\n\n\n\ndef np_sigmoid(x):\n\treturn 1 / (1 + np.exp(-x))\n\nprint('IMPORT OK')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-11T19:32:02.045595Z","iopub.execute_input":"2023-12-11T19:32:02.045988Z","iopub.status.idle":"2023-12-11T19:32:06.890363Z","shell.execute_reply.started":"2023-12-11T19:32:02.045958Z","shell.execute_reply":"2023-12-11T19:32:06.889494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"root_dir = \\\n\t'/kaggle/input/blood-vessel-segmentation'\n\t#'/home/user/share1/kaggle/2023/blood-vessel-segmentation'\n\n\nIMAGE_SIZE  = 640\nIMAGE_DEPTH = 32\n\n# Net\nclass MyDecoderBlock3d(nn.Module):\n\tdef __init__(\n\t\tself,\n\t\tin_channel,\n\t\tskip_channel,\n\t\tout_channel,\n\t):\n\t\tsuper().__init__()\n\t\tself.conv1 = nn.Sequential(\n\t\t\tnn.Conv3d(in_channel + skip_channel, out_channel, kernel_size=3, padding=1, bias=False),\n\t\t\tnn.BatchNorm3d(out_channel, eps=1e-4),\n\t\t\tnn.ReLU(inplace=True),\n\t\t)\n\t\tself.attention1 = nn.Identity()\n\t\tself.conv2 = nn.Sequential(\n\t\t\tnn.Conv3d(out_channel,out_channel,kernel_size=3, padding=1, bias=False),\n\t\t\tnn.BatchNorm3d(out_channel, eps=1e-4),\n\t\t\tnn.ReLU(inplace=True),\n\t\t)\n\t\tself.attention2 = nn.Identity()\n\n\tdef forward(self, x, skip=None):\n\t\tx = F.interpolate(x, scale_factor=2, mode='nearest')\n\t\tif skip is not None:\n\t\t\tx = torch.cat([x, skip], dim=1)\n\t\t\tx = self.attention1(x)\n\t\tx = self.conv1(x)\n\t\tx = self.conv2(x)\n\t\tx = self.attention2(x)\n\t\treturn x\n\nclass MyUnetDecoder3d(nn.Module):\n\tdef __init__(self,\n\t\t\t\t in_channel,\n\t\t\t\t skip_channel,\n\t\t\t\t out_channel,\n\t\t\t\t ):\n\t\tsuper().__init__()\n\t\tself.center = nn.Identity()\n\n\t\ti_channel = [in_channel, ] + out_channel[:-1]\n\t\ts_channel = skip_channel\n\t\to_channel = out_channel\n\t\tblock = [\n\t\t\tMyDecoderBlock3d(i, s, o)\n\t\t\tfor i, s, o in zip(i_channel, s_channel, o_channel)\n\t\t]\n\t\tself.block = nn.ModuleList(block)\n\n\tdef forward(self, feature, skip):\n\t\td = self.center(feature)\n\t\tdecode = []\n\t\tfor i, block in enumerate(self.block):\n\t\t\t#print(i, d.shape, skip[i].shape if skip[i] is not None else 'none')\n\t\t\t#print(block.conv1[0])\n\t\t\t#print('')\n\n\t\t\ts = skip[i]\n\t\t\td = block(d, s)\n\t\t\tdecode.append(d)\n\t\tlast = d\n\t\treturn last, decode\n\n#--------------------------------------\n\nclass Net(nn.Module):\n\tdef __init__(self, ):\n\t\tsuper().__init__()\n\t\tencoder_dim = [64, 256, 512, 1024, 2048]\n\t\tdecoder_dim = [256, 128, 128, 64, 32 ]\n\n\t\t#self.encoder = seresnext26d_32x4d(pretrained=True, in_chans=3)\n\t\tself.encoder = resnet50d(pretrained=False, in_chans=3)\n\t\tself.add_conv1 = nn.Sequential(\n\t\t\tnn.Conv3d(1, 32, kernel_size=3, padding=1, bias=False),\n\t\t\tnn.BatchNorm3d(32, eps=1e-4),\n\t\t\tnn.ReLU(inplace=True),\n\t\t\tnn.Conv3d(32, 32, kernel_size=3, padding=1, bias=False),\n\t\t\tnn.BatchNorm3d(32, eps=1e-4),\n\t\t\tnn.ReLU(inplace=True),\n\t\t\tnn.Conv3d(32, 32, kernel_size=3, padding=1, bias=False),\n\t\t\tnn.BatchNorm3d(32, eps=1e-4),\n\t\t\tnn.ReLU(inplace=True),\n\t\t)\n\n\n\t\tself.decoder = MyUnetDecoder3d(\n\t\t\tin_channel  = encoder_dim[-1],\n\t\t\tskip_channel= encoder_dim[:-1][::-1]+[32],\n\t\t\tout_channel = decoder_dim,\n\t\t)\n\t\tself.vessel = nn.Conv3d(decoder_dim[-1], 1, kernel_size=1)\n\n\t\t#just a simple demo. please improve this\n\t\tself.convert_3d = torch.nn.ModuleList([\n\t\t\t\tnn.Identity(),\n\t\t\t\tnn.Sequential(nn.Conv3d(  64,  64,kernel_size=( 2,1,1),stride=( 2,1,1),padding=(0,0,0)),  nn.BatchNorm3d(  64, eps=1e-4), nn.ReLU(inplace=True), ),\n\t\t\t\tnn.Sequential(nn.Conv3d( 256, 256,kernel_size=( 4,1,1),stride=( 4,1,1),padding=(0,0,0)),  nn.BatchNorm3d( 256, eps=1e-4), nn.ReLU(inplace=True), ),\n\t\t\t\tnn.Sequential(nn.Conv3d( 512, 512,kernel_size=( 8,1,1),stride=( 8,1,1),padding=(0,0,0)),  nn.BatchNorm3d( 512, eps=1e-4), nn.ReLU(inplace=True), ),\n\t\t\t\tnn.Sequential(nn.Conv3d(1024,1024,kernel_size=(16,1,1),stride=(16,1,1),padding=(0,0,0)),  nn.BatchNorm3d(1024, eps=1e-4), nn.ReLU(inplace=True), ),\n\t\t\t\tnn.Sequential(nn.Conv3d(2048,2048,kernel_size=(32,1,1),stride=(32,1,1),padding=(0,0,0)),  nn.BatchNorm3d(2048, eps=1e-4), nn.ReLU(inplace=True), ),\n\t\t])\n\n\tdef forward(self, subvolume):\n\t\txx = subvolume\n\t\tB, D, H, W = xx.shape\n\n\t\tx = xx.reshape(B * D, 1, H, W)\n\t\tx = x.expand(-1, 3, -1, -1)\n\n\t\tencode = []\n\t\txx = self.add_conv1(xx.unsqueeze(1)); encode.append(xx)\n\n\t\te = self.encoder\n\t\tx = e.conv1(x)\n\t\tx = e.bn1(x)\n\t\tx = e.act1(x);   encode.append(x)\n\t\tx = F.avg_pool2d(x, kernel_size=2, stride=2)\n\t\tx = e.layer1(x); encode.append(x)\n\t\tx = e.layer2(x); encode.append(x)\n\t\tx = e.layer3(x); encode.append(x)\n\t\tx = e.layer4(x); encode.append(x)\n\n\t\tfor i, x in enumerate(encode):\n\t\t\tif i == 0: continue\n\t\t\tx = encode[i]\n\t\t\tBD, c, h, w = x.shape\n\t\t\tx = x.reshape(B, D, c, h, w)\n\t\t\tx = x.transpose(1, 2)\n\t\t\tx = self.convert_3d[i](x)\n\t\t\tencode[i] = x\n\n\t\tlast, decode = self.decoder(\n\t\t\tfeature=encode[-1], skip=encode[:-1][::-1]\n\t\t)\n\n\t\tvessel = self.vessel(last).squeeze(1)\n\t\tvessel = torch.sigmoid(vessel.float())\n\t\treturn vessel\n\n#############################################################################33\n\ndef run_check_net():\n\theight, width = 480, 480\n\tdepth = 32\n\tbatch_size = 2\n\n\tsubvolume = torch.from_numpy(np.random.uniform(0, 1, (batch_size, depth, height, width))).float().cuda()\n\tnet = Net().cuda()\n\n\n\twith torch.no_grad():\n\t\twith torch.cuda.amp.autocast(enabled=True):\n\t\t\tvessel = net(subvolume)\n\tprint(subvolume.shape)\n\tprint(vessel.shape)\n\nrun_check_net()\nprint('NET OK!!!')","metadata":{"execution":{"iopub.status.busy":"2023-12-11T19:32:06.892118Z","iopub.execute_input":"2023-12-11T19:32:06.8924Z","iopub.status.idle":"2023-12-11T19:32:16.793835Z","shell.execute_reply.started":"2023-12-11T19:32:06.892376Z","shell.execute_reply":"2023-12-11T19:32:16.792923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Data\ndef norm_by_min_max(x, xmin, xmax, alpha=0.01):\n\txmin = float(xmin)\n\txmax = float(xmax)\n\tx = (x-xmin)/(xmax-xmin)\n\tif 1:\n\t\tx[x>1]=(x[x>1]-1)*alpha +1\n\t\tx[x<0]=(x[x<0])*alpha\n\t#x = np.clip(x,0,1)\n\treturn x\n\ndef load_dummy_data():\n\t#x,y,z = 740,470, 389\n\tname = 'kidney_3_sparse'\n\tx,y,z = 420, 150, 356\n\n\tsubvolume = []\n\ttruth = []\n\tfor i in range(z,z+IMAGE_DEPTH):\n\t\timage_file = f'{root_dir}/train/{name}/images/{i:04d}.tif'\n\t\tm = cv2.imread(image_file, cv2.IMREAD_UNCHANGED)\n\t\tsubvolume.append(m)\n\n\t\timage_file = f'{root_dir}/train/{name}/labels/{i:04d}.tif'\n\t\tv = cv2.imread(image_file, cv2.IMREAD_GRAYSCALE)\n\t\ttruth.append(v)\n\n\tsubvolume = np.stack(subvolume)\n\tsubvolume = subvolume[:,y:y+IMAGE_SIZE,x:x+IMAGE_SIZE]\n\txmin, xmax = (18806, 21903) #precompute low and high percentile\n\tsubvolume = norm_by_min_max(subvolume, xmin, xmax)\n\n\ttruth = np.stack(truth)\n\ttruth = truth[:,y:y+IMAGE_SIZE,x:x+IMAGE_SIZE]\n\ttruth = truth//255\n\n\treturn subvolume, truth\n\n\ndef run_demo():\n\n\tsubvolume, truth = load_dummy_data()\n\n\tnet = Net()\n\tcheckpoint_file ='/kaggle/input/2d-to-3d-demo-data/00000828.pth'\n\tstate_dict = torch.load(checkpoint_file, map_location=lambda storage, loc: storage)['state_dict']\n\tprint(net.load_state_dict(state_dict, strict=True))  # True\n\tnet = net.eval()\n\tnet = net.cuda()\n\n\ttensor = torch.from_numpy(subvolume).float().cuda()\n\twith torch.no_grad():\n\t\twith torch.cuda.amp.autocast(enabled=True):\n\t\t\tvessel = net(tensor.unsqueeze(0)).squeeze(0)\n\tvessel = vessel.float().data.cpu().numpy()\n\n\n\t#--------------------------------------------------------------\n\tif 0: #visualisation (run offline. does not work in kaggle\n\n\t\tD,H,W = subvolume.shape\n\t\tsubvolume = np.clip(subvolume,0,1)\n\t\tx_mean = subvolume.mean(0)\n\t\tp_mean = vessel.mean(0)\n\t\ty_mean = truth.mean(0)\n\n\t\t#---\n\t\tdef add_contrast_for_x(x_mean):\n\t\t\tx_mean = np_sigmoid((1-x_mean)*5)\n\t\t\tx_mean = 1- (x_mean-x_mean.min())/(x_mean.max()-x_mean.min())\n\t\t\treturn x_mean\n\n\t\tdef add_contrast_for_py(p_mean, y_mean):\n\t\t\tzero = np.zeros((H,W))\n\t\t\tpy_mean = np.dstack([zero,p_mean,y_mean])\n\t\t\tpy_mean = ((py_mean-py_mean.min())/(py_mean.max()-py_mean.min()))**0.25\n\t\t\treturn py_mean\n\t\t#---\n\n\t\tx_mean  = add_contrast_for_x(x_mean)\n\t\tpy_mean = add_contrast_for_py(p_mean, y_mean)\n\t\timage_show_norm('x_mean',x_mean, min=0,max=1,resize=1)\n\t\timage_show_norm('py_mean',py_mean, min=0,max=1,resize=1)\n\t\tcv2.waitKey(0)\n\telse:\n\t\tdisplay(Image(filename='../input/2d-to-3d-demo-data/Selection_999(4367).png'))\n\t\tdisplay(Image(filename='../input/2d-to-3d-demo-data/Selection_999(4368).png'))\n \n\n\t#--------------------------------------------------------------\n\tif 0: #3d visualisation (run offline. does not work in kaggle)\n\t\timport pyvista as pv\n\t\ty = truth>0\n\t\tp = vessel>0.5\n\t\thit = y*p\n\t\tfp = (1-y)*p\n\t\tmiss = y*(1-p)\n\n\t\tpl = pv.Plotter()\n\n\t\tmhit  = pv.PolyData(np.stack(np.where(hit > 0.1)).T).glyph(geom=pv.Cube())\n\t\tmfp   = pv.PolyData(np.stack(np.where(fp  > 0.1)).T).glyph(geom=pv.Cube())\n\t\tmmiss = pv.PolyData(np.stack(np.where(miss > 0.1)).T).glyph(geom=pv.Cube())\n\t\tpl.add_mesh(mhit,  color='yellow' )\n\t\tpl.add_mesh(mfp,   color='green' )\n\t\tpl.add_mesh(mmiss, color='red' )\n\t\tpl.show()\n\telse:\n\t\tprint('why cannot play GIF ???????')        \n\t\tdisplay(Image(filename='../input/2d-to-3d-demo-data/Peek 2023-12-12 02-38.gif', format='png'))\n\t\tdisplay(Image(filename='../input/2d-to-3d-demo-data/Peek 2023-12-12 02-36.gif', format='png'))\n    \n\nrun_demo()\nprint('DEMO OK!!!')","metadata":{"execution":{"iopub.status.busy":"2023-12-11T19:33:21.776675Z","iopub.execute_input":"2023-12-11T19:33:21.777278Z","iopub.status.idle":"2023-12-11T19:33:37.109962Z","shell.execute_reply.started":"2023-12-11T19:33:21.777231Z","shell.execute_reply":"2023-12-11T19:33:37.109091Z"},"trusted":true},"execution_count":null,"outputs":[]}]}