{"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":91249,"databundleVersionId":11294684,"sourceType":"competition"}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport random","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-20T13:42:22.094830Z","iopub.execute_input":"2025-03-20T13:42:22.095211Z","iopub.status.idle":"2025-03-20T13:42:22.099568Z","shell.execute_reply.started":"2025-03-20T13:42:22.095184Z","shell.execute_reply":"2025-03-20T13:42:22.098385Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def random_crop_patch_around_target(volume, target_coord, patch_size):\n    \"\"\"\n    Randomly crop a 3D patch ensuring the target is inside.\n\n    Args:\n        volume (torch.Tensor): Input volume of shape (C, D_full, H_full, W_full).\n        target_coord (tuple): (z, y, x) coordinates of the object in the full volume.\n        patch_size (tuple): (D, H, W) size of the patch.\n\n    Returns:\n        cropped_patch (torch.Tensor): Cropped volume of shape (C, D, H, W).\n        crop_start (tuple): (z_start, y_start, x_start) of the cropped patch.\n    \"\"\"\n    C, D_full, H_full, W_full = volume.shape\n    D, H, W = patch_size\n    z_target, y_target, x_target = target_coord\n\n    # Define valid range for cropping\n    z_min = max(0, z_target - D)  \n    z_max = min(D_full - D, z_target)  \n    y_min = max(0, y_target - H)  \n    y_max = min(H_full - H, y_target)  \n    x_min = max(0, x_target - W)  \n    x_max = min(W_full - W, x_target)  \n\n    # Randomly select a valid crop start position\n    z_start = random.randint(z_min, z_max)\n    y_start = random.randint(y_min, y_max)\n    x_start = random.randint(x_min, x_max)\n\n    # Crop the patch\n    cropped_patch = volume[:, z_start:z_start + D, y_start:y_start + H, x_start:x_start + W]\n\n    return cropped_patch, (z_start, y_start, x_start)\n    \n\ndef generate_3d_labels(global_target, roi_start, roi_size, stride=2):\n    \"\"\"\n    Generate class_map and offset_map for 3D object detection.\n\n    Args:\n        global_target (tuple): (z, y, x) coordinate in the full volume.\n        roi_start (tuple): (z_start, y_start, x_start) of the extracted patch in full volume.\n        roi_size (tuple): (depth, height, width) of the extracted patch.\n        class_map_size (tuple): (depth, height, width) of the class_map output.\n        stride (int): The stride factor, default is 2.\n\n    Returns:\n        class_map (torch.Tensor): Binary tensor of shape class_map_size.\n        offset_map (torch.Tensor): Offset tensor of shape (3, *class_map_size).\n    \"\"\"\n    #roi coordinate space\n    z_local = global_target[0] - roi_start[0]\n    y_local = global_target[1] - roi_start[1]\n    x_local = global_target[2] - roi_start[2]\n\n    #shrinkage coordinate\n    z_shrink = z_local//stride\n    y_shrink = y_local//stride\n    x_shrink = x_local//stride\n    \n    #relative coordinate \n    z_ratio = z_local/roi_size[0]\n    y_ratio = y_local/roi_size[1]\n    x_ratio = x_local/roi_size[2]\n\n    #\n    D, H, W = roi_size[0]//2, roi_size[1]//2, roi_size[2]//2\n    class_map = torch.zeros(1, D, H, W)\n    offset_map = torch.zeros(3, D, H, W)\n\n    for dz in range(2):\n        for dy in range(2):\n            for dx in range(2):\n                z_idx = z_shrink + dz\n                y_idx = y_shrink + dy\n                x_idx = x_shrink + dx\n                if z_idx<D and y_idx<H and x_idx<W:\n                    class_map[:, z_idx, y_idx, x_idx] = 1\n                    offset_map[0, z_idx, y_idx, x_idx] = z_ratio\n                    offset_map[1, z_idx, y_idx, x_idx] = y_ratio\n                    offset_map[2, z_idx, y_idx, x_idx] = x_ratio\n    return class_map, offset_map\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-20T13:42:22.261398Z","iopub.execute_input":"2025-03-20T13:42:22.261733Z","iopub.status.idle":"2025-03-20T13:42:22.272659Z","shell.execute_reply.started":"2025-03-20T13:42:22.261708Z","shell.execute_reply":"2025-03-20T13:42:22.271567Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Example usage\nvolume = torch.randn(1, 184, 630, 630)  # Simulated 3D volume (C, D, H, W)\ntarget_coord = (162, 23, 224)  # Object position in full volume\nroi_size = (96, 128, 128)  # Desired patch size\n\n# Step 1: Randomly crop a patch ensuring the object is inside\ncropped_patch, roi_start = random_crop_patch_around_target(volume, target_coord, roi_size)\n\nclass_map, offset_map = generate_3d_labels(target_coord, roi_start, roi_size)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-20T13:42:22.293672Z","iopub.execute_input":"2025-03-20T13:42:22.294004Z","iopub.status.idle":"2025-03-20T13:42:23.066976Z","shell.execute_reply.started":"2025-03-20T13:42:22.293979Z","shell.execute_reply":"2025-03-20T13:42:23.066003Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Class Map Shape:\", class_map.shape)\nprint(\"Offset Map Shape:\", offset_map.shape)\nprint(\"Class Map Non-Zero Indices:\", torch.nonzero(class_map))\nprint(\"Offset Values at Class Map Locations:\", offset_map[class_map.bool().repeat(3, 1, 1, 1)])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-20T13:42:23.068118Z","iopub.execute_input":"2025-03-20T13:42:23.068407Z","iopub.status.idle":"2025-03-20T13:42:23.078742Z","shell.execute_reply.started":"2025-03-20T13:42:23.068382Z","shell.execute_reply":"2025-03-20T13:42:23.077786Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_map.shape, offset_map.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-20T13:42:23.080381Z","iopub.execute_input":"2025-03-20T13:42:23.080694Z","iopub.status.idle":"2025-03-20T13:42:23.096995Z","shell.execute_reply.started":"2025-03-20T13:42:23.080669Z","shell.execute_reply":"2025-03-20T13:42:23.095847Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def recover_original_coordinate(class_map, offset_map, roi_start, roi_size, stride=2):\n    \"\"\"\n    Recover the original global target coordinate from class_map and offset_map.\n\n    Args:\n        class_map (torch.Tensor): Binary tensor of shape (1, D, H, W) indicating object presence.\n        offset_map (torch.Tensor): Offset tensor of shape (3, D, H, W).\n        roi_start (tuple): (z_start, y_start, x_start) of the extracted patch in full volume.\n        roi_size (tuple): (depth, height, width) of the extracted patch.\n        stride (int): The stride factor, default is 2.\n\n    Returns:\n        recovered_target (tuple): (z, y, x) coordinate in the full volume.\n    \"\"\"\n    # Get indices where class_map == 1\n    indices = torch.nonzero(class_map.squeeze(), as_tuple=True)\n    coords = []\n    # Take the first detected point (assuming 1 target object)\n    for z_idx, y_idx, x_idx in zip(*indices):\n        # Retrieve corresponding offsets\n        z_offset = offset_map[0, z_idx, y_idx, x_idx].item()\n        y_offset = offset_map[1, z_idx, y_idx, x_idx].item()\n        x_offset = offset_map[2, z_idx, y_idx, x_idx].item()\n        # Convert back to local coordinates within the ROI\n        z_local = (z_idx + z_offset) * stride\n        y_local = (y_idx + y_offset) * stride\n        x_local = (x_idx + x_offset) * stride\n    \n        # Convert back to global coordinates\n        z_global = int(z_local + roi_start[0])\n        y_global = int(y_local + roi_start[1])\n        x_global = int(x_local + roi_start[2])\n        coords.append((z_global, y_global, x_global))\n    return coords","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-20T13:42:23.098435Z","iopub.execute_input":"2025-03-20T13:42:23.098789Z","iopub.status.idle":"2025-03-20T13:42:23.118493Z","shell.execute_reply.started":"2025-03-20T13:42:23.098748Z","shell.execute_reply":"2025-03-20T13:42:23.117229Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Recover the target coordinate\nrecovered_coord = recover_original_coordinate(class_map, offset_map, roi_start, roi_size)\nprint(\"Recovered Target Coordinate:\", recovered_coord)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-20T13:42:23.119698Z","iopub.execute_input":"2025-03-20T13:42:23.120103Z","iopub.status.idle":"2025-03-20T13:42:23.146125Z","shell.execute_reply.started":"2025-03-20T13:42:23.120065Z","shell.execute_reply":"2025-03-20T13:42:23.145180Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}