{"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":"gpu","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":9867543,"sourceType":"datasetVersion","datasetId":6040935}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import 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\ntry :\n    import zarr\n    import cc3d\n    import torch_optimizer as optim\nexcept :\n    !pip install zarr\n    !pip install --no-index --find-links=/kaggle/input/hengck-czii-cryo-et-01/wheel_file connected-components-3d\n    import zarr\n    import cc3d","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-02-13T15:36:32.318928Z","iopub.execute_input":"2025-02-13T15:36:32.319419Z","iopub.status.idle":"2025-02-13T15:36:48.758057Z","shell.execute_reply.started":"2025-02-13T15:36:32.319384Z","shell.execute_reply":"2025-02-13T15:36:48.756970Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_KAGGLE_DIR = '/kaggle/input/czii-cryo-et-object-identification'\nTRAIN_DIR = f'{DATA_KAGGLE_DIR}/train'\nTEST_DIR = f'{DATA_KAGGLE_DIR}/test'\n\nTRAIN_EXP = [\"TS_5_4\",\"TS_69_2\",\"TS_6_4\",\"TS_6_6\",\"TS_73_6\",\"TS_86_3\",\"TS_99_9\"]\nTEST_EXP = [\"TS_5_4\",\"TS_69_2\",\"TS_6_4\"]\n\nscale = 10.012444196428572\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\ndef 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    # mean = volume.mean()\n    # std = volume.std()\n    # volume = (volume - mean) / std\n    return volume\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(reversed(list(json_data['points'][i]['location'].values())))  for i in range(num_point)]\n        location[p] = [[coo/scale for coo in coos] for coos in loc ]\n    return location\n\n\n\ndef visualize(data):\n    \"\"\"\n    Visualize multiple images or masks side by side for each slice in the dataset.\n    \n    Parameters:\n    data (dict): Dictionary where keys are the names of the data items,\n                 and values are either 3D images or masks.\n    \"\"\"\n    # Ensure all data items have the same depth (number of slices)\n    z_sizes = {name: item.shape[0] for name, item in data.items() if name!=\"center\"}\n    if len(set(z_sizes.values())) != 1:\n        raise ValueError(\"All items must have the same number of slices along the Z-axis.\")\n    \n    z_size = next(iter(z_sizes.values()))  # Depth of slices (Z-axis)\n    keys = list(data.keys())  # Get all keys for consistent ordering\n    # Loop over each slice (Z-index)\n    for z in range(z_size):\n        fig, axes = plt.subplots(1, len(keys), figsize=(3 * len(keys), 5))\n        for idx, key in enumerate(keys):\n            item = data[key]\n            cmap = \"gray\" if \"volume\" in key.lower() else \"jet\"\n            \n            # Display the slice\n            axes[idx].imshow(item[z, :, :], cmap=cmap)\n            axes[idx].set_title(f\"{key} - Slice (Z={z})\")\n            axes[idx].axis('off')\n        \n        plt.tight_layout()\n        plt.show()\n\ndef draw_sphere_in_image_fast(image, center, radius, radius_factor, value):\n    new_radius = radius * radius_factor\n    z_min = max(round(center[0] - new_radius), 0)\n    y_min = max(round(center[1] - new_radius), 0)\n    x_min = max(round(center[2] - new_radius), 0)\n\n    z_max = min(round(center[0] + new_radius) + 1, image.shape[0])\n    y_max = min(round(center[1] + new_radius) + 1, image.shape[1])\n    x_max = min(round(center[2] + new_radius) + 1, image.shape[2])\n    \n    local_region = image[z_min:z_max, y_min:y_max, x_min:x_max]\n    local_center = (new_radius, new_radius, new_radius)\n    \n    local_region = draw_sphere_in_local_image(local_region, local_center, new_radius, value)\n    image[z_min:z_max, y_min:y_max, x_min:x_max] = np.bitwise_or(local_region,image[z_min:z_max, y_min:y_max, x_min:x_max])\n\n    return image\n\ndef draw_sphere_in_local_image(image, center, radius, value):\n    shape = image.shape\n\n    z, y, x = np.indices(shape)\n\n    distance = (z - center[0])**2 + (y - center[1])**2 + (x - center[2])**2\n    \n    cylinder = distance <= radius**2\n\n    image[cylinder] = value\n\n    return image\n\ndef crop_with_center(image, patch_size, center):\n\n    center = [round (c) for c in center]\n    dims = image.shape\n    if len(dims) == 3:\n        start_indices = [max(0, c - p // 2) for c, p in zip(center, patch_size)]\n        end_indices = [min(dim, c + (p + 1) // 2) for c, p, dim in zip(center, patch_size, dims)]\n        \n        cropped_patch = image[\n            start_indices[0]:end_indices[0],\n            start_indices[1]:end_indices[1],\n            start_indices[2]:end_indices[2]\n        ]\n    else:\n        start_indices = [max(0, c - p // 2) for c, p in zip(center, patch_size)]\n        end_indices = [min(dim, c + (p + 1) // 2) for c, p, dim in zip(center, patch_size, dims[1:])]\n        \n        cropped_patch = image[\n            :,\n            start_indices[0]:end_indices[0],\n            start_indices[1]:end_indices[1],\n            start_indices[2]:end_indices[2]\n        ]\n    return cropped_patch\n\ndef one_hot(label):\n    return np.stack( [label==i for i in range(6)] ,  0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T15:36:57.806184Z","iopub.execute_input":"2025-02-13T15:36:57.806902Z","iopub.status.idle":"2025-02-13T15:36:57.824661Z","shell.execute_reply.started":"2025-02-13T15:36:57.806875Z","shell.execute_reply":"2025-02-13T15:36:57.823327Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data = {}\nfor exp_name in tqdm(TRAIN_EXP):\n    train_data[exp_name] = {}\n    train_data[exp_name][\"volume\"] = read_one_data(exp_name, static_dir=f'{TRAIN_DIR}/static/ExperimentRuns')\n    train_data[exp_name][\"truth\"] = read_one_truth(exp_name, overlay_dir=f'{TRAIN_DIR}/overlay/ExperimentRuns')\n    \n    train_data[exp_name][\"label\"] = np.zeros((184,630,630) , dtype = np.int8)\n    \n    for particle in train_data[exp_name][\"truth\"].keys():\n        radius = OBJECT_DICT[particle][\"radius\"]\n        radius_factor = np.log2(radius)/radius *.8\n        \n        label = OBJECT_DICT[particle][\"label\"]\n        \n        for point in train_data[exp_name][\"truth\"][particle]:\n            train_data[exp_name][\"label\"] = draw_sphere_in_image_fast(train_data[exp_name][\"label\"], point, radius, radius_factor = radius_factor, value = label)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T15:36:59.391801Z","iopub.execute_input":"2025-02-13T15:36:59.392182Z","iopub.status.idle":"2025-02-13T15:37:15.500597Z","shell.execute_reply.started":"2025-02-13T15:36:59.392155Z","shell.execute_reply":"2025-02-13T15:37:15.499820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"min_ = 0\nmax_ = 0\n\nfor k in train_data.keys():\n    pmin,pmax = np.percentile(train_data[k][\"volume\"],(5,99))\n    print(pmin,pmax)\n    min_ += pmin/7\n    max_ += pmax/7\n\nprint(min_,max_)\nfor k in train_data.keys():\n    train_data[k][\"volume\"] = (train_data[k][\"volume\"]-min_)/(max_-min_)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T15:37:15.501418Z","iopub.execute_input":"2025-02-13T15:37:15.501632Z","iopub.status.idle":"2025-02-13T15:37:24.943847Z","shell.execute_reply.started":"2025-02-13T15:37:15.501614Z","shell.execute_reply":"2025-02-13T15:37:24.942667Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def mean_std_shift (image,shift = 0.02):\n    factor = 1/(shift*2)\n    std = image.std()\n    mean = image.mean()\n    shift_mean = (torch.rand(1)/factor - shift).item()\n    shift_std = (torch.rand(1)/factor - shift).item()\n    new_mean = mean + mean * shift_mean\n    new_std = std + std * shift_std\n\n    new_image = (image-mean)/std*new_std+new_mean\n    return new_image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T15:37:24.945319Z","iopub.execute_input":"2025-02-13T15:37:24.945625Z","iopub.status.idle":"2025-02-13T15:37:24.950004Z","shell.execute_reply.started":"2025-02-13T15:37:24.945594Z","shell.execute_reply":"2025-02-13T15:37:24.949207Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn.functional as F\n\nclass SegmentationDataset(Dataset):\n    def __init__(self, patch_size,length ,experiments = [\"TS_6_4\"] , shift = 0.02):\n        self.patch_size = patch_size\n        self.experiments = experiments\n        self.length = length\n        self.shift = shift\n\n    def __len__(self):\n        return self.length\n\n    def augment(self,result):\n        \n        do_flip_z = torch.rand(1)<.5\n        do_flip_y = torch.rand(1)<.5\n        do_flip_x = torch.rand(1)<.5\n        rot_times = random.choice([0,1,2,3])\n\n        for key in result.keys():\n            if do_flip_z:\n                result[key] = np.flip(result[key], axis=-3)\n            if do_flip_y:\n                result[key] = np.flip(result[key], axis=-2)\n            if do_flip_x:\n                result[key] = np.flip(result[key], axis=-1)\n\n            if rot_times !=0:\n                result[key] = np.rot90(result[key] , k = rot_times, axes=(-2,-1))\n        return result\n        \n    def _to_tensor(self,result):\n        for k in result.keys():\n            if \"label\" in k or \"heat_map\" in k:\n                #result[k] = numpy_one_hot(result[k])\n                result[k] = torch.tensor(result[k].copy() , dtype = torch.long if \"label\" in k else torch.float32)\n                #result[k] = one_hot_encode_3d(result[k] ,2)\n            else : \n                result[k] = torch.tensor(result[k].copy() , dtype = torch.float32)\n\n        #result[\"label\"] = torch.stack([result[k] for k in result.keys() if \"label\" in k])\n        return result\n        \n    def __getitem__(self,idx):\n        zyx = [random.choice(range(self.patch_size[i]//2 ,dim-self.patch_size[i]//2)) for i,dim in enumerate((184,630,630))]\n        exp_name = random.choice(self.experiments)\n        \n        result = {}\n        for key in [\"label\", \"volume\"]:#,\"heat_map\"\n            result [key] = crop_with_center(train_data[exp_name][key], self.patch_size, zyx)\n\n        result = self.augment(result)\n\n        result[\"volume\"] = mean_std_shift(result[\"volume\"],self.shift)\n\n        # i left this so i get the same result (i forgot to delete it in original notebook i thaught it might interfere with the seed result if i don't keep it)\n        if torch.rand(1)> .8:\n            1\n\n\n        #result[\"label\"] = one_hot(result[\"label\"])\n        result = self._to_tensor(result)\n        return result\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T15:37:27.117848Z","iopub.execute_input":"2025-02-13T15:37:27.118173Z","iopub.status.idle":"2025-02-13T15:37:27.126544Z","shell.execute_reply.started":"2025-02-13T15:37:27.118149Z","shell.execute_reply":"2025-02-13T15:37:27.125532Z"}},"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            \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\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\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T15:38:22.573069Z","iopub.execute_input":"2025-02-13T15:38:22.573375Z","iopub.status.idle":"2025-02-13T15:38:22.585723Z","shell.execute_reply.started":"2025-02-13T15:38:22.573351Z","shell.execute_reply":"2025-02-13T15:38:22.584856Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def set_seed(seed):\n\n    random.seed(seed)  # Python's built-in random\n    np.random.seed(seed)  # NumPy random seed\n    torch.manual_seed(seed)  # Torch CPU random seed\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)  # Torch GPU random seed\n        torch.cuda.manual_seed_all(seed)  # All GPUs\n\n    # For deterministic behavior in CuDNN operations\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T15:38:23.453337Z","iopub.execute_input":"2025-02-13T15:38:23.453672Z","iopub.status.idle":"2025-02-13T15:38:23.458756Z","shell.execute_reply.started":"2025-02-13T15:38:23.453646Z","shell.execute_reply":"2025-02-13T15:38:23.457639Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.optim as optim\nfrom torch.cuda.amp import GradScaler, autocast\nfrom torch.nn.utils import clip_grad_norm_\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nbatch_size = 4\nnum_epochs = 36\npatch_size = (128, 128, 128)\n\nfolds = []\ntrain_experiments = [TRAIN_EXP[i] for i in range(7) if i not in folds]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-13T15:39:28.594072Z","iopub.execute_input":"2025-02-13T15:39:28.594381Z","iopub.status.idle":"2025-02-13T15:39:28.605122Z","shell.execute_reply.started":"2025-02-13T15:39:28.594358Z","shell.execute_reply":"2025-02-13T15:39:28.603822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_model():\n    \"\"\"Initialize and return the model.\"\"\"\n    model = Model().to(device)\n    return model\n\ndef get_loader():\n    \"\"\"Prepare and return the dataloaders.\"\"\"\n    train_dataset = SegmentationDataset(patch_size, 1024, experiments=train_experiments)\n    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=4)\n    \n    return train_loader\n\ndef train_one_epoch(model, train_loader, optimizer, scaler, max_norm=1.0):\n    \"\"\"Train the model for one epoch with gradient clipping and return the average loss.\"\"\"\n    model.train()\n    model.output_type = [\"loss\"]\n    train_loss = 0.0\n\n    for batch in tqdm(train_loader, desc=\"Training\", leave=False):\n        with autocast():\n            outputs = model(batch)\n            loss = outputs[\"loss\"]\n\n        optimizer.zero_grad()\n        scaler.scale(loss).backward()\n\n        # Unscales the gradients of optimizer's assigned params and clips gradients\n        scaler.unscale_(optimizer)  \n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)\n\n        scaler.step(optimizer)\n        scaler.update()\n\n        train_loss += loss.item()\n\n    avg_train_loss = train_loss / len(train_loader)\n    return avg_train_loss\n\ndef train_model(seed = 42):\n    \"\"\"Train the model across all epochs.\"\"\"\n    set_seed(seed)\n    model = get_model()\n    train_loader = get_loader()\n\n    # there is no logic behind using different shift factor i just forgot to change in different notebooks i keep here for same result\n    shift = 0.02 if seed == 42 else 0.03\n    train_loader.dataset.shift = shift\n    \n    scaler = GradScaler()\n\n    for epoch in range(num_epochs):\n\n\n        learning_rate = 1e-4\n        optimizer = optim.Adam(\n            model.parameters(), \n            lr=learning_rate, \n            betas=(0.9, 0.999)\n        )\n        \n        print(f\"Epoch [{epoch + 1}/{num_epochs}]\")\n        avg_train_loss = train_one_epoch(model, train_loader, optimizer, scaler)\n        print(f\"Train Loss: {avg_train_loss:.4f}\")\n\n        # Save model checkpoint\n        checkpoint_path = f\"model_all_{epoch}_{seed}.bin\"\n        torch.save(model.state_dict(), checkpoint_path)\n\n    # Free GPU memory\n    torch.cuda.empty_cache()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for seed in [80,120]:\n    train_model(seed)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}