{"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}],"dockerImageVersionId":30839,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"deps_path = '/kaggle/input/cziidependencies'","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:40:05.145772Z","iopub.execute_input":"2025-01-28T14:40:05.146088Z","iopub.status.idle":"2025-01-28T14:40:05.150598Z","shell.execute_reply.started":"2025-01-28T14:40:05.146055Z","shell.execute_reply":"2025-01-28T14:40:05.149175Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! cp -r /kaggle/input/cziidependencies/asciitree-0.3.3/ asciitree-0.3.3/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:40:05.151466Z","iopub.execute_input":"2025-01-28T14:40:05.151802Z","iopub.status.idle":"2025-01-28T14:40:05.335783Z","shell.execute_reply.started":"2025-01-28T14:40:05.151765Z","shell.execute_reply":"2025-01-28T14:40:05.334768Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! pip wheel asciitree-0.3.3/asciitree-0.3.3/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:40:05.336817Z","iopub.execute_input":"2025-01-28T14:40:05.337170Z","iopub.status.idle":"2025-01-28T14:40:08.872275Z","shell.execute_reply.started":"2025-01-28T14:40:05.337136Z","shell.execute_reply":"2025-01-28T14:40:08.871207Z"},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install asciitree-0.3.3-py3-none-any.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:40:08.875069Z","iopub.execute_input":"2025-01-28T14:40:08.875322Z","iopub.status.idle":"2025-01-28T14:40:12.958930Z","shell.execute_reply.started":"2025-01-28T14:40:08.875296Z","shell.execute_reply":"2025-01-28T14:40:12.958035Z"},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! pip install -q --no-index --find-links {deps_path} --requirement {deps_path}/requirements.txt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:40:12.960708Z","iopub.execute_input":"2025-01-28T14:40:12.961011Z","iopub.status.idle":"2025-01-28T14:40:20.863438Z","shell.execute_reply.started":"2025-01-28T14:40:12.960987Z","shell.execute_reply":"2025-01-28T14:40:20.862406Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip show pytorch-lightning\n!pip install --upgrade pytorch-lightning\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:40:20.864616Z","iopub.execute_input":"2025-01-28T14:40:20.864974Z","iopub.status.idle":"2025-01-28T14:40:28.139834Z","shell.execute_reply.started":"2025-01-28T14:40:20.864949Z","shell.execute_reply":"2025-01-28T14:40:28.138838Z"},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from typing import List, Tuple, Union\nimport gc\nimport numpy as np\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-28T14:40:28.140831Z","iopub.execute_input":"2025-01-28T14:40:28.141082Z","iopub.status.idle":"2025-01-28T14:40:54.957635Z","shell.execute_reply.started":"2025-01-28T14:40:28.141060Z","shell.execute_reply":"2025-01-28T14:40:54.956995Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Define some helper functions\n\n\n### Patching helper functions\n\nThese are mostly used to split large volumes into smaller ones and stitch them back together. ","metadata":{}},{"cell_type":"code","source":"def calculate_patch_starts(dimension_size: int,patch_size:int)->List[int]:\n    \"\"\"\n    Calculate the starting positions of patches along a single dimension\n    with minimal overlap to cover the entire dimension.\n    \n    Parameters:\n    -----------\n    dimension_size : int\n        Size of the dimension\n    patch_size : int\n        Size of the patch in this dimension\n        \n    Returns:\n    --------\n    List[int]\n        List of starting positions for patches\n    \"\"\"\n    if dimension_size <= patch_size:\n        return [0]\n\n    n_patches = np.ceil(dimension_size/patch_size)\n\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\ndef extract_3d_patches_minimal_overlap(arrays:List[np.ndarray],\n                                       patch_size: int) -> Tuple[List[np.ndarray],List[Tuple[int,int,int]]]:\n    \"\"\"\n    Extract 3D patches from multiple arrays with minimal overlap to cover the entire array.\n    \n    Parameters:\n    -----------\n    arrays : List[np.ndarray]\n        List of input arrays, each with shape (m, n, l)\n    patch_size : int\n        Size of cubic patches (a x a x a)\n        \n    Returns:\n    --------\n    patches : List[np.ndarray]\n        List of all patches from all input arrays\n    coordinates : List[Tuple[int, int, int]]\n        List of starting coordinates (x, y, z) for each patch\n    \"\"\"\n\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 the 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    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: List[np.ndarray],\n                      coordinates: List[Tuple[int,int,int]],\n                      original_shape: Tuple[int,int,int])->np.ndarray:\n    \"\"\"\n    Reconstruct array from patches.\n    \n    Parameters:\n    -----------\n    patches : List[np.ndarray]\n        List of patches to reconstruct from\n    coordinates : List[Tuple[int, int, int]]\n        Starting coordinates for each patch\n    original_shape : Tuple[int, int, int]\n        Shape of the original array\n        \n    Returns:\n    --------\n    np.ndarray\n        Reconstructed array\n    \"\"\"\n    reconstructed = np.zeros(original_shape,dtype=np.int64)\n    patch_size = patches[0].shape[0]\n\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        \n    return reconstructed","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:40:54.958501Z","iopub.execute_input":"2025-01-28T14:40:54.959554Z","iopub.status.idle":"2025-01-28T14:40:54.970170Z","shell.execute_reply.started":"2025-01-28T14:40:54.959513Z","shell.execute_reply":"2025-01-28T14:40:54.969173Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Submission helper functions\n\nThese help with getting the submission in the correct format","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\ndef dict_to_df(coord_dict,experiment_name):\n    \"\"\"\n    Convert dictionary of coordinates to pandas DataFrame.\n    \n    Parameters:\n    -----------\n    coord_dict : dict\n        Dictionary where keys are labels and values are Nx3 coordinate arrays\n        \n    Returns:\n    --------\n    pd.DataFrame\n        DataFrame with columns ['x', 'y', 'z', 'label']\n    \"\"\"\n    all_coords = []\n    all_labels = []\n    for label, coords in coord_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-28T14:40:54.971128Z","iopub.execute_input":"2025-01-28T14:40:54.971485Z","iopub.status.idle":"2025-01-28T14:40:55.019184Z","shell.execute_reply.started":"2025-01-28T14:40:54.971450Z","shell.execute_reply":"2025-01-28T14:40:55.018248Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Reading the data","metadata":{}},{"cell_type":"code","source":"TRAIN_DATA_DIR=\"/kaggle/input/cziinumpy-dataset-exp\"\nTEST_DATA_DIR=\"/kaggle/input/czii-cryo-et-object-identification\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:40:55.020120Z","iopub.execute_input":"2025-01-28T14:40:55.020495Z","iopub.status.idle":"2025-01-28T14:40:55.033717Z","shell.execute_reply.started":"2025-01-28T14:40:55.020463Z","shell.execute_reply":"2025-01-28T14:40:55.032993Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_names = ['TS_5_4','TS_69_2','TS_6_6','TS_73_6','TS_86_3','TS_99_9']\nvalid_names = ['TS_6_4']\n\ntrain_files = []\nvalid_files = []\n\nfor name in train_names:\n    image = np.load(f\"{TRAIN_DATA_DIR}/train_image_{name}.npy\")\n    label = np.load(f\"{TRAIN_DATA_DIR}/train_label_{name}.npy\")\n\n    train_files.append({\"image\": image, \"label\":label})\n\nfor name in valid_names:\n    image = np.load(f\"{TRAIN_DATA_DIR}/train_image_{name}.npy\")\n    label = np.load(f\"{TRAIN_DATA_DIR}/train_label_{name}.npy\")\n\n    valid_files.append({\"image\": image, \"label\":label})\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:40:55.034611Z","iopub.execute_input":"2025-01-28T14:40:55.035106Z","iopub.status.idle":"2025-01-28T14:41:14.251321Z","shell.execute_reply.started":"2025-01-28T14:40:55.035083Z","shell.execute_reply":"2025-01-28T14:41:14.250335Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(train_files[0]['label'].shape)\nprint(train_files[0]['image'].shape)\nprint(valid_files[0]['label'].shape)\nprint(valid_files[0]['image'].shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:41:14.252322Z","iopub.execute_input":"2025-01-28T14:41:14.252674Z","iopub.status.idle":"2025-01-28T14:41:14.259568Z","shell.execute_reply.started":"2025-01-28T14:41:14.252642Z","shell.execute_reply":"2025-01-28T14:41:14.258544Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Create the training data loader","metadata":{}},{"cell_type":"code","source":"non_random_transforms = Compose([\n    EnsureChannelFirstd(keys=[\"image\", \"label\"], channel_dim=\"no_channel\"),\n    NormalizeIntensityd(keys=\"image\"),\n    Orientationd(keys=[\"image\", \"label\"], axcodes=\"RAS\")\n])\n\n#train_images,train_labels = [dcts['image'] for dcts in train_files],[dcts['label'] for dcts in train_files]\n#train_image_patches, _ = extract_3d_patches_minimal_overlap(train_images,96)\n#train_label_patches, _ = extract_3d_patches_minimal_overlap(train_labels,96)\n#train_files = [{\"image\":img, \"label\":lbl} for img,lbl in zip(train_image_patches,train_label_patches)]\n#train_ds = CacheDataset(data= train_files, transform = non_random_transforms, cache_rate=1.0)\nraw_train_ds  = CacheDataset(data= train_files, transform = non_random_transforms, cache_rate=1.0)\n\nmy_num_samples= 32\ntrain_batch_size=16\n\nrandom_transforms = Compose([\n    RandCropByLabelClassesd(\n        keys = [\"image\",\"label\"],\n        label_key = \"label\",\n        spatial_size=[96,96,96],\n        num_classes = 7,\n        num_samples = my_num_samples\n    ),\n    RandRotate90d(keys=[\"image\",\"label\"],prob=0.5,spatial_axes=[0,2]),\n    RandFlipd(keys=[\"image\",\"label\"],prob=0.5, spatial_axis = 0),\n])\n#train_ds = Dataset(data = raw_train_ds, transform = random_transforms)\ntrain_ds = CacheDataset(data = raw_train_ds, transform = random_transforms, cache_rate=1.0)\ntrain_files = [\n    {\"image\": sub_file[\"image\"], \"label\": sub_file[\"label\"]}\n    for file in train_ds\n    for sub_file in file\n]\n\ntrain_loader = DataLoader(\n    train_files,\n    batch_size = train_batch_size,\n    shuffle = True,\n    num_workers = 4,\n    pin_memory = torch.cuda.is_available()\n)\n\ndel train_ds\ngc.collect","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:41:14.262871Z","iopub.execute_input":"2025-01-28T14:41:14.263110Z","iopub.status.idle":"2025-01-28T14:41:27.814714Z","shell.execute_reply.started":"2025-01-28T14:41:14.263083Z","shell.execute_reply":"2025-01-28T14:41:27.813733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(train_loader))\nfor i in train_loader:\n    print(i['label'].shape)\n    break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:41:27.816324Z","iopub.execute_input":"2025-01-28T14:41:27.816679Z","iopub.status.idle":"2025-01-28T14:41:29.147760Z","shell.execute_reply.started":"2025-01-28T14:41:27.816644Z","shell.execute_reply":"2025-01-28T14:41:29.146896Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_images,val_labels = [dcts['image'] for dcts in valid_files],[dcts['label'] for dcts in valid_files]\n\nval_image_patches, _ = extract_3d_patches_minimal_overlap(val_images,96) #98\nval_label_patches, _ = extract_3d_patches_minimal_overlap(val_labels,96)\n\nval_patched_data = [{\"image\":img, \"label\":lbl} for img,lbl in zip(val_image_patches,val_label_patches)]\n\nvalid_ds = CacheDataset(data=val_patched_data, transform = non_random_transforms,cache_rate=1.0)\n\nvalid_batch_size=16\nval_loader = DataLoader(\n    valid_ds,\n    batch_size = valid_batch_size,\n    shuffle=False,\n    num_workers=4,\n    pin_memory=torch.cuda.is_available()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:41:29.149029Z","iopub.execute_input":"2025-01-28T14:41:29.149405Z","iopub.status.idle":"2025-01-28T14:41:30.508073Z","shell.execute_reply.started":"2025-01-28T14:41:29.149373Z","shell.execute_reply":"2025-01-28T14:41:30.507142Z"}},"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)\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch['image'], batch['label']\n        y_hat = self(x)\n        loss = self.loss_fn(y_hat, y)\n        self.train_loss += loss\n        self.num_train_batch += 1\n        return loss\n        \n    def on_train_epoch_end(self):\n        loss_per_epoch = self.train_loss / self.num_train_batch\n        print(f\"Epoch {self.current_epoch} - Average Train Loss: {loss_per_epoch:.4f}\")\n        self.log('train_loss', loss_per_epoch, prog_bar=True)\n        self.train_loss = 0\n        self.num_train_batch = 0\n\n    def validation_step(self, batch, batch_idx):\n        with torch.no_grad():\n            x, y = batch['image'], batch['label']\n            y_hat = self(x)\n            metric_val_outputs = [AsDiscrete(argmax=True, to_onehot=self.hparams.out_channels)(i) for i in decollate_batch(y_hat)]\n            metric_val_labels = [AsDiscrete(to_onehot=self.hparams.out_channels)(i) for i in decollate_batch(y)]\n\n            self.metric_fn(y_pred=metric_val_outputs, y=metric_val_labels)\n            metrics = self.metric_fn.aggregate(reduction=\"mean_batch\")\n            val_metric = torch.mean(metrics)\n            self.val_metric += val_metric \n            self.num_val_batch += 1\n        return {'val_metric': val_metric}\n\n    def on_validation_epoch_end(self):\n        metric_per_epoch = self.val_metric / self.num_val_batch\n        current_lr = self.trainer.optimizers[0].param_groups[0]['lr']  # Get the current learning rate\n        print(f\"------Epoch {self.current_epoch} - Learning Rate: {current_lr:.6f} - Average Val Metric: {metric_per_epoch:.4f}\")\n        self.log('val_metric', metric_per_epoch, prog_bar=True)\n        self.val_metric = 0\n        self.num_val_batch = 0\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(self.parameters(), lr=self.hparams.lr)\n        scheduler = {\n            \"scheduler\": ReduceLROnPlateau(optimizer, mode=\"max\", factor=0.5, patience=5, verbose=True),\n            \"monitor\": \"val_metric\",\n            \"interval\": \"epoch\",\n            \"frequency\": 1,\n        }\n        return [optimizer], [scheduler]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:41:30.509105Z","iopub.execute_input":"2025-01-28T14:41:30.509483Z","iopub.status.idle":"2025-01-28T14:41:31.740492Z","shell.execute_reply.started":"2025-01-28T14:41:30.509448Z","shell.execute_reply":"2025-01-28T14:41:31.739762Z"}},"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 = 500\n\nmodel = Model(channels=channels, strides=strides_pattern, num_res_units=num_res_units, lr=learning_rate)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:41:31.741230Z","iopub.execute_input":"2025-01-28T14:41:31.741463Z","iopub.status.idle":"2025-01-28T14:41:31.767844Z","shell.execute_reply.started":"2025-01-28T14:41:31.741432Z","shell.execute_reply":"2025-01-28T14:41:31.766947Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.set_float32_matmul_precision('medium')\n\n# Check if CUDA is available and then count the GPUs\nif torch.cuda.is_available():\n    num_gpus = torch.cuda.device_count()\n    print(f\"Number of GPUs available: {num_gpus}\")\nelse:\n    print(\"No GPU available. Running on CPU.\")\ndevices = list(range(num_gpus))\nprint(devices)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:41:31.768836Z","iopub.execute_input":"2025-01-28T14:41:31.769157Z","iopub.status.idle":"2025-01-28T14:41:31.775132Z","shell.execute_reply.started":"2025-01-28T14:41:31.769134Z","shell.execute_reply":"2025-01-28T14:41:31.774274Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Create a custom callback to track metrics during training\nclass MetricPlotCallback(pl.Callback):\n    def __init__(self):\n        self.train_losses = []\n        self.val_metrics = []\n\n    def on_train_epoch_end(self, trainer, pl_module, outputs):\n        # Save training loss at the end of each epoch\n        self.train_losses.append(trainer.callback_metrics['train_loss'].item())\n\n    def on_validation_epoch_end(self, trainer, pl_module):\n        # Save validation loss at the end of each validation epoch\n        self.val_metrics.append(trainer.callback_metrics['val_metric'].item())\n\n    def plot_metrics(self):\n        # Plot training and validation losses\n        epochs = range(1, len(self.train_losses) + 1)\n        plt.figure(figsize=(12, 6))\n\n        # Plot training loss\n        plt.subplot(1, 2, 1)\n        plt.plot(epochs, self.train_losses, label='Train Loss')\n        plt.xlabel('Epochs')\n        plt.ylabel('Loss')\n        plt.title('Training Loss Over Epochs')\n        plt.legend()\n\n        # Plot validation metric\n        plt.subplot(1, 2, 2)\n        plt.plot(epochs, self.val_metrics, label='Validation Metric', color='orange')\n        plt.xlabel('Epochs')\n        plt.ylabel('Metric')\n        plt.title('Validation Metric Over Epochs')\n        plt.legend()\n\n        plt.tight_layout()\n        plt.savefig(\"taining_epoches.png\")\n        print(f\"Metrics plot saved.\")\n        plt.show()\n        \ncheckpoint_callback = ModelCheckpoint(\n    monitor=\"val_metric\",            # Metric to monitor\n    mode=\"max\",                      # Maximize the monitored metric (use \"min\" for loss)\n    save_top_k=1,                    # Save only the best model\n    filename=\"best-checkpoint\",  # Save format\n    verbose=True                     # Print save messages\n)\n\n# Instantiate the callback\nmetric_plot_callback = MetricPlotCallback()\n\n# Add this callback to your trainer\ntrainer = pl.Trainer(\n    callbacks=[checkpoint_callback],\n    max_epochs=num_epochs,\n    accelerator=\"gpu\",\n    devices=[0],\n    num_nodes=1,\n    log_every_n_steps=10,\n    enable_progress_bar=True,\n)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:41:31.775973Z","iopub.execute_input":"2025-01-28T14:41:31.776285Z","iopub.status.idle":"2025-01-28T14:41:31.838089Z","shell.execute_reply.started":"2025-01-28T14:41:31.776255Z","shell.execute_reply":"2025-01-28T14:41:31.837461Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer.fit(model, train_loader, val_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-28T14:41:31.838885Z","iopub.execute_input":"2025-01-28T14:41:31.839106Z","execution_failed":"2025-01-28T14:43:14.659Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Call the plot method after training finishes\n#metric_plot_callback.plot_metrics()\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-01-28T14:43:14.659Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save the model manually\ntorch.save(model.state_dict(), \"model_weights1000_epoches.pth\")\n\n# To load the model later:\n#model.load_state_dict(torch.load(\"model_weights.pth\"))\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-01-28T14:43:14.659Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}