{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":71698,"databundleVersionId":7906362,"sourceType":"competition"},{"sourceId":8301420,"sourceType":"datasetVersion","datasetId":4931836},{"sourceId":8523633,"sourceType":"datasetVersion","datasetId":5089654},{"sourceId":8524811,"sourceType":"datasetVersion","datasetId":5090495},{"sourceId":180090464,"sourceType":"kernelVersion"},{"sourceId":43019,"sourceType":"modelInstanceVersion","modelInstanceId":36136},{"sourceId":43079,"sourceType":"modelInstanceVersion","modelInstanceId":36183}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 1. Libraries","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.backends.cudnn as cudnn\nimport torchvision.transforms as transforms\nimport argparse\nimport os\nimport time\nimport math\nimport glob\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.datasets import VisionDataset\nfrom torchvision import transforms\nimport tifffile\n\nimport torch.nn.functional as F\nimport numpy as np","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-27T17:10:20.948659Z","iopub.execute_input":"2024-05-27T17:10:20.949484Z","iopub.status.idle":"2024-05-27T17:10:20.956947Z","shell.execute_reply.started":"2024-05-27T17:10:20.949444Z","shell.execute_reply":"2024-05-27T17:10:20.955856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define command line arguments\nclass Args:\n    def __init__(self):\n        self.save = 'logs'\n        self.batchSize = 32\n        self.learningRate = 0.01\n        self.learningRateDecay = 1e-7\n        self.weightDecay = 0.0005\n        self.momentum = 0.9\n        self.epoch_step = 20\n        self.gpu_index = 0\n        self.max_epoch = 25\n        self.jitter_step = 2\n        self.model = 'subvol'\n        self.train_data = 'data/modelnet40_60x/train_data.txt'\n        self.test_data = 'data/modelnet40_60x/test_data.txt'\n        self.num_classes = 12\n        self.dataset = \"siz\"","metadata":{"execution":{"iopub.status.busy":"2024-05-27T17:10:20.961891Z","iopub.execute_input":"2024-05-27T17:10:20.962650Z","iopub.status.idle":"2024-05-27T17:10:20.970481Z","shell.execute_reply.started":"2024-05-27T17:10:20.962616Z","shell.execute_reply":"2024-05-27T17:10:20.969480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2. Dataset creation","metadata":{}},{"cell_type":"code","source":"import SimpleITK as sitk\n\nclass CustomTiffDataset(VisionDataset):\n    def __init__(self, tiff_paths, transform=None):\n        super(CustomTiffDataset, self).__init__(root=None, transform=transform)\n        self.tiff_paths = tiff_paths\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.tiff_paths)\n\n    def __getitem__(self, index):\n        tiff_path = self.tiff_paths[index]\n        label = tiff_path.split(\"/\")[-2]\n        \n        encoded_label = one_hot_encode([label], classes)[0]\n        encoded_label = torch.tensor(encoded_label)\n        \n        # Load TIFF file as 3D volume\n        image = sitk.ReadImage(tiff_path)\n        image_array = sitk.GetArrayFromImage(image)\n#         print(tiff_path, image_array.shape, \"getting shape\")\n#         print(type(image_array), \"type\")\n        \n        # Reshape the image array to match [channel, depth, width, height] format\n        # Assuming your TIFF images represent 3D volumes, with channels as 1\n        image_array = image_array.reshape((1,) + image_array.shape)\n        \n#         print(\"getting index transformed\")\n        if self.transform:\n            # Apply transformations if any\n            image_array = self.transform(image_array)\n            \n#         print(\"got index transformed\")\n        \n        image_array = image_array.astype(np.float32)\n        image_array = torch.Tensor(image_array)\n        return image_array, encoded_label\n","metadata":{"execution":{"iopub.status.busy":"2024-05-27T17:10:20.972686Z","iopub.execute_input":"2024-05-27T17:10:20.973285Z","iopub.status.idle":"2024-05-27T17:10:20.983767Z","shell.execute_reply.started":"2024-05-27T17:10:20.973254Z","shell.execute_reply":"2024-05-27T17:10:20.982881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"paths = glob.glob('/kaggle/input/bugnist-slicing/train/*/*.tif')\n\nclasses = []\nclass_paths = glob.glob(\"/kaggle/input/bugnist-slicing/train/*\")\nfor folder_path in class_paths:\n    folder_name = folder_path.split(\"/\")[-1]\n    classes.append(folder_name)\nprint(classes)","metadata":{"execution":{"iopub.status.busy":"2024-05-27T17:10:20.984944Z","iopub.execute_input":"2024-05-27T17:10:20.985252Z","iopub.status.idle":"2024-05-27T17:10:21.033762Z","shell.execute_reply.started":"2024-05-27T17:10:20.985224Z","shell.execute_reply":"2024-05-27T17:10:21.032809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def one_hot_encode(labels, classes):\n    encoded_labels = []\n    for label in labels:\n        encoded = [0.0] * len(classes)\n        index = classes.index(label)\n        encoded[index] = 1.0\n        encoded_labels.append(encoded)\n    return encoded_labels\n\n# Example usage:\nlabels = ['CF', 'WO']\nencoded_labels = one_hot_encode(labels, classes)\nprint(encoded_labels)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-27T17:10:21.036233Z","iopub.execute_input":"2024-05-27T17:10:21.036599Z","iopub.status.idle":"2024-05-27T17:10:21.044246Z","shell.execute_reply.started":"2024-05-27T17:10:21.036569Z","shell.execute_reply":"2024-05-27T17:10:21.043284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\nimport os\n\ndef create_datasets(arg, paths_group, classes):\n    for c in classes:\n        if args.dataset == \"normal\":\n            c_paths = paths\n        else:\n            c_paths = glob.glob(f'/kaggle/input/bugnist-{arg.dataset}/{c}/*.tif')\n        c_paths.sort()  # Sort paths to ensure consistent splitting\n\n        tifCount = len(c_paths)\n        tifTrainCount = int(tifCount * 0.8)  # 80% for training\n        tifValCount = int(tifCount * 0.1)    # 10% for validation\n        tifTestCount = tifCount - tifTrainCount - tifValCount  # Remaining 10% for testing\n\n        c_groups = {\n            'train': c_paths[:tifTrainCount],\n            'validate': c_paths[tifTrainCount: tifTrainCount + tifValCount],\n            'test': c_paths[-tifTestCount:]\n        }\n        \n        for x in ['train', 'validate', 'test']:\n            paths_group[x].extend(c_groups[x])","metadata":{"execution":{"iopub.status.busy":"2024-05-27T17:10:21.045504Z","iopub.execute_input":"2024-05-27T17:10:21.046231Z","iopub.status.idle":"2024-05-27T17:10:21.055548Z","shell.execute_reply.started":"2024-05-27T17:10:21.046200Z","shell.execute_reply":"2024-05-27T17:10:21.054680Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3. Model setting","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass Block(nn.Module):\n    def __init__(self, in_channels, out_channels, kT, kW, kH, dT=1, dW=1, dH=1):\n        super(Block, self).__init__()\n        self.conv = nn.Conv3d(in_channels, out_channels, kernel_size=(kT, kW, kH), stride=(dT, dW, dH))\n        self.bn = nn.BatchNorm3d(out_channels)\n        self.relu = nn.ReLU(inplace=True)\n    \n    def forward(self, x):\n        x = self.conv(x)\n        x = self.bn(x)\n        x = self.relu(x)\n        return x\n    \ndef msr_init(layer):\n    if isinstance(layer, nn.Conv3d):\n        n = layer.kernel_size[0] * layer.kernel_size[1] * layer.kernel_size[2] * layer.out_channels\n        layer.weight.data.normal_(0, (2. / n) ** 0.5)\n        if layer.bias is not None:\n            layer.bias.data.zero_()\n\nclass SubVolNet(nn.Module):\n    def __init__(self, num_classes):\n        super(SubVolNet, self).__init__()\n        self.block0 = nn.Sequential(\n            Block(1, 48, 6, 6, 6, 2, 2, 2),\n            Block(48, 48, 1, 1, 1),\n            Block(48, 48, 1, 1, 1),\n            nn.Dropout(0.2)\n        )\n        \n        self.block1 = nn.Sequential(\n            Block(48, 48, 6, 6, 6, 2, 2, 2),\n            Block(48, 48, 1, 1, 1),\n            Block(48, 48, 1, 1, 1),\n            nn.Dropout(0.2)\n        )\n        \n        self.block2 = nn.Sequential(\n            Block(48, 160, 5, 5, 5, 2, 2, 2),\n            Block(160, 160, 1, 1, 1),\n            Block(160, 160, 1, 1, 1),\n            nn.Dropout(0.2)\n        )\n        \n        self.block3 = nn.Sequential(\n            Block(160, 512, 3, 3, 3, 2, 2, 2),\n            Block(512, 512, 1, 1, 1),\n            Block(512, 512, 1, 1, 1),\n            nn.Dropout(0.2)\n        )\n        \n        self.fc_blocks = nn.ModuleList([nn.Sequential(\n            nn.Linear(512, num_classes)\n        ) for _ in range(8)])\n        \n        self.w = nn.Sequential(\n#             nn.Flatten(),\n            nn.Linear(4096, 2048),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.5),\n            nn.Linear(2048, 2048),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.5),\n            nn.Linear(2048, num_classes)\n        )\n        \n    def forward(self, x):\n        x = self.block0(x)\n        x = self.block1(x)\n        x = self.block2(x)\n        x = self.block3(x)\n        \n        x = x.view(x.size(0), 512, 8)\n        \n        out = []\n        for i in range(8):\n            out.append(self.fc_blocks[i](x[:, :, i]))\n        \n        w_out = x.view(x.size(0), 4096)\n        w_out = self.w(w_out)\n        out.append(w_out)\n        \n        return out\n\ndef get_model(model_name, num_classes):\n    if (model_name == \"subvol\"):\n        model = SubVolNet(num_classes)\n         # Initialize the network\n        model.apply(msr_init)\n\n        # Define the criterion\n        criterion = [nn.CrossEntropyLoss() for _ in range(9)]\n        criterion = nn.ModuleList(criterion)\n        \n        return model, criterion\n#         model.MSRinit()\n        return model, criterion\n    elif (model_name == \"3dnin_fc\"):\n        model = Model3dninfc()\n        return model\n    return None","metadata":{"execution":{"iopub.status.busy":"2024-05-27T17:10:21.056672Z","iopub.execute_input":"2024-05-27T17:10:21.057438Z","iopub.status.idle":"2024-05-27T17:10:21.079544Z","shell.execute_reply.started":"2024-05-27T17:10:21.057402Z","shell.execute_reply":"2024-05-27T17:10:21.078360Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"args = Args()\n\n# Set GPU\ndevice = torch.device(\"cuda:\" + str(args.gpu_index) if torch.cuda.is_available() else \"cpu\")\ncudnn.benchmark = True\n\n# Define your transformations here if needed\ntransform = transforms.Compose([\n    # Add your transformations here\n    # e.g., ToTensor(), Normalize(), etc.\n])\n\npaths_group = {\n    'train': [],\n    'validate': [],\n    'test': []\n}\ncreate_datasets(args, paths_group, classes)\n\ntiff_files = {x: paths_group[x] for x in ['train', 'validate', 'test']}\ndatasets = {x: CustomTiffDataset(tiff_files[x], transform=transform) for x in ['train', 'validate', 'test']}\ndataset_sizes = {x: len(paths_group[x]) for x in ['train', 'validate', 'test']}\n\ndataloaders = {x: DataLoader(datasets[x], batch_size=args.batchSize, shuffle=True, num_workers=2)\n               for x in ['train', 'validate', 'test']}","metadata":{"execution":{"iopub.status.busy":"2024-05-27T17:10:21.105900Z","iopub.execute_input":"2024-05-27T17:10:21.106159Z","iopub.status.idle":"2024-05-27T17:10:21.166982Z","shell.execute_reply.started":"2024-05-27T17:10:21.106136Z","shell.execute_reply":"2024-05-27T17:10:21.166237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for inputs, labels in dataloaders['train']:\n    print(inputs.shape, labels)\n    break","metadata":{"execution":{"iopub.status.busy":"2024-05-27T17:10:21.168610Z","iopub.execute_input":"2024-05-27T17:10:21.168858Z","iopub.status.idle":"2024-05-27T17:10:22.964384Z","shell.execute_reply.started":"2024-05-27T17:10:21.168836Z","shell.execute_reply":"2024-05-27T17:10:22.963459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4. Model evaluation","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import f1_score, precision_score, recall_score, roc_auc_score\n\ndef evaluate_model(model_path, dataloader):\n    # Load the model\n    model, criterion = get_model(args.model, args.num_classes)  # Define your function to get the model\n    model.load_state_dict(torch.load(model_path))\n    model = model.to(device)\n    model.eval()  # Set the model to evaluation mode\n\n    running_corrects = 0\n    total_samples = 0\n    all_preds = []\n    all_labels = []\n\n    # Iterate over data.\n    for inputs, labels in dataloader:\n        label_indexes = torch.max(labels, 1).indices\n        label_indexes = label_indexes.to(device)\n        inputs = inputs.to(device)\n        labels = labels.to(device)\n        total_samples += inputs.size(0)\n\n        # Forward pass\n        with torch.no_grad():\n            outputs = model(inputs)\n            \n            losses = [criterion[i](outputs[i], labels) for i in range(len(outputs))]\n            loss = sum(losses)\n\n        # Calculate correct predictions\n        if labels.dim() > 1:  # Check if labels are one-hot encoded\n            _, labels = torch.max(labels, 1)  # Convert one-hot encoded labels to class indices\n        for i in range(len(outputs)):\n            _, preds = torch.max(outputs[i], 1)\n            running_corrects += torch.sum(preds == label_indexes)\n\n\n            # Store predictions and labels for additional metrics calculation\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n\n    # Calculate accuracy\n    accuracy = running_corrects.double() / (total_samples * len(outputs))\n    \n    # Calculate F1-score\n    f1 = f1_score(all_labels, all_preds, average='weighted')\n\n    # Calculate precision\n    precision = precision_score(all_labels, all_preds, average='weighted')\n    \n    # Calculate recall\n    recall = recall_score(all_labels, all_preds, average='weighted')\n    \n    # Calculate AUC-ROC\n#     auc_roc = roc_auc_score(all_labels, all_preds)\n\n    return accuracy.item(), f1, precision, recall\n\n# Usage:\nmodel_path = '/kaggle/input/subvolume-sup-training/subvol_siz.pth'\naccuracy, f1, precision, recall = evaluate_model(model_path, dataloaders['test'])\nprint(f'Test Accuracy: {accuracy:.4f}')\nprint(f'F1-score: {f1:.4f}')\nprint(f'Precision: {precision:.4f}')\nprint(f'Recall: {recall:.4f}')\n# print(f'AUC-ROC: {auc_roc:.4f}')\n","metadata":{"execution":{"iopub.status.busy":"2024-05-27T17:10:43.672266Z","iopub.execute_input":"2024-05-27T17:10:43.673163Z","iopub.status.idle":"2024-05-27T17:11:08.927105Z","shell.execute_reply.started":"2024-05-27T17:10:43.673131Z","shell.execute_reply":"2024-05-27T17:11:08.926065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}