{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":94689,"databundleVersionId":11605086,"sourceType":"competition"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"##### Setup","metadata":{}},{"cell_type":"code","source":"!pip install -q pylibraft-cu12==24.12.0 rmm-cu12==24.12.0 pylibcugraph-cu12==24.12.0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:07.586441Z","iopub.execute_input":"2025-05-31T11:29:07.587030Z","iopub.status.idle":"2025-05-31T11:29:10.633791Z","shell.execute_reply.started":"2025-05-31T11:29:07.586999Z","shell.execute_reply":"2025-05-31T11:29:10.633007Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q torchio","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:10.635390Z","iopub.execute_input":"2025-05-31T11:29:10.635683Z","iopub.status.idle":"2025-05-31T11:29:13.718074Z","shell.execute_reply.started":"2025-05-31T11:29:10.635660Z","shell.execute_reply":"2025-05-31T11:29:13.717197Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nimport plotly.figure_factory as ff\nfrom PIL import Image\nimport math\nimport numpy as np\nimport random\nfrom skimage import io\nimport re\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom collections import OrderedDict\nimport torch.nn.functional as F\nimport torchio as tio\nimport torch.optim as optim\nfrom torch.cuda.amp import autocast\nimport time\nimport warnings\nfrom torch.utils.data import random_split\nfrom torch.optim.lr_scheduler import LambdaLR\nfrom torch.utils.tensorboard import SummaryWriter\nfrom copy import deepcopy\nfrom tqdm import tqdm\nimport shutil\nfrom sklearn.model_selection import train_test_split","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:13.719126Z","iopub.execute_input":"2025-05-31T11:29:13.719345Z","iopub.status.idle":"2025-05-31T11:29:20.279768Z","shell.execute_reply.started":"2025-05-31T11:29:13.719323Z","shell.execute_reply":"2025-05-31T11:29:20.279221Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\" #optimise GPU usage\nwarnings.filterwarnings(\"ignore\") #ignore warnings","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:20.281648Z","iopub.execute_input":"2025-05-31T11:29:20.282456Z","iopub.status.idle":"2025-05-31T11:29:20.286141Z","shell.execute_reply.started":"2025-05-31T11:29:20.282432Z","shell.execute_reply":"2025-05-31T11:29:20.285240Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Exploratory Data Analysis","metadata":{}},{"cell_type":"code","source":"input_path = \"/kaggle/input/forams-classification-2025\"\n\n# Load the labelled data\ndf_labelled = pd.read_csv(f\"{input_path}/labelled.csv\")\ndf_unlabelled = pd.read_csv(f\"{input_path}/unlabelled.csv\")\n\nprint(\"Number of labelled samples: \", len(df_labelled))\nprint(\"Number of unlabelled samples: \", len(df_unlabelled)) \n\n# Some of the first labelled rows\nprint(\"First few rows of labelled dataset: \")\ndf_labelled.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:20.287028Z","iopub.execute_input":"2025-05-31T11:29:20.287287Z","iopub.status.idle":"2025-05-31T11:29:20.323473Z","shell.execute_reply.started":"2025-05-31T11:29:20.287256Z","shell.execute_reply":"2025-05-31T11:29:20.322965Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Labels distribution\nplt.figure(figsize=(10, 5))\nsns.countplot(data=df_labelled, x='label', order=df_labelled['label'].value_counts().index)\nplt.title(\"Distribution of Labels\")\nplt.show()","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-05-31T11:29:20.324239Z","iopub.execute_input":"2025-05-31T11:29:20.324431Z","iopub.status.idle":"2025-05-31T11:29:20.527294Z","shell.execute_reply.started":"2025-05-31T11:29:20.324414Z","shell.execute_reply":"2025-05-31T11:29:20.526553Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"From this, it can be seen that the number of labelled samples is same for each labels. 0-12 labels are known and, from the description of the data, 13th label is for the unknown volumetric images","metadata":{}},{"cell_type":"code","source":"# Defining path for labelled volumes\nlabelled_volumes = \"/kaggle/input/forams-classification-2025/volumes/volumes/labelled\"\n\n# Merging all images in one folder\nlabelled_images = f\"{input_path}/visualizations/visualizations/labelled\"\nunlabelled_images = f\"{input_path}/visualizations/visualizations/unlabelled\"\n\nimage_paths = []\n\n# Merge images in one folder\nfor image_id in df_labelled['id']:\n    image_filename = image_id.replace(\"labelled_\", \"labelled_foram_\") + \".jpg\"\n    full_path = os.path.join(labelled_images, image_filename)\n    image_paths.append(full_path)\n\nfor image_id in df_unlabelled['id']:\n    image_filename = f\"foram_{int(image_id):05}.jpg\"\n    full_path = os.path.join(unlabelled_images, image_filename)\n    image_paths.append(full_path)","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-05-31T11:29:20.528117Z","iopub.execute_input":"2025-05-31T11:29:20.528327Z","iopub.status.idle":"2025-05-31T11:29:20.560602Z","shell.execute_reply.started":"2025-05-31T11:29:20.528310Z","shell.execute_reply":"2025-05-31T11:29:20.559793Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Pick up 1 random labelles sample to visualise it\nidx = random.choice(range(len(df_labelled))) # random index from labelled data\nrow = df_labelled.iloc[idx]\nimage_id = row['id']\nlabel = row['label']\n\n# Image filename and image path\nlabelled_filename_base = f\"labelled_foram_{image_id.split('_')[1]}\" # base name for the labelled images (e.g labelled_foram_00129)\nimage_path = f\"{labelled_images}/{labelled_filename_base}.jpg\"\n\n# labelled_filename base is included in the name of volume, however scaling factor is also included in the name.\n# Hence, we need to match volume name using regular expressions. \n\nmatched_volume = next(\n    (f for f in os.listdir(labelled_volumes)\n     if f.startswith(labelled_filename_base) and f.endswith('.tif')),\n    None\n)\nvolume_path = f\"{labelled_volumes}/{matched_volume}\"\n\n# Read the volume using io (to get )\nvolume = io.imread(volume_path)\nprint(f\"Shape of the volume array: {volume.shape}\") # (Depth, Height, Width) = (128, 128, 128)\n\n# Get the scaling factor using regular expressions\nmatch = re.search(r'_sc_(\\d+_\\d+)', matched_volume)\nscaling_factor = float(match.group(1).replace('_', '.')) if match else 1.0\nprint(f\"Scaling Factor: {scaling_factor}\")\n\n# Open and plot visualisation\nimg = Image.open(image_path)\nplt.imshow(img) # (-0.5, 299.5, 599.5, -0.5) --> width: 300 pixels, height: 600 pixels\nplt.title(f\" Label is {label}\", fontsize=12)\nplt.axis('off')\n\n# Plot all 128 slices on 16x8 grid\nnum_slices = volume.shape[0]  # number of slices (128)\nrows, cols = 16, 8 \n\nfig, axes = plt.subplots(rows, cols, figsize=(20, 40))  # adjust figsize for readability\nfig.suptitle(\"Plot of the slices\", fontsize=18)\n\nfor i in range(num_slices):\n    r, c = divmod(i, cols)\n    axes[r, c].imshow(volume[i, :, :], cmap='gray')  # Show slice i\n    axes[r, c].set_title(f\"slice={i+1}\", fontsize=8)\n    axes[r, c].axis('off')\n\nplt.tight_layout()\nplt.subplots_adjust(top=0.95)  # adjust to fit suptitle\n\nplt.show()","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-05-31T11:29:20.561408Z","iopub.execute_input":"2025-05-31T11:29:20.561659Z","iopub.status.idle":"2025-05-31T11:29:30.205629Z","shell.execute_reply.started":"2025-05-31T11:29:20.561624Z","shell.execute_reply":"2025-05-31T11:29:30.204282Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"From plotting slices multiple times, It can be seen that in most cases at least first and last 10 slices do not contain any information about the shape of the sample. Hence, we can save computing time by ignoring these slices.\n\nFor now, I am ignoring the scaling factor and I will train the model without it. I will introduce scaling factor later if it will be needed to improve accuracy (Spoiler: I became a victim of exam session and didn't have time to do it).","metadata":{}},{"cell_type":"markdown","source":"# Building Dataset","metadata":{}},{"cell_type":"markdown","source":"It is important to preprocess data correctly. Since in MixMatch approcah two types of data - strongly and weakly augmented are used to train on the unlabelled data, we will return strongly and weakly augmented versions of unlabelled data (and one version of the labelled data)  ","metadata":{}},{"cell_type":"code","source":"# Helper function to extract file IDs and labels from the CSV.\ndef get_ids_and_labels(labelled, df):\n    file_ids = df[\"id\"].values\n    if not labelled:\n        # Unlabeled samples get a dummy label (-1)\n        labels = [-1] * len(file_ids)\n    else:\n        labels = df[\"label\"].values\n    return file_ids, labels\n    \n# Transforms\n# Weak augmentations \nweak_transform = tio.Compose([\n    tio.RandomFlip(axes=('LR', 'AP', 'IS')),\n    tio.RandomAffine(scales=(0.95, 1.05), degrees=15),\n    tio.RandomBlur(std=(0.3, 1.5)),  # Random Gaussian blur\n    tio.RandomGamma(log_gamma=(-0.8, 0.8)), \n    tio.ZNormalization(),\n])\n\n# Strong augmentations \nstrong_transform = tio.Compose([\n    tio.RandomFlip(axes=('LR', 'AP', 'IS')),  # Flip along any of the 3 axes\n    tio.RandomAffine(scales=(0.9, 1.1), degrees=30),  # Afine variations\n    tio.RandomBlur(std=(0.5, 1.7)),  # Random Gaussian blur\n    tio.RandomGamma(log_gamma=(-0.8, 0.8)),  # Adjust contrast via gamma correction\n    tio.RandomNoise(std=(0.02, 0.05)),  # Apply Gaussian noise to simulate acquisition artifacts\n    tio.ZNormalization(),  # Normalize to zero mean and unit variance\n])\n\n# Dataset\nclass ForamDataset3D(Dataset):\n    def __init__(self, volume_dir, csv_path, labelled=False, weak_transform=None, strong_transform=None):\n        self.volume_dir = volume_dir\n        self.labelled = labelled\n        self.weak_transform = weak_transform\n        self.strong_transform = strong_transform\n        self.df = pd.read_csv(csv_path)\n        self.file_ids, self.labels = get_ids_and_labels(labelled, self.df)\n\n    def __len__(self):\n        return len(self.file_ids)  # Return the total number of samples\n\n    def __getitem__(self, idx):\n        file_id = self.file_ids[idx]\n        label = self.labels[idx]\n        volume_dir = self.volume_dir\n        \n        if self.labelled:\n            # file_id might be 'labelled_00000', extract numeric part after 'labelled_'\n            numeric_id = file_id.split('_')[1]  # e.g. '00000'\n            prefix = \"labelled_foram_\"\n            # filenames have a suffix like _sc_0_752.tif - we need to find matching file\n            search_prefix = prefix + numeric_id\n        else:\n            # unlabelled IDs are like '1', so zero-pad to 5 digits (assuming 5-digit IDs)\n            numeric_id = str(file_id).zfill(5)  # '1' -> '00001'\n            prefix = \"foram_\"\n            search_prefix = prefix + numeric_id\n        \n        # Now find the file matching search_prefix (e.g. \"labelled_foram_00000\")\n        matched_volume = next(\n            (f for f in os.listdir(volume_dir) if f.startswith(search_prefix) and f.endswith('.tif')),\n            None\n        )\n        if matched_volume is None:\n            raise FileNotFoundError(f\"No file matching prefix {search_prefix} found in {volume_dir}\")\n        \n        filepath = os.path.join(volume_dir, matched_volume)\n        volume = io.imread(filepath)[25:-25]\n        volume = np.expand_dims(volume, axis=0)\n        subject = tio.Subject(volume=tio.ScalarImage(tensor=volume))\n        \n        if not self.labelled:\n            subject_weak = self.weak_transform(subject) if self.weak_transform else subject\n            subject_strong = self.strong_transform(subject) if self.strong_transform else subject\n            return subject_weak['volume'].data.float(), subject_strong['volume'].data.float(), label\n        else:\n            subject_weak = self.weak_transform(subject) if self.weak_transform else subject\n            return subject_weak['volume'].data.float(), label\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:30.207022Z","iopub.execute_input":"2025-05-31T11:29:30.207622Z","iopub.status.idle":"2025-05-31T11:29:30.224278Z","shell.execute_reply.started":"2025-05-31T11:29:30.207596Z","shell.execute_reply":"2025-05-31T11:29:30.223606Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Load your CSV with columns 'id' and 'label' (assuming 'label' is the class column)\ndf = pd.read_csv(\"/kaggle/input/forams-classification-2025/labelled.csv\")\n\n# Stratified split — keeps class proportions\ntrain_df, val_df = train_test_split(\n    df,\n    test_size=0.2,           # 20% validation, 80% training\n    stratify=df['label'],    # stratify on the label column\n    random_state=42          # for reproducibility\n)\n\nprint(\"Train class distribution:\\n\", train_df['label'].value_counts())\nprint(\"Test class distribution:\\n\", val_df['label'].value_counts())\n\n# Optional: save splits to CSV\ntrain_df.to_csv(\"train_split.csv\", index=False)\nval_df.to_csv(\"test_split.csv\", index=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:30.227559Z","iopub.execute_input":"2025-05-31T11:29:30.227890Z","iopub.status.idle":"2025-05-31T11:29:30.253346Z","shell.execute_reply.started":"2025-05-31T11:29:30.227867Z","shell.execute_reply":"2025-05-31T11:29:30.252700Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check if everything works as expected\n\n# directories\nlabelled_volumes_dir = '/kaggle/input/forams-classification-2025/volumes/volumes/labelled'\nlabelled_train_csv_path = \"/kaggle/working/train_split.csv\"\nlabelled_test_csv_path = \"/kaggle/working/test_split.csv\"\n\nunlabelled_volumes_dir = '/kaggle/input/forams-classification-2025/volumes/volumes/unlabelled'\nunlabelled_csv_path = \"/kaggle/input/forams-classification-2025/unlabelled.csv\"\n\nlabelled_train_dataset = ForamDataset3D(\n    labelled_volumes_dir, \n    labelled_train_csv_path, \n    labelled=True,\n    weak_transform=weak_transform  # Pass the weak transformation\n)\n\nlabelled_test_dataset = ForamDataset3D(\n    labelled_volumes_dir, \n    labelled_test_csv_path, \n    labelled=True,\n    weak_transform=weak_transform  # Pass the weak transformation\n)\n\nunlabelled_dataset = ForamDataset3D(\n    unlabelled_volumes_dir,\n    unlabelled_csv_path,\n    labelled=False,\n    weak_transform=weak_transform,   # Pass weak augmentation\n    strong_transform=strong_transform  # Pass strong augmentation\n)\n\nlabelled_volume, labelled_label = labelled_train_dataset[120]\nunlabelled_volume1, unlabelled_volume2, unlabelled_label = unlabelled_dataset[6072]\n\nprint(f\"unabelled dataset volume without augmentation: {unlabelled_volume1.shape}, volume for the input for augmentation: {unlabelled_volume2}, label: {unlabelled_label}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:30.254120Z","iopub.execute_input":"2025-05-31T11:29:30.254381Z","iopub.status.idle":"2025-05-31T11:29:30.958863Z","shell.execute_reply.started":"2025-05-31T11:29:30.254357Z","shell.execute_reply":"2025-05-31T11:29:30.958074Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plotting the strongly augmented data to see how the transforms look like\n\n# Ensure the volume has the correct shape\nunlabelled_volume = unlabelled_volume2.squeeze()\nprint(f\"Volume shape after squeeze: {unlabelled_volume.shape}\")\n\nnum_slices = unlabelled_volume.shape[0]  # Number of slices\nprint(f\"Number of slices: {num_slices}\")\n\nrows, cols = 13, 6  # Adjust grid size for readability\nfig, axes = plt.subplots(rows, cols, figsize=(20, 40))\nfig.suptitle(\"Unlabelled Volume Slices after srong augmnetation\", fontsize=18)\n\n# Loop through slices and display them\nfor i in range(min(num_slices, rows * cols)):  # Ensure we don't exceed available slices\n    r, c = divmod(i, cols)\n    axes[r, c].imshow(unlabelled_volume[i, :, :], cmap='gray')  # Show slice i\n    axes[r, c].set_title(f\"Slice {i+1}\", fontsize=8)\n    axes[r, c].axis('off')\n\nplt.tight_layout()\nplt.subplots_adjust(top=0.95)  # Adjust to fit suptitle\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:30.959653Z","iopub.execute_input":"2025-05-31T11:29:30.959879Z","iopub.status.idle":"2025-05-31T11:29:37.308000Z","shell.execute_reply.started":"2025-05-31T11:29:30.959862Z","shell.execute_reply":"2025-05-31T11:29:37.307114Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Developing Network","metadata":{}},{"cell_type":"code","source":"# Sourse: https://github.com/kenshohara/3D-ResNets-PyTorch/blob/master/models/densenet.py\nclass _DenseLayer(nn.Sequential):\n    \"\"\" Internal building block - Dense layer \n    Args: \n    n_layers (int) - number of layers\n    n_input_features (int) - number of input features\n    growth_rate (int) - growth rate (how many new feature maps are added after each layer)\n    drop_rate (float) - drop_rate (probability with which features will be dropped. This helps overcome overfitting)\n    bn_size - multiplicative factor for number of bottle neck layers\n          (i.e. bn_size * k features in the bottleneck layer)\n        \n    \"\"\"\n\n    def __init__(self, num_input_features, growth_rate, bn_size, drop_rate):\n        super().__init__()\n        self.add_module('norm1', nn.BatchNorm3d(num_input_features))\n        self.add_module('relu1', nn.ReLU(inplace=True))\n        self.add_module(\n            'conv1',\n            nn.Conv3d(num_input_features,\n                      bn_size * growth_rate,\n                      kernel_size=1,\n                      stride=1,\n                      bias=False))\n        self.add_module('norm2', nn.BatchNorm3d(bn_size * growth_rate))\n        self.add_module('relu2', nn.ReLU(inplace=True))\n        self.add_module(\n            'conv2',\n            nn.Conv3d(bn_size * growth_rate,\n                      growth_rate,\n                      kernel_size=3,\n                      stride=1,\n                      padding=1,\n                      bias=False))\n        self.drop_rate = drop_rate\n\n    def forward(self, x):\n        new_features = super().forward(x)\n        if self.drop_rate > 0:\n            new_features = F.dropout(new_features,\n                                     p=self.drop_rate,\n                                     training=self.training)\n        return torch.cat([x, new_features], 1)\n\n\nclass _DenseBlock(nn.Sequential):\n\n    def __init__(self, num_layers, num_input_features, bn_size, growth_rate,\n                 drop_rate):\n        super().__init__()\n        for i in range(num_layers):\n            layer = _DenseLayer(num_input_features + i * growth_rate,\n                                growth_rate, bn_size, drop_rate)\n            self.add_module('denselayer{}'.format(i + 1), layer)\n\n\nclass _Transition(nn.Sequential):\n    \"\"\" Transition Layer\n    Used to downsample the feature maps calculated by Dense Block and \n    to reduce computational load\n    \"\"\"\n\n    def __init__(self, num_input_features, num_output_features):\n        super().__init__()\n        self.add_module('norm', nn.BatchNorm3d(num_input_features))\n        self.add_module('relu', nn.ReLU(inplace=True))\n        self.add_module(\n            'conv',\n            nn.Conv3d(num_input_features,\n                      num_output_features,\n                      kernel_size=1,\n                      stride=1,\n                      bias=False))\n        self.add_module('pool', nn.AvgPool3d(kernel_size=2, stride=2))\n\n\nclass DenseNet(nn.Module):\n    \"\"\"DenseNet model class\n    Args:\n        n_input_channels (int) - how many input channels, 3 if RGB img\n        conv_1_t_size (int) - size of the first convolution along the depth/time dimension.\n        conv_1_t_stride (int) - stride for the first convolution along the depth/time dimension.\n        no_max_pool (boolean) - whether to skip the initial MaxPooling layer after the first conv\n            (False -> (default) maxpool after first conv layer -> downsampling of the input,\n             True -> no maxpool)\n        growth_rate (int) - how many new features are added after each layer\n        block_config (list of 4 ints) - how many layers in each pooling block\n        n_init_features (int) - number of features for first layer\n        bn_size (int) - multiplicative factor for number of bottle neck layers\n           (i.e. bn_size * k features in the bottleneck layer)\n        drop_rate (float) - dropout rate after each dense layer (should be between 0.2-0.5)\n        n_classes (int) - number of classification classes, in our case, 14   \n    \"\"\"\n\n    def __init__(self,\n                 n_input_channels=1,\n                 conv1_t_size=7,\n                 conv1_t_stride=1,\n                 no_max_pool=False,\n                 growth_rate=32,\n                 block_config=(6, 12, 24, 16),\n                 num_init_features=64,\n                 bn_size=4,\n                 drop_rate=0,\n                 num_classes=14):\n\n        super().__init__()\n\n        # First convolution\n        self.features = [('conv1',\n                          nn.Conv3d(n_input_channels,\n                                    num_init_features,\n                                    kernel_size=(conv1_t_size, 7, 7),\n                                    stride=(conv1_t_stride, 2, 2),\n                                    padding=(conv1_t_size // 2, 3, 3),\n                                    bias=False)),\n                         ('norm1', nn.BatchNorm3d(num_init_features)),\n                         ('relu1', nn.ReLU(inplace=True))]\n        if not no_max_pool:\n            self.features.append(\n                ('pool1', nn.MaxPool3d(kernel_size=3, stride=2, padding=1)))\n        self.features = nn.Sequential(OrderedDict(self.features))\n\n        # Each denseblock\n        num_features = num_init_features\n        for i, num_layers in enumerate(block_config):\n            block = _DenseBlock(num_layers=num_layers,\n                                num_input_features=num_features,\n                                bn_size=bn_size,\n                                growth_rate=growth_rate,\n                                drop_rate=drop_rate)\n            self.features.add_module('denseblock{}'.format(i + 1), block)\n            num_features = num_features + num_layers * growth_rate\n            if i != len(block_config) - 1:\n                trans = _Transition(num_input_features=num_features,\n                                    num_output_features=num_features // 2)\n                self.features.add_module('transition{}'.format(i + 1), trans)\n                num_features = num_features // 2\n\n        # Final batch norm\n        self.features.add_module('norm5', nn.BatchNorm3d(num_features))\n\n        for m in self.modules():\n            if isinstance(m, nn.Conv3d):\n                m.weight = nn.init.kaiming_normal(m.weight, mode='fan_out')\n            elif isinstance(m, nn.BatchNorm3d) or isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()\n\n        # Linear layer\n        self.classifier = nn.Linear(num_features, num_classes)\n\n        for m in self.modules():\n            if isinstance(m, nn.Conv3d):\n                nn.init.kaiming_normal_(m.weight,\n                                        mode='fan_out',\n                                        nonlinearity='relu')\n            elif isinstance(m, nn.BatchNorm3d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.Linear):\n                nn.init.constant_(m.bias, 0)\n\n    def forward(self, x):\n        features = self.features(x)\n        out = F.relu(features, inplace=True)\n        out = F.adaptive_avg_pool3d(out,\n                                    output_size=(1, 1,\n                                                 1)).view(features.size(0), -1)\n        out = self.classifier(out)\n        return out\n\n\ndef generate_DenseNet121(model_depth=121, **kwargs):\n    assert model_depth == 121 \n\n    # Call the DenseNet model constructor\n    model = DenseNet(num_init_features=64,\n                     growth_rate=32,\n                     block_config=(6, 12, 24, 16),  \n                     **kwargs)  # Pass additional kwargs (like num_classes)\n\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:37.308893Z","iopub.execute_input":"2025-05-31T11:29:37.309256Z","iopub.status.idle":"2025-05-31T11:29:37.330006Z","shell.execute_reply.started":"2025-05-31T11:29:37.309225Z","shell.execute_reply":"2025-05-31T11:29:37.329169Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training model - FixMatch approach","metadata":{}},{"cell_type":"code","source":"# Source: https://github.com/kekmodel/FixMatch-pytorch/tree/master","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:37.330837Z","iopub.execute_input":"2025-05-31T11:29:37.331064Z","iopub.status.idle":"2025-05-31T11:29:37.345311Z","shell.execute_reply.started":"2025-05-31T11:29:37.331041Z","shell.execute_reply":"2025-05-31T11:29:37.344558Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Utility functions\n\ndef get_mean_and_std(dataset):\n    '''Compute the mean and std value of dataset.'''\n    dataloader = torch.utils.data.DataLoader(\n        dataset, batch_size=1, shuffle=False, num_workers=4)\n\n    mean = torch.zeros(3)\n    std = torch.zeros(3)\n    logger.info('==> Computing mean and std..')\n    for inputs, targets in dataloader:\n        for i in range(3):\n            mean[i] += inputs[:, i, :, :].mean()\n            std[i] += inputs[:, i, :, :].std()\n    mean.div_(len(dataset))\n    std.div_(len(dataset))\n    return mean, std\n\n\ndef accuracy(output, target, topk=(1,)):\n    \"\"\"Computes the precision@k for the specified values of k\"\"\"\n    maxk = max(topk)\n    batch_size = target.size(0)\n\n    _, pred = output.topk(maxk, 1, True, True)\n    pred = pred.t()\n    correct = pred.eq(target.reshape(1, -1).expand_as(pred))\n\n    res = []\n    for k in topk:\n        correct_k = correct[:k].reshape(-1).float().sum(0)\n        res.append(correct_k.mul_(100.0 / batch_size))\n    return res\n\n\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\n       Imported from https://github.com/pytorch/examples/blob/master/imagenet/main.py#L247-L262\n    \"\"\"\n\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:37.346271Z","iopub.execute_input":"2025-05-31T11:29:37.346476Z","iopub.status.idle":"2025-05-31T11:29:37.357346Z","shell.execute_reply.started":"2025-05-31T11:29:37.346459Z","shell.execute_reply":"2025-05-31T11:29:37.356683Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ModelEMA(object):\n    def __init__(self, model, decay):\n        self.ema = deepcopy(model)\n        self.ema.eval()\n        self.decay = decay\n        self.ema_has_module = hasattr(self.ema, 'module')\n        # Fix EMA. https://github.com/valencebond/FixMatch_pytorch thank you!\n        self.param_keys = [k for k, _ in self.ema.named_parameters()]\n        self.buffer_keys = [k for k, _ in self.ema.named_buffers()]\n        for p in self.ema.parameters():\n            p.requires_grad_(False)\n\n    def update(self, model):\n        needs_module = hasattr(model, 'module') and not self.ema_has_module\n        with torch.no_grad():\n            msd = model.state_dict()\n            esd = self.ema.state_dict()\n            for k in self.param_keys:\n                if needs_module:\n                    j = 'module.' + k\n                else:\n                    j = k\n                model_v = msd[j].detach()\n                ema_v = esd[k]\n                esd[k].copy_(ema_v * self.decay + (1. - self.decay) * model_v)\n\n            for k in self.buffer_keys:\n                if needs_module:\n                    j = 'module.' + k\n                else:\n                    j = k\n                esd[k].copy_(msd[j])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:37.358058Z","iopub.execute_input":"2025-05-31T11:29:37.358284Z","iopub.status.idle":"2025-05-31T11:29:37.375308Z","shell.execute_reply.started":"2025-05-31T11:29:37.358268Z","shell.execute_reply":"2025-05-31T11:29:37.374527Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def save_checkpoint(state, is_best, checkpoint, filename='checkpoint.pth.tar'):\n    filepath = os.path.join(checkpoint, filename)\n    torch.save(state, filepath)\n    if is_best:\n        shutil.copyfile(filepath, os.path.join(checkpoint,\n                                               'model_best.pth.tar'))\n\n\ndef set_seed(args):\n    random.seed(args.seed)\n    np.random.seed(args.seed)\n    torch.manual_seed(args.seed)\n    torch.cuda.manual_seed_all(args.seed)\n\n\ndef get_cosine_schedule_with_warmup(optimizer,\n                                    num_warmup_steps,\n                                    num_training_steps,\n                                    num_cycles=7./16.,\n                                    last_epoch=-1):\n    def _lr_lambda(current_step):\n        if current_step < num_warmup_steps:\n            return float(current_step) / float(max(1, num_warmup_steps))\n        no_progress = float(current_step - num_warmup_steps) / \\\n            float(max(1, num_training_steps - num_warmup_steps))\n        return max(0., math.cos(math.pi * num_cycles * no_progress))\n\n    return LambdaLR(optimizer, _lr_lambda, last_epoch)\n\n\ndef interleave(x, size):\n    s = list(x.shape)\n    return x.reshape([-1, size] + s[1:]).transpose(0, 1).reshape([-1] + s[1:])\n\n\ndef de_interleave(x, size):\n    s = list(x.shape)\n    return x.reshape([size, -1] + s[1:]).transpose(0, 1).reshape([-1] + s[1:])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:37.376083Z","iopub.execute_input":"2025-05-31T11:29:37.376250Z","iopub.status.idle":"2025-05-31T11:29:37.392299Z","shell.execute_reply.started":"2025-05-31T11:29:37.376235Z","shell.execute_reply":"2025-05-31T11:29:37.391709Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Args:\n    # hardware\n    world_size    = 1\n    local_rank    = 0\n    num_workers   = 4\n    device        = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    writer        = None\n    out           = 'result'\n\n    # data / dataset\n    dataset       = 'Foram3D'     \n    num_labeled   = 210\n    expand_labels = False\n\n    # model / arch\n    use_ema       = True\n    ema_decay     = 0.999\n    num_classes   = 14\n\n    # training schedule\n    total_steps           = 2**20\n    eval_step             = 1024\n    start_epoch           = 0\n    batch_size            = 7\n    mu                    = 3 # unlabeled to labeled batch ratio\n    lambda_u              = 1.0\n    threshold             = 0.95\n    T                     = 1.0 # pseudo-label temperature\n    epochs                = 10\n\n    # optimizer / lr\n    lr              = 0.001\n    wdecay          = 5e-4\n    nesterov        = True\n    warmup          = 0            # warmup steps\n    label_smoothing = 0.1\n\n    # reproducibility & logging\n    seed          = 1234\n    out           = 'result'\n    amp           = False\n    opt_level     = 'O1'\n    no_progress   = False\n    resume        = ''           # e.g. 'result/checkpoint.pth.tar\n\nargs = Args()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:37.393025Z","iopub.execute_input":"2025-05-31T11:29:37.393191Z","iopub.status.idle":"2025-05-31T11:29:37.456533Z","shell.execute_reply.started":"2025-05-31T11:29:37.393177Z","shell.execute_reply":"2025-05-31T11:29:37.455853Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_acc = 0.0\nif args.seed is not None:\n    set_seed(args)\n\nargs.out = \"/kaggle/working/runs\"\nos.makedirs(args.out, exist_ok=True)\nargs.writer = SummaryWriter(log_dir=args.out)\n\ncriterion = nn.CrossEntropyLoss(label_smoothing=args.label_smoothing)\n\n# DataLoaders\nlabelled_trainloader = DataLoader(\n    labelled_train_dataset,\n    batch_size=args.batch_size,\n    shuffle=True,\n    num_workers= args.num_workers,  \n    drop_last=False,\n)\n\ntest_loader = DataLoader(\n    labelled_test_dataset,\n    batch_size=args.batch_size,\n    shuffle=True,\n    num_workers= args.num_workers,  \n    drop_last=False,\n)\n\n\nnum_unlabelled =  18216 # Total samples\nsubset_size = 5000  # Desired subset size\n\n# Split dataset into a subset (5000 samples) and the rest\nsubset_unlabelled, _ = random_split(unlabelled_dataset, [subset_size, num_unlabelled - subset_size])\n\n# Create DataLoader for the subset\nunlabelled_trainloader = DataLoader(\n    subset_unlabelled,\n    batch_size=args.batch_size,\n    shuffle=True,\n    num_workers=args.num_workers,\n    drop_last=True,\n)\n\n# Create the primary model and an EMA model.\nmodel = generate_DenseNet121(121, num_classes=args.num_classes)\nmodel = nn.DataParallel(model)\nmodel.to(args.device)\nema_model = ModelEMA(model, args.ema_decay)\n\nno_decay = ['bias', 'bn']\ngrouped_parameters = [\n    {'params': [p for n, p in model.named_parameters() if not any(\n        nd in n for nd in no_decay)], 'weight_decay': args.wdecay},\n    {'params': [p for n, p in model.named_parameters() if any(\n        nd in n for nd in no_decay)], 'weight_decay': 0.0}\n]\noptimizer = optim.SGD(grouped_parameters, lr=args.lr,\n                      momentum=0.9, nesterov=args.nesterov)\n\n#args.epochs = math.ceil(args.total_steps / args.eval_step)\nscheduler = get_cosine_schedule_with_warmup(\n    optimizer, args.warmup, args.total_steps)\n\nargs.start_epoch = 0\n\nif args.resume:\n    logger.info(\"==> Resuming from checkpoint..\")\n    assert os.path.isfile(\n        args.resume), \"Error: no checkpoint directory found!\"\n    args.out = os.path.dirname(args.resume)\n    checkpoint = torch.load(args.resume)\n    best_acc = checkpoint['best_acc']\n    args.start_epoch = checkpoint['epoch']\n    model.load_state_dict(checkpoint['state_dict'])\n    if args.use_ema:\n        ema_model.ema.load_state_dict(checkpoint['ema_state_dict'])\n    optimizer.load_state_dict(checkpoint['optimizer'])\n    scheduler.load_state_dict(checkpoint['scheduler'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:37.457246Z","iopub.execute_input":"2025-05-31T11:29:37.457429Z","iopub.status.idle":"2025-05-31T11:29:38.739455Z","shell.execute_reply.started":"2025-05-31T11:29:37.457413Z","shell.execute_reply":"2025-05-31T11:29:38.738864Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train(args, labeled_trainloader, unlabeled_trainloader, test_loader,\n          model, optimizer, ema_model, scheduler, criterion):\n    if args.amp:\n        from apex import amp\n    global best_acc\n    test_accs = []\n    end = time.time()\n\n    labeled_iter = iter(labeled_trainloader)\n    unlabeled_iter = iter(unlabeled_trainloader)\n\n    model.train()\n    for epoch in range(args.start_epoch, args.epochs):\n        batch_time = AverageMeter()\n        data_time = AverageMeter()\n        losses = AverageMeter()\n        losses_x = AverageMeter()\n        losses_u = AverageMeter()\n        mask_probs = AverageMeter()\n        if not args.no_progress:\n            p_bar = tqdm(range(args.eval_step),\n                         disable=args.local_rank not in [-1, 0])\n        for batch_idx in range(args.eval_step):\n            try:\n                inputs_x, targets_x = next(labeled_iter)\n            except:\n                if args.world_size > 1:\n                    labeled_epoch += 1\n                    labeled_trainloader.sampler.set_epoch(labeled_epoch)\n                labeled_iter = iter(labeled_trainloader)\n                inputs_x, targets_x = next(labeled_iter)\n\n            try:\n                inputs_u_w, inputs_u_s, _ = next(unlabeled_iter)\n            except:\n                if args.world_size > 1:\n                    unlabeled_epoch += 1\n                    unlabeled_trainloader.sampler.set_epoch(unlabeled_epoch)\n                unlabeled_iter = iter(unlabeled_trainloader)\n                inputs_u_w, inputs_u_s, _ = next(unlabeled_iter)\n\n            data_time.update(time.time() - end)\n            batch_size = inputs_x.shape[0]\n            #print(\"inputs_x shape:\", inputs_x.shape)\n            #print(\"inputs_u_w shape:\", inputs_u_w.shape)\n            #print(\"inputs_u_s shape:\", inputs_u_s.shape)\n            #print(\"args.mu:\", args.mu)\n            #print(\"Concatenated shape:\", torch.cat((inputs_x, inputs_u_w, inputs_u_s)).shape)\n            inputs = interleave(\n                torch.cat((inputs_x, inputs_u_w, inputs_u_s)), 2*args.mu+1).to(args.device)\n            targets_x = targets_x.to(args.device)\n            logits = model(inputs)\n            logits = de_interleave(logits, 2*args.mu+1)\n            logits_x = logits[:batch_size]\n            logits_u_w, logits_u_s = logits[batch_size:].chunk(2)\n            del logits\n\n            #Lx = F.cross_entropy(logits_x, targets_x, reduction='mean')\n            Lx = criterion(logits_x, targets_x)\n\n            pseudo_label = torch.softmax(logits_u_w.detach()/args.T, dim=-1)\n            max_probs, targets_u = torch.max(pseudo_label, dim=-1)\n            mask = max_probs.ge(args.threshold).float()\n\n            Lu = (F.cross_entropy(logits_u_s, targets_u,\n                                  reduction='none') * mask).mean()\n\n            loss = Lx + args.lambda_u * Lu\n\n            if args.amp:\n                with amp.scale_loss(loss, optimizer) as scaled_loss:\n                    scaled_loss.backward()\n            else:\n                loss.backward()\n\n            losses.update(loss.item())\n            losses_x.update(Lx.item())\n            losses_u.update(Lu.item())\n            optimizer.step()\n            scheduler.step()\n            if args.use_ema:\n                ema_model.update(model)\n            model.zero_grad()\n\n            batch_time.update(time.time() - end)\n            end = time.time()\n            mask_probs.update(mask.mean().item())\n            if not args.no_progress:\n                p_bar.set_description(\"Train Epoch: {epoch}/{epochs:4}. Iter: {batch:4}/{iter:4}. LR: {lr:.4f}. Data: {data:.3f}s. Batch: {bt:.3f}s. Loss: {loss:.4f}. Loss_x: {loss_x:.4f}. Loss_u: {loss_u:.4f}. Mask: {mask:.2f}. \".format(\n                    epoch=epoch + 1,\n                    epochs=args.epochs,\n                    batch=batch_idx + 1,\n                    iter=args.eval_step,\n                    lr=scheduler.get_last_lr()[0],\n                    data=data_time.avg,\n                    bt=batch_time.avg,\n                    loss=losses.avg,\n                    loss_x=losses_x.avg,\n                    loss_u=losses_u.avg,\n                    mask=mask_probs.avg))\n                p_bar.update()\n\n        if not args.no_progress:\n            p_bar.close()\n\n        if args.use_ema:\n            test_model = ema_model.ema\n        else:\n            test_model = model\n\n        if args.local_rank in [-1, 0]:\n            test_loss, test_acc = test(args, test_loader, test_model, epoch)\n\n            args.writer.add_scalar('train/1.train_loss', losses.avg, epoch)\n            args.writer.add_scalar('train/2.train_loss_x', losses_x.avg, epoch)\n            args.writer.add_scalar('train/3.train_loss_u', losses_u.avg, epoch)\n            args.writer.add_scalar('train/4.mask', mask_probs.avg, epoch)\n            args.writer.add_scalar('test/1.test_acc', test_acc, epoch)\n            args.writer.add_scalar('test/2.test_loss', test_loss, epoch)\n\n            is_best = test_acc > best_acc\n            best_acc = max(test_acc, best_acc)\n\n            model_to_save = model.module if hasattr(model, \"module\") else model\n            if args.use_ema:\n                ema_to_save = ema_model.ema.module if hasattr(\n                    ema_model.ema, \"module\") else ema_model.ema\n            save_checkpoint({\n                'epoch': epoch + 1,\n                'state_dict': model_to_save.state_dict(),\n                'ema_state_dict': ema_to_save.state_dict() if args.use_ema else None,\n                'acc': test_acc,\n                'best_acc': best_acc,\n                'optimizer': optimizer.state_dict(),\n                'scheduler': scheduler.state_dict(),\n            }, is_best, args.out)\n\n            test_accs.append(test_acc)\n            print('Best top-1 acc: {:.2f}'.format(best_acc))\n            print('Mean top-1 acc: {:.2f}\\n'.format(\n                np.mean(test_accs[-20:])))\n\n    if args.local_rank in [-1, 0]:\n        args.writer.close()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:38.740271Z","iopub.execute_input":"2025-05-31T11:29:38.740901Z","iopub.status.idle":"2025-05-31T11:29:38.757090Z","shell.execute_reply.started":"2025-05-31T11:29:38.740879Z","shell.execute_reply":"2025-05-31T11:29:38.756401Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def test(args, test_loader, model, epoch):\n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    top1 = AverageMeter()\n    top5 = AverageMeter()\n    end = time.time()\n    \n    if not args.no_progress:\n        test_loader = tqdm(test_loader,\n                           disable=args.local_rank not in [-1, 0])\n    \n    with torch.no_grad():\n        for batch_idx, (inputs, targets) in enumerate(test_loader):\n            data_time.update(time.time() - end)\n            model.eval()\n    \n            inputs = inputs.to(args.device)\n            targets = targets.to(args.device)\n            outputs = model(inputs)\n            loss = F.cross_entropy(outputs, targets)\n\n            prec1, prec5 = accuracy(outputs, targets, topk=(1, 5))\n            losses.update(loss.item(), inputs.shape[0])\n            top1.update(prec1.item(), inputs.shape[0])\n            top5.update(prec5.item(), inputs.shape[0])\n            batch_time.update(time.time() - end)\n            end = time.time()\n            if not args.no_progress:\n                test_loader.set_description(\"Test Iter: {batch:4}/{iter:4}. Data: {data:.3f}s. Batch: {bt:.3f}s. Loss: {loss:.4f}. top1: {top1:.2f}. top5: {top5:.2f}. \".format(\n                    batch=batch_idx + 1,\n                    iter=len(test_loader),\n                    data=data_time.avg,\n                    bt=batch_time.avg,\n                    loss=losses.avg,\n                    top1=top1.avg,\n                    top5=top5.avg,\n                ))\n        if not args.no_progress:\n            test_loader.close()\n    \n    print(\"top-1 acc: {:.2f}\".format(top1.avg))\n    print(\"top-5 acc: {:.2f}\".format(top5.avg))\n    return losses.avg, top1.avg\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:38.757863Z","iopub.execute_input":"2025-05-31T11:29:38.758701Z","iopub.status.idle":"2025-05-31T11:29:38.779905Z","shell.execute_reply.started":"2025-05-31T11:29:38.758674Z","shell.execute_reply":"2025-05-31T11:29:38.779238Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.zero_grad()\ntrain(args, labelled_trainloader, unlabelled_trainloader, test_loader,\n      model, optimizer, ema_model, scheduler, criterion)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T11:29:38.780580Z","iopub.execute_input":"2025-05-31T11:29:38.780773Z","iopub.status.idle":"2025-05-31T18:37:20.787655Z","shell.execute_reply.started":"2025-05-31T11:29:38.780757Z","shell.execute_reply":"2025-05-31T18:37:20.786866Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(model.state_dict(), \"/kaggle/working/model.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T18:37:20.788885Z","iopub.execute_input":"2025-05-31T18:37:20.789185Z","iopub.status.idle":"2025-05-31T18:37:20.953882Z","shell.execute_reply.started":"2025-05-31T18:37:20.789159Z","shell.execute_reply":"2025-05-31T18:37:20.953040Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(ema_model.ema.state_dict(), \"/kaggle/working/ema_model.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T18:37:20.954878Z","iopub.execute_input":"2025-05-31T18:37:20.955124Z","iopub.status.idle":"2025-05-31T18:37:21.116132Z","shell.execute_reply.started":"2025-05-31T18:37:20.955106Z","shell.execute_reply":"2025-05-31T18:37:21.115330Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Make a submission","metadata":{}},{"cell_type":"code","source":"model = generate_DenseNet121(121, num_classes=args.num_classes)\nmodel = nn.DataParallel(model)\nmodel.load_state_dict(torch.load(\"/kaggle/working/model.pth\"))\nmodel.to(args.device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T18:54:30.105418Z","iopub.execute_input":"2025-05-31T18:54:30.106187Z","iopub.status.idle":"2025-05-31T18:54:30.582720Z","shell.execute_reply.started":"2025-05-31T18:54:30.106143Z","shell.execute_reply":"2025-05-31T18:54:30.582060Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ema_model = ModelEMA(model, args.ema_decay)\nema_model.ema.load_state_dict(torch.load(\"/kaggle/working/ema_model.pth\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T18:37:21.620217Z","iopub.execute_input":"2025-05-31T18:37:21.620896Z","iopub.status.idle":"2025-05-31T18:37:21.834183Z","shell.execute_reply.started":"2025-05-31T18:37:21.620866Z","shell.execute_reply":"2025-05-31T18:37:21.833420Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# make a class for the inference dataset\ndef get_ids( df):\n    file_ids = df[\"id\"].values\n    return file_ids\n\n\nclass InferenceDataset(Dataset):\n    def __init__(self, volume_dir, csv_path, transform=None):\n        self.volume_dir = volume_dir\n        self.transform = transform\n        self.df = pd.read_csv(csv_path)\n        self.file_ids = get_ids(self.df)\n\n    def __len__(self):\n        return len(self.file_ids)  # Return the total number of samples\n\n    def __getitem__(self, idx):\n        file_id = self.file_ids[idx]\n        volume_dir = self.volume_dir\n        numeric_id = str(file_id).zfill(5)  # '1' -> '00001'\n        prefix = \"foram_\"\n        search_prefix = prefix + numeric_id\n\n        matched_volume = next(\n            (f for f in os.listdir(volume_dir) if f.startswith(search_prefix) and f.endswith('.tif')),\n            None\n        )\n        if matched_volume is None:\n            raise FileNotFoundError(f\"No file matching prefix {search_prefix} found in {volume_dir}\")\n  \n        filepath = os.path.join(volume_dir, matched_volume)\n        volume = io.imread(filepath)[25:-25]\n        volume = np.expand_dims(volume, axis=0)\n        subject = tio.Subject(volume=tio.ScalarImage(tensor=volume))\n        return subject['volume'].data.float(), file_id\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T18:37:21.835021Z","iopub.execute_input":"2025-05-31T18:37:21.835208Z","iopub.status.idle":"2025-05-31T18:37:21.842369Z","shell.execute_reply.started":"2025-05-31T18:37:21.835192Z","shell.execute_reply":"2025-05-31T18:37:21.841665Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"unlabelled_dataset_inference = InferenceDataset(\n    unlabelled_volumes_dir,\n    unlabelled_csv_path,\n    transform=weak_transform\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T18:37:21.846248Z","iopub.execute_input":"2025-05-31T18:37:21.846607Z","iopub.status.idle":"2025-05-31T18:37:21.865117Z","shell.execute_reply.started":"2025-05-31T18:37:21.846591Z","shell.execute_reply":"2025-05-31T18:37:21.864458Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"inference_volume, inference_id = unlabelled_dataset_inference[1008]\n\nprint(f\"inference dataset volume: {unlabelled_volume.shape}, id: {inference_id}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T18:37:21.866156Z","iopub.execute_input":"2025-05-31T18:37:21.866423Z","iopub.status.idle":"2025-05-31T18:37:21.971076Z","shell.execute_reply.started":"2025-05-31T18:37:21.866400Z","shell.execute_reply":"2025-05-31T18:37:21.970466Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create DataLoader for the subset\ninference_loader = DataLoader(\n    unlabelled_dataset_inference,\n    batch_size=8,\n    shuffle=False,\n    num_workers=8,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T18:37:21.971771Z","iopub.execute_input":"2025-05-31T18:37:21.972032Z","iopub.status.idle":"2025-05-31T18:37:21.975861Z","shell.execute_reply.started":"2025-05-31T18:37:21.972006Z","shell.execute_reply":"2025-05-31T18:37:21.975177Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_and_save_predictions(model, inference_loader, device=args.device, output_csv=\"submission.csv\", temperature=1.0, suppress_class=12, suppress_strength=150):\n    model.eval()\n    predictions = []\n    ids = []\n    \n    print(f\"Starting evaluation... Number of batches: {len(inference_loader)}\")\n    \n    with torch.no_grad():\n        for batch_idx, (volumes, file_ids) in enumerate(inference_loader):\n            print(f\"Batch {batch_idx} - volumes shape: {volumes.shape}, file_ids: {file_ids}\")\n            volumes = volumes.to(device)\n            \n            outputs = model(volumes)  # raw logits: shape (batch_size, num_classes)\n            print(f\"Outputs shape: {outputs.shape}\")\n            \n            # Apply temperature scaling\n            scaled_outputs = outputs / temperature\n\n            # ↓↓↓ Suppress logits for class 12 ↓↓↓\n            scaled_outputs[:, suppress_class] -= suppress_strength\n\n            # Print logits of first sample in batch\n            print(\"Scaled & suppressed logits for first sample:\", scaled_outputs[0])\n            \n            # Compute softmax\n            probs = torch.softmax(scaled_outputs, dim=1)\n            print(\"Softmax probabilities for first sample:\", probs[0])\n            \n            max_prob, pred_label = torch.max(probs, dim=1)\n            print(\"Max probabilities:\", max_prob)\n            print(\"Predicted labels:\", pred_label)\n            \n            predictions.extend(pred_label.cpu().numpy())\n            ids.extend(file_ids.cpu().numpy())\n    \n    # Save to CSV\n    submission_df = pd.DataFrame({\n        \"id\": ids,\n        \"label\": predictions\n    })\n    submission_df.to_csv(output_csv, index=False)\n    print(f\"Saved predictions to {output_csv}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T19:06:20.167030Z","iopub.execute_input":"2025-05-31T19:06:20.167706Z","iopub.status.idle":"2025-05-31T19:06:20.173976Z","shell.execute_reply.started":"2025-05-31T19:06:20.167683Z","shell.execute_reply":"2025-05-31T19:06:20.173360Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"evaluate_and_save_predictions(model, inference_loader, device=args.device, temperature=2.0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T19:06:49.419838Z","iopub.execute_input":"2025-05-31T19:06:49.420162Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x = torch.randn(1, 1, 78, 128, 128).to(args.device)\nwith torch.no_grad():\n    logits = model(x)\n    print(logits)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T18:49:11.296137Z","iopub.status.idle":"2025-05-31T18:49:11.296416Z","shell.execute_reply.started":"2025-05-31T18:49:11.296285Z","shell.execute_reply":"2025-05-31T18:49:11.296297Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# References\n\n","metadata":{}},{"cell_type":"markdown","source":"1. Ilesanmi, A. E., Ilesanmi, T. O., & Ajayi, B. O. (2024). Reviewing 3D convolutional neural network approaches for medical image segmentation. Heliyon, 10(6), e27398. https://doi.org/10.1016/j.heliyon.2024.e27398\n\n2. Zhang, K., Guo, Y., Wang, X., Yuan, J., & Ding, Q. (2019). Multiple Feature Reweight DenseNet for Image Classification. IEEE Access, 7, 9872–9880. https://doi.org/10.1109/ACCESS.2018.2890127\n\n3. https://www.geeksforgeeks.org/densenet-explained/\n\n4. Huang, G., Liu, Z., Pleiss, G., Maaten, L. van der, & Weinberger, K. Q. (2022). Convolutional Networks with Dense Connectivity. IEEE Transactions on Pattern Analysis and Machine Intelligence, 44(12), 8704–8716. https://doi.org/10.1109/TPAMI.2019.2918284\n\n5. https://medium.com/@reh.yawar2/how-the-bottleneck-layers-in-the-deep-networks-work-and-how-do-those-layers-reduce-computational-7bc99c0d1e96\n\n6. Zhou, T., Ye, X., Lu, H., Zheng, X., Qiu, S., Liu, Y., & Li, C. (2022). Dense Convolutional Network and Its Application in Medical Image Analysis. BioMed Research International, 2022(1), 2384830–2384830. https://doi.org/10.1155/2022/2384830\n\n7. https://paperswithcode.com/method/he-initialization#:~:text=Kaiming%20Initialization%2C%20or%20He%20Initialization,functions%2C%20such%20as%20ReLU%20activations.\n\n8. Li, S., Kou, P., Ma, M., Yang, H., Huang, S., & Yang, Z. (2024). Application of Semi-supervised Learning in Image Classification: Research on Fusion of Labeled and Unlabeled Data. IEEE Access, 12, 1–1. https://doi.org/10.1109/ACCESS.2024.3367772\n\n9. Chen, Z., Jing, L., Yang, L., Li, Y., & Li, B. (2023). Class-Level Confidence Based 3D Semi-Supervised Learning. 2023 IEEE/CVF Winter Conference on Applications of Computer Vision (WACV), 633–642. https://doi.org/10.1109/WACV56688.2023.00070\n\n10. https://github.com/YU1ut/MixMatch-pytorch\n\n11. Chen, Z., Jing, L., Liang, Y., Tian, Y., & Li, B. (2021). Multimodal Semi-Supervised Learning for 3D Objects. https://doi.org/10.48550/arxiv.2110.11601\n","metadata":{}}]}