{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":36363,"databundleVersionId":4050810,"sourceType":"competition"},{"sourceId":2768643,"sourceType":"datasetVersion","datasetId":1688447},{"sourceId":4407134,"sourceType":"datasetVersion","datasetId":2556595},{"sourceId":4454029,"sourceType":"datasetVersion","datasetId":2607864},{"sourceId":7910796,"sourceType":"datasetVersion","datasetId":4647323},{"sourceId":7911066,"sourceType":"datasetVersion","datasetId":4647528},{"sourceId":7921148,"sourceType":"datasetVersion","datasetId":4654903},{"sourceId":7921218,"sourceType":"datasetVersion","datasetId":4654958},{"sourceId":7921255,"sourceType":"datasetVersion","datasetId":4654987},{"sourceId":8124472,"sourceType":"datasetVersion","datasetId":4730514},{"sourceId":8130848,"sourceType":"datasetVersion","datasetId":4753262},{"sourceId":8240821,"sourceType":"datasetVersion","datasetId":4808693},{"sourceId":104036025,"sourceType":"kernelVersion"}],"dockerImageVersionId":30683,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# !pip install /kaggle/input/rsna-2022-whl/{pydicom-2.3.0-py3-none-any.whl,pylibjpeg-1.4.0-py3-none-any.whl,python_gdcm-3.0.15-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl}\n# !pip install /kaggle/input/rsna-2022-whl/{torch-1.12.1-cp37-cp37m-manylinux1_x86_64.whl,torchvision-0.13.1-cp37-cp37m-manylinux1_x86_64.whl}","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install /kaggle/input/rsna-weights/timm-0.5.4-py3-none-any.whl\n# !pip install /kaggle/input/rsna-weights/tifffile-2022.8.8-py3-none-any.whl\n# !pip install /kaggle/input/rsna-weights/einops-0.5.0-py3-none-any.whl\n\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install pydicom gdcm pylibjpeg pylibjpeg-libjpeg","metadata":{"execution":{"iopub.status.busy":"2024-04-16T02:58:45.573008Z","iopub.execute_input":"2024-04-16T02:58:45.573881Z","iopub.status.idle":"2024-04-16T02:59:00.033409Z","shell.execute_reply.started":"2024-04-16T02:58:45.573844Z","shell.execute_reply":"2024-04-16T02:59:00.032312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install python_gdcm\n# !pip install pylibjpeg\n# !pip install pydicom\n# !pip install einops\n!pip install monai\nimport time","metadata":{"execution":{"iopub.status.busy":"2024-04-16T02:59:00.035673Z","iopub.execute_input":"2024-04-16T02:59:00.036414Z","iopub.status.idle":"2024-04-16T02:59:27.329905Z","shell.execute_reply.started":"2024-04-16T02:59:00.036356Z","shell.execute_reply":"2024-04-16T02:59:27.328816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import os\n# os.listdir('/kaggle/input/mmdetection-2-17-offline')\n\n# !pip install /kaggle/input/mmdetection-2-17-offline/mmcv_full-1.3.14-cp37-cp37m-linux_x86_64.whl --no-deps\n# !pip install /kaggle/input/mmdetection-2-17-offline/pycocotools-2.0.2-cp37-cp37m-linux_x86_64.whl --no-deps\n# !pip install /kaggle/input/mmdetection-2-17-offline/terminaltables-3.1.0-py3-none-any.whl --no-deps\n# !pip install /kaggle/input/mmdetection-2-17-offline/pytest_runner-5.3.1-py3-none-any.whl --no-deps\n# !pip install /kaggle/input/mmdetection-2-17-offline/mmpycocotools-12.0.3-cp37-cp37m-linux_x86_64.whl --no-deps\n# !pip install /kaggle/input/mmdetection-2-17-offline/terminal-0.4.0-py3-none-any.whl --no-deps\n# !pip install /kaggle/input/mmdetection-2-17-offline/mmdet-2.17.0-py3-none-any.whl --no-deps\n# !pip install /kaggle/input/mmdetection-2-17-offline/addict-2.4.0-py3-none-any.whl --no-deps\n# !pip install /kaggle/input/mmdetection-2-17-offline/yapf-0.31.0-py2.py3-none-any.whl --no-deps","metadata":{"execution":{"iopub.status.busy":"2024-04-16T02:59:27.331407Z","iopub.execute_input":"2024-04-16T02:59:27.331716Z","iopub.status.idle":"2024-04-16T02:59:27.336296Z","shell.execute_reply.started":"2024-04-16T02:59:27.331688Z","shell.execute_reply":"2024-04-16T02:59:27.335495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import sys\n# sys.path.append('/kaggle/input/rsnazoopublic')","metadata":{"execution":{"iopub.status.busy":"2024-04-16T02:59:27.338264Z","iopub.execute_input":"2024-04-16T02:59:27.338535Z","iopub.status.idle":"2024-04-16T02:59:27.345936Z","shell.execute_reply.started":"2024-04-16T02:59:27.338509Z","shell.execute_reply":"2024-04-16T02:59:27.345208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport random\nimport re\nfrom dataclasses import dataclass\nfrom typing import Dict\nfrom typing import List\n\nimport albumentations\nimport cv2\nimport numpy as np\nimport pydicom\nimport tifffile\nimport torch\nimport torch.hub\nfrom albumentations import ReplayCompose\nfrom skimage import measure\nfrom torch.functional import Tensor\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2024-04-16T03:02:48.516515Z","iopub.execute_input":"2024-04-16T03:02:48.517432Z","iopub.status.idle":"2024-04-16T03:02:48.828124Z","shell.execute_reply.started":"2024-04-16T03:02:48.517396Z","shell.execute_reply":"2024-04-16T03:02:48.827126Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchvision import models","metadata":{"execution":{"iopub.status.busy":"2024-04-16T03:02:50.492087Z","iopub.execute_input":"2024-04-16T03:02:50.492836Z","iopub.status.idle":"2024-04-16T03:02:50.497193Z","shell.execute_reply.started":"2024-04-16T03:02:50.492804Z","shell.execute_reply":"2024-04-16T03:02:50.496284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from monai.networks.nets import densenet, resnet, senet, autoencoder, attentionunet\n# from monai.networks.nets import densenet, resnet, senet\nclass MonaiModelWithClassification(nn.Module):\n    def __init__(self):\n        super(MonaiModelWithClassification, self).__init__()\n#         self.features = densenet.DenseNet(spatial_dims = 3, in_channels = 3, out_channels = 8)\n\n#         self.features = resnet.ResNet(spatial_dims=3, num_classes = 8, block = \"basic\", layers = [3, 4, 23, 3], block_inplanes = [64, 64, 64, 64])\n\n#         self.features = senet.SEResNext101(spatial_dims=3, num_classes = 8, in_channels = 3)\n\n        self.features = autoencoder.AutoEncoder(spatial_dims=3, in_channels=3, out_channels=8, channels=(2, 4, 8), strides=(2, 2, 2))\n        self.final_layer = nn.Linear(8 * 40 * 256 * 256, 8)\n#         self.fc1 = nn.Linear(8 * 40 * 256 * 256, 1024)  # Add first fully connected layer\n#         self.fc2 = nn.Linear(1024, 512)  # Add second fully connected layer\n#         self.fc3 = nn.Linear(512, 8) \n#         num_ftrs = self.features.fc.in_features\n#         self.features.fc = nn.Linear(num_ftrs, 8)\n \n    def forward(self, x):\n        x = self.features(x)\n        x = torch.sigmoid(x)  # Apply sigmoid activation for multi-label classification\n        x = x.view(x.size(0), -1)  # Flatten the output tensor\n        x = self.final_layer(x)  # Apply the final layer\n        return x\n    \nmodel = MonaiModelWithClassification()\n# input=torch.rand(1,3,40,256,256)\n# output=model(input)\nprint(model)","metadata":{"execution":{"iopub.status.busy":"2024-04-16T03:02:51.786353Z","iopub.execute_input":"2024-04-16T03:02:51.787083Z","iopub.status.idle":"2024-04-16T03:02:53.493864Z","shell.execute_reply.started":"2024-04-16T03:02:51.78704Z","shell.execute_reply":"2024-04-16T03:02:53.492853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # from torchvision import models\n# from monai.networks.nets import densenet, resnet, senet\n# class MonaiModelWithClassification(nn.Module):\n#     def __init__(self):\n#         super(MonaiModelWithClassification, self).__init__()\n# #         self.features = densenet.DenseNet(spatial_dims = 3, in_channels = 3, out_channels = 8)\n\n# #         self.features = resnet.ResNet(spatial_dims=3, num_classes = 8, block = \"basic\", layers = [3, 4, 23, 3], block_inplanes = [64, 64, 64, 64])\n\n# #         self.features = senet.SEResNext101(spatial_dims=3, num_classes = 8, in_channels = 3)\n\n#         self.features = senet.SEResNet152(spatial_dims=3, num_classes = 8, in_channels = 3)\n \n#     def forward(self, x):\n#         x = self.features(x)\n#         x = torch.sigmoid(x)  # Apply sigmoid activation for multi-label classification\n#         return x\n    \n# model = MonaiModelWithClassification()\n# # input=torch.rand(1,3,40,256,256)\n# # output=model(input)\n# # print(output)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@dataclass\nclass BatchSlice:\n    i_from: int\n    i_to: int\n    i_start: int\n\n\ndef get_slices(batch: Tensor, dim=1, window: int = 16, overlap: int = 8) -> List[BatchSlice]:\n    num_imgs = batch.size(dim)\n    if num_imgs <= window:\n        return [BatchSlice(0, num_imgs, 0)]\n    stride = window - overlap\n    result = []\n    current_idx = 0\n    while True:\n        next_idx = current_idx + window\n\n        if next_idx >= num_imgs:\n            current_idx = num_imgs - window\n            offset = overlap // 2 if current_idx > 0 else 0\n            next_idx = num_imgs\n            result.append(BatchSlice(current_idx, next_idx, offset))\n            break\n        else:\n            offset = overlap // 2 if current_idx > 0 else 0\n            result.append(BatchSlice(current_idx, next_idx, offset))\n        current_idx += stride\n    return result","metadata":{"execution":{"iopub.status.busy":"2024-04-16T03:03:05.502092Z","iopub.execute_input":"2024-04-16T03:03:05.502861Z","iopub.status.idle":"2024-04-16T03:03:05.511971Z","shell.execute_reply.started":"2024-04-16T03:03:05.50283Z","shell.execute_reply":"2024-04-16T03:03:05.511001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\ndef _read_labels():\n    labels_df = pd.read_csv(\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\n    labels_dict = {}\n    for index, row in labels_df.iterrows():\n        cube_id = row['StudyInstanceUID']\n        overall_patient = row['patient_overall']\n        c1, c2, c3, c4, c5, c6, c7 = row['C1'], row['C2'], row['C3'], row['C4'], row['C5'], row['C6'], row['C7']\n        labels_dict[cube_id] = [overall_patient, c1, c2, c3, c4, c5, c6, c7]\n    return labels_dict\n\nlabels_dict = _read_labels()\n# print(labels_dict)","metadata":{"execution":{"iopub.status.busy":"2024-04-16T03:03:06.908059Z","iopub.execute_input":"2024-04-16T03:03:06.909059Z","iopub.status.idle":"2024-04-16T03:03:07.129734Z","shell.execute_reply.started":"2024-04-16T03:03:06.909024Z","shell.execute_reply":"2024-04-16T03:03:07.128858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def combine_scan(scan_dir: str, size=512, fix_monochrome: bool = True) -> np.ndarray:\n    num_files = len(os.listdir(scan_dir))\n    images = []\n    offset = 0\n    first = None\n    last = None\n    files = []\n    for i in range(num_files):\n        dpath = os.path.join(scan_dir, f\"{i + offset}.dcm\")\n        if i == 0:\n            while not os.path.exists(dpath):\n                offset += 1\n                dpath = os.path.join(scan_dir, f\"{i + offset}.dcm\")\n        files.append(dpath)\n\n    for dpath in files[::2]:\n        ds = pydicom.dcmread(dpath)\n        if not first:\n            first = ds\n        last = ds\n\n        data = ds.pixel_array\n        data = cv2.resize(data, (size, size))\n        if fix_monochrome and ds.PhotometricInterpretation == \"MONOCHROME1\":\n            data = np.amax(data) - data\n        images.append(data)\n\n    if first and last:\n        if last.ImagePositionPatient[2] > first.ImagePositionPatient[2]:\n            images = images[::-1]\n    return np.array(images)\n\nclass DatasetSeg(Dataset):\n    def __init__(\n            self,\n            dataset_dir: str,\n            cases: List[str],\n    ):\n        self.dataset_dir = dataset_dir\n        self.cases = cases\n\n    def __getitem__(self, i):\n        cube_id = self.cases[i]\n        image_cube = combine_scan(os.path.join(self.dataset_dir, cube_id), size=256)\n        image_mean = image_cube.mean()\n        image_std = image_cube.std()\n        h = image_cube.shape[0]\n\n        images = image_cube\n        if h % 32 > 0:\n            tmp = np.zeros(((h // 32 + 1) * 32, 256, 256))\n            tmp[:h] = images\n            images = tmp\n        images = (images - image_mean) / image_std\n        images = np.expand_dims(images, 0)\n        sample = {}\n        sample['image'] = torch.from_numpy(images).float()\n        sample['cube_id'] = cube_id\n        sample['h'] = h\n        return sample\n\n    def __len__(self):\n        return len(self.cases)\n\n\ncrop_augs =  albumentations.ReplayCompose([\n            albumentations.LongestMaxSize(256),\n            albumentations.PadIfNeeded(256, 256, border_mode=cv2.BORDER_CONSTANT),\n        ])\n\nclass DatasetCrops(Dataset):\n    def __init__(\n            self,\n            dataset_dir: str,\n            cases: List[str],\n            transforms=crop_augs,\n            slice_size=40,\n    ):\n        self.dataset_dir = dataset_dir\n        self.transforms = transforms\n        self.slice_size = slice_size\n        self.cases = cases\n\n    def __getitem__(self, i):\n        cube_id = self.cases[i]\n#         mask_cube = tifffile.imread(os.path.join(\"seg_preds\", f\"{cube_id}.tif\"))\n        mask_cube_path = os.path.join(\"/kaggle/input/masks-part1\", f\"{cube_id}.tif\")\n        if os.path.exists(mask_cube_path):\n            mask_cube = tifffile.imread(mask_cube_path)\n        else:\n            # If the file is not found in the first directory, try the second directory\n            mask_cube_path = os.path.join(\"/kaggle/input/masks-part2\", f\"{cube_id}.tif\")\n            if os.path.exists(mask_cube_path):\n                mask_cube = tifffile.imread(mask_cube_path)\n            else:\n                mask_cube_path = os.path.join(\"/kaggle/input/masks-part3\", f\"{cube_id}.tif\")\n                if os.path.exists(mask_cube_path):\n                    mask_cube = tifffile.imread(mask_cube_path)\n                else:\n                    mask_cube_path = os.path.join(\"/kaggle/input/masks-part4\", f\"{cube_id}.tif\")\n                    if os.path.exists(mask_cube_path):\n                        mask_cube = tifffile.imread(mask_cube_path)\n                    else:\n                        mask_cube_path = os.path.join(\"/kaggle/input/masks-part5\", f\"{cube_id}.tif\")\n                        if os.path.exists(mask_cube_path):\n                            mask_cube = tifffile.imread(mask_cube_path)\n                        else:\n                            mask_cube = torch.rand(256, 256, 256)\n                            mask_cube = mask_cube.cpu().numpy().astype(np.int)\n                            print(\"Error: mask_cube.tif not found in both directories.\")\n\n        image_cube = combine_scan(os.path.join(self.dataset_dir, cube_id) ,size=512)\n        boxes = {}\n        for rprop in measure.regionprops(mask_cube):\n            boxes[rprop.label] = rprop.bbox, rprop.area\n\n        image_mean = image_cube.mean()\n        image_std = image_cube.std()\n        slice_size = self.slice_size\n        all_images = []\n#         labels = np.zeros((8,))\n        labels = labels_dict[cube_id]\n        for li in range(1, 8):\n            if li not in boxes:\n                all_images.append(np.zeros((3, self.slice_size, 256, 256)))\n            else:\n                bbox, area = boxes[li]\n                z1, z2 = bbox[0], bbox[3]\n                y1, y2 = max(bbox[1] - 16, 0), min(bbox[4] + 16, 256)\n                x1, x2 = max(bbox[2] - 16, 0), min(bbox[5] + 16, 256)\n                # if z2 - z1 < slice_size:\n                #     z1 = random.randint(max(z2 - slice_size, 0), z1)\n                #     z2 = z1 + slice_size\n                # todo: verify\n                if z2 - z1 < slice_size:\n                    diff = (slice_size - z2 + z1) // 2\n                    z1 = max(0, z1 - diff)\n                    z2 = z1 + slice_size\n                images = image_cube[z1:z2, y1 * 2:y2 * 2, x1 * 2:x2 * 2].copy()\n                masks = mask_cube[z1:z2, y1:y2, x1:x2].copy()\n                slice_size = self.slice_size\n\n                replay = None\n                image_crops = []\n                mask_crops = []\n                for i in range(images.shape[0]):\n                    image = images[i]\n                    mask = masks[i]\n                    h, w, = mask.shape\n                    mask = cv2.resize(mask, (w * 2, h * 2), interpolation=cv2.INTER_NEAREST)\n                    if replay is None:\n                        sample = self.transforms(image=image, mask=mask)\n                        replay = sample[\"replay\"]\n                    else:\n                        sample = ReplayCompose.replay(replay, image=image, mask=mask)\n                    image_ = sample[\"image\"]\n                    image_crops.append(image_)\n                    mask_crops.append(sample[\"mask\"])\n                images = np.array(image_crops).astype(np.float32)\n                masks = np.array(mask_crops).astype(np.float32)\n                images = np.expand_dims(images, -1)\n                masks = np.expand_dims(masks, -1)\n                images = (images - image_mean) / image_std\n\n                images = np.concatenate([images, images, masks], axis=-1)\n                h = images.shape[0]\n                if h > slice_size:\n                    images = images[: slice_size]\n                    all_images.append(np.moveaxis(images, -1, 0))\n                    images = images[-slice_size:]\n                    all_images.append(np.moveaxis(images, -1, 0))\n                else:\n                    if h != slice_size:\n                        tmp = np.zeros((slice_size, *images.shape[1:]))\n                        tmp[:h] = images\n                        images = tmp\n                    all_images.append(np.moveaxis(images, -1, 0))\n\n        sample = {}\n        sample['image'] = torch.from_numpy(np.array(all_images)).float()\n        sample['label'] = torch.from_numpy(np.array(labels)).float()\n        sample['cube_id'] = cube_id\n        return sample\n\n    def __len__(self):\n        return len(self.cases)","metadata":{"execution":{"iopub.status.busy":"2024-04-16T03:03:08.125198Z","iopub.execute_input":"2024-04-16T03:03:08.125625Z","iopub.status.idle":"2024-04-16T03:03:08.163055Z","shell.execute_reply.started":"2024-04-16T03:03:08.125581Z","shell.execute_reply":"2024-04-16T03:03:08.161995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import random_split\nimport numpy as np\nfrom torch.utils.data import Subset \n\nvalidation_split = 0.1\ntest_split = 0.1\n\ndataset_dir = \"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train_images/\"\n\ncases = os.listdir(dataset_dir)\n# Create dataset\ndataset = DatasetCrops(dataset_dir=dataset_dir, cases=cases)\n\n# Compute sizes\ntrain_size = int(len(dataset) * (1 - validation_split - test_split))\nval_test_size = len(dataset) - train_size\nval_size = int(val_test_size / 2)\ntest_size = val_test_size - val_size\n\n# Random split into train, validation, and test datasets\ntrain_dataset, val_test_dataset = random_split(dataset, [train_size, val_test_size])\nval_dataset, test_dataset = random_split(val_test_dataset, [val_size, test_size])\n\n# train_dataset = Subset(train_dataset, range(5))\n# val_dataset = Subset(val_dataset, range(5))\n# test_dataset = Subset(test_dataset, range(5))\n\n# Create data loaders\ntrain_loader = DataLoader(train_dataset, batch_size=1, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=1, shuffle=False)\ntest_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\nprint(len(train_loader), len(val_loader), len(test_loader)) \n\n# train_loader = train_loader[:100]","metadata":{"execution":{"iopub.status.busy":"2024-04-16T03:03:09.30334Z","iopub.execute_input":"2024-04-16T03:03:09.304247Z","iopub.status.idle":"2024-04-16T03:03:09.435184Z","shell.execute_reply.started":"2024-04-16T03:03:09.304213Z","shell.execute_reply":"2024-04-16T03:03:09.434254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import csv\ncsv_file_path = '/kaggle/working/Auto_Encoder_new_data.csv'\n\n# Column names\nfieldnames = ['Epoch','Train_Loss', 'Train_Accuracy', 'Val_Loss', 'Val_Accuracy']\n\n\n# Open the file outside of the with statement\ncsvfile = open(csv_file_path, 'w', newline='')\ncsv_writer = csv.DictWriter(csvfile, fieldnames=fieldnames)\n\n# Write the header\ncsv_writer.writeheader()\n\n# with open(csv_file_path, 'r') as csvfile:\n#     csv_reader = csv.reader(csvfile)\n#     for row in csv_reader:\n#         print(row)","metadata":{"execution":{"iopub.status.busy":"2024-04-16T03:03:10.550838Z","iopub.execute_input":"2024-04-16T03:03:10.551508Z","iopub.status.idle":"2024-04-16T03:03:10.560332Z","shell.execute_reply.started":"2024-04-16T03:03:10.551476Z","shell.execute_reply":"2024-04-16T03:03:10.559335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\ndef weighted_cross_entropy(predicted, label):\n    num_samples = predicted.shape[0]\n    \n    # Calculate element-wise losses\n    losses = -label * torch.log(predicted) - (1 - label) * torch.log(1 - predicted)\n#     print(\"losses are \",losses)\n    \n    # Sum the total_loss for all inputs and divide by the number of samples\n    final_loss = torch.sum(losses) / num_samples\n#     print(final_loss)\n    \n    return final_loss","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import nn\nimport torch.optim as optim\nfrom sklearn.metrics import accuracy_score\nfrom torch.utils.data import random_split\n\nlearning_rate = 0.001  # Define the learning rate\n\ndef train_model(model, imgs):\n    preds = []\n    with torch.no_grad():\n        for i in range(len(imgs)):\n            with torch.cuda.amp.autocast():\n                output = model(imgs[i:i + 1])\n            pred_slice = (output.float()).cpu().numpy().astype(np.float32)\n            with torch.cuda.amp.autocast():\n                output = model(torch.flip(imgs[i:i + 1], dims=(-1,)))\n            pred_slice += (output.float()).cpu().numpy().astype(np.float32)\n            pred_slice /= 2\n            preds.append(pred_slice)\n    preds = np.max(np.array(preds), axis=0)\n    preds[np.isnan(preds)] = 0.01\n    return preds\n\ndef train_classification(model, cases: List, num_epochs: int = 1):\n    # Define loss criterion for multi-label classification\n    optimizer = optim.Adam(model.parameters(), lr=learning_rate)  # Define optimizer\n    loss_criterion = weighted_cross_entropy  # You can adjust the loss function as needed\n    best_metric = -1\n    best_metric_epoch = -1\n    best_metrics_epochs_and_time = [[], [], []]\n    total_start = time.time()\n    total_images = 0\n    for epoch in range(start_epoch, num_epochs):\n        total_images = 0  # Reset total_images for each epoch\n        print(f\"Epoch [{epoch + 1}/{num_epochs}]\")\n        # Training phase\n        train_loss = 0.0\n        train_preds = []\n        train_labels = []\n\n        for sample in tqdm(train_loader, desc=\"Training\"):\n            imgs = sample[\"image\"].cuda().float()[0]\n            cube_id = sample[\"cube_id\"][0]\n            with torch.no_grad():\n                preds = []\n                model.train()  # Set the model to training mode\n                x = train_model(model, imgs)\n                preds.append(x.squeeze())\n                preds = np.average(np.array(preds), axis=0)\n                preds = np.clip(preds, 0.01, 0.99)\n\n                loss = loss_criterion(torch.tensor(preds).squeeze(), sample['label'].squeeze())  # Calculate loss\n                loss.requires_grad = True\n                optimizer.zero_grad()  # Zero the gradients\n                loss.backward()  # Backpropagate the gradients\n                optimizer.step()  # Update model parameters\n\n                train_loss += loss.item() * imgs.size(0)\n                total_images += imgs.size(0)\n\n                train_preds.extend(preds)\n                train_labels.extend(sample['label'].squeeze().cpu().numpy())\n\n        train_loss /= total_images\n        train_accuracy = accuracy_score(np.array(train_labels), np.array(train_preds) >= 0.48)\n        print(f\"Train Loss: {train_loss:.4f} | Train Accuracy: {train_accuracy:.4f}\")\n\n        # Validation phase\n        val_loss = 0.0\n        val_preds = []\n        val_labels = []\n        model.eval()  # Set the model to evaluation mode\n        total_val_images = 0  # Track total validation images\n        with torch.no_grad():\n            for val_sample in tqdm(val_loader, desc=\"Validation\"):\n                val_imgs = val_sample[\"image\"].cuda().float()[0]\n                total_val_images += val_imgs.size(0)  # Increment total validation images count\n                val_preds_batch = []\n                val_x = train_model(model, val_imgs)\n                val_preds_batch.append(val_x.squeeze())\n                val_preds_batch = np.average(np.array(val_preds_batch), axis=0)\n                val_preds_batch = np.clip(val_preds_batch, 0.01, 0.99)\n                val_loss += loss_criterion(torch.tensor(val_preds_batch), val_sample['label'].squeeze()).item() * val_imgs.size(0)\n                val_preds.extend(val_preds_batch)\n                val_labels.extend(val_sample['label'].squeeze().cpu().numpy())\n\n        val_loss /= total_val_images  # Calculate average loss per image\n        val_accuracy = accuracy_score(np.array(val_labels), np.array(val_preds) >= 0.5)\n        print(f\"Validation Loss: {val_loss:.4f} | Validation Accuracy: {val_accuracy:.4f}\")\n\n        with open(csv_file_path, 'a', newline='') as csvfile:\n            csv_writer = csv.DictWriter(csvfile, fieldnames=fieldnames)\n            csv_writer.writerow({\n                'Epoch': epoch+1,\n                'Train_Loss': f\"{train_loss:.4f}\",\n                'Train_Accuracy': f\"{train_accuracy:.4f}\",\n                'Val_Loss': f\"{val_loss:.4f}\",\n                'Val_Accuracy': f\"{val_accuracy:.4f}\"            \n            })\n\n        if val_accuracy > best_metric:\n            best_metric = val_accuracy\n            best_metric_epoch = epoch + 1\n            best_metrics_epochs_and_time[0].append(best_metric)\n            best_metrics_epochs_and_time[1].append(best_metric_epoch)\n            best_metrics_epochs_and_time[2].append(time.time() - total_start)\n            torch.save(\n                model.state_dict(),\n                os.path.join(\"/kaggle/working/\", f\"Auto_Encoder_new_best_metric_model_{epoch+1}.pth\"),\n            )\n            print(\"saved new best metric model\")\n        else:\n            chk_file_name = \"Auto_Encoder_new_epoch_\" + str(epoch+1) + \"_model.pth\"    \n            torch.save(\n                model.state_dict(),\n                os.path.join(\"/kaggle/working/\", chk_file_name),\n            )","metadata":{"execution":{"iopub.status.busy":"2024-04-16T03:03:34.90472Z","iopub.execute_input":"2024-04-16T03:03:34.905065Z","iopub.status.idle":"2024-04-16T03:03:34.931056Z","shell.execute_reply.started":"2024-04-16T03:03:34.905041Z","shell.execute_reply":"2024-04-16T03:03:34.930062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"output = []\n\nstart_epoch = 0\nlatest_checkpoint_path = \"/kaggle/input/auto-encoder-new/Auto_Encoder_new_best_metric_model_42.pth\"\nif os.path.exists(latest_checkpoint_path):\n    checkpoint = torch.load(latest_checkpoint_path)\n    model.load_state_dict(checkpoint)\n    start_epoch = 46\n    print(\"checkpoint loaded\")\nelse:\n    print(\"No checkpoint. Starting from scratch\")\n# print(model)\ndevice=torch.device(\"cuda:0\")\nmodel = model.to(device)\ntrain_classification(model, cases, 51)","metadata":{"execution":{"iopub.status.busy":"2024-04-16T03:03:36.42998Z","iopub.execute_input":"2024-04-16T03:03:36.430514Z","iopub.status.idle":"2024-04-16T03:04:14.493967Z","shell.execute_reply.started":"2024-04-16T03:03:36.430478Z","shell.execute_reply":"2024-04-16T03:04:14.492756Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import List\nimport torch\nimport numpy as np\nfrom sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score, roc_auc_score\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\n\ndef predict_classification(models: List[nn.Module]):\n#     test_dataset = DatasetCrops(dataset_dir=test_dataset_dir, cases=cases)\n#     dataloader = DataLoader(\n#         test_dataset, batch_size=1, sampler=None, shuffle=False, num_workers=1, pin_memory=False\n#     )\n    \n    def predict_model(model, imgs):\n        preds = []\n        with torch.no_grad():\n            for i in range(len(imgs)):\n                with torch.cuda.amp.autocast():\n                    output = model(imgs[i:i + 1])\n                pred_slice = (output.float()).cpu().numpy().astype(np.float32)\n                with torch.cuda.amp.autocast():\n                    output = model(torch.flip(imgs[i:i + 1], dims=(-1,)))\n                pred_slice += (output.float()).cpu().numpy().astype(np.float32)\n                pred_slice /= 2\n                preds.append(pred_slice)\n        preds = np.max(np.array(preds), axis=0)\n        preds[np.isnan(preds)] = 0.01\n        return preds\n        \n    output_list = []  # Initialize the output list\n    all_labels = []\n    all_predictions = []\n#     print(\".......................................................................\", len(test_loader))\n    for sample in tqdm(test_loader):\n        imgs = sample[\"image\"].cuda().float()[0]\n        cube_id = sample[\"cube_id\"][0]\n        labels = sample[\"label\"].cpu().numpy()\n        # Option 1: Remove outer list\n        labels = labels[0]\n#         print(\"len(labels), labels\", len(labels), labels)\n        all_labels.extend(labels)\n        \n        with torch.no_grad():\n            preds = []\n            \n            preds.append(predict_model(model, imgs))  # Pass imgs to predict_model\n            preds = np.average(np.array(preds), axis=0)\n            preds = np.clip(preds, 0.01, 0.99)\n#             print(preds)\n            all_predictions.extend(preds.squeeze())\n#             output_list.append([cube_id, preds])\n    \n    # Calculate metrics\n    \n    accuracy = accuracy_score(all_labels, (np.array(all_predictions) >= 0.5).astype(int))\n    precision = precision_score(all_labels, (np.array(all_predictions) >= 0.5).astype(int))\n    recall = recall_score(all_labels, (np.array(all_predictions) >= 0.5).astype(int))\n    f1 = f1_score(all_labels, (np.array(all_predictions) >= 0.5).astype(int))\n    print(\"accuracy, precision, recall, f1\", accuracy, precision, recall, f1)\n   # auc = roc_auc_score(all_labels, (np.array(all_predictions)))\n    \n    return output_list, accuracy, precision, recall, f1","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load model from the specified path\nmodel_path = \"/kaggle/input/auto-encoder-new/Auto_Encoder_new_best_metric_model_42.pth\"\nif os.path.exists(model_path):\n    model.load_state_dict(torch.load(model_path))\n    print(\"Model loaded for testing.\")\nelse:\n    print(\"Model path not found. Ensure the correct path is provided.\")\n\noutput = []\n# cases = os.listdir(train_dataset_dir)\npredict_classification(model)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import shutil\n# shutil.rmtree('/kaggle/working/seg_preds')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}