{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":10130907,"sourceType":"datasetVersion","datasetId":6251933},{"sourceId":10479616,"sourceType":"datasetVersion","datasetId":6251769},{"sourceId":10671331,"sourceType":"datasetVersion","datasetId":6252074}],"dockerImageVersionId":30840,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%capture\n\ntry:\n    import zarr\n    import monai\nexcept: \n    !cp -r '/kaggle/input/cryo-wheels-2/wheels' '/kaggle/working/'\n    !pip install /kaggle/working/wheels/asciitree-0.3.3/asciitree-0.3.3\n    !pip install --no-index --find-links=/kaggle/working/wheels zarr\n    !pip install --no-index --find-links=/kaggle/working/wheels connected-components-3d\n    !pip install --no-index --find-links=/kaggle/working/wheels monai\n    \n\nprint('PIP INSTALL OK!!!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:02:58.496094Z","iopub.execute_input":"2025-02-05T16:02:58.496545Z","iopub.status.idle":"2025-02-05T16:03:38.858041Z","shell.execute_reply.started":"2025-02-05T16:02:58.496504Z","shell.execute_reply":"2025-02-05T16:03:38.856789Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!rm -rf /kaggle/working/wheels","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:03:38.859335Z","iopub.execute_input":"2025-02-05T16:03:38.859590Z","iopub.status.idle":"2025-02-05T16:03:39.605071Z","shell.execute_reply.started":"2025-02-05T16:03:38.859555Z","shell.execute_reply":"2025-02-05T16:03:39.604038Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\ndirectory_path = \"/kaggle/input/cryo-utils\"\nsys.path.append(directory_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:03:39.607079Z","iopub.execute_input":"2025-02-05T16:03:39.607323Z","iopub.status.idle":"2025-02-05T16:03:39.611206Z","shell.execute_reply.started":"2025-02-05T16:03:39.607297Z","shell.execute_reply":"2025-02-05T16:03:39.610368Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEBUG = False\n\nclass Config:\n    DEPTH = 96\n    HEIGHT = 128\n    WIDTH = 128\n    SKIP_RATIO = 0.85 # Was 0.85\n    BATCH_SIZE = 8\n\ncfg = Config()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:03:39.612487Z","iopub.execute_input":"2025-02-05T16:03:39.612799Z","iopub.status.idle":"2025-02-05T16:03:39.625516Z","shell.execute_reply.started":"2025-02-05T16:03:39.612773Z","shell.execute_reply":"2025-02-05T16:03:39.624812Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # BEST\n# valid_thresholds = {'TS_6_4': {'apo-ferritin': 0.8749999999999999, 'beta-galactosidase': 0.975, 'ribosome': 0.95, 'thyroglobulin': 0.975, 'virus-like-particle': 0.95},\n\n# 'TS_5_4': {'apo-ferritin': 0.975, 'beta-galactosidase': 0.975, 'ribosome': 0.95, 'thyroglobulin': 0.95, 'virus-like-particle': 0.975},\n\n# 'TS_69_2': {'apo-ferritin': 0.95, 'beta-galactosidase': 0.95, 'ribosome': 0.7749999999999998, 'thyroglobulin': 0.9249999999999999, 'virus-like-particle': 0.6249999999999997}}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:03:39.626249Z","iopub.execute_input":"2025-02-05T16:03:39.626428Z","iopub.status.idle":"2025-02-05T16:03:39.641844Z","shell.execute_reply.started":"2025-02-05T16:03:39.626411Z","shell.execute_reply":"2025-02-05T16:03:39.641046Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # 01-29-25\n# valid_thresholds = {'TS_6_4': {'apo-ferritin': 0.8999999999999999, 'beta-galactosidase': 0.95, 'ribosome': 0.975, 'thyroglobulin': 0.975, 'virus-like-particle': 0.8749999999999999},\n\n# 'TS_5_4': {'apo-ferritin': 0.95, 'beta-galactosidase': 0.975, 'ribosome': 0.9249999999999999, 'thyroglobulin': 0.975, 'virus-like-particle': 0.5249999999999996},\n\n# 'TS_69_2': {'apo-ferritin': 0.95, 'beta-galactosidase': 0.975, 'ribosome': 0.95, 'thyroglobulin': 0.95, 'virus-like-particle': 0.8999999999999999}}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:03:39.642658Z","iopub.execute_input":"2025-02-05T16:03:39.642897Z","iopub.status.idle":"2025-02-05T16:03:39.656627Z","shell.execute_reply.started":"2025-02-05T16:03:39.642878Z","shell.execute_reply":"2025-02-05T16:03:39.655892Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# UNET 128 128 96 020425\n\nvalid_thresholds = {'TS_6_4': {'apo-ferritin': 0.975, 'beta-galactosidase': 0.975, 'ribosome': 0.8999999999999999, 'thyroglobulin': 0.975, 'virus-like-particle': 0.6249999999999997},\n\n'TS_5_4': {'apo-ferritin': 0.975, 'beta-galactosidase': 0.95, 'ribosome': 0.7249999999999998, 'thyroglobulin': 0.975, 'virus-like-particle': 0.5249999999999996},\n\n'TS_69_2': {'apo-ferritin': 0.975, 'beta-galactosidase': 0.9249999999999999, 'ribosome': 0.49999999999999956, 'thyroglobulin': 0.95, 'virus-like-particle': 0.49999999999999956}}\n\n\nvalid_radii = {'TS_6_4': {'apo-ferritin': 50, 'beta-galactosidase': 80, 'ribosome': 140, 'thyroglobulin': 120, 'virus-like-particle': 130},\n\n'TS_5_4': {'apo-ferritin': 50, 'beta-galactosidase': 80, 'ribosome': 140, 'thyroglobulin': 40, 'virus-like-particle': 130},\n\n'TS_69_2': {'apo-ferritin': 50, 'beta-galactosidase': 80, 'ribosome': 90, 'thyroglobulin': 120, 'virus-like-particle': 130}}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:49:33.997944Z","iopub.execute_input":"2025-02-05T16:49:33.998241Z","iopub.status.idle":"2025-02-05T16:49:34.003521Z","shell.execute_reply.started":"2025-02-05T16:49:33.998220Z","shell.execute_reply":"2025-02-05T16:49:34.002554Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # 01-30-25\n\n# valid_thresholds = {'TS_6_4': {'apo-ferritin': 0.975, 'beta-galactosidase': 0.8749999999999999, 'ribosome': 0.95, 'thyroglobulin': 0.975, 'virus-like-particle': 0.8749999999999999},\n\n# 'TS_5_4': {'apo-ferritin': 0.975, 'beta-galactosidase': 0.975, 'ribosome': 0.95, 'thyroglobulin': 0.975, 'virus-like-particle': 0.1999999999999993},\n\n# 'TS_69_2': {'apo-ferritin': 0.975, 'beta-galactosidase': 0.95, 'ribosome': 0.975, 'thyroglobulin': 0.975, 'virus-like-particle': 0.49999999999999956}}\n\n\n# valid_radii = {'TS_6_4': {'apo-ferritin': 30, 'beta-galactosidase': 80, 'ribosome': 140, 'thyroglobulin': 120, 'virus-like-particle': 70},\n\n# 'TS_5_4': {'apo-ferritin': 20, 'beta-galactosidase': 20, 'ribosome': 140, 'thyroglobulin': 110, 'virus-like-particle': 130},\n\n# 'TS_69_2': {'apo-ferritin': 50, 'beta-galactosidase': 80, 'ribosome': 140, 'thyroglobulin': 90, 'virus-like-particle': 130}}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:03:39.657324Z","iopub.execute_input":"2025-02-05T16:03:39.657600Z","iopub.status.idle":"2025-02-05T16:03:39.669728Z","shell.execute_reply.started":"2025-02-05T16:03:39.657580Z","shell.execute_reply":"2025-02-05T16:03:39.668923Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # 01-28-2025\n\n# valid_thresholds = {'TS_6_4': {'apo-ferritin': 0.8999999999999999, 'beta-galactosidase': 0.9249999999999999, 'ribosome': 0.975, 'thyroglobulin': 0.8999999999999999, 'virus-like-particle': 0.8999999999999999},\n\n# 'TS_5_4': {'apo-ferritin': 0.975, 'beta-galactosidase': 0.975, 'ribosome': 0.975, 'thyroglobulin': 0.95, 'virus-like-particle': 0.2999999999999994},\n\n# 'TS_69_2': {'apo-ferritin': 0.7999999999999998, 'beta-galactosidase': 0.8999999999999999, 'ribosome': 0.975, 'thyroglobulin': 0.9249999999999999, 'virus-like-particle': 0.8749999999999999}}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:03:39.672464Z","iopub.execute_input":"2025-02-05T16:03:39.672706Z","iopub.status.idle":"2025-02-05T16:03:39.683841Z","shell.execute_reply.started":"2025-02-05T16:03:39.672675Z","shell.execute_reply":"2025-02-05T16:03:39.683206Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PARTICLE= [\n    {\n        \"name\": \"apo-ferritin\",\n        \"difficulty\": 'easy',\n        \"pdb_id\": \"4V1W\",\n        \"label\": 1,\n        \"color\": [0, 255, 0, 0],\n        \"radius\": 60,\n        \"map_threshold\": 0.0418\n    },\n    # {\n    #     \"name\": \"beta-amylase\",\n    #     \"difficulty\": 'ignore',\n    #     \"pdb_id\": \"1FA2\",\n    #     \"label\": 2,\n    #     \"color\": [0, 0, 255, 255],\n    #     \"radius\": 65,\n    #     \"map_threshold\": 0.035\n    # },\n    {\n        \"name\": \"beta-galactosidase\",\n        \"difficulty\": 'hard',\n        \"pdb_id\": \"6X1Q\",\n        \"label\": 2,\n        \"color\": [0, 255, 0, 255],\n        \"radius\": 90,\n        \"map_threshold\": 0.0578\n    },\n    {\n        \"name\": \"ribosome\",\n        \"difficulty\": 'easy',\n        \"pdb_id\": \"6EK0\",\n        \"label\": 3,\n        \"color\": [0, 0, 255, 0],\n        \"radius\": 150,\n        \"map_threshold\": 0.0374\n    },\n    {\n        \"name\": \"thyroglobulin\",\n        \"difficulty\": 'hard',\n        \"pdb_id\": \"6SCJ\",\n        \"label\": 4,\n        \"color\": [0, 255, 255, 0],\n        \"radius\": 130,\n        \"map_threshold\": 0.0278\n    },\n    {\n        \"name\": \"virus-like-particle\",\n        \"difficulty\": 'easy',\n        \"pdb_id\": \"6N4V\",\n        \"label\": 5,\n        \"color\": [0, 0, 0, 255],\n        \"radius\": 135,\n        \"map_threshold\": 0.201\n    }\n]\n\nPARTICLE_NAME=['none']+[\n    PARTICLE[i]['name'] for i in range(5)\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:03:39.685138Z","iopub.execute_input":"2025-02-05T16:03:39.685380Z","iopub.status.idle":"2025-02-05T16:03:39.697781Z","shell.execute_reply.started":"2025-02-05T16:03:39.685356Z","shell.execute_reply":"2025-02-05T16:03:39.696954Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport plotly.graph_objects as go\nimport zarr # NEED\n# import copick # NEED\nfrom matplotlib import pyplot as plt\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport os\nimport shutil\nimport os\nimport shutil\n# from copick_utils.segmentation import segmentation_from_picks # NEED\n# import copick_utils.writers.write as write # Need\nfrom collections import defaultdict\nfrom tqdm import tqdm\nfrom torch.utils.data import DataLoader, Dataset\n# import segmentation_models_pytorch as smp\nimport sys\n# import neptune\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport yaml\n# from dataset_gaus_3d import CryoDataset\nimport monai.transforms as transforms\n# from monai.transforms.spatial.functional import resize\n# from data_utils import visualize, split_data, create_loaders, build_augmentations, calculate_loss_weights, probability_to_location, do_one_eval, plot_3d_points, convert_mask_to_one_hot, PARTICLE\nfrom monai.networks.nets import AHNet, SegResNetDS, VISTA3D, SegResNetDS2, UNet\nfrom scipy.optimize import linear_sum_assignment\nfrom sklearn.metrics import fbeta_score\nimport gc\nimport cc3d\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:03:39.698494Z","iopub.execute_input":"2025-02-05T16:03:39.698771Z","iopub.status.idle":"2025-02-05T16:04:04.963790Z","shell.execute_reply.started":"2025-02-05T16:03:39.698742Z","shell.execute_reply":"2025-02-05T16:04:04.963047Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class dotdict(dict):\n    __setattr__ = dict.__setitem__\n    __delattr__ = dict.__delitem__\n\n    def __getattr__(self, name):\n        try:\n            return self[name]\n        except KeyError:\n            raise AttributeError(name)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:04.964608Z","iopub.execute_input":"2025-02-05T16:04:04.965534Z","iopub.status.idle":"2025-02-05T16:04:04.969864Z","shell.execute_reply.started":"2025-02-05T16:04:04.965495Z","shell.execute_reply":"2025-02-05T16:04:04.968903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport torch\nfrom torch.utils.data import Dataset\nimport math\n\n# Instead of precomputing, compute at runtime\nclass CryoDataset(Dataset):\n    def __init__(self, data, num_slices=64, height=64, width=64, skip_ratio=0.5, transform=None):\n        self.num_slices = num_slices\n        self.skip_ratio = skip_ratio\n        self.height = height\n        self.width = width\n        self.data = data\n        # self.experiment = experiment\n        self.data['image'] = self.normalise_by_percentile(self.data['image'])\n        # self.data[i]['mask'] = np.where(self.data[i]['mask'] > 2, self.data[i]['mask'] - 1, self.data[i]['mask']) # Remove 2 class and readjust labels\n        # self.transform = transform\n\n        # Calculate the total number of possible images\n        self.total_images = 0\n        self.D, self.H, self.W = self.data['image'].shape\n        D, H, W = self.data['image'].shape\n        num_i = math.ceil((D - self.num_slices) / int(self.skip_ratio * self.num_slices)) + 1\n        num_h = math.ceil((H - self.height) / int(self.skip_ratio * self.height)) + 1\n        num_w = math.ceil((W - self.width) / int(self.skip_ratio * self.width)) + 1\n        self.total_images += num_i * num_h * num_w\n\n        self.num_i = num_i\n        self.num_h = num_h\n        self.num_w = num_w\n\n    def __len__(self):\n        return self.total_images\n    \n    def normalise_by_percentile(self, data, min=1, max=99):\n        min = np.percentile(data, min)\n        max = np.percentile(data, max)\n        data = np.clip(data, min, max)\n        data = (data - min) / (max - min)\n        return data\n\n    def __getitem__(self, idx):\n        data = self.data\n        D, H, W = self.D, self.H, self.W\n        \n        # Calculate number of steps as before\n        num_i = math.ceil((D - self.num_slices) / int(self.skip_ratio * self.num_slices)) + 1\n        num_h = math.ceil((H - self.height) / int(self.skip_ratio * self.height)) + 1\n        num_w = math.ceil((W - self.width) / int(self.skip_ratio * self.width)) + 1\n        total = num_i * num_h * num_w\n\n        i_idx = idx // (num_h * num_w)\n        h_idx = (idx % (num_h * num_w)) // num_w\n        w_idx = idx % num_w\n\n        # Calculate regular step positions\n        i = i_idx * int(self.skip_ratio * self.num_slices)\n        h = h_idx * int(self.skip_ratio * self.height)\n        w = w_idx * int(self.skip_ratio * self.width)\n\n        # Adjust final positions to ensure coverage of ends\n        if i_idx == num_i - 1:  # Last position in depth\n            i = D - self.num_slices\n        if h_idx == num_h - 1:  # Last position in height\n            h = H - self.height\n        if w_idx == num_w - 1:  # Last position in width\n            w = W - self.width\n\n        # Remove the min() calls since we now handle edge cases explicitly\n        image = data['image'][i:i+self.num_slices, h:h+self.height, w:w+self.width]\n        image = np.expand_dims(image, axis=0)\n        image = np.expand_dims(image, axis=0)\n\n        xyz = (w, h, i)\n        return {'image': image, 'xyz': xyz, 'experiment': data['experiment']}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:04.970619Z","iopub.execute_input":"2025-02-05T16:04:04.970874Z","iopub.status.idle":"2025-02-05T16:04:05.030148Z","shell.execute_reply.started":"2025-02-05T16:04:04.970853Z","shell.execute_reply":"2025-02-05T16:04:05.029445Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nD_SIZE, H_SIZE, W_SIZE = 184, 630, 630\n\nBASE_DIR = '/kaggle/input/czii-cryo-et-object-identification/test/static/ExperimentRuns/'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:05.031055Z","iopub.execute_input":"2025-02-05T16:04:05.031348Z","iopub.status.idle":"2025-02-05T16:04:05.046406Z","shell.execute_reply.started":"2025-02-05T16:04:05.031320Z","shell.execute_reply":"2025-02-05T16:04:05.045576Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def collate_fn(batch):\n    images = [torch.from_numpy(b['image']) for b in batch]\n    \n    xyzs = [b['xyz'] for b in batch]\n    experiment = [b['experiment'] for b in batch]\n    # Stack the images and masks so that a new dimension is added at the beginning\n    return {'image':torch.concat(images, dim=0), 'xyz':xyzs, 'experiment':experiment}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:05.047231Z","iopub.execute_input":"2025-02-05T16:04:05.047502Z","iopub.status.idle":"2025-02-05T16:04:05.058393Z","shell.execute_reply.started":"2025-02-05T16:04:05.047474Z","shell.execute_reply":"2025-02-05T16:04:05.057691Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def probability_to_location_old(probability, threshold):\n    _,D,H,W = probability.shape\n\n    location={}\n    for p in PARTICLE:\n        p = dotdict(p)\n        l = p.label\n\n        cc, P = cc3d.connected_components(probability[l]>threshold[p.name], return_N=True)\n        stats = cc3d.statistics(cc)\n        zyx=stats['centroids'][1:]*10\n        xyz = np.ascontiguousarray(zyx[:,::-1])\n        # Cluster the coordinates before storing\n        # xyz = cluster_coordinates(xyz, eps=40, min_samples=1)\n        location[p.name]=xyz\n        \n    return location\n\n\nif DEBUG:\n    \n    def probability_to_location(probability, threshold, radius_dict):\n        _,D,H,W = probability.shape\n    \n        location={}\n        for p in PARTICLE:\n            p = dotdict(p)\n            l = p.label\n    \n            cc, P = cc3d.connected_components(probability[l]>threshold[p.name], return_N=True, connectivity=26)\n            stats = cc3d.statistics(cc)\n            \n            # Filter components based on volume threshold\n            radius = radius_dict[p['name']] / 10\n            min_volume = (0.35 * radius)**3 * 4/3 * np.pi  # Minimum volume as fraction of particle volume\n            valid_components = stats['voxel_counts'][1:] >= min_volume\n            \n            # Only keep centroids of components that meet volume threshold\n            zyx = stats['centroids'][1:][valid_components] * 10\n            xyz = np.ascontiguousarray(zyx[:,::-1])\n            # Cluster the coordinates before storing\n            # xyz = cluster_coordinates(xyz, eps=40, min_samples=1)\n            location[p.name]=xyz\n            \n        return location\n\nif not DEBUG:\n\n    def probability_to_location(probability, radius_dict):\n        _,D,H,W = probability.shape\n    \n        location={}\n        for p in PARTICLE:\n            p = dotdict(p)\n            l = p.label\n    \n            cc, P = cc3d.connected_components(probability[l], return_N=True, connectivity=26)\n            stats = cc3d.statistics(cc)\n            \n            # Filter components based on volume threshold\n            radius = radius_dict[p['name']] / 10\n            min_volume = (0.35 * radius)**3 * 4/3 * np.pi  # Minimum volume as fraction of particle volume\n            valid_components = stats['voxel_counts'][1:] >= min_volume\n            \n            # Only keep centroids of components that meet volume threshold\n            zyx = stats['centroids'][1:][valid_components] * 10\n            xyz = np.ascontiguousarray(zyx[:,::-1])\n            # Cluster the coordinates before storing\n            # xyz = cluster_coordinates(xyz, eps=40, min_samples=1)\n            location[p.name]=xyz\n            \n        return location","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:05.059153Z","iopub.execute_input":"2025-02-05T16:04:05.059368Z","iopub.status.idle":"2025-02-05T16:04:05.068394Z","shell.execute_reply.started":"2025-02-05T16:04:05.059349Z","shell.execute_reply":"2025-02-05T16:04:05.067722Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_pred_df(pred_heat_maps, thresholds, optimal_blob_radii, experiment):\n    pred_df = pd.DataFrame(columns=['id', 'experiment', 'particle_type', 'x', 'y', 'z'])\n    idx = 0\n    # {p['name']: threshold for p in PARTICLE}\n    locs = probability_to_location(pred_heat_maps, threshold=thresholds, radius_dict=optimal_blob_radii)\n    for p in PARTICLE:\n        temp_locs = locs[p['name']]\n        for l in temp_locs:\n            pred_df.at[idx, 'id'] = idx\n            pred_df.at[idx, 'experiment'] = experiment\n            pred_df.at[idx, 'particle_type'] = p['name']\n            pred_df.at[idx, 'x'] = l[0]\n            pred_df.at[idx, 'y'] = l[1]\n            pred_df.at[idx, 'z'] = l[2]\n            idx += 1\n    return pred_df","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.DataFrame()\ndf_idx = 0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:05.069182Z","iopub.execute_input":"2025-02-05T16:04:05.069409Z","iopub.status.idle":"2025-02-05T16:04:05.087707Z","shell.execute_reply.started":"2025-02-05T16:04:05.069390Z","shell.execute_reply":"2025-02-05T16:04:05.087054Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TEST_RUN = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:05.088429Z","iopub.execute_input":"2025-02-05T16:04:05.088711Z","iopub.status.idle":"2025-02-05T16:04:05.101079Z","shell.execute_reply.started":"2025-02-05T16:04:05.088680Z","shell.execute_reply":"2025-02-05T16:04:05.100389Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if not DEBUG:\n    models_list = [\n        # '/kaggle/input/cryo-models/UNET_128_128_96_ADAMW_20250205-070947_20250205-070947-20250205T152105Z-001/UNET_128_128_96_ADAMW_20250205-070947_20250205-070947/Fold_1_TS_5_4'\n        # '/kaggle/input/cryo-models/UNET_128_128_96_ADAMW_20250205-070947_20250205-070947-20250205T152105Z-001/UNET_128_128_96_ADAMW_20250205-070947_20250205-070947/Fold_2_TS_69_2'\n        '/kaggle/input/cryo-models/UNET_128_128_96_ADAMW_20250205-070947_20250205-070947-20250205T152105Z-001/UNET_128_128_96_ADAMW_20250205-070947_20250205-070947/Fold_3_TS_6_4'\n    ]\n    use_tta = [False]\n    # Load models\n    models = []\n    for weights_dir in models_list:\n        model_file = min(\n            [file for file in os.listdir(weights_dir) if file.endswith(\".pth\")],\n            key=lambda x: float(x.split(\"_\")[-1].replace(\".pth\", \"\"))\n        )\n        weights_dir = os.path.join(weights_dir, model_file)\n    \n        model_fold_name = weights_dir.split('/')[-2][7:]\n    \n        print(f'Using model directory {weights_dir} with fold name {model_fold_name}')\n    \n        # thresh_val = thresholds[model_fold_name]\n        # optimal_thresholds = {p['name']: thresh_val for p in PARTICLE}\n    \n        optimal_thresholds = valid_thresholds[model_fold_name]\n\n        optimal_radii = valid_radii[model_fold_name]\n    \n        state_dict = torch.load(weights_dir, weights_only=True)\n        net = UNet(\n            spatial_dims=3,            # 3D\n            in_channels=1,             # Adjust if you have >1 input channels\n            out_channels=6,            # e.g. 6 classes/particles\n            channels=(64, 128, 256, 512, 1024),  # increased channels\n            strides=(2, 2, 2, 2),      # Downsampling at each level\n            num_res_units=4,           # Residual blocks for deeper capacity\n            norm='batch',              # Or instance/group norm if you prefer\n            # dropout=0.2                # Could add small dropout if needed\n        ).to(device)\n        net.load_state_dict(state_dict)\n        net = net.to(device)\n        net = torch.nn.DataParallel(net, device_ids=[0, 1])\n        net.eval()\n        net = torch.compile(net)\n        models.append(net)\n        \n        print('Loaded model weights')\n    # thresholds = {'TS_5_4': 0.5, 'TS_69_2': 0.65, 'TS_6_4': 0.15, 'TS_6_6': 0.75, 'TS_73_6': 0.6, 'TS_86_3': 0.75, 'TS_99_9': 0.65}\n    # for sample in data:\n    for experiment_dir in tqdm(os.listdir(BASE_DIR)):\n        sample = {}\n        zarr_dir = os.path.join(BASE_DIR, experiment_dir, 'VoxelSpacing10.000/denoised.zarr')\n        z = zarr.open(zarr_dir, mode='r')\n        sample['image'] = z[0]\n        sample['experiment'] = experiment_dir\n        \n        ensemble_probs = np.zeros((6, D_SIZE, H_SIZE, W_SIZE), dtype=np.float32)\n        # ensemble_counts = np.zeros((6, D_SIZE, H_SIZE, W_SIZE), dtype=np.float32)\n        exp = sample['experiment']\n        print(f'Running experiment {exp}')\n        ds = CryoDataset(sample, num_slices=cfg.DEPTH, height=cfg.HEIGHT, width=cfg.WIDTH, skip_ratio=cfg.SKIP_RATIO, transform=None)\n        dataloader = DataLoader(ds, batch_size=16, shuffle=False, num_workers=4, persistent_workers=True, prefetch_factor=2, pin_memory=True, collate_fn=collate_fn)\n        for index_m, net in enumerate(models):\n            experiment_probs = np.zeros((6, D_SIZE, H_SIZE, W_SIZE), dtype=np.float32)\n            experiment_counts = np.zeros((6, D_SIZE, H_SIZE, W_SIZE), dtype=np.float32)\n            # with torch.amp.autocast('cuda', enabled=True):\n            with torch.no_grad():\n                for batch in dataloader:\n                    images = batch['image'].to('cuda')\n                    xyz = batch['xyz']\n                    \n                    # Original predictions\n                    preds = net(images)\n\n                    if use_tta[index_m]:\n                        # Rotate predictions\n                        flipped_images = torch.flip(images, dims=(2,))\n                        flipped_preds = net(flipped_images)\n                        # flipped_preds = F.softmax(flipped_preds, dim=1).detach().cpu().numpy()\n                        # Flip back the predictions\n                        flipped_preds = torch.flip(flipped_preds, dims=(2,))\n                        averaged_preds = preds + flipped_preds\n                    else:\n                        averaged_preds = preds\n\n                    averaged_preds = F.softmax(averaged_preds, dim=1).detach().cpu().numpy()\n    \n                    # Update probabilities and counts\n                    for i, pred in enumerate(averaged_preds):\n                        W, H, D = xyz[i]\n                        experiment_probs[:, D:D+cfg.DEPTH, H:H+cfg.HEIGHT, W:W+cfg.WIDTH] += pred\n                        experiment_counts[:, D:D+cfg.DEPTH, H:H+cfg.HEIGHT, W:W+cfg.WIDTH] += 1\n\n            if TEST_RUN:\n                assert not np.any(experiment_counts == 0), \"Counts Array contains values equal to 0\"\n            experiment_probs = experiment_probs / experiment_counts\n            \n            # For each particle, any values under threshold should be set to 0\n            for p in PARTICLE:\n                experiment_probs[p['label']] = np.where(\n                    experiment_probs[p['label']] < optimal_thresholds[p['name']], \n                    0, \n                    1\n                )\n            # Add to ensemble probabilities\n            ensemble_probs += experiment_probs\n        # assert not np.any(experiment_probs == 0), \"Probs Array contains values equal to 0\"\n        # Create temporary optimal thresholds that check if a majority of the models in the ensemble agree\n        majority_threshold = len(models) // 2 + 1\n        majority_probs = ensemble_probs >= majority_threshold\n        locs = probability_to_location(majority_probs, optimal_radii)\n        rows = []  # Temporary list to store rows\n        for particle, particle_locs in locs.items():\n            for p_loc in particle_locs:\n                rows.append({\n                    'id': df_idx,\n                    'experiment': exp,\n                    'particle_type': particle,\n                    'x': p_loc[0],\n                    'y': p_loc[1],\n                    'z': p_loc[2],\n                })\n                df_idx += 1\n        \n        # Append all rows to the DataFrame at once\n        df = pd.concat([df, pd.DataFrame(rows)], ignore_index=True)\n        del sample\n        del experiment_probs\n        del experiment_counts\n        del locs\n        del ensemble_probs\n        gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:05.101820Z","iopub.execute_input":"2025-02-05T16:04:05.102177Z","iopub.status.idle":"2025-02-05T16:04:05.117453Z","shell.execute_reply.started":"2025-02-05T16:04:05.102143Z","shell.execute_reply":"2025-02-05T16:04:05.116680Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nif DEBUG:\n    import json\n    gt_locs = {}\n    experiments = os.listdir(\"/kaggle/input/czii-cryo-et-object-identification/train/overlay/ExperimentRuns\")\n    \n    def read_one_truth(id, overlay_dir):\n        location={}\n    \n        json_dir = f'{overlay_dir}/{id}/Picks'\n        for p in PARTICLE_NAME[1:]:\n            json_file = f'{json_dir}/{p}.json'\n    \n            with open(json_file, 'r') as f:\n                json_data = json.load(f)\n    \n            num_point = len(json_data['points'])\n            loc = np.array([list(json_data['points'][i]['location'].values()) for i in range(num_point)])\n            location[p] = loc\n    \n        return location\n\n    for exp in experiments:\n        gt_locations = read_one_truth(exp, '/kaggle/input/czii-cryo-et-object-identification/train/overlay/ExperimentRuns')\n        gt_locs[exp] = gt_locations\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:05.118238Z","iopub.execute_input":"2025-02-05T16:04:05.118530Z","iopub.status.idle":"2025-02-05T16:04:05.266579Z","shell.execute_reply.started":"2025-02-05T16:04:05.118503Z","shell.execute_reply":"2025-02-05T16:04:05.266011Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Scoring functions\n\nfrom scipy.spatial import KDTree\n\n\nclass ParticipantVisibleError(Exception):\n    pass\n\n\ndef compute_metrics(reference_points, reference_radius, candidate_points):\n    num_reference_particles = len(reference_points)\n    num_candidate_particles = len(candidate_points)\n\n    if len(reference_points) == 0:\n        return 0, num_candidate_particles, 0\n\n    if len(candidate_points) == 0:\n        return 0, 0, num_reference_particles\n\n    ref_tree = KDTree(reference_points)\n    candidate_tree = KDTree(candidate_points)\n    raw_matches = candidate_tree.query_ball_tree(ref_tree, r=reference_radius)\n    matches_within_threshold = []\n    for match in raw_matches:\n        matches_within_threshold.extend(match)\n    # Prevent submitting multiple matches per particle.\n    # This won't be be strictly correct in the (extremely rare) case where true particles\n    # are very close to each other.\n    matches_within_threshold = set(matches_within_threshold)\n    tp = int(len(matches_within_threshold))\n    fp = int(num_candidate_particles - tp)\n    fn = int(num_reference_particles - tp)\n    return tp, fp, fn\n\ndef score(\n        solution: pd.DataFrame,\n        submission: pd.DataFrame,\n        row_id_column_name: str,\n        distance_multiplier: float,\n        beta: int) -> float:\n    '''\n    F_beta\n      - a true positive occurs when\n         - (a) the predicted location is within a threshold of the particle radius, and\n         - (b) the correct `particle_type` is specified\n      - raw results (TP, FP, FN) are aggregated across all experiments for each particle type\n      - f_beta is calculated for each particle type\n      - individual f_beta scores are weighted by particle type for final score\n    '''\n\n    particle_radius = {\n        'apo-ferritin': 60,\n        'beta-amylase': 65,\n        'beta-galactosidase': 90,\n        'ribosome': 150,\n        'thyroglobulin': 130,\n        'virus-like-particle': 135,\n    }\n\n    weights = {\n        'apo-ferritin': 1,\n        'beta-amylase': 0,\n        'beta-galactosidase': 2,\n        'ribosome': 1,\n        'thyroglobulin': 2,\n        'virus-like-particle': 1,\n    }\n\n    particle_radius = {k: v * distance_multiplier for k, v in particle_radius.items()}\n\n    # Filter submission to only contain experiments found in the solution split\n    split_experiments = set(solution['experiment'].unique())\n    submission = submission.loc[submission['experiment'].isin(split_experiments)]\n\n    # Only allow known particle types\n    if not set(submission['particle_type'].unique()).issubset(set(weights.keys())):\n        raise ParticipantVisibleError('Unrecognized `particle_type`.')\n\n    assert solution.duplicated(subset=['experiment', 'x', 'y', 'z']).sum() == 0\n    assert particle_radius.keys() == weights.keys()\n\n    results = {}\n    for particle_type in solution['particle_type'].unique():\n        results[particle_type] = {\n            'total_tp': 0,\n            'total_fp': 0,\n            'total_fn': 0,\n        }\n\n    for experiment in split_experiments:\n        for particle_type in solution['particle_type'].unique():\n            reference_radius = particle_radius[particle_type]\n            select = (solution['experiment'] == experiment) & (solution['particle_type'] == particle_type)\n            reference_points = solution.loc[select, ['x', 'y', 'z']].values\n\n            select = (submission['experiment'] == experiment) & (submission['particle_type'] == particle_type)\n            candidate_points = submission.loc[select, ['x', 'y', 'z']].values\n\n            if len(reference_points) == 0:\n                reference_points = np.array([])\n                reference_radius = 1\n\n            if len(candidate_points) == 0:\n                candidate_points = np.array([])\n\n            tp, fp, fn = compute_metrics(reference_points, reference_radius, candidate_points)\n\n            results[particle_type]['total_tp'] += tp\n            results[particle_type]['total_fp'] += fp\n            results[particle_type]['total_fn'] += fn\n\n    aggregate_fbeta = 0.0\n    for particle_type, totals in results.items():\n        tp = totals['total_tp']\n        fp = totals['total_fp']\n        fn = totals['total_fn']\n\n        precision = tp / (tp + fp) if tp + fp > 0 else 0\n        recall = tp / (tp + fn) if tp + fn > 0 else 0\n        fbeta = (1 + beta**2) * (precision * recall) / (beta**2 * precision + recall) if (precision + recall) > 0 else 0.0\n        aggregate_fbeta += fbeta * weights.get(particle_type, 1.0)\n\n    if weights:\n        aggregate_fbeta = aggregate_fbeta / sum(weights.values())\n    else:\n        aggregate_fbeta = aggregate_fbeta / len(results)\n    return aggregate_fbeta","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:05.267262Z","iopub.execute_input":"2025-02-05T16:04:05.267450Z","iopub.status.idle":"2025-02-05T16:04:05.279258Z","shell.execute_reply.started":"2025-02-05T16:04:05.267433Z","shell.execute_reply":"2025-02-05T16:04:05.278466Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# valid_thresholds = {\"TS_6_4\": {'apo-ferritin': 0.2, 'beta-galactosidase': 0.9500000000000001, 'ribosome': 0.15000000000000002, 'thyroglobulin': 0.3, 'virus-like-particle': 0.45},\n# \"TS_5_4\": {'apo-ferritin': 0.4, 'beta-galactosidase': 0.05, 'ribosome': 0.1, 'thyroglobulin': 0.5, 'virus-like-particle': 0.5},\n# \"TS_69_2\": {'apo-ferritin': 0.8, 'beta-galactosidase': 0.05, 'ribosome': 0.2, 'thyroglobulin': 0.25, 'virus-like-particle': 0.6500000000000001}}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:05.280050Z","iopub.execute_input":"2025-02-05T16:04:05.280330Z","iopub.status.idle":"2025-02-05T16:04:05.295351Z","shell.execute_reply.started":"2025-02-05T16:04:05.280298Z","shell.execute_reply":"2025-02-05T16:04:05.294598Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if DEBUG:\n    DIR = '/kaggle/input/cryo-models/UNET_128_128_96_ADAMW_20250205-070947_20250205-070947-20250205T152105Z-001/UNET_128_128_96_ADAMW_20250205-070947_20250205-070947'\n    models = []\n    valid_ids = ['TS_6_4', 'TS_5_4', 'TS_69_2']\n    for vid in valid_ids:\n        # Load model\n        model_dir = None\n        for m in os.listdir(DIR):\n            if vid in m:\n                model_dir = m\n        weights_dir = os.path.join(DIR, model_dir)\n        \n        model_file = min(\n            [file for file in os.listdir(weights_dir) if file.endswith(\".pth\")],\n            key=lambda x: float(x.split(\"_\")[-1].replace(\".pth\", \"\"))\n        )\n        print(f\"Loading {model_file}\")\n        weights_dir = os.path.join(weights_dir, model_file)\n\n        model_fold_name = weights_dir.split('/')[-2][7:]\n\n        # thresh_val = thresholds[model_fold_name]\n        # optimal_thresholds = {p['name']: thresh_val for p in PARTICLE}\n    \n        state_dict = torch.load(weights_dir, weights_only=True)\n        net = UNet(\n            spatial_dims=3,            # 3D\n            in_channels=1,             # Adjust if you have >1 input channels\n            out_channels=6,            # e.g. 6 classes/particles\n            channels=(64, 128, 256, 512, 1024),  # increased channels\n            strides=(2, 2, 2, 2),      # Downsampling at each level\n            num_res_units=4,           # Residual blocks for deeper capacity\n            norm='batch',              # Or instance/group norm if you prefer\n            dropout=0.2                # Could add small dropout if needed\n        ).to(device)\n        net.load_state_dict(state_dict)\n        # net = torch.nn.DataParallel(net, device_ids=[0, 1])\n        net.eval()\n        \n        # net = net.to(device)\n        models.append(net)\n        \n        print('Loaded model weights')\n        \n    import pandas as pd\n    def build_gt_df(gt_locs):\n        experiments = gt_locs.keys()\n        gt_df = pd.DataFrame()\n        idx = 0\n        for experiment in experiments:\n            true_locs = gt_locs[experiment]\n            for p in PARTICLE:\n                temp_locs = true_locs[p['name']]\n                for l in temp_locs:\n                    gt_df.at[idx, 'id'] = idx\n                    gt_df.at[idx, 'experiment'] = experiment\n                    gt_df.at[idx, 'particle_type'] = p['name']\n                    gt_df.at[idx, 'x'] = l[0]\n                    gt_df.at[idx, 'y'] = l[1]\n                    gt_df.at[idx, 'z'] = l[2]\n                    idx += 1\n        return gt_df\n    \n    def build_pred_df(pred_heat_maps, thresholds, optimal_blob_radii, experiment):\n        pred_df = pd.DataFrame(columns=['id', 'experiment', 'particle_type', 'x', 'y', 'z'])\n        idx = 0\n        # {p['name']: threshold for p in PARTICLE}\n        locs = probability_to_location(pred_heat_maps, threshold=thresholds, radius_dict=optimal_blob_radii)\n        for p in PARTICLE:\n            temp_locs = locs[p['name']]\n            for l in temp_locs:\n                pred_df.at[idx, 'id'] = idx\n                pred_df.at[idx, 'experiment'] = experiment\n                pred_df.at[idx, 'particle_type'] = p['name']\n                pred_df.at[idx, 'x'] = l[0]\n                pred_df.at[idx, 'y'] = l[1]\n                pred_df.at[idx, 'z'] = l[2]\n                idx += 1\n        return pred_df\n\n    def find_best_threshold_per_particle(pred_heat_maps, gt_df, experiment):\n        temp_gt = gt_df[gt_df.experiment == experiment]\n        optimal_thresholds = {p['name']: 0.0 for p in PARTICLE}\n        threshold_options = [x for x in np.arange(1.0, 0.0, -0.025)]\n        threshold_options.append(0.98)\n        threshold_options.append(0.985)\n        threshold_options.append(0.99)\n        threshold_options.append(0.995)\n        for p in tqdm(PARTICLE):\n            cur_best_score = -1\n            for i, to in enumerate(threshold_options):\n                optimal_thresholds[p['name']] = to\n                df_pred = build_pred_df(pred_heat_maps, optimal_thresholds, valid_radii[experiment], experiment)\n                lb_score = score(temp_gt, df_pred, row_id_column_name='id', distance_multiplier=0.5, beta=4)\n                if lb_score < cur_best_score:\n                    optimal_thresholds[p['name']] = threshold_options[i-1]\n                    break\n                else:\n                    cur_best_score = lb_score\n        return optimal_thresholds\n\n    def find_best_blob_size_per_particle(pred_heat_maps, gt_df, experiment, optimal_thresholds):\n        temp_gt = gt_df[gt_df.experiment == experiment]\n        optimal_blob_radii = {p['name']: 0.0 for p in PARTICLE}\n        for p in tqdm(PARTICLE):\n            blob_radius = [x for x in np.arange(0, p['radius'], 10)]\n            cur_best_score = -1\n            for i, to in enumerate(blob_radius):\n                optimal_blob_radii[p['name']] = to\n                df_pred = build_pred_df(pred_heat_maps, optimal_thresholds, optimal_blob_radii, experiment)\n                lb_score = score(temp_gt, df_pred, row_id_column_name='id', distance_multiplier=0.5, beta=4)\n                if lb_score < cur_best_score:\n                    optimal_blob_radii[p['name']] = blob_radius[i-1]\n                    break\n                else:\n                    cur_best_score = lb_score\n        return optimal_blob_radii\n                \n\n    dfs = []\n    # thresholds = {'TS_5_4': 0.5, 'TS_69_2': 0.65, 'TS_6_4': 0.15, 'TS_6_6': 0.75, 'TS_73_6': 0.6, 'TS_86_3': 0.75, 'TS_99_9': 0.65}\n    gt_df = build_gt_df(gt_locs)\n    sub_gt_df = gt_df[gt_df.experiment.isin(['TS_6_4', 'TS_5_4', 'TS_69_2'])]\n    # Run predictions on validation sets\n    for vid, experiment_dir in tqdm(enumerate(os.listdir(\"/kaggle/input/czii-cryo-et-object-identification/test/static/ExperimentRuns/\")), total=len(valid_ids)):\n        sample = {}\n        zarr_dir = os.path.join(BASE_DIR, experiment_dir, 'VoxelSpacing10.000/denoised.zarr')\n        z = zarr.open(zarr_dir, mode='r')\n        sample['image'] = z[0]\n        sample['experiment'] = experiment_dir\n        \n        experiment_probs = np.zeros((6, D_SIZE, H_SIZE, W_SIZE), dtype=np.float32)\n        experiment_counts = np.zeros((6, D_SIZE, H_SIZE, W_SIZE), dtype=np.float32)\n        exp = sample['experiment']\n        print(f'Running experiment {exp}')\n        ds = CryoDataset(sample, num_slices=cfg.DEPTH, height=cfg.HEIGHT, width=cfg.WIDTH, skip_ratio=cfg.SKIP_RATIO, transform=None)\n        dataloader = DataLoader(ds, batch_size=16, shuffle=False, num_workers=0, pin_memory=True, collate_fn=collate_fn)\n        net = models[vid]\n        net = net.to(device)\n        net.eval()\n        # with torch.amp.autocast('cuda', enabled=True):\n        with torch.no_grad():\n            for batch in tqdm(dataloader, leave=True, position=0):\n                images = batch['image'].to('cuda')\n                xyz = batch['xyz']\n                \n                # --- 1) Original Predictions ---\n                preds = net(images)\n                # print(preds)\n                # preds = F.softmax(preds, dim=1).detach().cpu().numpy()\n                \n                # # --- 2) Rotate Predictions (k=1) ---\n                # rotated_images = torch.rot90(images, k=1, dims=(3, 4))\n                # rotated_preds = net(rotated_images)\n                # # rotated_preds = F.softmax(rotated_preds, dim=1).detach().cpu().numpy()\n                # # rotate predictions back\n                # rotated_preds = torch.rot90(rotated_preds, k=-1, dims=(3, 4))\n                \n                # # --- 3) Rotate Predictions (k=-1) ---\n                # rotated_images2 = torch.rot90(images, k=-1, dims=(3, 4))\n                # rotated_preds2 = net(rotated_images2)\n                # # rotated_preds2 = F.softmax(rotated_preds2, dim=1).detach().cpu().numpy()\n                # # rotate predictions back\n                # rotated_preds2 = torch.rot90(rotated_preds2, k=1, dims=(3, 4))\n                \n                # --- 4) Flip TTA (Horizontal Flip) ---\n                # Flip along the last dimension (width). If your width is at index 4, use dims=(4,).\n                flipped_images = torch.flip(images, dims=(2,))\n                flipped_preds = net(flipped_images)\n                # flipped_preds = F.softmax(flipped_preds, dim=1).detach().cpu().numpy()\n                # Flip back the predictions\n                flipped_preds = torch.flip(flipped_preds, dims=(2,))\n                \n                # --- 5) Average Everything ---\n                averaged_preds = preds + flipped_preds # + rotated_preds + rotated_preds2 + flipped_preds\n                averaged_preds = F.softmax(averaged_preds, dim=1).detach().cpu().numpy()\n\n                # print(averaged_preds)\n                # Update probabilities and counts\n                for i, pred in enumerate(averaged_preds):\n                    W, H, D = xyz[i]\n                    experiment_probs[:, D:D+cfg.DEPTH, H:H+cfg.HEIGHT, W:W+cfg.WIDTH] += pred\n                    experiment_counts[:, D:D+cfg.DEPTH, H:H+cfg.HEIGHT, W:W+cfg.WIDTH] += 1\n\n        \n        experiment_probs = experiment_probs / experiment_counts\n        # experiment_probs = F.interpolate(experiment_probs, scale_factor=0.5, mode='bilinear', align_corners=False)\n        # Make sure no values in count are 0 (were not hit)\n        assert not np.any(experiment_counts == 0), \"Counts Array contains values equal to 0\"\n        # assert not np.any(experiment_probs == 0), \"Probs Array contains values equal to 0\"\n        # experiment_probs = np.nan_to_num(experiment_probs, nan=0)\n        # locs = probability_to_location(experiment_probs, optimal_thresholds)\n        optimal_th = valid_thresholds[experiment_dir]\n        optimal_blob_radii = valid_radii[experiment_dir]\n        if False:\n            optimal_blob_radii = {p['name']: p['radius'] * 0.35 for p in PARTICLE}\n        if False:\n            optimal_th = find_best_threshold_per_particle(experiment_probs, gt_df, experiment_dir)\n        if True:\n            optimal_blob_radii = find_best_blob_size_per_particle(experiment_probs, gt_df, experiment_dir, optimal_th)\n        print(f'Optimal threshold values for experiment {experiment_dir}', optimal_th)\n        print(f'Optimal blob radii values for experiment {experiment_dir}', optimal_blob_radii)\n        df_pred = build_pred_df(experiment_probs, optimal_th, optimal_blob_radii, experiment_dir)\n        dfs.append(df_pred)\n        # del sample\n        # del experiment_probs\n        # del experiment_counts\n        # del locs\n        gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:34:42.897219Z","iopub.execute_input":"2025-02-05T16:34:42.897625Z","iopub.status.idle":"2025-02-05T16:49:18.883897Z","shell.execute_reply.started":"2025-02-05T16:34:42.897598Z","shell.execute_reply":"2025-02-05T16:49:18.883165Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# TTA NEW BLOBS\n\nif DEBUG:\n\n    pred_df = pd.concat(dfs)\n    \n    print(pred_df.experiment.unique())\n    \n    sub_gt_df = gt_df[gt_df.experiment.isin(['TS_6_4', 'TS_5_4', 'TS_69_2'])]\n    \n    sub_gt_df.experiment.unique()\n    \n    print(score(sub_gt_df, pred_df, row_id_column_name='id', distance_multiplier=0.5, beta=4))\n    \n    valid_ids = ['TS_6_4', 'TS_5_4', 'TS_69_2']\n    for vid in valid_ids:\n        print(vid, score(sub_gt_df[sub_gt_df.experiment == vid], pred_df[pred_df.experiment == vid], row_id_column_name='id', distance_multiplier=0.5, beta=4))\n\n    plt.imshow(experiment_probs[0,92,:,:])\n    plt.colorbar()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:49:44.065927Z","iopub.execute_input":"2025-02-05T16:49:44.066246Z","iopub.status.idle":"2025-02-05T16:49:44.365665Z","shell.execute_reply.started":"2025-02-05T16:49:44.066219Z","shell.execute_reply":"2025-02-05T16:49:44.364822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# TTA\n\nif DEBUG:\n\n    pred_df = pd.concat(dfs)\n    \n    print(pred_df.experiment.unique())\n    \n    sub_gt_df = gt_df[gt_df.experiment.isin(['TS_6_4', 'TS_5_4', 'TS_69_2'])]\n    \n    sub_gt_df.experiment.unique()\n    \n    print(score(sub_gt_df, pred_df, row_id_column_name='id', distance_multiplier=0.5, beta=4))\n    \n    valid_ids = ['TS_6_4', 'TS_5_4', 'TS_69_2']\n    for vid in valid_ids:\n        print(vid, score(sub_gt_df[sub_gt_df.experiment == vid], pred_df[pred_df.experiment == vid], row_id_column_name='id', distance_multiplier=0.5, beta=4))\n\n    plt.imshow(experiment_probs[0,92,:,:])\n    plt.colorbar()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:31:28.088039Z","iopub.execute_input":"2025-02-05T16:31:28.088340Z","iopub.status.idle":"2025-02-05T16:31:28.399444Z","shell.execute_reply.started":"2025-02-05T16:31:28.088319Z","shell.execute_reply":"2025-02-05T16:31:28.398607Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# UNET 128 128 96 020425\n\n# TS_6_4 {'apo-ferritin': 0.975, 'beta-galactosidase': 0.975, 'ribosome': 0.8999999999999999, 'thyroglobulin': 0.975, 'virus-like-particle': 0.6249999999999997}\n\n# TS_5_4 {'apo-ferritin': 0.975, 'beta-galactosidase': 0.95, 'ribosome': 0.7249999999999998, 'thyroglobulin': 0.975, 'virus-like-particle': 0.5249999999999996}\n\n# TS_69_2 {'apo-ferritin': 0.975, 'beta-galactosidase': 0.9249999999999999, 'ribosome': 0.49999999999999956, 'thyroglobulin': 0.95, 'virus-like-particle': 0.49999999999999956}\n\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:25:31.644398Z","iopub.execute_input":"2025-02-05T16:25:31.644749Z","iopub.status.idle":"2025-02-05T16:25:31.648131Z","shell.execute_reply.started":"2025-02-05T16:25:31.644721Z","shell.execute_reply":"2025-02-05T16:25:31.647221Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 01-30-25\n\n# TS_6_4 {'apo-ferritin': 30, 'beta-galactosidase': 30, 'ribosome': 140, 'thyroglobulin': 120, 'virus-like-particle': 130}\n\n# TS_5_4 {'apo-ferritin': 20, 'beta-galactosidase': 20, 'ribosome': 140, 'thyroglobulin': 110, 'virus-like-particle': 130}\n\n# TS_69_2 {'apo-ferritin': 50, 'beta-galactosidase': 80, 'ribosome': 140, 'thyroglobulin': 90, 'virus-like-particle': 130}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:28.160098Z","iopub.status.idle":"2025-02-05T16:04:28.160384Z","shell.execute_reply":"2025-02-05T16:04:28.160272Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 01-30-25\n\n# TS_6_4 {'apo-ferritin': 0.975, 'beta-galactosidase': 0.95, 'ribosome': 0.975, 'thyroglobulin': 0.975, 'virus-like-particle': 0.975}\n\n# TS_5_4 {'apo-ferritin': 0.975, 'beta-galactosidase': 0.975, 'ribosome': 0.95, 'thyroglobulin': 0.975, 'virus-like-particle': 0.1999999999999993}\n\n# TS_69_2 {'apo-ferritin': 0.975, 'beta-galactosidase': 0.95, 'ribosome': 0.975, 'thyroglobulin': 0.975, 'virus-like-particle': 0.49999999999999956}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:28.161316Z","iopub.status.idle":"2025-02-05T16:04:28.161661Z","shell.execute_reply":"2025-02-05T16:04:28.161507Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 01-30-25\n\nif DEBUG:\n\n    pred_df = pd.concat(dfs)\n    \n    print(pred_df.experiment.unique())\n    \n    sub_gt_df = gt_df[gt_df.experiment.isin(['TS_6_4', 'TS_5_4', 'TS_69_2'])]\n    \n    sub_gt_df.experiment.unique()\n    \n    print(score(sub_gt_df, pred_df, row_id_column_name='id', distance_multiplier=0.5, beta=4))\n    \n    valid_ids = ['TS_6_4', 'TS_5_4', 'TS_69_2']\n    for vid in valid_ids:\n        print(vid, score(sub_gt_df[sub_gt_df.experiment == vid], pred_df[pred_df.experiment == vid], row_id_column_name='id', distance_multiplier=0.5, beta=4))\n\n    plt.imshow(experiment_probs[0,92,:,:])\n    plt.colorbar()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:25:34.410004Z","iopub.execute_input":"2025-02-05T16:25:34.410446Z","iopub.status.idle":"2025-02-05T16:25:34.767473Z","shell.execute_reply.started":"2025-02-05T16:25:34.410408Z","shell.execute_reply":"2025-02-05T16:25:34.766548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 0.8068723226284283","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:28.164007Z","iopub.status.idle":"2025-02-05T16:04:28.164338Z","shell.execute_reply":"2025-02-05T16:04:28.164189Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 01-29-25\n# TS_6_4 {'apo-ferritin': 0.8999999999999999, 'beta-galactosidase': 0.95, 'ribosome': 0.975, 'thyroglobulin': 0.975, 'virus-like-particle': 0.8749999999999999}\n\n# TS_5_4 {'apo-ferritin': 0.95, 'beta-galactosidase': 0.975, 'ribosome': 0.9249999999999999, 'thyroglobulin': 0.975, 'virus-like-particle': 0.5249999999999996}\n\n# TS_69_2 {'apo-ferritin': 0.95, 'beta-galactosidase': 0.975, 'ribosome': 0.95, 'thyroglobulin': 0.95, 'virus-like-particle': 0.8999999999999999}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:28.165251Z","iopub.status.idle":"2025-02-05T16:04:28.165599Z","shell.execute_reply":"2025-02-05T16:04:28.165435Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 01-28-25\n\nif DEBUG:\n\n    pred_df = pd.concat(dfs)\n    \n    print(pred_df.experiment.unique())\n    \n    sub_gt_df = gt_df[gt_df.experiment.isin(['TS_6_4', 'TS_5_4', 'TS_69_2'])]\n    \n    sub_gt_df.experiment.unique()\n    \n    print(score(sub_gt_df, pred_df, row_id_column_name='id', distance_multiplier=0.5, beta=4))\n    \n    valid_ids = ['TS_6_4', 'TS_5_4', 'TS_69_2']\n    for vid in valid_ids:\n        print(vid, score(sub_gt_df[sub_gt_df.experiment == vid], pred_df[pred_df.experiment == vid], row_id_column_name='id', distance_multiplier=0.5, beta=4))\n\n    plt.imshow(experiment_probs[0,92,:,:])\n    plt.colorbar()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:28.166542Z","iopub.status.idle":"2025-02-05T16:04:28.166868Z","shell.execute_reply":"2025-02-05T16:04:28.166761Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 01282025\n\n# TS_6_4 {'apo-ferritin': 0.8999999999999999, 'beta-galactosidase': 0.9249999999999999, 'ribosome': 0.975, 'thyroglobulin': 0.8999999999999999, 'virus-like-particle': 0.8999999999999999}\n\n# TS_5_4 {'apo-ferritin': 0.975, 'beta-galactosidase': 0.975, 'ribosome': 0.975, 'thyroglobulin': 0.95, 'virus-like-particle': 0.2999999999999994}\n\n# TS_69_2 {'apo-ferritin': 0.7999999999999998, 'beta-galactosidase': 0.8999999999999999, 'ribosome': 0.975, 'thyroglobulin': 0.9249999999999999, 'virus-like-particle': 0.8749999999999999}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:28.167846Z","iopub.status.idle":"2025-02-05T16:04:28.168240Z","shell.execute_reply":"2025-02-05T16:04:28.168070Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if DEBUG:\n\n    pred_df = pd.concat(dfs)\n    \n    print(pred_df.experiment.unique())\n    \n    sub_gt_df = gt_df[gt_df.experiment.isin(['TS_6_4', 'TS_5_4', 'TS_69_2'])]\n    \n    sub_gt_df.experiment.unique()\n    \n    print(score(sub_gt_df, pred_df, row_id_column_name='id', distance_multiplier=0.5, beta=4))\n    \n    valid_ids = ['TS_6_4', 'TS_5_4', 'TS_69_2']\n    for vid in valid_ids:\n        print(vid, score(sub_gt_df[sub_gt_df.experiment == vid], pred_df[pred_df.experiment == vid], row_id_column_name='id', distance_multiplier=0.5, beta=4))\n\n    plt.imshow(experiment_probs[0,92,:,:])\n    plt.colorbar()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:28.169169Z","iopub.status.idle":"2025-02-05T16:04:28.169500Z","shell.execute_reply":"2025-02-05T16:04:28.169330Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# No tta\n# 0.6693667828728712\n# TS_6_4 0.8260117849858227\n# TS_5_4 0.6819642395251888\n# TS_69_2 0.5892248591547758","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:28.170323Z","iopub.status.idle":"2025-02-05T16:04:28.170832Z","shell.execute_reply":"2025-02-05T16:04:28.170620Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# With tta\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:28.171618Z","iopub.status.idle":"2025-02-05T16:04:28.171983Z","shell.execute_reply":"2025-02-05T16:04:28.171870Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df.to_csv('./submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:28.172572Z","iopub.status.idle":"2025-02-05T16:04:28.172888Z","shell.execute_reply":"2025-02-05T16:04:28.172743Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print('Done')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T16:04:28.173542Z","iopub.status.idle":"2025-02-05T16:04:28.173849Z","shell.execute_reply":"2025-02-05T16:04:28.173718Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}