{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":91249,"databundleVersionId":11294684,"sourceType":"competition"},{"sourceId":11483969,"sourceType":"datasetVersion","datasetId":7197682}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Let's make a 3d volume encoder that will output 3d feature map. It is constructed from 2d imagenet encoder","metadata":{}},{"cell_type":"code","source":"#model.py\nimport timm\nprint('timm.__version__',timm.__version__)\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\n\ndef encode_for_pvtv2(e, x, B, depth_scaling=[2,2,2,2,]):\n\t#poor man's attention = avg + max pool\n    def pool_in_depth(x, depth_scaling):\n        bd, c, h, w = x.shape\n        x1 = x.reshape(B, -1, c, h, w).permute(0, 2, 1, 3, 4)\n        x1 = F.avg_pool3d(x1, kernel_size=(depth_scaling, 1, 1), stride=(depth_scaling, 1, 1), padding=0) \\\n                + F.gelu(F.max_pool3d(x1, kernel_size=(depth_scaling, 1, 1), stride=(depth_scaling, 1, 1), padding=0))\n        x = x1.permute(0, 2, 1, 3, 4).reshape(-1, c, h, w)\n        return x, x1\n\n    encode=[]  #x = 1,256,512,512\n    x = e.patch_embed(x) # x seq: 1024, 128, 128, 64\n\n    x = e.stages_0(x)   #4, 64, 128, 128\n    x, x1 = pool_in_depth(x, depth_scaling[0])\n    #encode.append(x1)\n    x = e.stages_1(x)   #4, 128, 64, 64\n    x, x1 = pool_in_depth(x, depth_scaling[1])\n    #encode.append(x1)\n    x = e.stages_2(x)   #4, 320, 32, 32\n    x, x1 = pool_in_depth(x, depth_scaling[2])\n    #encode.append(x1)\n    x = e.stages_3(x)   #4, 512, 16, 16\n    x, x1 = pool_in_depth(x, depth_scaling[3])\n    encode.append(x1)\n\n    return encode\n\n\nclass Net(nn.Module):\n    def __init__(self, pretrained=False, cfg=None):\n        super(Net, self).__init__()\n        self.output_type = ['infer', 'loss', ]\n        self.register_buffer('D', torch.tensor(0))\n\n        self.arch = 'pvt_v2_b1'\n        encoder_dim = {\n            'resnet34d': [64, 64, 128, 256, 512, ],\n            'resnet50d': [64, 256, 512, 1024, 2048, ],\n            'seresnext26d_32x4d': [64, 256, 512, 1024, 2048, ],\n            'convnext_small.fb_in22k': [96, 192, 384, 768],\n            'pvt_v2_b1': [64, 128, 320, 512],\n            'pvt_v2_b2': [64, 128, 320, 512],\n        }.get(self.arch, [1024])\n\n        self.encoder = timm.create_model(\n            model_name=self.arch, pretrained=pretrained, in_chans=3, num_classes=0, global_pool='', features_only=True,\n        )\n        self.mask = nn.Conv3d(encoder_dim[-1],1, kernel_size=1)\n\n    def forward(self, batch):\n        device = self.D.device\n\n        image = batch['image'].to(device)\n        B, D, H, W = image.shape\n        image = image.reshape(B*D, 1,H, W)\n\n        x = (image.half() - 128) / 128\n        x = x.expand(-1, 3, -1, -1)\n\n        encode = encode_for_pvtv2(self.encoder, x, B)\n        last = encode[-1] #this is the feature map !!!!\n        logit = self.mask(last) .squeeze(1)\n\n        \n        #print(f'last', last.shape)\n        #[print(f'encode_{i}', e.shape) for i,e in enumerate(encode)]\n        #print('logit', logit.shape)\n\n        output = {} \n\n        #loss for pretraining 2d-3d encoder\n        if 'loss' in self.output_type:\n            truth = batch['truth'].to(device)\n            output['mask_loss'] = F.binary_cross_entropy_with_logits(logit,truth)\n\n        if 'infer' in self.output_type:\n            output['motor'] = torch.sigmoid(logit)\n        return output\n\n\n###---------------------------------------------\n# run some dummy data\ndef run_check_net():\n\n    B = 1\n    slice_shape = (192,384,384) \n    mask_shape  = (12,12,12)\n\n    batch = {\n        'image': torch.from_numpy(np.random.uniform(0,1, (B, *slice_shape))).byte(),\n        'truth': torch.from_numpy(np.random.choice(2, (B, *mask_shape))).half(),\n    }\n    net = Net(pretrained=False, cfg=None).cuda()\n\n    with torch.no_grad():\n        with torch.amp.autocast('cuda',enabled=True):\n            output = net(batch)\n    # ---\n    print('batch')\n    for k, v in batch.items():\n        if k == 'D':\n            print(f'{k:>32} : {v} ')\n        else:\n            print(f'{k:>32} : {v.shape} ')\n\n    print('output')\n    for k, v in output.items():\n        if 'loss' not in k:\n            print(f'{k:>32} : {v.shape} ')\n    print('loss')\n    for k, v in output.items():\n        if 'loss' in k:\n            print(f'{k:>32} : {v.item()} ')\n\nrun_check_net()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-20T08:58:41.714047Z","iopub.execute_input":"2025-04-20T08:58:41.714389Z","iopub.status.idle":"2025-04-20T08:58:42.813466Z","shell.execute_reply.started":"2025-04-20T08:58:41.714370Z","shell.execute_reply":"2025-04-20T08:58:42.812888Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"let's load a trained model and show prediction results on show validation tomographs","metadata":{}},{"cell_type":"code","source":"import matplotlib\nimport matplotlib.pyplot as plt \nimport cv2\n\n\ncheckpoint ='/kaggle/input/example-2d-3d-encoder-weight/00003808.pth'\n\n\nnet = Net(pretrained=False)\nnet.cuda()\nnet.output_type = ['infer']\nf = torch.load(checkpoint, map_location=lambda storage, loc: storage, weights_only=False)\nstate_dict = f.get('state_dict', f) \nprint(net.load_state_dict(state_dict, strict=False))\n\nvalid_data=[\n    {\n        'tomo_id':'tomo_00e047',\n        'z0z1'  :[73,265],\n        'label' :1,  \n    },\n    {\n        'tomo_id':'tomo_17143f',\n        'z0z1'  :[0,256],\n        'label' :0,  \n    },\n    {\n        'tomo_id':'tomo_0fe63f',\n        'z0z1'  :[101,293],\n        'label' :1,  \n    },\n]\n\n#helper function\nKAGGLE_DATA_DIR ='/kaggle/input/byu-locating-bacterial-flagellar-motors-2025'\ndef read_image_stack(tomo_id, start_no, end_no, resize=-1, mode=cv2.IMREAD_GRAYSCALE):\n    image = []\n    for z in range(start_no,end_no):\n        jpg_file = f'{KAGGLE_DATA_DIR}/train/{tomo_id}/slice_{z:04d}.jpg'\n        m=cv2.imread(jpg_file, mode)\n        if resize>0:\n            m=cv2.resize(m, (resize, resize),cv2.INTER_LINEAR)\n        image.append(m)\n    image = np.stack(image)\n    return image\n\ndef make_overlay(image, prob, axis=0): \n    m = image.mean(axis)\n    p = prob.max(axis)\n    #t = truth.max(axis)\n    \n    overlay = np.stack([m, m, m], 2)\n    op = p[..., None] * [[[1, 0, 0]]]\n    overlay = 255 - (255 - overlay) * (1 - op)\n    #ot = t[...,None]*[[[0,1,0]]]\n    #overlay = 255-(255-overlay)*(1-ot)\n    overlay = overlay.astype(np.uint8)\n    return overlay\n\n\nfor r in valid_data:\n    tomo_id = r['tomo_id']\n    label = r['label']\n    z0,z1 = r['z0z1']\n    image = read_image_stack(tomo_id, z0, z1, resize=384, mode=cv2.IMREAD_GRAYSCALE)\n    batch = {\n        'image': torch.from_numpy(image).unsqueeze(0).byte(),\n    }\t\t\n    with torch.amp.autocast('cuda', dtype=torch.float16):\n        with torch.no_grad():\n            output = net(batch)\n\n    prob = F.interpolate(\n\t    output['motor'].unsqueeze(1),\n\t    scale_factor=(16, 32, 32), mode='nearest',  \n\t)[0,0].float().data.cpu().numpy()\n\n    #---draw the results ---\n    print('tomo_id', tomo_id)\n    print('label', label)\n    print('image', image.shape)\n    print('prob', prob.shape)\n\n    overlay0 = make_overlay(image, prob, axis=0)\n    overlay1 = make_overlay(image, prob, axis=1)\n    plt.imshow(image.mean(0), cmap='gray')\n    plt.show()\n    plt.imshow(overlay0)\n    plt.show()\n    plt.imshow(overlay1)\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-20T08:59:02.572252Z","iopub.execute_input":"2025-04-20T08:59:02.572532Z","iopub.status.idle":"2025-04-20T08:59:10.920195Z","shell.execute_reply.started":"2025-04-20T08:59:02.572509Z","shell.execute_reply":"2025-04-20T08:59:10.919425Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"How to add DETR head? This is left as an exercise for you. you can ask chatgpt:\n\nQ: \"do you know about DETR for object detection?\"\n\nQ: \"in my application, i have a 3d volume as input. i need to predict if the volume contains the target object or not. If yes, the most likely object location as well. now I have 3d volume feature encoder that will make a feature, e.g. of size 12x12x12. i am thinking of using DETR head for my task. can you help?\n\neg: https://chatgpt.com/share/6804b1fd-a31c-800b-9f6c-5a36d9fbdde2","metadata":{}}]}