{"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"}],"dockerImageVersionId":30698,"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-03T04:13:18.919935Z","iopub.execute_input":"2024-05-03T04:13:18.920450Z","iopub.status.idle":"2024-05-03T04:13:26.731716Z","shell.execute_reply.started":"2024-05-03T04:13:18.920408Z","shell.execute_reply":"2024-05-03T04:13:26.730632Z"},"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 Block2(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel_size, stride=1, padding=0):\n        super(Block2, self).__init__()\n        self.conv = nn.Conv3d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding)\n        self.batchnorm = nn.BatchNorm3d(out_channels)\n    \n    def forward(self, x):\n        x = self.conv(x)\n        x = self.batchnorm(x)\n        x = F.relu(x, inplace=True)\n        return x\n\nclass Model3dninfc(nn.Module):\n    def __init__(self):\n        super(Model3dninfc, self).__init__()\n        self.net = nn.Sequential()\n\n        self.net.add_module('Block1', Block2(1, 48, kernel_size=(6,6,6), stride=(2,2,2)))\n        self.net.add_module('Block2', Block2(48, 48, kernel_size=(1,1,1)))\n        self.net.add_module('Block3', Block2(48, 48, kernel_size=(1,1,1)))\n        self.net.add_module('Dropout1', nn.Dropout3d(p=0.2))\n\n        self.net.add_module('Block4', Block2(48, 96, kernel_size=(5,5,5), stride=(2,2,2)))\n        self.net.add_module('Block5', Block2(96, 96, kernel_size=(1,1,1)))\n        self.net.add_module('Block6', Block2(96, 96, kernel_size=(1,1,1)))\n        self.net.add_module('Dropout2', nn.Dropout3d(p=0.2))\n\n        self.net.add_module('Block7', Block2(96, 512, kernel_size=(3,3,3), stride=(2,2,2)))\n        self.net.add_module('Block8', Block2(512, 512, kernel_size=(1,1,1)))\n        self.net.add_module('Block9', Block2(512, 6, kernel_size=(1,1,1)))\n        self.net.add_module('Dropout3', nn.Dropout3d(p=0.2))\n\n        self.net.add_module('View', nn.Flatten())\n        self.net.add_module('Linear1', nn.Linear(3024, 512))\n        self.net.add_module('ReLU', nn.ReLU(inplace=True))\n        self.net.add_module('Dropout4', nn.Dropout(p=0.5))\n        self.net.add_module('Linear2', nn.Linear(512, 12))\n\n        self.init_weights()\n\n    def forward(self, x):\n        return self.net(x)\n\n    def init_weights(self):\n        def init(m):\n            if isinstance(m, nn.Conv3d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n        self.apply(init)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T04:13:26.749606Z","iopub.execute_input":"2024-05-03T04:13:26.750191Z","iopub.status.idle":"2024-05-03T04:13:26.768118Z","shell.execute_reply.started":"2024-05-03T04:13:26.750160Z","shell.execute_reply":"2024-05-03T04:13:26.766733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_model(model_name):\n    if (model_name == \"3dnin\"):\n        model = Model3dnin()\n        model.MSRinit()\n        return model\n    elif (model_name == \"3dnin_fc\"):\n        model = Model3dninfc()\n        return model\n    return None","metadata":{"execution":{"iopub.status.busy":"2024-05-03T04:13:27.224013Z","iopub.execute_input":"2024-05-03T04:13:27.224398Z","iopub.status.idle":"2024-05-03T04:13:27.230554Z","shell.execute_reply.started":"2024-05-03T04:13:27.224367Z","shell.execute_reply":"2024-05-03T04:13:27.229234Z"},"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        # 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        if self.transform:\n            # Apply transformations if any\n            image_array = self.transform(image_array)\n            \n#         print(image_array.shape, \"getting shape\")\n        return image_array, encoded_label\n","metadata":{"execution":{"iopub.status.busy":"2024-05-03T04:13:26.771172Z","iopub.execute_input":"2024-05-03T04:13:26.772148Z","iopub.status.idle":"2024-05-03T04:13:27.222700Z","shell.execute_reply.started":"2024-05-03T04:13:26.772097Z","shell.execute_reply":"2024-05-03T04:13:27.221405Z"},"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-03T04:13:27.247232Z","iopub.execute_input":"2024-05-03T04:13:27.247606Z","iopub.status.idle":"2024-05-03T04:13:31.782118Z","shell.execute_reply.started":"2024-05-03T04:13:27.247573Z","shell.execute_reply":"2024-05-03T04:13:31.780873Z"},"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-03T04:13:31.783642Z","iopub.execute_input":"2024-05-03T04:13:31.784031Z","iopub.status.idle":"2024-05-03T04:13:31.791533Z","shell.execute_reply.started":"2024-05-03T04:13:31.783993Z","shell.execute_reply":"2024-05-03T04:13:31.790530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_datasets(paths_group, classes):\n    for c in classes:\n        c_paths = glob.glob(f'/kaggle/input/bugnist2024fgvc/BugNIST_DATA/train/{c}/*.tif')\n        \n        tifCount = len(c_paths)\n        tifTrainCount = (tifCount * 8) // 10\n        tifValCount = tifCount // 10\n\n        c_groups = {\n            'train': paths[:tifTrainCount],\n            'validate': paths[tifTrainCount: tifTrainCount + tifValCount],\n            'test': paths[-tifValCount:]\n        }\n        \n        for x in ['train', 'validate', 'test']:\n            paths_group[x].extend(c_groups[x])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Setting up","metadata":{}},{"cell_type":"code","source":"# Define command line arguments\nclass Args:\n    def __init__(self):\n        self.save = 'logs'\n        self.batchSize = 4\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 = '3dnin_fc'\n        self.train_data = 'data/modelnet40_60x/train_data.txt'\n        self.test_data = 'data/modelnet40_60x/test_data.txt'\n\nargs = Args()\n\n# Print chosen options\nprint(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# Load model\nmodel = get_model(args.model)  # 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\ncriterion = nn.CrossEntropyLoss().cuda()\n\n# Define transformations\ntransform = transforms.Compose([\n    # Add your transformations here\n    # e.g., ToTensor(), Normalize(), etc.\n])\n\n# Loading datasets\npaths_group = {\n    'train': [],\n    'validate': [],\n    'test': []\n}\ncreate_datasets(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\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\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-03T04:13:31.792880Z","iopub.execute_input":"2024-05-03T04:13:31.793556Z","iopub.status.idle":"2024-05-03T04:13:31.941893Z","shell.execute_reply.started":"2024-05-03T04:13:31.793526Z","shell.execute_reply":"2024-05-03T04:13:31.940515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training function","metadata":{}},{"cell_type":"code","source":"# Training the model\ndef train_model(model, criterion, optimizer, scheduler, 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#                 print(inputs.shape)\n#                 print(label_indexes)\n                inputs = inputs.to(device)\n                labels = labels.to(device)\n                label_indexes = label_indexes.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#                     print(inputs.shape)\n                    outputs = model(inputs)\n                    _, preds = torch.max(outputs, 1)\n#                     print(preds, outputs, labels)\n                    loss = criterion(outputs, labels)\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                \n                running_corrects += torch.sum(preds == label_indexes)\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]\n\n            print(f'{phase} Loss: {epoch_loss:.4f} Acc: {epoch_acc:.4f}')\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2024-05-03T04:34:52.428908Z","iopub.execute_input":"2024-05-03T04:34:52.429345Z","iopub.status.idle":"2024-05-03T04:34:52.440015Z","shell.execute_reply.started":"2024-05-03T04:34:52.429298Z","shell.execute_reply":"2024-05-03T04:34:52.439078Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_trained = train_model(model, criterion, optimizer, scheduler, num_epochs=args.max_epoch)\n# Save the model\ntorch.save(model_trained.state_dict(), '/kaggle/working/3dninfc.pth')","metadata":{"execution":{"iopub.status.busy":"2024-05-03T04:34:54.966913Z","iopub.execute_input":"2024-05-03T04:34:54.967388Z","iopub.status.idle":"2024-05-03T04:35:03.618216Z","shell.execute_reply.started":"2024-05-03T04:34:54.967350Z","shell.execute_reply":"2024-05-03T04:35:03.616395Z"},"trusted":true},"execution_count":null,"outputs":[]}]}