{"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"}],"dockerImageVersionId":30823,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -U -q lightning zarr mrcfile patchify","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-24T14:23:02.713795Z","iopub.execute_input":"2024-12-24T14:23:02.714107Z","iopub.status.idle":"2024-12-24T14:23:11.115353Z","shell.execute_reply.started":"2024-12-24T14:23:02.714076Z","shell.execute_reply":"2024-12-24T14:23:11.114539Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!git clone -q https://github.com/wolny/pytorch-3dunet.git\nimport sys\nsys.path.append('/kaggle/working/pytorch-3dunet/')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T14:23:11.116136Z","iopub.execute_input":"2024-12-24T14:23:11.116359Z","iopub.status.idle":"2024-12-24T14:23:19.063268Z","shell.execute_reply.started":"2024-12-24T14:23:11.116332Z","shell.execute_reply":"2024-12-24T14:23:19.062131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Basic\nimport zarr\nimport numpy as np\nimport os\nimport pandas as pd\nfrom glob import glob\n# Data and visualization\nfrom sklearn.model_selection import train_test_split\nfrom scipy.spatial.transform import Rotation as R\nfrom scipy.ndimage import map_coordinates\nfrom skimage.measure import block_reduce\nimport h5py, mrcfile, os, warnings\nimport tqdm\nfrom patchify import patchify, unpatchify\nimport plotly.graph_objects as go\nfrom plotly.subplots import make_subplots\n# Model\nfrom pytorch3dunet.unet3d.model import UNet3D, ResidualUNet3D, ResidualUNetSE3D\nfrom pytorch3dunet.unet3d.losses import DiceLoss, WeightedCrossEntropyLoss, GeneralizedDiceLoss, BCEDiceLoss\nfrom pytorch3dunet.unet3d.metrics import MeanIoU, DiceCoefficient, BoundaryAveragePrecision, AdaptedRandError\nfrom pytorch3dunet.unet3d.seg_metrics import AveragePrecision\nfrom torchinfo import summary\n# Training\nimport torch\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch import optim\nfrom torch.utils.data import DataLoader, Subset, Dataset\nimport lightning as L\nfrom lightning.pytorch import Trainer, seed_everything\nfrom lightning.pytorch.callbacks import DeviceStatsMonitor, ModelCheckpoint, StochasticWeightAveraging\nfrom lightning.pytorch.tuner import Tuner\nfrom lightning.pytorch.loggers import TensorBoardLogger\nseed_everything(472, workers=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T14:23:19.064276Z","iopub.execute_input":"2024-12-24T14:23:19.064590Z","iopub.status.idle":"2024-12-24T14:23:25.857493Z","shell.execute_reply.started":"2024-12-24T14:23:19.064568Z","shell.execute_reply":"2024-12-24T14:23:25.856839Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"PATCH_SHAPE=[23, 64, 64]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T14:23:25.859297Z","iopub.execute_input":"2024-12-24T14:23:25.859802Z","iopub.status.idle":"2024-12-24T14:23:25.863195Z","shell.execute_reply.started":"2024-12-24T14:23:25.859779Z","shell.execute_reply":"2024-12-24T14:23:25.862404Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CZIIDataset(Dataset):\n    def __init__(self, mini_tomos, mini_masks):\n        self.mini_tomos = mini_tomos\n        self.mini_masks = mini_masks\n        \n        self.data = [(torch.from_numpy(t).reshape(tuple([1] + PATCH_SHAPE)).to(torch.float32), F.one_hot(torch.tensor(m).to(torch.int64), num_classes=6).permute(3, 0, 1, 2).to(torch.float32)) for t, m in zip(self.mini_tomos, self.mini_masks)]\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        return self.data[idx]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T14:23:25.864466Z","iopub.execute_input":"2024-12-24T14:23:25.864797Z","iopub.status.idle":"2024-12-24T14:23:25.881338Z","shell.execute_reply.started":"2024-12-24T14:23:25.864760Z","shell.execute_reply":"2024-12-24T14:23:25.880763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CZIIDataModule(L.LightningDataModule):\n\n    def __init__(self, \n                 batch_size,\n                 data_dir='/kaggle/input/czii-cryo-et-object-identification/train/static/ExperimentRuns',\n                 split_strategy='holdout'):\n        \n        super().__init__()\n        \n        self.save_hyperparameters(ignore=['data_dir'])\n        \n        self.data_dir = data_dir\n        \n        self.split_strategy=self.hparams.split_strategy\n        self.batch_size=self.hparams.batch_size\n        \n        self.radius = {\n            'apo-ferritin':60,\n            'beta-galactosidase':90,\n            'ribosome':150,\n            'thyroglobulin':130,\n            'virus-like-particle':135\n        }\n        self.scale = 10.012444196428572\n        \n    def setup(self, stage=None):\n        print('Preparing the dataset...')\n        paths = list()\n        for run_id in os.listdir(self.data_dir):\n            paths += glob(f\"{self.data_dir.replace('static', 'overlay')}/{run_id}/Picks/*.json\")\n        df = pd.concat([pd.read_json(x) for x in paths]).reset_index(drop=True)\n        for axis in \"x\", \"y\", \"z\":\n            df[axis] = df.points.apply(lambda x: x[\"location\"][axis])\n        df.rename(columns={'pickable_object_name':'class'}, inplace=True)\n        df.drop(columns=['user_id', 'session_id', 'unit', 'trust_orientation', 'points'], inplace=True)\n        df = df[df['class']!='beta-amylase']\n        df['voxel_spacing'] = 10.\n        label={\n            'apo-ferritin':1,\n            'beta-galactosidase':2,\n            'ribosome':3,\n            'thyroglobulin':4,\n            'virus-like-particle':5\n        }\n        df['label']=df['class'].map(label)\n        self.df=df\n\n        mini_tomos=list()\n        mini_masks=list()\n        for run_id in tqdm.tqdm(os.listdir(self.data_dir)):\n            tomo, mask = self.generate_datapoint(run_id)\n            tomo = np.pad(tomo, [[0, 0], [0, 10], [0, 10]], mode='constant', constant_values=0)\n            mask = np.pad(mask, [[0, 0], [0, 10], [0, 10]], mode='constant', constant_values=0)\n            tomo_patches = patchify(tomo, PATCH_SHAPE, step=PATCH_SHAPE)\n            mask_patches = patchify(mask, PATCH_SHAPE, step=PATCH_SHAPE)\n            mini_tomos += list(np.reshape(tomo_patches, tuple([-1] + PATCH_SHAPE)))\n            mini_masks += list(np.reshape(mask_patches, tuple([-1] + PATCH_SHAPE)))\n\n        print('Data Preparation Done...')\n        \n        print('Starting Holdout (5/2) setup...')\n        \n        self.data = CZIIDataset(mini_tomos, mini_masks)\n        self.train = Subset(self.data, list(range(4000)))\n        self.val = Subset(self.data, list(range(4000, 5600)))\n\n        print('Completed Holdout (5/2) setup...')\n    \n    # 5 - 2 split\n    def train_dataloader(self):\n        return DataLoader(self.train, batch_size=self.batch_size, num_workers=4, shuffle=False)\n\n    def val_dataloader(self):\n        return DataLoader(self.val, batch_size=self.batch_size, num_workers=4, shuffle=False)\n\n    def generate_datapoint(self, run_id):\n        target_zeros=np.zeros((184,630,630), dtype=np.uint32)\n        radius_list=np.array([i/self.scale for i in self.radius.values()], dtype=np.uint32)\n        Rmax = np.max(radius_list)\n        dim = [2*Rmax, 2*Rmax, 2*Rmax]\n        ref_list = []\n        for idx in range(len(radius_list)):\n            ref_list.append(self.create_sphere(dim, radius_list[idx]))\n        temp = self.df[self.df['run_name']==run_id][['x', 'y', 'z', 'label']]\n        # the x, y, z values are also divided by 10 to adjust to matrix\n        temp['x']/=self.scale\n        temp['y']/=self.scale\n        temp['z']/=self.scale\n        temp['phi']=self.scale\n        temp['psi']=self.scale\n        temp['the']=self.scale\n        objl=temp.to_dict(orient=\"records\")\n        tomo = self.get_tomogram(self.data_dir, run_id)\n        mask = self.generate_with_shapes(objl, target_zeros, ref_list)\n        return tomo, mask\n        \n    def get_tomogram(self, dir, run_id):\n        zarr_file=f'{dir}/{run_id}/VoxelSpacing10.000/denoised.zarr'\n        zarr_array=zarr.open(zarr_file, mode='r')[0][:]\n        vol = (zarr_array - zarr_array.min()) / (zarr_array.max() - zarr_array.min())\n        return vol.astype(np.float32)\n    \n    def create_sphere(self, dim, R):\n        C = np.floor((dim[0] / 2, dim[1] / 2, dim[2] / 2))\n        x, y, z = np.meshgrid(range(dim[0]), range(dim[1]), range(dim[2]))\n        sphere = ((x - C[0]) / R) ** 2 + ((y - C[1]) / R) ** 2 + ((z - C[2]) / R) ** 2\n        sphere = np.int32(sphere <= 1)\n        return sphere\n    \n    def rotate_array(self, array, orient):\n        phi = orient[0]\n        psi = orient[1]\n        the = orient[2]\n    \n        new_phi = -phi\n        new_psi = -the\n        new_the = -psi\n    \n        dim = array.shape\n        ax = np.arange(dim[0])\n        ay = np.arange(dim[1])\n        az = np.arange(dim[2])\n        coords = np.meshgrid(ax, ay, az)\n    \n        xyz = np.vstack(\n            [\n                coords[0].reshape(-1) - float(dim[0]) / 2,\n                coords[1].reshape(-1) - float(dim[1]) / 2,\n                coords[2].reshape(-1) - float(dim[2]) / 2,\n            ],\n        )\n    \n        r = R.from_euler(\"YZY\", [new_phi, new_psi, new_the], degrees=True)\n        mat = r.as_matrix()\n    \n        transformed_xyz = np.dot(mat, xyz)\n    \n        x = transformed_xyz[0, :] + float(dim[0]) / 2\n        y = transformed_xyz[1, :] + float(dim[1]) / 2\n        z = transformed_xyz[2, :] + float(dim[2]) / 2\n    \n        x = x.reshape((dim[1], dim[0], dim[2]))\n        y = y.reshape((dim[1], dim[0], dim[2]))\n        z = z.reshape((dim[1], dim[0], dim[2]))\n    \n        new_xyz = [y, x, z]\n        arrayR = map_coordinates(array, new_xyz, order=1)\n        return arrayR\n    \n    def generate_with_shapes(self, objl, target_array, ref_list):\n    \n        N = len(objl)\n        dim = target_array.shape       \n        for p in range(len(objl)):\n            lbl = int(objl[p]['label'])\n            x = int(objl[p]['x'])\n            y = int(objl[p]['y'])\n            z = int(objl[p]['z'])\n            phi = objl[p]['phi']\n            psi = objl[p]['psi']\n            the = objl[p]['the']\n    \n            ref = ref_list[lbl - 1]\n            centeroffset = np.int_(np.floor(ref.shape[0] / 2))\n    \n            if phi!=None and psi!=None and the!=None:\n                ref = self.rotate_array(ref, (phi, psi, the))\n                ref = np.int32(np.round(ref))\n    \n            obj_voxels = np.nonzero(ref == 1)\n            x_vox = obj_voxels[2] + x - centeroffset\n            y_vox = obj_voxels[1] + y - centeroffset\n            z_vox = obj_voxels[0] + z - centeroffset\n    \n            for idx in range(x_vox.size):\n                xx = x_vox[idx]\n                yy = y_vox[idx]\n                zz = z_vox[idx]\n                if xx >= 0 and xx < dim[2] and yy >= 0 and yy < dim[1] and zz >= 0 and zz < dim[0]: \n                    target_array[zz, yy, xx] = lbl\n    \n        return np.int32(target_array)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T14:23:25.882019Z","iopub.execute_input":"2024-12-24T14:23:25.882203Z","iopub.status.idle":"2024-12-24T14:23:25.906200Z","shell.execute_reply.started":"2024-12-24T14:23:25.882187Z","shell.execute_reply":"2024-12-24T14:23:25.905586Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CZIILightningModule(L.LightningModule):\n\n    def __init__(self, \n                 model, \n                 batch_size,\n                 lr, \n                 normalization='sigmoid', \n                 weight_decay=0.,\n                 betas=(0.9, 0.999)):\n        \n        super().__init__()\n        \n        self.save_hyperparameters(ignore=['model'])\n\n        # load hyperparameters\n        self.batch_size = self.hparams.batch_size\n        self.lr = self.hparams.lr\n        self.normalization=self.hparams.normalization\n        self.weight_decay = self.hparams.weight_decay\n        self.betas = self.hparams.betas\n        \n        # model\n        self.model = model\n        \n        # loss function\n        self.weight=torch.tensor([0., 1., 2., 1., 2., 1.]) #These weights are the same weights which will be assigned during evaluation\n        self.loss = DiceLoss(weight=self.weight, normalization=self.normalization)\n\n    def forward(self, x):\n        return self.model(x)\n        \n    def common_step(self, x, y, mode):\n\n        # forward pass\n        y_hat = self(x)\n\n        # calculate loss\n        dice_loss = self.loss(y_hat, y)\n        \n        return dice_loss\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        loss = self.common_step(x, y, 'train')\n        self.log('train_loss', loss, prog_bar=True, logger=True, on_step=True, on_epoch=True, enable_graph=True, sync_dist=True)\n        return loss\n    \n    def validation_step(self, batch, batch_idx):\n        x, y = batch\n        loss = self.common_step(x, y, 'val')\n        self.log('val_loss', loss, prog_bar=True, logger=True, on_step=False, on_epoch=True, enable_graph=True, sync_dist=True)\n        return loss\n        \n    def configure_optimizers(self):\n        optimizer = optim.Adam(self.model.parameters(), lr=self.lr, betas=self.betas, weight_decay=self.weight_decay)\n        lr_scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min')\n        return {\n                    \"optimizer\": optimizer,\n                    \"lr_scheduler\": {\n                        \"scheduler\": lr_scheduler,\n                        \"monitor\": \"train_loss\"\n                    }\n               }\n\n    def configure_callbacks(self):\n\n        return [ModelCheckpoint(monitor=\"val_loss\", mode='min'), \n                StochasticWeightAveraging(swa_lrs=self.lr),\n                DeviceStatsMonitor(cpu_stats=True)]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T14:23:25.907053Z","iopub.execute_input":"2024-12-24T14:23:25.907330Z","iopub.status.idle":"2024-12-24T14:23:25.924434Z","shell.execute_reply.started":"2024-12-24T14:23:25.907302Z","shell.execute_reply":"2024-12-24T14:23:25.923752Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dm = CZIIDataModule(batch_size=33)\n#dm.setup()\n#train_loader = dm.train_dataloader()\n#val_loader = dm.val_dataloader()\nMODEL = UNet3D(in_channels=1, out_channels=6)\nlm = CZIILightningModule(model=MODEL, batch_size=33, lr=0.0005248074602497723)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T14:23:25.925110Z","iopub.execute_input":"2024-12-24T14:23:25.925359Z","iopub.status.idle":"2024-12-24T14:23:26.104487Z","shell.execute_reply.started":"2024-12-24T14:23:25.925339Z","shell.execute_reply":"2024-12-24T14:23:26.103866Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!mkdir \"lightning_logs\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T14:23:26.105236Z","iopub.execute_input":"2024-12-24T14:23:26.105514Z","iopub.status.idle":"2024-12-24T14:23:26.237010Z","shell.execute_reply.started":"2024-12-24T14:23:26.105485Z","shell.execute_reply":"2024-12-24T14:23:26.236103Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%reload_ext tensorboard\n%tensorboard --logdir=lightning_logs/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T14:23:26.237925Z","iopub.execute_input":"2024-12-24T14:23:26.238185Z","iopub.status.idle":"2024-12-24T14:23:36.294850Z","shell.execute_reply.started":"2024-12-24T14:23:26.238155Z","shell.execute_reply":"2024-12-24T14:23:36.293890Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"logger = TensorBoardLogger(save_dir='/kaggle/working/lightning_logs/', name='lightning_logs')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T14:23:36.295759Z","iopub.execute_input":"2024-12-24T14:23:36.296058Z","iopub.status.idle":"2024-12-24T14:23:36.303167Z","shell.execute_reply.started":"2024-12-24T14:23:36.296033Z","shell.execute_reply":"2024-12-24T14:23:36.302266Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer = Trainer(accelerator='gpu',\n                 strategy='ddp_notebook',\n                 devices=1,\n                 precision='32',\n                 gradient_clip_val=None, \n                 logger=logger,\n                 max_epochs=30,\n                 enable_checkpointing=True,\n                 enable_progress_bar=True,\n                 enable_model_summary=False,\n                 inference_mode=True,\n                 default_root_dir='/kaggle/working/',\n                 num_sanity_val_steps=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T14:23:36.304194Z","iopub.execute_input":"2024-12-24T14:23:36.304497Z","iopub.status.idle":"2024-12-24T14:23:36.387054Z","shell.execute_reply.started":"2024-12-24T14:23:36.304468Z","shell.execute_reply":"2024-12-24T14:23:36.386315Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer.fit(model=lm, \n            #train_dataloaders=train_loader, \n            #val_dataloaders=val_loader, \n            datamodule=dm, \n            ckpt_path=None)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T14:23:36.389744Z","iopub.execute_input":"2024-12-24T14:23:36.390017Z","execution_failed":"2024-12-24T14:36:32.031Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def probability_to_location(probability,cfg):\n    _,D,H,W = probability.shape\n\n    location={}\n    for p in PARTICLE:\n        p = dotdict(p)\n        l = p.label\n\n        cc, P = cc3d.connected_components(probability[l]>cfg.threshold[p.name], return_N=True)\n        stats = cc3d.statistics(cc)\n        zyx=stats['centroids'][1:]*10\n        xyz = np.ascontiguousarray(zyx[:,::-1]) \n        location[p.name]=xyz\n        '''\n            j=1\n            z,y,x = np.where(cc==j)\n            z=z.mean()\n            y=y.mean()\n            x=x.mean()\n            print([x,y,z])\n        '''\n    return location\n\ndef location_to_df(location):\n    location_df = []\n    for p in PARTICLE:\n        p = dotdict(p)\n        xyz = location[p.name]\n        if len(xyz)>0:\n            df = pd.DataFrame(data=xyz, columns=['x','y','z'])\n            #df.loc[:,'particle_type']= p.name\n            df.insert(loc=0, column='particle_type', value=p.name)\n            location_df.append(df)\n    location_df = pd.concat(location_df)\n    return location_df","metadata":{"trusted":true,"execution":{"execution_failed":"2024-12-24T14:36:32.032Z"}},"outputs":[],"execution_count":null}]}