{"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":8523633,"sourceType":"datasetVersion","datasetId":5089654},{"sourceId":8524811,"sourceType":"datasetVersion","datasetId":5090495},{"sourceId":179890551,"sourceType":"kernelVersion"}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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# from models import *  # Import your model definitions\n# from provider import *  # Import your data loading functions\n\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.datasets import VisionDataset\nfrom torchvision import transforms\nimport tifffile","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-27T18:36:15.562604Z","iopub.execute_input":"2024-05-27T18:36:15.563012Z","iopub.status.idle":"2024-05-27T18:36:23.437102Z","shell.execute_reply.started":"2024-05-27T18:36:15.562983Z","shell.execute_reply":"2024-05-27T18:36:23.436225Z"},"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 = \"ess\"","metadata":{"execution":{"iopub.status.busy":"2024-05-27T18:36:23.438840Z","iopub.execute_input":"2024-05-27T18:36:23.439247Z","iopub.status.idle":"2024-05-27T18:36:23.445929Z","shell.execute_reply.started":"2024-05-27T18:36:23.439218Z","shell.execute_reply":"2024-05-27T18:36:23.444855Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model definition and functions","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","metadata":{"execution":{"iopub.status.busy":"2024-05-27T18:36:23.447079Z","iopub.execute_input":"2024-05-27T18:36:23.447432Z","iopub.status.idle":"2024-05-27T18:36:23.467000Z","shell.execute_reply.started":"2024-05-27T18:36:23.447401Z","shell.execute_reply":"2024-05-27T18:36:23.466081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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-27T18:36:23.469238Z","iopub.execute_input":"2024-05-27T18:36:23.469517Z","iopub.status.idle":"2024-05-27T18:36:23.481660Z","shell.execute_reply.started":"2024-05-27T18:36:23.469487Z","shell.execute_reply":"2024-05-27T18:36:23.480736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset classes and functions","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        \n        if self.transform:\n            # Apply transformations if any\n            image_array = self.transform(image_array)\n        \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-27T18:36:23.483223Z","iopub.execute_input":"2024-05-27T18:36:23.483850Z","iopub.status.idle":"2024-05-27T18:36:23.897093Z","shell.execute_reply.started":"2024-05-27T18:36:23.483816Z","shell.execute_reply":"2024-05-27T18:36:23.896342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"paths = glob.glob('/kaggle/input/bugnist2024fgvc/BugNIST_DATA/train/*/*.tif')\n\nclasses = []\nclass_paths = glob.glob(\"/kaggle/input/bugnist2024fgvc/BugNIST_DATA/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-27T18:36:23.898123Z","iopub.execute_input":"2024-05-27T18:36:23.898376Z","iopub.status.idle":"2024-05-27T18:36:25.567416Z","shell.execute_reply.started":"2024-05-27T18:36:23.898354Z","shell.execute_reply":"2024-05-27T18:36:25.566518Z"},"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-27T18:36:25.568701Z","iopub.execute_input":"2024-05-27T18:36:25.568987Z","iopub.status.idle":"2024-05-27T18:36:25.575437Z","shell.execute_reply.started":"2024-05-27T18:36:25.568963Z","shell.execute_reply":"2024-05-27T18:36:25.574391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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-27T18:36:25.576812Z","iopub.execute_input":"2024-05-27T18:36:25.577188Z","iopub.status.idle":"2024-05-27T18:36:25.591143Z","shell.execute_reply.started":"2024-05-27T18:36:25.577155Z","shell.execute_reply":"2024-05-27T18:36:25.590259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training function","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\n\ndef train_model(model, criterion, optimizer, scheduler, dataloaders, dataset_sizes, device, num_epochs=25):\n    for epoch in range(num_epochs):\n        print(f'Epoch {epoch}/{num_epochs - 1}')\n        print('-' * 10)\n\n        # Each epoch has a training and validation phase\n        for phase in ['train', 'validate']:\n            if phase == 'train':\n                model.train()  # Set model to training mode\n            else:\n                model.eval()   # Set model to evaluate mode\n\n            running_loss = 0.0\n            running_corrects = 0\n\n            # Iterate over data.\n            for inputs, labels in dataloaders[phase]:\n#                 inputs = inputs.to(torch.float32)\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\n                # Zero the parameter gradients\n                optimizer.zero_grad()\n\n                # Forward\n                # track history if only in train\n                with torch.set_grad_enabled(phase == 'train'):\n                    outputs = model(inputs)\n#                     print(inputs.shape, outputs[0].shape, labels.shape)\n                    losses = [criterion[i](outputs[i], labels) for i in range(len(outputs))]\n                    loss = sum(losses)\n\n                    # Backward + optimize only if in training phase\n                    if phase == 'train':\n                        loss.backward()\n                        optimizer.step()\n\n                # Statistics\n                running_loss += loss.item() * inputs.size(0)\n                for i in range(len(outputs)):\n                    _, preds = torch.max(outputs[i], 1)\n#                     print(preds, labels)\n                    running_corrects += torch.sum(preds == label_indexes)\n\n            if phase == 'train':\n                scheduler.step()\n\n            epoch_loss = running_loss / dataset_sizes[phase]\n            epoch_acc = running_corrects.double() / (dataset_sizes[phase] * len(outputs))\n\n            print(f'{phase} Loss: {epoch_loss:.4f} Acc: {epoch_acc:.4f}')\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2024-05-27T18:36:25.592492Z","iopub.execute_input":"2024-05-27T18:36:25.592870Z","iopub.status.idle":"2024-05-27T18:36:25.605497Z","shell.execute_reply.started":"2024-05-27T18:36:25.592840Z","shell.execute_reply":"2024-05-27T18:36:25.604507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transformation","metadata":{}},{"cell_type":"code","source":"import torch\nimport torchvision.transforms as transforms\nimport numpy as np\n\nimport numpy as np\nimport cv2\n\ndef depth_rotation(volume, angle):\n    \"\"\"\n    Apply depth rotation augmentation to a 3D volumetric object.\n    \n    Parameters:\n        volume (np.ndarray): 3D volumetric object represented as a numpy array.\n        angle(float): Angle in degrees for augmentation.\n    \n    Returns:\n        np.ndarray: Augmented 3D volumetric object.\n    \"\"\"\n    # Get dimensions of the volumetric object\n    depth, height, width = volume.shape\n    \n    # Calculate center of the object\n    center = (width // 2, height // 2)\n    \n    # Initialize empty array for augmented volume\n    augmented_volume = np.zeros_like(volume)\n    \n    # Iterate over each slice of the volumetric object\n    for z in range(depth):\n        # Create rotation matrix for the given angle (in degrees)\n        rotation_matrix = cv2.getRotationMatrix2D(center, angle, 1.0)\n        \n        # Rotate slice around the center\n        rotated_slice = cv2.warpAffine(volume[z], rotation_matrix, (width, height))\n        \n        # Assign the rotated slice to the corresponding place in the augmented volume\n        augmented_volume[z] = rotated_slice\n    \n    return augmented_volume\n\nimport numpy as np\nimport cv2\n\ndef width_rotation(volume, angle):\n    \"\"\"\n    Apply width rotation augmentation to a 3D volumetric object.\n    \n    Parameters:\n        volume (np.ndarray): 3D volumetric object represented as a numpy array.\n        angle (float): Angle in degrees for augmentation.\n    \n    Returns:\n        np.ndarray: Augmented 3D volumetric object.\n    \"\"\"\n    # Get dimensions of the volumetric object\n    depth, height, width = volume.shape\n    \n    # Calculate center of the object\n    center = (height // 2, depth // 2)\n    \n    # Initialize empty array for augmented volume\n    augmented_volume = np.zeros_like(volume)\n    \n    # Iterate over each width slice of the volumetric object\n    for x in range(width):\n        # Create rotation matrix for the given angle (in degrees)\n        rotation_matrix = cv2.getRotationMatrix2D(center, angle, 1.0)\n        \n        # Rotate slice around the center\n        rotated_slice = cv2.warpAffine(volume[:, :, x], rotation_matrix, (height, depth))\n        \n        # Assign the rotated slice to the corresponding place in the augmented volume\n        augmented_volume[:, :, x] = rotated_slice\n    \n    return augmented_volume\n\ndef height_rotation(volume, angle):\n    \"\"\"\n    Apply height rotation augmentation to a 3D volumetric object.\n    \n    Parameters:\n        volume (np.ndarray): 3D volumetric object represented as a numpy array.\n        angle (float): Angle in degrees for augmentation.\n    \n    Returns:\n        np.ndarray: Augmented 3D volumetric object.\n    \"\"\"\n    # Get dimensions of the volumetric object\n    depth, height, width = volume.shape\n    \n    # Calculate center of the object\n    center = (width // 2, depth // 2)\n    \n    # Initialize empty array for augmented volume\n    augmented_volume = np.zeros_like(volume)\n    \n    # Iterate over each height slice of the volumetric object\n    for y in range(height):\n        # Create rotation matrix for the given angle (in degrees)\n        rotation_matrix = cv2.getRotationMatrix2D(center, angle, 1.0)\n        \n        # Rotate slice around the center\n        rotated_slice = cv2.warpAffine(volume[:, y, :], rotation_matrix, (width, depth))\n        \n        # Assign the rotated slice to the corresponding place in the augmented volume\n        augmented_volume[:, y, :] = rotated_slice\n    \n    return augmented_volume\n\ndef rotate(volume, d_angle = 0, h_angle = 0, w_angle = 0):\n    aug_vol = depth_rotation(volume, d_angle)\n    aug_vol = width_rotation(aug_vol, h_angle)\n    aug_vol = height_rotation(aug_vol, w_angle)\n    return aug_vol","metadata":{"execution":{"iopub.status.busy":"2024-05-27T18:36:25.609134Z","iopub.execute_input":"2024-05-27T18:36:25.609486Z","iopub.status.idle":"2024-05-27T18:36:25.853700Z","shell.execute_reply.started":"2024-05-27T18:36:25.609441Z","shell.execute_reply":"2024-05-27T18:36:25.852794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\n\nclass RandomRotation3D:\n    def __init__(self, d_angle_range=(-10, 10), h_angle_range=(-10, 10), w_angle_range=(-10, 10)):\n        self.d_angle_range = d_angle_range\n        self.h_angle_range = h_angle_range\n        self.w_angle_range = w_angle_range\n\n    def __call__(self, volume):\n        d_angle = random.uniform(*self.d_angle_range)\n        h_angle = random.uniform(*self.h_angle_range)\n        w_angle = random.uniform(*self.w_angle_range)\n        return rotate(volume, d_angle, h_angle, w_angle)","metadata":{"execution":{"iopub.status.busy":"2024-05-27T18:36:25.854830Z","iopub.execute_input":"2024-05-27T18:36:25.855180Z","iopub.status.idle":"2024-05-27T18:36:25.862305Z","shell.execute_reply.started":"2024-05-27T18:36:25.855147Z","shell.execute_reply":"2024-05-27T18:36:25.861349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Setting up","metadata":{}},{"cell_type":"code","source":"args = Args()\n\n# Print chosen options\n# print(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 transformations\ntransform = transforms.Compose([\n    RandomRotation3D(d_angle_range=(-90, 90), h_angle_range=(-90, 90), w_angle_range=(-90, 90)),\n])\n\n# Loading datasets\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']}\n","metadata":{"execution":{"iopub.status.busy":"2024-05-27T18:36:25.863538Z","iopub.execute_input":"2024-05-27T18:36:25.863827Z","iopub.status.idle":"2024-05-27T18:36:27.850029Z","shell.execute_reply.started":"2024-05-27T18:36:25.863802Z","shell.execute_reply":"2024-05-27T18:36:27.849258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n# Load model\nmodel, criterion = get_model(args.model, args.num_classes)  # Define your function to get the model\nassert(model != None)\nmodel = model.to(device)\nmodel.zero_grad()\nparameters = list(model.parameters())\n# print(\"parameters\", parameters)\n\n# Set criterion\n# criterion = nn.CrossEntropyLoss().cuda()\n\n# Config for SGD solver and scheduler\noptimizer = optim.SGD(parameters, lr=args.learningRate, weight_decay=args.weightDecay, momentum=args.momentum)\nscheduler = optim.lr_scheduler.StepLR(optimizer, step_size=7, gamma=0.1)\n\n\n\nmodel_trained = train_model(model, criterion, optimizer, scheduler, dataloaders, dataset_sizes, device, num_epochs=args.max_epoch)\n# Save the model\ntorch.save(model_trained.state_dict(), f'/kaggle/working/subvol_{args.dataset}.pth')","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2024-05-27T18:36:27.851032Z","iopub.execute_input":"2024-05-27T18:36:27.851291Z","iopub.status.idle":"2024-05-27T18:39:58.238429Z","shell.execute_reply.started":"2024-05-27T18:36:27.851268Z","shell.execute_reply":"2024-05-27T18:39:58.236936Z"},"trusted":true},"execution_count":null,"outputs":[]}]}