{"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":10493780,"sourceType":"datasetVersion","datasetId":6497225},{"sourceId":10534032,"sourceType":"datasetVersion","datasetId":6517838},{"sourceId":10625614,"sourceType":"datasetVersion","datasetId":6531827}],"dockerImageVersionId":30839,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"class CONFIG:\n    DEPS_PATH = '/kaggle/input/cziidependencies'\n    TRAIN_DATA_DIR=\"/kaggle/input/cziinumpy-dataset-exp\"\n    TEST_DATA_DIR=\"/kaggle/input/czii-cryo-et-object-identification/test/static\"\n    MODEL_DIR=\"/kaggle/input/cziiunet-light/model_weights1000_epoches.pth\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-01-28T08:35:54.852131Z","iopub.execute_input":"2025-01-28T08:35:54.852470Z","iopub.status.idle":"2025-01-28T08:35:54.856854Z","shell.execute_reply.started":"2025-01-28T08:35:54.852446Z","shell.execute_reply":"2025-01-28T08:35:54.855862Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! cp -r /kaggle/input/cziidependencies/asciitree-0.3.3/ asciitree-0.3.3/\n! pip wheel asciitree-0.3.3/asciitree-0.3.3/\n!pip install asciitree-0.3.3-py3-none-any.whl\n! pip install -q --no-index --find-links {CONFIG.DEPS_PATH} --requirement {CONFIG.DEPS_PATH}/requirements.txt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T08:35:54.858206Z","iopub.execute_input":"2025-01-28T08:35:54.858448Z","iopub.status.idle":"2025-01-28T08:36:08.249206Z","shell.execute_reply.started":"2025-01-28T08:35:54.858429Z","shell.execute_reply":"2025-01-28T08:36:08.248301Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\nimport numpy as np\nimport pandas as pd\nfrom typing import List, Tuple, Union\n\nimport torch\nfrom monai.data import DataLoader, Dataset, CacheDataset, decollate_batch\nfrom monai.transforms import(\n    Compose,\n    EnsureChannelFirstd,\n    Orientationd,\n    AsDiscrete,\n    RandFlipd,\n    RandRotate90d,\n    NormalizeIntensityd,\n    RandCropByLabelClassesd,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T08:36:08.251103Z","iopub.execute_input":"2025-01-28T08:36:08.251367Z","iopub.status.idle":"2025-01-28T08:36:33.149352Z","shell.execute_reply.started":"2025-01-28T08:36:08.251344Z","shell.execute_reply":"2025-01-28T08:36:33.148667Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_patch_starts(dimension_size,patch_size):\n    if dimension_size <= patch_size:\n        return[0]\n\n    # number og patch\n    n_patches = np.ceil(dimension_size/patch_size)\n    if n_patches ==1:\n        return[0]\n\n    total_overlap = (n_patches*patch_size - dimension_size)/(n_patches-1)\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:\n            positions.append(pos)\n    return positions\n\n\ndef extract_3d_patches_minimal_overlap(arrays,patch_size):\n    if not arrays or not isinstance(arrays,list):\n        raise ValueError(\"Input must be a non-empty list of arrays\")\n\n    shape = arrays[0].shape\n    if not all(arr.shape == shape for arr in arrays):\n        raise ValueError(\"All input arrays must have yhe same shape\")\n\n    if patch_size> min(shape):\n        raise ValueError(f\"patch_size({patch_size}) must be smaller than smallest dimension {min(shape)}\")\n\n    m,n,l = shape\n    patches= []\n    coordinates = []\n\n    x_starts = calculate_patch_starts(m,patch_size)\n    y_starts = calculate_patch_starts(n,patch_size)\n    z_starts = calculate_patch_starts(l,patch_size)\n\n    for arr in arrays:\n        for x in x_starts:\n            for y in y_starts:\n                for z in z_starts:\n                    patch=arr[\n                    x:x+patch_size,\n                    y:y+patch_size,\n                    z:z+patch_size\n                    ]\n                    patches.append(patch)\n                    coordinates.append(((x,y,z)))\n    return patches, coordinates\n\ndef reconstruct_array(patches, coordinates,original_shape):\n    reconstructed= np.zeros(original_shape,dtype = np.int64)\n    patch_size = patches[0].shape[0]\n    for patch,(x,y,z) in zip(patches,coordinates):\n        reconstructed[\n            x:x+patch_size,\n            y:y+patch_size,\n            z:z+patch_size\n        ]= patch\n    return reconstructed\n\ndef dict_to_df(coor_dict, experiment_name):\n    all_coords = []\n    all_labels = []\n    for label, coords in coor_dict.items():\n        all_coords.append(coords)\n        all_labels.extend([label]*len(coords))\n\n    all_coords = np.vstack(all_coords)\n\n    df = pd.DataFrame({\n        'experiment': experiment_name,\n        'particle_type': all_labels,\n        'x':all_coords[:,0],\n        'y':all_coords[:,1],\n        'z':all_coords[:,2]\n    })\n    \n    return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T08:36:33.150698Z","iopub.execute_input":"2025-01-28T08:36:33.151510Z","iopub.status.idle":"2025-01-28T08:36:33.160279Z","shell.execute_reply.started":"2025-01-28T08:36:33.151486Z","shell.execute_reply":"2025-01-28T08:36:33.159365Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"copick_config_path = CONFIG.TRAIN_DATA_DIR + \"/copick.config\"\n\nwith open(copick_config_path,'r') as f:\n    copick_config = json.load(f)\n\ncopick_config['static_root'] = CONFIG.TEST_DATA_DIR\n\ncopick_test_config_path = 'copick_test.config'\n\nwith open(copick_test_config_path,'w') as outfile:\n    json.dump(copick_config,outfile)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T08:36:33.161218Z","iopub.execute_input":"2025-01-28T08:36:33.161453Z","iopub.status.idle":"2025-01-28T08:36:33.458988Z","shell.execute_reply.started":"2025-01-28T08:36:33.161422Z","shell.execute_reply":"2025-01-28T08:36:33.458117Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import copick\n\nroot = copick.from_file(copick_test_config_path)\ncopick_user_name = \"copickUtils\"\ncopick_segmentation_name = \"paintedPicks\"\nvoxel_size = 10\ntomo_type = \"denoised\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T08:36:33.459847Z","iopub.execute_input":"2025-01-28T08:36:33.460160Z","iopub.status.idle":"2025-01-28T08:36:35.010404Z","shell.execute_reply.started":"2025-01-28T08:36:33.460138Z","shell.execute_reply":"2025-01-28T08:36:35.009506Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Non-random transform\n\ninference_transforms = Compose([\n    EnsureChannelFirstd(keys=[\"image\"], channel_dim = \"no_channel\"),\n    NormalizeIntensityd(keys = 'image'),\n    Orientationd(keys = [\"image\"], axcodes=\"RAS\")\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T08:36:35.011277Z","iopub.execute_input":"2025-01-28T08:36:35.012206Z","iopub.status.idle":"2025-01-28T08:36:35.016358Z","shell.execute_reply.started":"2025-01-28T08:36:35.012181Z","shell.execute_reply":"2025-01-28T08:36:35.015392Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cc3d\n\nid_to_name = {\n    1: \"apo-ferritin\", \n    2: \"beta-amylase\",\n    3: \"beta-galactosidase\", \n    4: \"ribosome\", \n    5: \"thyroglobulin\", \n    6: \"virus-like-particle\"\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T08:36:35.017382Z","iopub.execute_input":"2025-01-28T08:36:35.017720Z","iopub.status.idle":"2025-01-28T08:36:35.038184Z","shell.execute_reply.started":"2025-01-28T08:36:35.017684Z","shell.execute_reply":"2025-01-28T08:36:35.037344Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pytorch_lightning as pl\nfrom monai.networks.nets import UNet\nfrom monai.losses import TverskyLoss\nfrom monai.metrics import DiceMetric\nfrom pytorch_lightning.callbacks import ModelCheckpoint\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport matplotlib.pyplot as plt\nfrom typing import Union, Tuple, List\n\nclass Model(pl.LightningModule):\n    def __init__(\n        self,\n        spatial_dims: int = 3,\n        in_channels: int = 1,\n        out_channels: int = 7,\n        channels: Union[Tuple[int, ...], List[int]] = (48, 64, 80, 80),\n        strides: Union[Tuple[int, ...], List[int]] = (2, 2, 1),\n        num_res_units: int = 1,\n        lr: float = 1e-3,\n    ):\n        super().__init__()\n        self.save_hyperparameters()\n        self.model = UNet(\n            spatial_dims=self.hparams.spatial_dims,\n            in_channels=self.hparams.in_channels,\n            out_channels=self.hparams.out_channels,\n            channels=self.hparams.channels,\n            strides=self.hparams.strides,\n            num_res_units=self.hparams.num_res_units,\n        )\n        self.loss_fn = TverskyLoss(include_background=True, to_onehot_y=True, softmax=True)\n        self.metric_fn = DiceMetric(include_background=False, reduction=\"mean\", ignore_empty=True)\n\n        self.train_loss = 0\n        self.val_metric = 0\n        self.num_train_batch = 0\n        self.num_val_batch = 0\n        \n    def forward(self, x):\n        return self.model(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T08:36:35.040478Z","iopub.execute_input":"2025-01-28T08:36:35.041139Z","iopub.status.idle":"2025-01-28T08:36:36.217124Z","shell.execute_reply.started":"2025-01-28T08:36:35.041116Z","shell.execute_reply":"2025-01-28T08:36:36.216248Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"channels = (48, 64, 80, 80)\nstrides_pattern = (2, 2, 1)       \nnum_res_units = 1\nlearning_rate = 1e-3\nnum_epochs = 1000\n\nif str(CONFIG.MODEL_DIR).split(\".\")[1] =='ckpt':\n    model = Model.load_from_checkpoint(CONFIG.MODEL_DIR,channels=channels, strides=strides_pattern, num_res_units=num_res_units, lr=learning_rate)\nelif str(CONFIG.MODEL_DIR).split(\".\")[1] =='pth':\n    model = Model(channels=channels, strides=strides_pattern, num_res_units=num_res_units, lr=learning_rate)\n    model.load_state_dict(torch.load(CONFIG.MODEL_DIR))\n\n\nmodel.eval()\nmodel.to(\"cuda\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T08:36:36.218317Z","iopub.execute_input":"2025-01-28T08:36:36.218627Z","iopub.status.idle":"2025-01-28T08:36:37.187593Z","shell.execute_reply.started":"2025-01-28T08:36:36.218597Z","shell.execute_reply":"2025-01-28T08:36:37.186696Z"},"_kg_hide-output":true,"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BLOB_THRESHOLD = 500\nCERTAINTY_THRESHOLD=0.5\n\nclasses = [1,2,3,4,5,6]\nwith torch.no_grad():\n    location_df = []\n    for run in root.runs:\n        print(run)\n\n        tomo = run.get_voxel_spacing(10)\n        tomo =tomo.get_tomogram(tomo_type).numpy()\n\n        tomo_patches, coordinates = extract_3d_patches_minimal_overlap([tomo],96)\n        tomo_patched_data = [{\"image\":img} for img in tomo_patches]\n\n        tomo_ds = CacheDataset(data=tomo_patched_data,transform = inference_transforms, cache_rate=1.0)\n\n        pred_masks=[]\n\n        for i in range(len(tomo_ds)):\n            input_tensor = tomo_ds[i]['image'].unsqueeze(0).to(\"cuda\")\n            model_output = model(input_tensor)\n            probs = torch.softmax(model_output[0],dim=0)\n            thresh_probs = probs > CERTAINTY_THRESHOLD\n            _,max_classes = thresh_probs.max(dim=0)\n            pred_masks.append(max_classes.cpu().numpy())\n        reconstructed_mask = reconstruct_array(pred_masks, coordinates, tomo.shape)\n        location ={}\n        \n        for c in classes:\n            cc = cc3d.connected_components(reconstructed_mask == c)\n            stats = cc3d.statistics(cc)\n            zyx=stats['centroids'][1:]*10 #.012444 #https://www.kaggle.com/competitions/czii-cryo-et-object-identification/discussion/544895#3040071\n            zyx_large = zyx[stats['voxel_counts'][1:] > BLOB_THRESHOLD]\n            xyz =np.ascontiguousarray(zyx_large[:,::-1])\n\n            location[id_to_name[c]] = xyz\n\n\n        df = dict_to_df(location, run.name)\n        location_df.append(df)\n    \n    location_df = pd.concat(location_df)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T08:36:37.188533Z","iopub.execute_input":"2025-01-28T08:36:37.188842Z","iopub.status.idle":"2025-01-28T08:37:13.143420Z","shell.execute_reply.started":"2025-01-28T08:36:37.188806Z","shell.execute_reply":"2025-01-28T08:37:13.142753Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"location_df.insert(loc=0, column='id', value=np.arange(len(location_df)))\nlocation_df.to_csv(\"submission.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T08:37:13.144187Z","iopub.execute_input":"2025-01-28T08:37:13.144407Z","iopub.status.idle":"2025-01-28T08:37:13.163765Z","shell.execute_reply.started":"2025-01-28T08:37:13.144388Z","shell.execute_reply":"2025-01-28T08:37:13.162953Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"location_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T08:37:13.165159Z","iopub.execute_input":"2025-01-28T08:37:13.165393Z","iopub.status.idle":"2025-01-28T08:37:13.184931Z","shell.execute_reply.started":"2025-01-28T08:37:13.165373Z","shell.execute_reply":"2025-01-28T08:37:13.184348Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}