{"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":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":220875180,"sourceType":"kernelVersion"}],"dockerImageVersionId":30840,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"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\nfrom IPython.display import display, clear_output\nfrom torch.optim.lr_scheduler import ExponentialLR\nimport glob\ntry :\n    import zarr\n    import monai\n    import cc3d\n    from monai.networks.blocks import MaxAvgPool\nexcept :\n    !pip install zarr\n    !pip install monai\n    !pip install segmentation_models_pytorch\n    !pip install --no-index --find-links=/kaggle/input/hengck-czii-cryo-et-01/wheel_file connected-components-3d\n    import zarr\n    import monai\n    import cc3d\n    from monai.networks.blocks import MaxAvgPool\n\nprint('PIP INSTALL OK!!!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:07:35.701139Z","iopub.execute_input":"2025-02-05T18:07:35.701439Z","iopub.status.idle":"2025-02-05T18:08:22.280522Z","shell.execute_reply.started":"2025-02-05T18:07:35.701404Z","shell.execute_reply":"2025-02-05T18:08:22.279696Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"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 = 2\n\nblob_factor = 6\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:08:22.281750Z","iopub.execute_input":"2025-02-05T18:08:22.282605Z","iopub.status.idle":"2025-02-05T18:08:22.371031Z","shell.execute_reply.started":"2025-02-05T18:08:22.282580Z","shell.execute_reply":"2025-02-05T18:08:22.370210Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:08:23.928796Z","iopub.execute_input":"2025-02-05T18:08:23.929239Z","iopub.status.idle":"2025-02-05T18:08:23.933511Z","shell.execute_reply.started":"2025-02-05T18:08:23.929210Z","shell.execute_reply":"2025-02-05T18:08:23.932826Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:08:24.973396Z","iopub.execute_input":"2025-02-05T18:08:24.973695Z","iopub.status.idle":"2025-02-05T18:08:24.977566Z","shell.execute_reply.started":"2025-02-05T18:08:24.973675Z","shell.execute_reply":"2025-02-05T18:08:24.976684Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"OBJECT_DICT","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:08:25.237996Z","iopub.execute_input":"2025-02-05T18:08:25.238270Z","iopub.status.idle":"2025-02-05T18:08:25.244230Z","shell.execute_reply.started":"2025-02-05T18:08:25.238248Z","shell.execute_reply":"2025-02-05T18:08:25.243339Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:08:26.083313Z","iopub.execute_input":"2025-02-05T18:08:26.083679Z","iopub.status.idle":"2025-02-05T18:08:26.094065Z","shell.execute_reply.started":"2025-02-05T18:08:26.083653Z","shell.execute_reply":"2025-02-05T18:08:26.093135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"6*6*3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:08:27.439978Z","iopub.execute_input":"2025-02-05T18:08:27.440320Z","iopub.status.idle":"2025-02-05T18:08:27.445005Z","shell.execute_reply.started":"2025-02-05T18:08:27.440295Z","shell.execute_reply":"2025-02-05T18:08:27.444134Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"patch_size = (128,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 np.ceil(dimension_size / patch_size) \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        \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            '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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:08:27.653188Z","iopub.execute_input":"2025-02-05T18:08:27.653399Z","iopub.status.idle":"2025-02-05T18:08:27.670618Z","shell.execute_reply.started":"2025-02-05T18:08:27.653381Z","shell.execute_reply":"2025-02-05T18:08:27.669643Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:08:29.784895Z","iopub.execute_input":"2025-02-05T18:08:29.785243Z","iopub.status.idle":"2025-02-05T18:08:29.791379Z","shell.execute_reply.started":"2025-02-05T18:08:29.785215Z","shell.execute_reply":"2025-02-05T18:08:29.790442Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:08:30.270962Z","iopub.execute_input":"2025-02-05T18:08:30.271381Z","iopub.status.idle":"2025-02-05T18:08:30.275102Z","shell.execute_reply.started":"2025-02-05T18:08:30.271349Z","shell.execute_reply":"2025-02-05T18:08:30.274092Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:08:31.000339Z","iopub.execute_input":"2025-02-05T18:08:31.000636Z","iopub.status.idle":"2025-02-05T18:08:31.016399Z","shell.execute_reply.started":"2025-02-05T18:08:31.000614Z","shell.execute_reply":"2025-02-05T18:08:31.015144Z"}},"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        \n        batch[\"volume\"] = torch.cat([volume ,x_flip , y_flip ,z_flip , rot_1] , 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(5):\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'] /= 5\n        return output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:08:33.112781Z","iopub.execute_input":"2025-02-05T18:08:33.113242Z","iopub.status.idle":"2025-02-05T18:08:33.123862Z","shell.execute_reply.started":"2025-02-05T18:08:33.113201Z","shell.execute_reply":"2025-02-05T18:08:33.122574Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"paths = [\n    \"/kaggle/input/deepfinder-120-seed\",\n    \"/kaggle/input/deepfinder-111-seed\",\n    \"/kaggle/input/deepfinder-128\",\n    \"/kaggle/input/deep-finder-128-90\",\n    \"/kaggle/input/deep-finder-128-80-seed\"\n]\n\ndef get_deep(fold,epoch ,path):\n    model = Model()\n\n    state_dict = torch.load(f\"\"\"{path}/model_{fold}_{epoch}.bin\"\"\",weights_only=True , map_location=\"cpu\")\n    model.load_state_dict(state_dict)\n    model.eval()\n    \n    return model\n\ndef get_fintuned(epoch,seed):\n    model = Model()\n\n    state_dict = torch.load(f\"\"\"/kaggle/input/finetuning/model_all_{epoch}_{seed}.bin\"\"\",weights_only=True , map_location=\"cpu\")\n    model.load_state_dict(state_dict)\n    model.eval()\n    \n    return model\nseeds = [111,120,42,80,90]\n\nmodel = {\n    \"cuda:0\": AVGModel([get_fintuned(EPOCH,seed).to(\"cuda:0\") for seed in seeds]),\n    \"cuda:1\": AVGModel([get_fintuned(EPOCH,seed).to(\"cuda:1\") for seed in seeds]),\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:08:50.999591Z","iopub.execute_input":"2025-02-05T18:08:50.999891Z","iopub.status.idle":"2025-02-05T18:08:53.296283Z","shell.execute_reply.started":"2025-02-05T18:08:50.999869Z","shell.execute_reply":"2025-02-05T18:08:53.295332Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#stats,evals = get_stats(probs,loaders[\"cuda:0\"])\nbest_eval = {'apo-ferritin': {'truth': 46,\n   'predict': 63,\n   'hit': 41,\n   'fp': 22,\n   'miss': 5,\n   'f_b': 0.8723404255319148,\n   'thresh': 0.0575000000000000005},\n  'beta-galactosidase': {'truth': 12,\n   'predict': 16,\n   'hit': 8,\n   'fp': 8,\n   'miss': 4,\n   'f_b': 0.6538461538461539,\n   'thresh': 0.04750000000000004},\n  'ribosome': {'truth': 31,\n   'predict': 49,\n   'hit': 29,\n   'fp': 20,\n   'miss': 2,\n   'f_b': 0.9045871559633026,\n   'thresh': 0.0575000000000000005},\n  'thyroglobulin': {'truth': 30,\n   'predict': 97,\n   'hit': 22,\n   'fp': 75,\n   'miss': 8,\n   'f_b': 0.551802426343154,\n   'thresh': 0.0475},\n  'virus-like-particle': {'truth': 11,\n   'predict': 11,\n   'hit': 11,\n   'fp': 0,\n   'miss': 0,\n   'f_b': 1.0,\n   'thresh': 0.125}}\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:08:53.297342Z","iopub.execute_input":"2025-02-05T18:08:53.297705Z","iopub.status.idle":"2025-02-05T18:08:53.311745Z","shell.execute_reply.started":"2025-02-05T18:08:53.297667Z","shell.execute_reply":"2025-02-05T18:08:53.310883Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-05T18:08:53.699638Z","iopub.execute_input":"2025-02-05T18:08:53.699960Z","iopub.status.idle":"2025-02-05T18:11:21.725122Z","shell.execute_reply.started":"2025-02-05T18:08:53.699935Z","shell.execute_reply":"2025-02-05T18:11:21.724200Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":".7927742226431762","metadata":{}},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-03T01:26:37.065265Z","iopub.execute_input":"2025-02-03T01:26:37.065611Z","iopub.status.idle":"2025-02-03T01:26:42.743077Z","shell.execute_reply.started":"2025-02-03T01:26:37.065576Z","shell.execute_reply":"2025-02-03T01:26:42.742265Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-03T01:17:23.934001Z","iopub.execute_input":"2025-02-03T01:17:23.934272Z","iopub.status.idle":"2025-02-03T01:17:23.950583Z","shell.execute_reply.started":"2025-02-03T01:17:23.934252Z","shell.execute_reply":"2025-02-03T01:17:23.949866Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-03T01:17:23.951546Z","iopub.execute_input":"2025-02-03T01:17:23.951865Z","iopub.status.idle":"2025-02-03T01:19:50.484559Z","shell.execute_reply.started":"2025-02-03T01:17:23.951833Z","shell.execute_reply":"2025-02-03T01:19:50.483476Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-03T01:19:50.48582Z","iopub.execute_input":"2025-02-03T01:19:50.486154Z","iopub.status.idle":"2025-02-03T01:19:50.491796Z","shell.execute_reply.started":"2025-02-03T01:19:50.486117Z","shell.execute_reply":"2025-02-03T01:19:50.490916Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submit_df.to_csv(\"submission.csv\",index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-03T01:19:50.492458Z","iopub.execute_input":"2025-02-03T01:19:50.49277Z","iopub.status.idle":"2025-02-03T01:19:50.519462Z","shell.execute_reply.started":"2025-02-03T01:19:50.492748Z","shell.execute_reply":"2025-02-03T01:19:50.518588Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submit_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-03T01:19:50.5203Z","iopub.execute_input":"2025-02-03T01:19:50.520512Z","iopub.status.idle":"2025-02-03T01:19:50.541055Z","shell.execute_reply.started":"2025-02-03T01:19:50.520493Z","shell.execute_reply":"2025-02-03T01:19:50.540341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}