{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":9862305,"sourceType":"datasetVersion","datasetId":6052780},{"sourceId":9867543,"sourceType":"datasetVersion","datasetId":6040935},{"sourceId":218555756,"sourceType":"kernelVersion"},{"sourceId":218685046,"sourceType":"kernelVersion"},{"sourceId":219125142,"sourceType":"kernelVersion"},{"sourceId":219258944,"sourceType":"kernelVersion"},{"sourceId":219270994,"sourceType":"kernelVersion"},{"sourceId":219271883,"sourceType":"kernelVersion"},{"sourceId":219272252,"sourceType":"kernelVersion"},{"sourceId":219371331,"sourceType":"kernelVersion"},{"sourceId":219580016,"sourceType":"kernelVersion"},{"sourceId":219580126,"sourceType":"kernelVersion"},{"sourceId":219729281,"sourceType":"kernelVersion"},{"sourceId":219729422,"sourceType":"kernelVersion"},{"sourceId":220068208,"sourceType":"kernelVersion"},{"sourceId":220390556,"sourceType":"kernelVersion"},{"sourceId":220650147,"sourceType":"kernelVersion"},{"sourceId":220650312,"sourceType":"kernelVersion"},{"sourceId":220772186,"sourceType":"kernelVersion"},{"sourceId":220820654,"sourceType":"kernelVersion"},{"sourceId":222388629,"sourceType":"kernelVersion"},{"sourceId":222388796,"sourceType":"kernelVersion"}],"dockerImageVersionId":30887,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":373.918993,"end_time":"2025-02-04T19:38:27.465221","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-02-04T19:32:13.546228","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"try:\n    import zarr\nexcept: \n    !cp -r '/kaggle/input/hengck-czii-cryo-et-01/wheel_file' '/kaggle/working/'\n    !pip install /kaggle/working/wheel_file/asciitree-0.3.3/asciitree-0.3.3\n    !pip install --no-index --find-links=/kaggle/working/wheel_file zarr\n    !pip install --no-index --find-links=/kaggle/working/wheel_file connected-components-3d\n\n\nfrom typing import List, Tuple, Union\ndeps_path = '/kaggle/input/czii-cryoet-dependencies'\n! pip install -q --no-index --find-links {deps_path} --requirement {deps_path}/requirements.txt\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport json\nimport torch\nimport torch.nn as nn\nimport gc\nimport random\nfrom torch.utils.data import Dataset, DataLoader\nimport os\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom scipy.optimize import linear_sum_assignment\nimport glob\n\nimport zarr\nimport cc3d\n\nprint('PIP INSTALL OK!!!')","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2025-02-14T20:38:56.288790Z","iopub.execute_input":"2025-02-14T20:38:56.289072Z","iopub.status.idle":"2025-02-14T20:39:20.678488Z","shell.execute_reply.started":"2025-02-14T20:38:56.289040Z","shell.execute_reply":"2025-02-14T20:39:20.677522Z"},"papermill":{"duration":47.269918,"end_time":"2025-02-04T19:33:03.396875","exception":false,"start_time":"2025-02-04T19:32:16.126957","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{"papermill":{"duration":0.005534,"end_time":"2025-02-04T19:33:03.408372","exception":false,"start_time":"2025-02-04T19:33:03.402838","status":"completed"},"tags":[]}},{"cell_type":"code","source":"DATA_KAGGLE_DIR = '/kaggle/input/czii-cryo-et-object-identification'\nTRAIN_DIR = f'{DATA_KAGGLE_DIR}/train'\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nscale = 10.012444196428572\n\nEPOCH = 35\n\nblob_factor = 7\n\ndo_xy = False\n\nOBJECT_DICT = {\n    'apo-ferritin': {'label': 1, 'radius': 60/scale}, \n    'beta-galactosidase': {'label': 2, 'radius': 90/scale}, \n    'ribosome': {'label': 3, 'radius': 150/scale}, \n    'thyroglobulin': {'label': 4, 'radius': 130/scale}, \n    'virus-like-particle': {'label': 5, 'radius': 135/scale},\n    #'beta-amylase' : {'label': 6, 'radius':65/scale},\n}\n\nMODE='submit'\n\nif MODE=='local':\n    valid_dir =f'{DATA_KAGGLE_DIR}/train'\n    valid_id = [\"TS_5_4\"]\n    \nif MODE=='submit':\n    valid_dir =f'{DATA_KAGGLE_DIR}/test'\n    valid_id = glob.glob(f'{valid_dir}/static/ExperimentRuns/*')\n    valid_id = [f.split('/')[-1] for f in valid_id]\n\nprint('valid_id:',len(valid_id), valid_id)","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:20.679558Z","iopub.execute_input":"2025-02-14T20:39:20.680162Z","iopub.status.idle":"2025-02-14T20:39:20.780911Z","shell.execute_reply.started":"2025-02-14T20:39:20.680125Z","shell.execute_reply":"2025-02-14T20:39:20.780076Z"},"papermill":{"duration":0.09265,"end_time":"2025-02-04T19:33:03.506675","exception":false,"start_time":"2025-02-04T19:33:03.414025","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_sphere_volume(radius):\n    \"\"\"\n    Calculate the volume of a sphere given its radius.\n\n    Parameters:\n    radius (float): The radius of the sphere.\n\n    Returns:\n    float: The volume of the sphere.\n    \"\"\"\n    if radius < 0:\n        raise ValueError(\"Radius cannot be negative.\")\n    volume = (4 / 3) * np.pi * np.power(radius * 0.8, 3)\n    return volume","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:20.782713Z","iopub.execute_input":"2025-02-14T20:39:20.783031Z","iopub.status.idle":"2025-02-14T20:39:20.800734Z","shell.execute_reply.started":"2025-02-14T20:39:20.782997Z","shell.execute_reply":"2025-02-14T20:39:20.799890Z"},"papermill":{"duration":0.01098,"end_time":"2025-02-04T19:33:03.523691","exception":false,"start_time":"2025-02-04T19:33:03.512711","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for k in OBJECT_DICT.keys():\n    OBJECT_DICT[k][\"blob\"] = calculate_sphere_volume(np.log2(OBJECT_DICT[k][\"radius\"]))/blob_factor","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:20.802092Z","iopub.execute_input":"2025-02-14T20:39:20.802405Z","iopub.status.idle":"2025-02-14T20:39:20.819924Z","shell.execute_reply.started":"2025-02-14T20:39:20.802373Z","shell.execute_reply":"2025-02-14T20:39:20.819043Z"},"papermill":{"duration":0.010597,"end_time":"2025-02-04T19:33:03.540253","exception":false,"start_time":"2025-02-04T19:33:03.529656","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"OBJECT_DICT","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:20.820754Z","iopub.execute_input":"2025-02-14T20:39:20.821042Z","iopub.status.idle":"2025-02-14T20:39:20.838622Z","shell.execute_reply.started":"2025-02-14T20:39:20.821010Z","shell.execute_reply":"2025-02-14T20:39:20.837901Z"},"papermill":{"duration":0.013951,"end_time":"2025-02-04T19:33:03.559922","exception":false,"start_time":"2025-02-04T19:33:03.545971","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def read_one_data(id, static_dir):\n    zarr_dir = f'{static_dir}/{id}/VoxelSpacing10.000'\n    zarr_file = f'{zarr_dir}/denoised.zarr'\n    zarr_data = zarr.open(zarr_file, mode='r')\n    volume = zarr_data[0][:]\n\n    pmin,pmax = -1.1769942479337e-05 , 1.2801160441345688e-05\n    # mean = volume.mean()\n    # std = volume.std()\n    # volume = (volume - mean) / std\n    return ((volume-pmin)/(pmax-pmin)).astype(np.float16)\n\n\ndef read_one_truth(id, overlay_dir):\n    location={}\n    json_dir = f'{overlay_dir}/{id}/Picks'\n    for p in OBJECT_DICT.keys():\n        json_file = f'{json_dir}/{p}.json'\n        with open(json_file, 'r') as f:\n            json_data = json.load(f)\n\n        num_point = len(json_data['points'])\n        loc = [list(json_data['points'][i]['location'].values())  for i in range(num_point)]\n        location[p] = [[coo for coo in coos] for coos in loc ]\n    return location\n\ndef do_one_eval(truth, predict, threshold = 3):\n    P=len(predict)\n    T=len(truth)\n\n    if P==0:\n        hit=[[],[]]\n        miss=np.arange(T).tolist()\n        fp=[]\n        metric = [P,T,len(hit[0]),len(miss),len(fp)]\n        return hit, fp, miss, metric\n\n    if T==0:\n        hit=[[],[]]\n        fp=np.arange(P).tolist()\n        miss=[]\n        metric = [P,T,len(hit[0]),len(miss),len(fp)]\n        return hit, fp, miss, metric\n\n    #---\n    distance = predict.reshape(P,1,3)-truth.reshape(1,T,3)\n    distance = distance**2\n    distance = distance.sum(axis=2)\n    distance = np.sqrt(distance)\n    p_index, t_index = linear_sum_assignment(distance)\n\n    valid = distance[p_index, t_index] <= threshold\n    p_index = p_index[valid]\n    t_index = t_index[valid]\n    hit = [p_index.tolist(), t_index.tolist()]\n    miss = np.arange(T)\n    miss = miss[~np.isin(miss,t_index)].tolist()\n    fp = np.arange(P)\n    fp = fp[~np.isin(fp,p_index)].tolist()\n\n    metric = [P,T,len(hit[0]),len(miss),len(fp)] #for lb metric F-beta copmutation\n    return hit, fp, miss, metric","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:20.839581Z","iopub.execute_input":"2025-02-14T20:39:20.839865Z","iopub.status.idle":"2025-02-14T20:39:20.855947Z","shell.execute_reply.started":"2025-02-14T20:39:20.839840Z","shell.execute_reply":"2025-02-14T20:39:20.855057Z"},"papermill":{"duration":0.016322,"end_time":"2025-02-04T19:33:03.582082","exception":false,"start_time":"2025-02-04T19:33:03.565760","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"patch_size = (184,128,128)\ndef calculate_patch_starts(dimension_size: int, patch_size: int):\n    if dimension_size <= patch_size:\n        return [0]\n        \n\n    n_patches = np.ceil(dimension_size / patch_size)  +1 if dimension_size>300 else 1\n\n    \n    if n_patches == 1:\n        return [0]\n    \n    # Calculate overlap\n    total_overlap = (n_patches * patch_size - dimension_size) / (n_patches - 1)\n    \n    # Generate starting positions\n    positions = []\n    for i in range(int(n_patches)):\n        pos = int(i * (patch_size - total_overlap))\n        if pos + patch_size > dimension_size:\n            pos = dimension_size - patch_size\n        if pos not in positions:  # Avoid duplicates\n            positions.append(pos)\n        \n    return positions\n\n\n    \nclass PredDataset(Dataset):\n    def __init__(self,experiment ,patch_size = patch_size):\n        self.is_local = MODE == \"local\"\n        self.patch_size = patch_size\n        #self.zyx = np.ones((3, 184+patch_size , 630+patch_size, 630+patch_size))*-1\n        pad_size = [patch_size[i]//3 for i in range(3)]\n        #self.zyx [:,pad_size:184+pad_size, pad_size:630+pad_size, pad_size:630+pad_size] = np.indices((184,630,630))\n        \n        self.volume = read_one_data(experiment, static_dir=f'{valid_dir}/static/ExperimentRuns')\n        \n        self.locations = read_one_truth(experiment, overlay_dir=f'{TRAIN_DIR}/overlay/ExperimentRuns') if self.is_local else None\n\n        self.indexes = [[z,y,x] \n                       for z in calculate_patch_starts(184,patch_size[0])\n                       for y in calculate_patch_starts(630, patch_size[1])\n                       for x in calculate_patch_starts(630, patch_size[2])]\n    def __len__(self):\n        return len(self.indexes)\n\n    def __getitem__(self,idx):\n\n        zyx = self.indexes [idx]\n        patch = self.volume[zyx[0]:zyx[0]+self.patch_size[0],zyx[1]:zyx[1]+self.patch_size[1],zyx[2]:zyx[2]+self.patch_size[2]]\n\n        x_flip = np.flip(patch,-1)\n        y_flip = np.flip(patch,-2)\n        z_flip = np.flip(patch,-3)\n        patch = torch.tensor(patch,dtype = torch.float32)\n        rot_1 = torch.rot90(patch , k = 1, dims = (-1,-2))\n        rot_2 = torch.rot90(patch , k = 2, dims = (-1,-2))\n        rot_3 = torch.rot90(patch , k = 3, dims = (-1,-2))\n        \n        return {\n            \"volume\": patch,\n            \"x_flip\":torch.tensor(x_flip.copy(),dtype = torch.float32),\n            \"y_flip\":torch.tensor(y_flip.copy(),dtype = torch.float32),\n            \"z_flip\":torch.tensor(z_flip.copy(),dtype = torch.float32),\n            \"rot_1\":rot_1,\n            \"rot_2\":rot_2,\n            \"rot_3\":rot_3,\n            'zyx':  torch.tensor(zyx,dtype = torch.long)}\n\ndef evaluate_predictions(stats, pred_loader, distance_threshold=3, beta=4 , particle_name = None):\n    best_f_beta = 0\n    best_metric = None\n    # Filter predictions based on voxel count\n    pred = np.array([centroid for i, centroid in enumerate(stats[particle_name][\"centroids\"]) if i != 0 and stats[particle_name][\"voxel_counts\"][i] > OBJECT_DICT[particle_name][\"blob\"]])\n    pred *= scale\n    if len(pred)==0:\n        return {\n            \"truth\": 0,\n            \"predict\": 0,\n            \"hit\": 0,\n            \"fp\": 0,\n            \"miss\": 0,\n            \"f_b\": 0\n        }\n    pred = pred[:,::-1]\n    truth_locations = np.array(pred_loader.dataset.locations[particle_name])\n    # Perform evaluation\n    hit, fp, miss, metric = do_one_eval(truth_locations, pred, distance_threshold)\n\n    # Calculate precision, recall, and F-beta score\n    precision = len(hit[0]) / (len(hit[0]) + len(fp)) if (len(hit[0]) + len(fp)) > 0 else 0\n    recall = len(hit[0]) / (len(hit[0]) + len(miss)) if (len(hit[0]) + len(miss)) > 0 else 0\n\n    beta_squared = beta ** 2\n    f_beta = (1 + beta_squared) * (precision * recall) / (beta_squared * precision + recall) if (precision + recall) > 0 else 0\n    if f_beta>= best_f_beta:\n        best_f_beta = f_beta\n        best_metric = {\n            \"truth\": len(truth_locations),\n            \"predict\": len(pred),\n            \"hit\": len(hit[0]),\n            \"fp\": len(fp),\n            \"miss\": len(miss),\n            \"f_b\": f_beta\n        }\n    # Return results as JSON-like dictionary\n    return best_metric\n    \ndef do_one_eval(truth, predict, threshold = 3):\n    P=len(predict)\n    T=len(truth)\n\n    if P==0:\n        hit=[[],[]]\n        miss=np.arange(T).tolist()\n        fp=[]\n        metric = [P,T,len(hit[0]),len(miss),len(fp)]\n        return hit, fp, miss, metric\n\n    if T==0:\n        hit=[[],[]]\n        fp=np.arange(P).tolist()\n        miss=[]\n        metric = [P,T,len(hit[0]),len(miss),len(fp)]\n        return hit, fp, miss, metric\n\n    #---\n    distance = predict.reshape(P,1,3)-truth.reshape(1,T,3)\n    distance = distance**2\n    distance = distance.sum(axis=2)\n    distance = np.sqrt(distance)\n    p_index, t_index = linear_sum_assignment(distance)\n\n    valid = distance[p_index, t_index] <= threshold\n    p_index = p_index[valid]\n    t_index = t_index[valid]\n    hit = [p_index.tolist(), t_index.tolist()]\n    miss = np.arange(T)\n    miss = miss[~np.isin(miss,t_index)].tolist()\n    fp = np.arange(P)\n    fp = fp[~np.isin(fp,p_index)].tolist()\n\n    metric = [P,T,len(hit[0]),len(miss),len(fp)] #for lb metric F-beta copmutation\n    return hit, fp, miss, metric\n","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:20.856928Z","iopub.execute_input":"2025-02-14T20:39:20.857188Z","iopub.status.idle":"2025-02-14T20:39:20.876483Z","shell.execute_reply.started":"2025-02-14T20:39:20.857161Z","shell.execute_reply":"2025-02-14T20:39:20.875704Z"},"papermill":{"duration":0.024436,"end_time":"2025-02-04T19:33:03.659289","exception":false,"start_time":"2025-02-04T19:33:03.634853","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"maps = {\n    \"cuda:0\":[],\n    \"cuda:1\":[]\n}\nfor i,exp_name in enumerate(valid_id):\n    if i%2 == 0:\n        maps[\"cuda:0\"].append(exp_name)\n    else:\n        maps[\"cuda:1\"].append(exp_name)\nprint(maps)\n\nloaders = {\n    \"cuda:0\":None,\n    \"cuda:1\":None\n}\nprint(loaders)","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:20.879716Z","iopub.execute_input":"2025-02-14T20:39:20.879962Z","iopub.status.idle":"2025-02-14T20:39:20.895188Z","shell.execute_reply.started":"2025-02-14T20:39:20.879943Z","shell.execute_reply":"2025-02-14T20:39:20.894226Z"},"papermill":{"duration":0.012163,"end_time":"2025-02-04T19:33:03.677243","exception":false,"start_time":"2025-02-04T19:33:03.665080","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#maps[\"cuda:0\"] = [valid_id[0] for _ in range(250)]\n#maps[\"cuda:1\"] = [valid_id[0] for _ in range(250)]","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:20.896417Z","iopub.execute_input":"2025-02-14T20:39:20.896604Z","iopub.status.idle":"2025-02-14T20:39:20.911446Z","shell.execute_reply.started":"2025-02-14T20:39:20.896588Z","shell.execute_reply":"2025-02-14T20:39:20.910621Z"},"papermill":{"duration":0.010235,"end_time":"2025-02-04T19:33:03.693403","exception":false,"start_time":"2025-02-04T19:33:03.683168","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.006051,"end_time":"2025-02-04T19:33:03.705443","exception":false,"start_time":"2025-02-04T19:33:03.699392","status":"completed"},"tags":[],"trusted":true},"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)\n\nimport torch.nn as nn\n\nclass ConvBNReLU2D(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel_size=3, stride=1, padding=1):\n        super(ConvBNReLU2D, self).__init__()\n        if kernel_size == 5:\n            padding = 2\n\n        if kernel_size == 7:\n            padding = 3\n            \n        self.block = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size, stride, padding, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.block(x)\n        \nclass ConvBNReLU3D(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel_size=3, stride=1, padding=1):\n        super(ConvBNReLU3D, self).__init__()\n        self.block = nn.Sequential(\n            nn.Conv3d(in_channels, out_channels, kernel_size, stride, padding, bias=False),\n            nn.BatchNorm3d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.block(x)\n            \n\nimport torch.nn.functional as F\n\nclass EncoderBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, skip_channels = 0, num_conv3d=1, \n                 do_up = True , do_down = True ,use_transpose = False):\n        super(EncoderBlock, self).__init__()\n        self.do_up = do_up\n        \n        if self.do_up:\n            self.upsample = lambda x: F.interpolate(x, scale_factor=2, mode='trilinear')\n            \n        if use_transpose:\n            self.upsample = nn.Sequential(\n                nn.ConvTranspose3d(out_channels, out_channels, kernel_size=2, stride=2),\n                nn.BatchNorm3d(out_channels),\n                nn.ReLU(inplace=True)\n            )\n            \n        self.do_down = do_down \n        if self.do_down:\n            self.downsample = lambda x: F.interpolate(x, scale_factor=.5, mode='trilinear')\n\n        \n        self.conv3d_layers = nn.Sequential(\n            *[ConvBNReLU3D(out_channels if i!=0 else in_channels + skip_channels, out_channels, stride= (1, 1, 1)) \n              for i in range(num_conv3d)]\n        )\n\n    def forward(self, x , xskip = None):\n        if xskip is not None:\n            x = torch.cat([x, xskip], dim=1)\n\n        out = self.conv3d_layers(x)\n        output = {\n            \"out\": out,\n            \"up\":None,\n            \"down\":None\n        }\n\n        if self.do_up:\n            output[\"up\"] = self.upsample(out)   \n        if self.do_down:\n            output[\"down\"] = self.downsample(out)   \n        return dotdict(output)\n    \nclass Model(nn.Module):\n    def __init__(self , channels = [28,32,36]):\n        super(Model, self).__init__()\n        self.register_buffer('D', torch.tensor(0))\n        self.output_type = ['particle', 'loss']\n\n        self.norm = nn.BatchNorm3d(1)\n        \n        self.encoder1 = EncoderBlock(in_channels = 1, out_channels = channels[0], num_conv3d=2 , do_up = False, do_down=True)\n        self.encoder2 = EncoderBlock(in_channels = channels[0], out_channels = channels[1], num_conv3d=2 , do_up = True, do_down=True)\n        \n        self.decoder1 = EncoderBlock(in_channels = channels[1], out_channels = channels[2], num_conv3d=4 , do_up = True, do_down=False)\n        self.decoder2 = EncoderBlock(in_channels = channels[2], out_channels = channels[1], num_conv3d=2 , skip_channels = channels[1], do_up = True, do_down=False , use_transpose = True)\n\n        self.pre = EncoderBlock(in_channels = channels[1], out_channels = channels[0], num_conv3d=2 , do_up = False, do_down=False)\n\n        self.mask = nn.Conv3d(channels[0], 6, 1, 1, bias=False)\n        \n        \n\n    def forward(self,batch):\n        device = self.D.device\n        volume = batch[\"volume\"].to(device).unsqueeze(1)\n\n        input_ = self.norm(volume)\n        \n        encode1 = self.encoder1(input_)\n        encode2 = self.encoder2(encode1.down)\n        \n        decode1 = self.decoder1(encode2.down)\n        #print(encode2.out.shape , decode1.up.shape)\n        decode2 = self.decoder2(encode2.out , decode1.up)\n\n        pre = self.pre(decode2.up)\n\n        logit = self.mask(pre.out)\n        #print(mask.shape)\n\n        output = {}\n        \n        if \"loss\" in self.output_type and \"label\" in batch.keys():\n        \n            # Apply weighted cross-entropy loss\n            output[\"loss\"] = F.cross_entropy(\n                logit, \n                batch['label'].to(device), \n                label_smoothing=0.01,\n            )\n\n        if \"particle\" in self.output_type:\n            output['particle'] = F.softmax(logit,1)\n            \n        return output","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:20.912371Z","iopub.execute_input":"2025-02-14T20:39:20.912663Z","iopub.status.idle":"2025-02-14T20:39:20.937556Z","shell.execute_reply.started":"2025-02-14T20:39:20.912620Z","shell.execute_reply":"2025-02-14T20:39:20.936542Z"},"papermill":{"duration":0.022,"end_time":"2025-02-04T19:33:03.733387","exception":false,"start_time":"2025-02-04T19:33:03.711387","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AVGModel(nn.Module):\n    def __init__(self, models):\n        super(AVGModel, self).__init__()\n        self.models = nn.ModuleList(models)\n\n    def forward(self,batch):\n        output = {\"particle\": 0}\n        volume = batch[\"volume\"].to(device)\n        b = len(volume)\n        z_flip = batch[\"z_flip\"].to(device)\n        y_flip = batch[\"y_flip\"].to(device)\n        x_flip = batch[\"x_flip\"].to(device)\n        rot_1 = batch[\"rot_1\"].to(device)\n        rot_2 = batch[\"rot_2\"].to(device)\n        rot_3 = batch[\"rot_3\"].to(device)\n        \n        batch[\"volume\"] = torch.cat([volume ,x_flip , y_flip ,z_flip , rot_1, rot_2, rot_3] , 0)\n        all_ = 0\n        for model in self.models:\n            all_ += model(batch)[\"particle\"]\n        all_ /= len(self.models)\n        \n        for i in range(7):\n            if i ==0:\n                output['particle'] += all_[b*i:b*(i+1)]\n            elif i<4 :\n                output['particle'] += torch.flip(all_[b*i:b*(i+1)], dims = [-i])\n            else :\n                rot = i-3\n                output['particle'] += torch.rot90(all_[b*i:b*(i+1)], k = rot, dims = (-2,-1))\n        output['particle'] /= 7\n        return output","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:20.938532Z","iopub.execute_input":"2025-02-14T20:39:20.938789Z","iopub.status.idle":"2025-02-14T20:39:20.958321Z","shell.execute_reply.started":"2025-02-14T20:39:20.938768Z","shell.execute_reply":"2025-02-14T20:39:20.957543Z"},"papermill":{"duration":0.013645,"end_time":"2025-02-04T19:33:03.753104","exception":false,"start_time":"2025-02-04T19:33:03.739459","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"paths = [\n    \"/kaggle/input/deepfinder-120-seed/model_all_35.bin\",\n    \"/kaggle/input/deepfinder-128/model_all_35.bin\",\n    \"/kaggle/input/deep-finder-128-80-seed/model_all_35.bin\",\n    \"/kaggle/input/deep-finder-128-90/model_all_35.bin\"\n]\n\npaths = [\n    #\"/kaggle/input/deepfinder-111-seed\",\n    \"/kaggle/input/train-script-42-90/model_all_35_42.bin\",\n    \"/kaggle/input/train-script-42-90/model_all_35_90.bin\",\n    \"/kaggle/input/train-script-80-120/model_all_35_120.bin\",\n    \"/kaggle/input/train-script-80-120/model_all_35_80.bin\"\n]\n\n\n\ndef get_deep(path):\n    model = Model()\n\n    state_dict = torch.load(path,weights_only=True , map_location=\"cpu\")\n    model.load_state_dict(state_dict)\n    model.eval()\n    \n    return model\n\n\n\nmodel = {\n    \"cuda:0\": AVGModel([get_deep(path).to(\"cuda:0\") for path in paths]),\n    \"cuda:1\": AVGModel([get_deep(path).to(\"cuda:1\") for path in paths]),\n}\n\nweight = torch.zeros((patch_size[0], patch_size[1], patch_size[2]) , dtype =torch.float16)\nweight[8:patch_size[0]-8, 8:patch_size[1]-8, 8:patch_size[2]-8] += 1\nweight += .1\n\nweights = {\n    \"cuda:0\": weight.to(\"cuda:0\"),\n    \"cuda:1\": weight.to(\"cuda:1\")\n}\n\nlogits = {\n    \"cuda:0\": torch.zeros((6, 184, 630, 630) , dtype =torch.float16).to(\"cuda:0\"),\n    \"cuda:1\": torch.zeros((6, 184, 630, 630 ) , dtype =torch.float16).to(\"cuda:1\")\n}\n\ncount = {\n    \"cuda:0\": torch.zeros((184 , 630 , 630 ) , dtype =torch.float16).to(\"cuda:0\"),\n    \"cuda:1\": torch.zeros((184 , 630 , 630 ) , dtype =torch.float16).to(\"cuda:1\")\n}","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:20.959110Z","iopub.execute_input":"2025-02-14T20:39:20.959324Z","iopub.status.idle":"2025-02-14T20:39:23.176196Z","shell.execute_reply.started":"2025-02-14T20:39:20.959307Z","shell.execute_reply":"2025-02-14T20:39:23.175243Z"},"papermill":{"duration":2.003155,"end_time":"2025-02-04T19:33:05.762157","exception":false,"start_time":"2025-02-04T19:33:03.759002","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#stats,evals = get_stats(probs,loaders[\"cuda:0\"])\nbest_eval = {\n    'apo-ferritin': {'thresh': 0.0575000000000000005},\n    'beta-galactosidase': {'thresh': 0.050000000000004},\n    'ribosome': {'thresh': 0.0575000000000000005},\n    'thyroglobulin':{'thresh': 0.05000000000000},\n    'virus-like-particle': {'thresh': 0.125}\n}\n\ndef get_probs(exp_name, device):\n    global loaders, weight, logits, count\n    #set the loader\n    del loaders[device]\n    gc.collect()\n    loaders[device] = DataLoader(PredDataset(experiment = exp_name ),batch_size=1,shuffle=False,num_workers=2 )\n    logits[device] = logits[device].zero_()\n    count[device] = count[device].zero_()\n\n    with torch.no_grad():\n        with torch.amp.autocast(device):\n            for batch in tqdm(loaders[device]):\n                local_logits = model[device](batch)[\"particle\"]\n                for i, logits_patch in enumerate(local_logits):\n                    \n                    z, y, x = batch[\"zyx\"][i]\n                    z_slice = slice(z, z + patch_size[0])\n                    y_slice = slice(y, y + patch_size[1])\n                    x_slice = slice(x, x + patch_size[2])\n                    \n                    count[device][z_slice, y_slice, x_slice] += weights[device]\n                    logits[device][:, z_slice, y_slice, x_slice] += logits_patch * weights[device]\n\n\n            probs = (logits[device]/count[device]).detach().cpu().numpy()\n    return probs\n\ndef get_stats(probs,pred_loader,search_thresh = False):\n    stats = {}\n    evals = {}\n    for particle in OBJECT_DICT.keys():\n        label = OBJECT_DICT[particle][\"label\"]\n        if MODE == \"local\" and search_thresh:\n            for prob_thresh in np.arange(0.05, 0.1, 0.01):\n                thresh = OBJECT_DICT[particle][\"radius\"]/2*scale\n                labels_out = cc3d.connected_components(probs[label, :, :, :] > prob_thresh, connectivity=18)\n                stats[particle] = cc3d.statistics(labels_out)\n                \n                eval_ = evaluate_predictions(stats, pred_loader, distance_threshold = thresh, particle_name = particle)\n                if particle not in evals.keys():\n                    evals[particle] = eval_\n                    evals[particle][\"thresh\"] = prob_thresh\n                elif eval_[\"f_b\"]>=evals[particle][\"f_b\"]:\n                    evals[particle] = eval_\n                    evals[particle][\"thresh\"] = prob_thresh\n\n        elif MODE == \"local\":\n            \n            thresh = OBJECT_DICT[particle][\"radius\"]/2*scale\n            labels_out = cc3d.connected_components(probs[label, :, :, :] > best_eval[particle][\"thresh\"], connectivity=18)\n            stats[particle] = cc3d.statistics(labels_out)\n            \n            evals[particle] = evaluate_predictions(stats, pred_loader, distance_threshold = thresh, particle_name = particle)\n        \n        else :\n            labels_out = cc3d.connected_components(probs[label, :, :, :] > best_eval[particle][\"thresh\"], connectivity=18)\n            stats[particle] = cc3d.statistics(labels_out)\n            \n    return stats,evals\n\ndef stats_to_df(stats,exp_name):\n    result = pd.DataFrame(columns=[\"x\",\"y\",\"z\",\"particle_type\",\"experiment\"])\n    for particle_name in OBJECT_DICT.keys():\n        pred = np.array([centroid for i, centroid in enumerate(stats[particle_name][\"centroids\"]) if i != 0 and stats[particle_name][\"voxel_counts\"][i] > OBJECT_DICT[particle_name][\"blob\"]])\n        if len(pred)==0:\n            continue\n\n\n        pred *= scale\n        pred = pred[:,::-1]\n        df = pd.DataFrame(pred, columns=[\"x\",\"y\",\"z\"])\n        df[\"experiment\"] = [exp_name for _ in range(len(df))]\n        df[\"particle_type\"] = [particle_name for _ in range(len(df))]\n        result = pd.concat([result,df])\n    return result","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:23.177193Z","iopub.execute_input":"2025-02-14T20:39:23.177509Z","iopub.status.idle":"2025-02-14T20:39:23.189990Z","shell.execute_reply.started":"2025-02-14T20:39:23.177477Z","shell.execute_reply":"2025-02-14T20:39:23.189244Z"},"papermill":{"duration":0.022492,"end_time":"2025-02-04T19:33:05.797195","exception":false,"start_time":"2025-02-04T19:33:05.774703","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.006324,"end_time":"2025-02-04T19:33:05.810048","exception":false,"start_time":"2025-02-04T19:33:05.803724","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if MODE == \"local\":\n    for k in OBJECT_DICT.keys():\n        OBJECT_DICT[k][\"blob\"] = calculate_sphere_volume(np.log2(OBJECT_DICT[k][\"radius\"]))/100\n    probs = get_probs(valid_id[0],\"cuda:0\")\n    s,e = get_stats(probs, loaders[\"cuda:0\"])\n    fb = 0\n    s,e = get_stats(probs, loaders[\"cuda:0\"])\n    print(e)\n    for p in e.keys():\n        if p == \"beta-galactosidase\" or p == \"thyroglobulin\":\n            fb+= 2*e[p][\"f_b\"]\n    \n        else :\n            fb+= e[p][\"f_b\"]\n    \n        print(fb)\n    \n    print(fb/7)","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:23.190742Z","iopub.execute_input":"2025-02-14T20:39:23.190939Z","iopub.status.idle":"2025-02-14T20:39:23.211022Z","shell.execute_reply.started":"2025-02-14T20:39:23.190922Z","shell.execute_reply":"2025-02-14T20:39:23.210208Z"},"papermill":{"duration":0.012233,"end_time":"2025-02-04T19:33:05.828161","exception":false,"start_time":"2025-02-04T19:33:05.815928","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"0.7933213556664216","metadata":{"papermill":{"duration":0.005708,"end_time":"2025-02-04T19:33:05.839859","exception":false,"start_time":"2025-02-04T19:33:05.834151","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if MODE == \"local\":\n    fb = 0\n    for k in OBJECT_DICT.keys():\n        OBJECT_DICT[k][\"blob\"] = calculate_sphere_volume(np.log2(OBJECT_DICT[k][\"radius\"]))/3\n    s,e = get_stats(probs, loaders[\"cuda:0\"])\n    print(e)\n    for p in e.keys():\n        if p == \"beta-galactosidase\" or p == \"thyroglobulin\":\n            fb+= 2*e[p][\"f_b\"]\n    \n        else :\n            fb+= e[p][\"f_b\"]\n    \n        print(fb)\n    \n    print(fb/7)","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:23.211853Z","iopub.execute_input":"2025-02-14T20:39:23.212129Z","iopub.status.idle":"2025-02-14T20:39:23.230157Z","shell.execute_reply.started":"2025-02-14T20:39:23.212107Z","shell.execute_reply":"2025-02-14T20:39:23.229402Z"},"papermill":{"duration":0.011574,"end_time":"2025-02-04T19:33:05.857329","exception":false,"start_time":"2025-02-04T19:33:05.845755","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dfs = {\n    \"cuda:0\":pd.DataFrame(columns=[\"x\",\"y\",\"z\",\"particle_type\",\"experiment\"]),\n    \"cuda:1\":pd.DataFrame(columns=[\"x\",\"y\",\"z\",\"particle_type\",\"experiment\"])\n}\n\ndef predict(exp_name,device, search_thresh=False):\n    probs = get_probs(exp_name,device)\n    stats,evals = get_stats(probs , loaders[device], search_thresh)\n    if MODE == \"local\":\n        print(exp_name,evals)\n    prediction_df = stats_to_df(stats,exp_name)\n    if prediction_df is not None:\n        dfs[device] = pd.concat([dfs[device],prediction_df])\n\ndef predict_all(device):\n    \n    print(\"predicting\" , device)\n    for exp_name in maps[device] :\n        predict(exp_name,device)","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:23.230935Z","iopub.execute_input":"2025-02-14T20:39:23.231172Z","iopub.status.idle":"2025-02-14T20:39:23.251119Z","shell.execute_reply.started":"2025-02-14T20:39:23.231128Z","shell.execute_reply":"2025-02-14T20:39:23.250288Z"},"papermill":{"duration":0.017069,"end_time":"2025-02-04T19:33:05.880343","exception":false,"start_time":"2025-02-04T19:33:05.863274","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from concurrent.futures import ThreadPoolExecutor\n\n# Initialize DataFrames\ndfs = {\n    \"cuda:0\": pd.DataFrame(columns=[\"x\", \"y\", \"z\", \"particle_type\", \"experiment\"]),\n    \"cuda:1\": pd.DataFrame(columns=[\"x\", \"y\", \"z\", \"particle_type\", \"experiment\"]),\n}\n\nwith ThreadPoolExecutor(max_workers=2) as executor:\n    futures = [\n        executor.submit(predict_all, \"cuda:0\"),\n        executor.submit(predict_all, \"cuda:1\"),\n    ]\n\n# Wait for all threads to finish\nfor future in futures:\n    future.result()  # Ensures exceptions in threads are raised\n\nprint(\"Processing complete!\")","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:39:23.251987Z","iopub.execute_input":"2025-02-14T20:39:23.252206Z","iopub.status.idle":"2025-02-14T20:44:54.431220Z","shell.execute_reply.started":"2025-02-14T20:39:23.252187Z","shell.execute_reply":"2025-02-14T20:44:54.430067Z"},"papermill":{"duration":318.043189,"end_time":"2025-02-04T19:38:23.929459","exception":false,"start_time":"2025-02-04T19:33:05.886270","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submit_df = pd.concat([dfs[\"cuda:0\"],dfs[\"cuda:1\"]])\nsubmit_df[\"id\"] = [i for i in range(len(submit_df))]","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:44:54.432373Z","iopub.execute_input":"2025-02-14T20:44:54.432806Z","iopub.status.idle":"2025-02-14T20:44:54.438791Z","shell.execute_reply.started":"2025-02-14T20:44:54.432765Z","shell.execute_reply":"2025-02-14T20:44:54.437896Z"},"papermill":{"duration":0.028512,"end_time":"2025-02-04T19:38:23.979423","exception":false,"start_time":"2025-02-04T19:38:23.950911","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submit_df.to_csv(\"submission.csv\",index=False)","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:44:54.439701Z","iopub.execute_input":"2025-02-14T20:44:54.440077Z","iopub.status.idle":"2025-02-14T20:44:54.472318Z","shell.execute_reply.started":"2025-02-14T20:44:54.440045Z","shell.execute_reply":"2025-02-14T20:44:54.471422Z"},"papermill":{"duration":0.037402,"end_time":"2025-02-04T19:38:24.038140","exception":false,"start_time":"2025-02-04T19:38:24.000738","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submit_df","metadata":{"execution":{"iopub.status.busy":"2025-02-14T20:44:54.473423Z","iopub.execute_input":"2025-02-14T20:44:54.473764Z","iopub.status.idle":"2025-02-14T20:44:54.499223Z","shell.execute_reply.started":"2025-02-14T20:44:54.473735Z","shell.execute_reply":"2025-02-14T20:44:54.498240Z"},"papermill":{"duration":0.041142,"end_time":"2025-02-04T19:38:24.100449","exception":false,"start_time":"2025-02-04T19:38:24.059307","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.022161,"end_time":"2025-02-04T19:38:24.144342","exception":false,"start_time":"2025-02-04T19:38:24.122181","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null}]}