{"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":10256284,"sourceType":"datasetVersion","datasetId":6344547}],"dockerImageVersionId":30823,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 2D Unet for czii-cryo-et-object-identification\n\nThe pipeline is the following:\n\n1. create segmentation mask using copick\n2. create 2D dataset of the tomograph along the depth dimension.\n3. use UNet model and cross-entropy loss to train the model\n4. During inference we will slice the tomograph and then predict it's corresponding heatmap and then stack all the heatmaps together to find the centeroids.","metadata":{}},{"cell_type":"code","source":"import os\nos.environ['CUDA_LAUNCH_BLOCKING'] = '1'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:14:00.375307Z","iopub.execute_input":"2024-12-20T16:14:00.375682Z","iopub.status.idle":"2024-12-20T16:14:00.379807Z","shell.execute_reply.started":"2024-12-20T16:14:00.375653Z","shell.execute_reply":"2024-12-20T16:14:00.378827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install git+https://github.com/copick/copick-utils.git matplotlib tqdm copick \n!pip install -q \"monai-weekly[mlflow]\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:14:00.784782Z","iopub.execute_input":"2024-12-20T16:14:00.785055Z","iopub.status.idle":"2024-12-20T16:14:27.323457Z","shell.execute_reply.started":"2024-12-20T16:14:00.785034Z","shell.execute_reply":"2024-12-20T16:14:27.322449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install zarr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:14:27.324642Z","iopub.execute_input":"2024-12-20T16:14:27.324870Z","iopub.status.idle":"2024-12-20T16:14:30.675814Z","shell.execute_reply.started":"2024-12-20T16:14:27.324852Z","shell.execute_reply":"2024-12-20T16:14:30.674569Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:14:30.678368Z","iopub.execute_input":"2024-12-20T16:14:30.678780Z","iopub.status.idle":"2024-12-20T16:14:38.984365Z","shell.execute_reply.started":"2024-12-20T16:14:30.678747Z","shell.execute_reply":"2024-12-20T16:14:38.983518Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install copick","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:14:38.985845Z","iopub.execute_input":"2024-12-20T16:14:38.986148Z","iopub.status.idle":"2024-12-20T16:14:42.552634Z","shell.execute_reply.started":"2024-12-20T16:14:38.986117Z","shell.execute_reply":"2024-12-20T16:14:42.551800Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Following code snippet is taken from: https://www.kaggle.com/code/ahsuna123/3d-u-net-training-only\n\n# Make a copick project\nimport os\nimport shutil\n\nconfig_blob = \"\"\"{\n    \"name\": \"czii_cryoet_mlchallenge_2024\",\n    \"description\": \"2024 CZII CryoET ML Challenge training data.\",\n    \"version\": \"1.0.0\",\n\n    \"pickable_objects\": [\n        {\n            \"name\": \"apo-ferritin\",\n            \"is_particle\": true,\n            \"pdb_id\": \"4V1W\",\n            \"label\": 1,\n            \"color\": [  0, 117, 220, 128],\n            \"radius\": 60,\n            \"map_threshold\": 0.0418\n        },\n        {\n            \"name\": \"beta-galactosidase\",\n            \"is_particle\": true,\n            \"pdb_id\": \"6X1Q\",\n            \"label\": 2,\n            \"color\": [ 76,   0,  92, 128],\n            \"radius\": 90,\n            \"map_threshold\": 0.0578\n        },\n        {\n            \"name\": \"ribosome\",\n            \"is_particle\": true,\n            \"pdb_id\": \"6EK0\",\n            \"label\": 3,\n            \"color\": [  0,  92,  49, 128],\n            \"radius\": 150,\n            \"map_threshold\": 0.0374\n        },\n        {\n            \"name\": \"thyroglobulin\",\n            \"is_particle\": true,\n            \"pdb_id\": \"6SCJ\",\n            \"label\": 4,\n            \"color\": [ 43, 206,  72, 128],\n            \"radius\": 130,\n            \"map_threshold\": 0.0278\n        },\n        {\n            \"name\": \"virus-like-particle\",\n            \"is_particle\": true,\n            \"label\": 5,\n            \"color\": [255, 204, 153, 128],\n            \"radius\": 135,\n            \"map_threshold\": 0.201\n        },\n        {\n            \"name\": \"membrane\",\n            \"is_particle\": false,\n            \"label\": 8,\n            \"color\": [100, 100, 100, 128]\n        },\n        {\n            \"name\": \"background\",\n            \"is_particle\": false,\n            \"label\": 9,\n            \"color\": [10, 150, 200, 128]\n        }\n    ],\n\n    \"overlay_root\": \"/kaggle/working/overlay\",\n\n    \"overlay_fs_args\": {\n        \"auto_mkdir\": true\n    },\n\n    \"static_root\": \"/kaggle/input/czii-cryo-et-object-identification/train/static\"\n}\"\"\"\n\ncopick_config_path = \"/kaggle/working/copick.config\"\noutput_overlay = \"/kaggle/working/overlay\"\n\nwith open(copick_config_path, \"w\") as f:\n    f.write(config_blob)\n    \n# Update the overlay\n# Define source and destination directories\nsource_dir = '/kaggle/input/czii-cryo-et-object-identification/train/overlay'\ndestination_dir = '/kaggle/working/overlay'\n\n# Walk through the source directory\nfor root, dirs, files in os.walk(source_dir):\n    # Create corresponding subdirectories in the destination\n    relative_path = os.path.relpath(root, source_dir)\n    target_dir = os.path.join(destination_dir, relative_path)\n    os.makedirs(target_dir, exist_ok=True)\n    \n    # Copy and rename each file\n    for file in files:\n        if file.startswith(\"curation_0_\"):\n            new_filename = file\n        else:\n            new_filename = f\"curation_0_{file}\"\n            \n        \n        # Define full paths for the source and destination files\n        source_file = os.path.join(root, file)\n        destination_file = os.path.join(target_dir, new_filename)\n        \n        # Copy the file with the new name\n        shutil.copy2(source_file, destination_file)\n        print(f\"Copied {source_file} to {destination_file}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:14:42.553578Z","iopub.execute_input":"2024-12-20T16:14:42.553812Z","iopub.status.idle":"2024-12-20T16:14:42.772502Z","shell.execute_reply.started":"2024-12-20T16:14:42.553792Z","shell.execute_reply":"2024-12-20T16:14:42.771679Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Following code snippet is taken from: https://www.kaggle.com/code/ahsuna123/3d-u-net-training-only\n\nimport copick\nfrom tqdm import tqdm\nfrom copick_utils.segmentation import segmentation_from_picks\nimport copick_utils.writers.write as write\nfrom collections import defaultdict\nimport numpy as np\n\nroot = copick.from_file(copick_config_path)\n\ncopick_user_name = \"copickUtils\"\ncopick_segmentation_name = \"paintedPicks\"\nvoxel_size = 10\ntomo_type = \"denoised\"\n\n\n\n\n\n# Just do this once\ngenerate_masks = True\n\nif generate_masks:\n    target_objects = defaultdict(dict)\n    for object in root.pickable_objects:\n        if object.is_particle:\n            target_objects[object.name]['label'] = object.label\n            target_objects[object.name]['radius'] = object.radius\n\n\n    for run in tqdm(root.runs):\n        tomo = run.get_voxel_spacing(10)\n        tomo = tomo.get_tomogram(tomo_type).numpy()\n        target = np.zeros(tomo.shape, dtype=np.uint8)\n        for pickable_object in root.pickable_objects:\n            pick = run.get_picks(object_name=pickable_object.name, user_id=\"curation\")\n            if len(pick):  \n                target = segmentation_from_picks.from_picks(pick[0], \n                                                            target, \n                                                            target_objects[pickable_object.name]['radius'] * 0.8,\n                                                            target_objects[pickable_object.name]['label']\n                                                            )\n        write.segmentation(run, target, copick_user_name, name=copick_segmentation_name)\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:14:42.773422Z","iopub.execute_input":"2024-12-20T16:14:42.773670Z","iopub.status.idle":"2024-12-20T16:15:01.999315Z","shell.execute_reply.started":"2024-12-20T16:14:42.773649Z","shell.execute_reply":"2024-12-20T16:15:01.998639Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\ndata_dicts = []\nfor run in tqdm(root.runs):\n    tomogram = run.get_voxel_spacing(voxel_size).get_tomogram(tomo_type).numpy()\n    segmentation = run.get_segmentations(name=copick_segmentation_name, user_id=copick_user_name, voxel_size=voxel_size, is_multilabel=True)[0].numpy()\n    data_dicts.append({\"image\": tomogram, \"label\": segmentation})\n    \nprint(np.unique(data_dicts[0]['label']))\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:15:02.000136Z","iopub.execute_input":"2024-12-20T16:15:02.000777Z","iopub.status.idle":"2024-12-20T16:15:09.252530Z","shell.execute_reply.started":"2024-12-20T16:15:02.000748Z","shell.execute_reply":"2024-12-20T16:15:09.251454Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Create 2D Dataset","metadata":{}},{"cell_type":"code","source":"import cv2\nimport numpy as np\nimport zarr\nfrom omegaconf import OmegaConf\n\nimport glob\nimport os\n\n\ndef convert_to_8bit(x):\n    lower, upper = np.percentile(x, (0.5, 99.5))\n    x = np.clip(x, lower, upper)\n    x = (x - x.min()) / (x.max() - x.min() + 1e-12) * 255\n    return x.round().astype(\"uint8\")\n\n\ndef make_2d_dataset(run_name):\n\n    print(f\"{run_name=}\")\n\n    pixelspacing = 10\n    imgsize = 640\n    depth = 3\n\n    run_root_path = f\"/kaggle/working/{run_name}\"\n\n    os.makedirs(run_root_path, exist_ok=True)\n    os.makedirs(f\"{run_root_path}/images\", exist_ok=True)\n    os.makedirs(f\"{run_root_path}/labels\", exist_ok=True)\n\n    # read a volume\n    vol = zarr.open(f'/kaggle/input/czii-cryo-et-object-identification/train/static/ExperimentRuns/{run_name}/VoxelSpacing10.000/denoised.zarr/0', mode='r')\n\n    masks = zarr.open(f\"/kaggle/working/overlay/ExperimentRuns/{run_name}/Segmentations/10.000_copickUtils_0_paintedPicks-multilabel.zarr\", mode='r')\n    masks = masks[0]\n\n    # normalize [0, 255]\n    vol = convert_to_8bit(vol)\n\n    print(f\"{vol.shape=}, {masks.shape=}, {type(vol)=}\")\n\n\n    for j in range(vol.shape[0]):\n\n        newvols = []\n\n        # Here we take neighbouring images as different channels of the image instead of replicating the same image thrice.\n        for k in range(depth):\n            if (imgidx := j - k + 1) < vol.shape[0] and imgidx >= 0:\n                newvols.append(vol[j - k + 1])\n            else:\n                newvols.append(vol[j])\n\n        newvolf = np.stack(newvols, axis=-1)\n        mask = masks[j]\n\n        newvolf = cv2.resize(newvolf, (imgsize,imgsize))\n        mask = cv2.resize(mask, (imgsize,imgsize))\n\n        np.savez(f\"{run_root_path}/{run_name}_{j*10}.npy\", img=newvolf, mask=mask)\n\n\n\n\nruns = [\"TS_6_4\", \"TS_6_6\", \"TS_69_2\", \"TS_73_6\", \"TS_86_3\", \"TS_99_9\", \"TS_5_4\"]\n\nparticle_names = []\n\nfor i, r in enumerate(runs):\n    make_2d_dataset(r)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:15:09.254112Z","iopub.execute_input":"2024-12-20T16:15:09.254336Z","iopub.status.idle":"2024-12-20T16:15:53.728448Z","shell.execute_reply.started":"2024-12-20T16:15:09.254314Z","shell.execute_reply":"2024-12-20T16:15:53.727707Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Unet","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\nfrom pytorch_lightning import LightningModule, LightningDataModule, Trainer, seed_everything\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping, LearningRateMonitor\n\nimport torch\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nimport cv2\n\nimport segmentation_models_pytorch as smp\n\n\nimport numpy as np\n\nfrom timm.optim import create_optimizer_v2\n\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\n\n# \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:15:53.729558Z","iopub.execute_input":"2024-12-20T16:15:53.729815Z","iopub.status.idle":"2024-12-20T16:16:03.576488Z","shell.execute_reply.started":"2024-12-20T16:15:53.729793Z","shell.execute_reply":"2024-12-20T16:16:03.575621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BATCH_SIZE = 4\nNUM_WORKERS = 4\nSEED = 12\nBACKBONE = \"efficientnet-b7\"\nIN_CHANNELS = 3\nNUM_CLASSES = 5\nMAX_EPOCHS = 128\nDEVICE = \"cuda\"\n\nTRAIN = False\n\n# Class weights for cross entropy loss background has a weight of 0.5 and rest of the particles have weight as 32\nclass_weights = torch.tensor([0.5,32,32,32,32,32]).to(DEVICE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:16:03.577349Z","iopub.execute_input":"2024-12-20T16:16:03.577679Z","iopub.status.idle":"2024-12-20T16:16:03.810242Z","shell.execute_reply.started":"2024-12-20T16:16:03.577642Z","shell.execute_reply":"2024-12-20T16:16:03.809270Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def dice_score_fn(ytrue, ypred):\n    \"\"\"\n        count number of correctly predicted pixels that are not background pixels.\n    \"\"\"\n\n    # We are not intereseted in how many background pixels the model predicts correctly\n    # The idea behind this is generally the image will have 95%+ background pixels\n    bg_preds = ~((ytrue == 0) & (ypred == 0))\n\n    ytrue = ytrue[bg_preds]\n    ypred = ypred[bg_preds]\n   \n    return torch.count_nonzero(ytrue == ypred) / len(ytrue)\n\ndef loss_fn(ytrue, ypred):\n    b, c, h, w = ypred.shape\n\n    ytrue = ytrue.reshape(b, -1)\n    ypred = ypred.permute(0,2,3,1).reshape(b,h * w, c)\n\n    return F.cross_entropy(ypred.reshape(-1, c), ytrue.reshape(-1), weight=class_weights)\n\n\ndef seed_everything(seed):\n    print(f\"seeding code with: {seed=}\")\n    import random, os\n    import numpy as np\n    import torch\n\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:16:03.812060Z","iopub.execute_input":"2024-12-20T16:16:03.812285Z","iopub.status.idle":"2024-12-20T16:16:03.818208Z","shell.execute_reply.started":"2024-12-20T16:16:03.812267Z","shell.execute_reply":"2024-12-20T16:16:03.817471Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Define Dataset and Dataloader","metadata":{}},{"cell_type":"code","source":"class MyDataset(Dataset):\n    def __init__(self, ids, mode):\n        assert mode in [\"train\", \"val\"]\n\n        self.ids = ids\n        self.transforms = get_train_transforms() if mode == \"train\" else get_val_transforms()\n\n    def __len__(self):\n        return len(self.ids)\n\n    def __getitem__(self, idx):\n        sample_path = self.ids[idx]\n\n        sample = np.load(sample_path)\n        img, mask = sample[\"img\"], sample[\"mask\"]\n\n        sample = self.transforms(image=img, mask=mask)\n        sample[\"image\"] = torch.from_numpy(sample[\"image\"]) / 255.0\n        sample[\"mask\"] = torch.from_numpy(sample[\"mask\"])\n\n        return sample[\"image\"].transpose(0,2), sample[\"mask\"].long()\n\ndef get_train_transforms():\n    return A.Compose(\n        [\n            A.ShiftScaleRotate(p=0.5, border_mode=cv2.BORDER_CONSTANT, shift_limit=0.1, scale_limit=0.2, value=0,\n                               rotate_limit=30, mask_value=0),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n        ],\n        p=1.0,\n    )\n\n\ndef get_val_transforms():\n    return A.Compose(\n        [ \n        ],\n        p=1.0,\n    )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:16:03.819323Z","iopub.execute_input":"2024-12-20T16:16:03.819583Z","iopub.status.idle":"2024-12-20T16:16:03.844974Z","shell.execute_reply.started":"2024-12-20T16:16:03.819564Z","shell.execute_reply":"2024-12-20T16:16:03.844122Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MyDataModule(LightningDataModule):\n    def __init__(self):\n        super().__init__()\n        self.train_dataset = None\n        self.val_dataset = None\n        self.test_dataset = None\n\n    def setup(self, stage=None):\n\n        train_experiments = [\"TS_6_4\", \"TS_6_6\", \"TS_69_2\", \"TS_73_6\", \"TS_86_3\", \"TS_99_9\"]\n        val_experiments = [\"TS_5_4\"]\n\n        root_dir = Path(\"/kaggle/working\")\n\n        \n        train_ids = []\n\n        for exp in train_experiments:\n            exppath = root_dir.joinpath(exp)\n            train_ids.extend(exppath.glob(\"*.npz\"))\n\n        self.train_dataset = MyDataset(train_ids, \"train\")\n\n        val_ids = []\n\n        for exp in val_experiments:\n            exppath = root_dir.joinpath(exp)\n            val_ids.extend(exppath.glob(\"*.npz\"))\n\n        self.val_dataset = MyDataset(val_ids, \"val\")\n\n        print(f\"train: {len(train_ids)}, val: {len(val_ids)}\")\n\n\n    def train_dataloader(self):\n\n        return DataLoader(self.train_dataset, batch_size=BATCH_SIZE, collate_fn=None,\n                          shuffle=True, drop_last=True, num_workers=NUM_WORKERS)\n\n    def val_dataloader(self):\n          return DataLoader(self.val_dataset, batch_size=BATCH_SIZE, collate_fn=None,\n                            shuffle=False, drop_last=False, num_workers=NUM_WORKERS)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:16:03.845932Z","iopub.execute_input":"2024-12-20T16:16:03.846155Z","iopub.status.idle":"2024-12-20T16:16:03.864018Z","shell.execute_reply.started":"2024-12-20T16:16:03.846138Z","shell.execute_reply":"2024-12-20T16:16:03.863318Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Define Trainer","metadata":{}},{"cell_type":"code","source":"class SegmentationModel(LightningModule):\n    def __init__(self):\n        super().__init__()\n\n        self.model = smp.Unet(\n            encoder_name=BACKBONE,\n            encoder_weights=\"imagenet\",\n            in_channels=IN_CHANNELS,\n            classes=NUM_CLASSES + 1,\n        )\n\n        self.loss_fn = loss_fn\n        self.metric_fn = dice_score_fn\n\n\n    def forward(self, x):\n        return self.model(x)\n\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch\n\n        y_pred = self.model(x)\n        loss = self.loss_fn(y,y_pred)\n\n        dice_score = self.metric_fn(y, y_pred.softmax(1).argmax(1))\n\n        self.log('train_loss', loss, on_step=True, on_epoch=True, prog_bar=True, logger=True)\n        self.log(\"train_dice_score\", dice_score, on_step=True, on_epoch=True, prog_bar=True, logger=True)\n        return dict(loss=loss)\n\n    def validation_step(self, batch, batch_idx):\n\n        x, y = batch\n\n        y_pred = self.model(x)\n\n        loss = self.loss_fn(y.squeeze(),y_pred.squeeze())\n        dice_score = self.metric_fn(y, y_pred.softmax(1).argmax(1))\n\n\n        self.log(\"val_loss\", loss, on_step=False, on_epoch=True, prog_bar=True, logger=True)\n        self.log(\"val_dice_score\", dice_score, on_step=False, on_epoch=True, prog_bar=True, logger=True)\n\n    def configure_optimizers(self):\n        optimizer = create_optimizer_v2(model_or_params=self.model, opt=\"AdamW\", lr=2e-4, weight_decay=0.01)\n\n        batch_size = BATCH_SIZE\n        updates_per_epoch = len(self.trainer.datamodule.train_dataset) // batch_size\n\n        scheduler = CosineAnnealingLR(optimizer, T_max=updates_per_epoch*MAX_EPOCHS, eta_min=0.0)\n\n        lr_dict = dict(\n            scheduler=scheduler,\n            interval=\"step\",\n            frequency=1,  # same as default\n        )\n        return dict(optimizer=optimizer, lr_scheduler=lr_dict)\n\n\n    def lr_scheduler_step(self, scheduler, metric):\n        scheduler.step()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:16:03.864781Z","iopub.execute_input":"2024-12-20T16:16:03.864992Z","iopub.status.idle":"2024-12-20T16:16:03.882223Z","shell.execute_reply.started":"2024-12-20T16:16:03.864973Z","shell.execute_reply":"2024-12-20T16:16:03.881359Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"seed_everything(SEED)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:16:03.883197Z","iopub.execute_input":"2024-12-20T16:16:03.883448Z","iopub.status.idle":"2024-12-20T16:16:03.904206Z","shell.execute_reply.started":"2024-12-20T16:16:03.883427Z","shell.execute_reply":"2024-12-20T16:16:03.903533Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lr_monitor = LearningRateMonitor()\ncheckpoint_callback = ModelCheckpoint(dirpath=\"./saved_models\", monitor=\"val_dice_score\", mode=\"max\", filename=f\"2DCNN\" + \"_{epoch:03d}_{val_dice_score:.4f}\", save_weights_only=True, save_top_k=1)\nearly_stop_callback = EarlyStopping( monitor='val_dice_score', min_delta=0.00, patience=32, verbose=True, mode='max' )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:16:03.905052Z","iopub.execute_input":"2024-12-20T16:16:03.905297Z","iopub.status.idle":"2024-12-20T16:16:03.915409Z","shell.execute_reply.started":"2024-12-20T16:16:03.905277Z","shell.execute_reply":"2024-12-20T16:16:03.914744Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dm = MyDataModule()\nmodel = SegmentationModel()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:16:03.916159Z","iopub.execute_input":"2024-12-20T16:16:03.916475Z","iopub.status.idle":"2024-12-20T16:16:09.069992Z","shell.execute_reply.started":"2024-12-20T16:16:03.916446Z","shell.execute_reply":"2024-12-20T16:16:09.069236Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer = Trainer(\n    max_epochs=MAX_EPOCHS\n    , strategy=\"auto\"\n    , check_val_every_n_epoch=1\n    , sync_batchnorm=False\n    , accelerator=DEVICE\n    , precision=16\n    , gradient_clip_val = None\n    , deterministic=False\n    , log_every_n_steps=40\n    , callbacks=[checkpoint_callback, early_stop_callback, lr_monitor])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:16:09.070822Z","iopub.execute_input":"2024-12-20T16:16:09.071111Z","iopub.status.idle":"2024-12-20T16:16:09.119001Z","shell.execute_reply.started":"2024-12-20T16:16:09.071082Z","shell.execute_reply":"2024-12-20T16:16:09.118155Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if TRAIN:\n    trainer.fit(model, datamodule=dm)\nelse:\n    model = SegmentationModel.load_from_checkpoint(\"/kaggle/input/czii-cryo-et-object-identification-2d-unet/2DCNN_epoch100_val_dice_score0.2597.ckpt\")\n    model.eval()\n    print(\"loaded the model\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:20:40.619685Z","iopub.execute_input":"2024-12-20T16:20:40.620011Z","iopub.status.idle":"2024-12-20T16:20:41.990949Z","shell.execute_reply.started":"2024-12-20T16:20:40.619984Z","shell.execute_reply":"2024-12-20T16:20:41.990241Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Let's look at the couple of validation predictions","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:20:43.371667Z","iopub.execute_input":"2024-12-20T16:20:43.372072Z","iopub.status.idle":"2024-12-20T16:20:43.376535Z","shell.execute_reply.started":"2024-12-20T16:20:43.372037Z","shell.execute_reply":"2024-12-20T16:20:43.375371Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dm.setup()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:20:43.821539Z","iopub.execute_input":"2024-12-20T16:20:43.821862Z","iopub.status.idle":"2024-12-20T16:20:43.832245Z","shell.execute_reply.started":"2024-12-20T16:20:43.821837Z","shell.execute_reply":"2024-12-20T16:20:43.831530Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x,y = dm.val_dataset[12]\ny = y.numpy()\nwith torch.no_grad():\n    ypred = model.model(x.cuda().unsqueeze(0)).squeeze().cpu().softmax(0).argmax(0).numpy()\n\nfig, ax = plt.subplots(1, 2, figsize=(10, 4))  # 1 row, 2 columns\n\n\nax[0].imshow(y)\nax[0].set_title('Ground Truth')\nax[0].legend()\n\n# Second plot\nax[1].imshow(ypred)\nax[1].set_title('Prediction')\nax[1].legend()\n\n# Adjust layout to prevent overlap\nplt.tight_layout()\n\n# Display the plots\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:23:15.820353Z","iopub.execute_input":"2024-12-20T16:23:15.820703Z","iopub.status.idle":"2024-12-20T16:23:16.541218Z","shell.execute_reply.started":"2024-12-20T16:23:15.820677Z","shell.execute_reply":"2024-12-20T16:23:16.540473Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x,y = dm.val_dataset[14]\ny = y.numpy()\nwith torch.no_grad():\n    ypred = model.model(x.cuda().unsqueeze(0)).squeeze().cpu().softmax(0).argmax(0).numpy()\n\nfig, ax = plt.subplots(1, 2, figsize=(10, 4))  # 1 row, 2 columns\n\n\nax[0].imshow(y)\nax[0].set_title('Ground Truth')\nax[0].legend()\n\n# Second plot\nax[1].imshow(ypred)\nax[1].set_title('Prediction')\nax[1].legend()\n\n# Adjust layout to prevent overlap\nplt.tight_layout()\n\n# Display the plots\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:23:53.721213Z","iopub.execute_input":"2024-12-20T16:23:53.721598Z","iopub.status.idle":"2024-12-20T16:23:54.339696Z","shell.execute_reply.started":"2024-12-20T16:23:53.721562Z","shell.execute_reply":"2024-12-20T16:23:54.338789Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x,y = dm.val_dataset[30]\ny = y.numpy()\nwith torch.no_grad():\n    ypred = model.model(x.cuda().unsqueeze(0)).squeeze().cpu().softmax(0).argmax(0).numpy()\n\nfig, ax = plt.subplots(1, 2, figsize=(10, 4))  # 1 row, 2 columns\n\n\nax[0].imshow(y)\nax[0].set_title('Ground Truth')\nax[0].legend()\n\n# Second plot\nax[1].imshow(ypred)\nax[1].set_title('Prediction')\nax[1].legend()\n\n# Adjust layout to prevent overlap\nplt.tight_layout()\n\n# Display the plots\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T16:24:17.250653Z","iopub.execute_input":"2024-12-20T16:24:17.250936Z","iopub.status.idle":"2024-12-20T16:24:17.894373Z","shell.execute_reply.started":"2024-12-20T16:24:17.250915Z","shell.execute_reply":"2024-12-20T16:24:17.893406Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}