{"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":"none","dataSources":[{"sourceId":24800,"databundleVersionId":1831594,"sourceType":"competition"},{"sourceId":386525,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":318712,"modelId":339285}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nfrom pathlib import Path\n\n# Load full metadata\ntrain_df = pd.read_csv(\"/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train.csv\")\n\n# Get first 1000 DICOM image paths\ninput_dir = Path(\"/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train\")\nimage_paths = list(input_dir.glob(\"*.dicom\"))[:1000]\n\n# Extract image IDs (without .dicom extension)\nimage_ids_1000 = [p.stem for p in image_paths]\n\n# Filter metadata for those 1000 images\nsubset_df = train_df[train_df[\"image_id\"].isin(image_ids_1000)]\n\n# Display\nprint(subset_df.shape)\nsubset_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:11.814641Z","iopub.execute_input":"2025-05-10T10:04:11.814909Z","iopub.status.idle":"2025-05-10T10:04:12.77677Z","shell.execute_reply.started":"2025-05-10T10:04:11.814887Z","shell.execute_reply":"2025-05-10T10:04:12.775877Z"}},"outputs":[{"name":"stdout","text":"(4416, 8)\n","output_type":"stream"},{"name":"stderr","text":"/usr/local/lib/python3.11/dist-packages/pandas/io/formats/format.py:1458: RuntimeWarning: invalid value encountered in greater\n  has_large_values = (abs_vals > 1e6).any()\n/usr/local/lib/python3.11/dist-packages/pandas/io/formats/format.py:1459: RuntimeWarning: invalid value encountered in less\n  has_small_values = ((abs_vals < 10 ** (-self.digits)) & (abs_vals > 0)).any()\n/usr/local/lib/python3.11/dist-packages/pandas/io/formats/format.py:1459: RuntimeWarning: invalid value encountered in greater\n  has_small_values = ((abs_vals < 10 ** (-self.digits)) & (abs_vals > 0)).any()\n","output_type":"stream"},{"execution_count":2,"output_type":"execute_result","data":{"text/plain":"                            image_id    class_name  class_id rad_id   x_min  \\\n12  5550a493b1c4554da469a072fdfab974    No finding        14     R9     NaN   \n15  f55460fccf2d3c591f57f9c0de2c37c2    No finding        14     R6     NaN   \n28  9b85b7ef757927db44393d03083a757c    No finding        14    R17     NaN   \n36  be1bb194dfb986bf7554b491852b8901  Lung Opacity         7     R9  2233.0   \n45  25f2c7b53a6ed09a9aaf73c30357aaf6  Cardiomegaly         3     R8   707.0   \n\n     y_min   x_max   y_max  \n12     NaN     NaN     NaN  \n15     NaN     NaN     NaN  \n28     NaN     NaN     NaN  \n36  1536.0  2518.0  1827.0  \n45  1316.0  2026.0  1686.0  ","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>image_id</th>\n      <th>class_name</th>\n      <th>class_id</th>\n      <th>rad_id</th>\n      <th>x_min</th>\n      <th>y_min</th>\n      <th>x_max</th>\n      <th>y_max</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>12</th>\n      <td>5550a493b1c4554da469a072fdfab974</td>\n      <td>No finding</td>\n      <td>14</td>\n      <td>R9</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n    </tr>\n    <tr>\n      <th>15</th>\n      <td>f55460fccf2d3c591f57f9c0de2c37c2</td>\n      <td>No finding</td>\n      <td>14</td>\n      <td>R6</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n    </tr>\n    <tr>\n      <th>28</th>\n      <td>9b85b7ef757927db44393d03083a757c</td>\n      <td>No finding</td>\n      <td>14</td>\n      <td>R17</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n    </tr>\n    <tr>\n      <th>36</th>\n      <td>be1bb194dfb986bf7554b491852b8901</td>\n      <td>Lung Opacity</td>\n      <td>7</td>\n      <td>R9</td>\n      <td>2233.0</td>\n      <td>1536.0</td>\n      <td>2518.0</td>\n      <td>1827.0</td>\n    </tr>\n    <tr>\n      <th>45</th>\n      <td>25f2c7b53a6ed09a9aaf73c30357aaf6</td>\n      <td>Cardiomegaly</td>\n      <td>3</td>\n      <td>R8</td>\n      <td>707.0</td>\n      <td>1316.0</td>\n      <td>2026.0</td>\n      <td>1686.0</td>\n    </tr>\n  </tbody>\n</table>\n</div>"},"metadata":{}}],"execution_count":2},{"cell_type":"code","source":"subset_df[\"image_id\"].nunique()  # Should be 1000","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:12.777986Z","iopub.execute_input":"2025-05-10T10:04:12.778502Z","iopub.status.idle":"2025-05-10T10:04:12.784951Z","shell.execute_reply.started":"2025-05-10T10:04:12.778482Z","shell.execute_reply":"2025-05-10T10:04:12.784085Z"}},"outputs":[{"execution_count":3,"output_type":"execute_result","data":{"text/plain":"1000"},"metadata":{}}],"execution_count":3},{"cell_type":"code","source":"subset_df_unique = subset_df.drop_duplicates(subset=\"image_id\")\nprint(len(subset_df_unique))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:12.785777Z","iopub.execute_input":"2025-05-10T10:04:12.786066Z","iopub.status.idle":"2025-05-10T10:04:12.798527Z","shell.execute_reply.started":"2025-05-10T10:04:12.78604Z","shell.execute_reply":"2025-05-10T10:04:12.797822Z"}},"outputs":[{"name":"stdout","text":"1000\n","output_type":"stream"}],"execution_count":4},{"cell_type":"code","source":"print(\"Unique image_ids in subset:\", subset_df[\"image_id\"].nunique())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:12.800491Z","iopub.execute_input":"2025-05-10T10:04:12.800691Z","iopub.status.idle":"2025-05-10T10:04:12.810573Z","shell.execute_reply.started":"2025-05-10T10:04:12.800666Z","shell.execute_reply":"2025-05-10T10:04:12.809932Z"}},"outputs":[{"name":"stdout","text":"Unique image_ids in subset: 1000\n","output_type":"stream"}],"execution_count":5},{"cell_type":"code","source":"# Step 1: Create a mapping from image_id to image_path\nimage_path_map = {p.stem: str(p) for p in image_paths}  # {image_id: full_path}\n\n# Step 2: Add image_path column to the subset_df using map()\nsubset_df_unique.loc[:, \"image_path\"] = subset_df_unique[\"image_id\"].map(image_path_map)\n\n# Step 3: Select only required columns\nmetadata = subset_df_unique[[\"image_id\", \"class_name\", \"class_id\", \"image_path\"]]\n\n# Optional: Reset index\nmetadata = metadata.reset_index(drop=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:12.811254Z","iopub.execute_input":"2025-05-10T10:04:12.811477Z","iopub.status.idle":"2025-05-10T10:04:12.825674Z","shell.execute_reply.started":"2025-05-10T10:04:12.811456Z","shell.execute_reply":"2025-05-10T10:04:12.824792Z"}},"outputs":[{"name":"stderr","text":"/tmp/ipykernel_31/856752283.py:5: SettingWithCopyWarning: \nA value is trying to be set on a copy of a slice from a DataFrame.\nTry using .loc[row_indexer,col_indexer] = value instead\n\nSee the caveats in the documentation: https://pandas.pydata.org/pandas-docs/stable/user_guide/indexing.html#returning-a-view-versus-a-copy\n  subset_df_unique.loc[:, \"image_path\"] = subset_df_unique[\"image_id\"].map(image_path_map)\n","output_type":"stream"}],"execution_count":6},{"cell_type":"code","source":"print(metadata[\"image_id\"].iloc[333])\nprint(metadata[\"image_path\"].iloc[333])\nprint(len(metadata))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:12.826698Z","iopub.execute_input":"2025-05-10T10:04:12.826947Z","iopub.status.idle":"2025-05-10T10:04:12.839245Z","shell.execute_reply.started":"2025-05-10T10:04:12.826923Z","shell.execute_reply":"2025-05-10T10:04:12.838691Z"}},"outputs":[{"name":"stdout","text":"72a61ca0140d994daccd2dc3ad4f3905\n/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train/72a61ca0140d994daccd2dc3ad4f3905.dicom\n1000\n","output_type":"stream"}],"execution_count":7},{"cell_type":"code","source":"unique_classes = metadata[\"class_name\"].unique()\nprint(unique_classes)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:12.840053Z","iopub.execute_input":"2025-05-10T10:04:12.840285Z","iopub.status.idle":"2025-05-10T10:04:12.852039Z","shell.execute_reply.started":"2025-05-10T10:04:12.840262Z","shell.execute_reply":"2025-05-10T10:04:12.851373Z"}},"outputs":[{"name":"stdout","text":"['No finding' 'Lung Opacity' 'Cardiomegaly' 'Nodule/Mass'\n 'Pulmonary fibrosis' 'Other lesion' 'Aortic enlargement'\n 'Pleural thickening' 'Infiltration' 'Pleural effusion' 'Atelectasis'\n 'Consolidation' 'ILD' 'Calcification' 'Pneumothorax']\n","output_type":"stream"}],"execution_count":8},{"cell_type":"code","source":"class_counts = metadata[\"class_name\"].value_counts()\nprint(class_counts)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:12.852773Z","iopub.execute_input":"2025-05-10T10:04:12.852968Z","iopub.status.idle":"2025-05-10T10:04:12.863549Z","shell.execute_reply.started":"2025-05-10T10:04:12.852943Z","shell.execute_reply":"2025-05-10T10:04:12.862732Z"}},"outputs":[{"name":"stdout","text":"class_name\nNo finding            714\nAortic enlargement     70\nCardiomegaly           55\nPleural thickening     36\nPulmonary fibrosis     30\nLung Opacity           23\nOther lesion           19\nInfiltration           15\nPleural effusion       11\nNodule/Mass             9\nILD                     6\nCalcification           6\nAtelectasis             4\nConsolidation           1\nPneumothorax            1\nName: count, dtype: int64\n","output_type":"stream"}],"execution_count":9},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport pydicom\nfrom PIL import Image\nimport numpy as np\nimport pandas as pd\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom pathlib import Path\nimport os\nfrom sklearn.model_selection import train_test_split\nfrom torch import optim\nfrom torch.utils.data import WeightedRandomSampler\nfrom transformers import ViTForImageClassification, SwinForImageClassification\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:12.86432Z","iopub.execute_input":"2025-05-10T10:04:12.865048Z","iopub.status.idle":"2025-05-10T10:04:38.590864Z","shell.execute_reply.started":"2025-05-10T10:04:12.865025Z","shell.execute_reply":"2025-05-10T10:04:38.590244Z"}},"outputs":[{"name":"stderr","text":"2025-05-10 10:04:27.292631: E external/local_xla/xla/stream_executor/cuda/cuda_fft.cc:477] Unable to register cuFFT factory: Attempting to register factory for plugin cuFFT when one has already been registered\nWARNING: All log messages before absl::InitializeLog() is called are written to STDERR\nE0000 00:00:1746871467.492267      31 cuda_dnn.cc:8310] Unable to register cuDNN factory: Attempting to register factory for plugin cuDNN when one has already been registered\nE0000 00:00:1746871467.559380      31 cuda_blas.cc:1418] Unable to register cuBLAS factory: Attempting to register factory for plugin cuBLAS when one has already been registered\n","output_type":"stream"}],"execution_count":10},{"cell_type":"code","source":"class DICOMDataset(Dataset):\n    def __init__(self, metadata_df, transform=None):\n        self.metadata_df = metadata_df\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.metadata_df)\n\n    def __getitem__(self, idx):\n        image_path = self.metadata_df.iloc[idx]['image_path']\n        label = self.metadata_df.iloc[idx]['class_id']\n\n        # Load DICOM image\n        dicom_image = pydicom.dcmread(image_path).pixel_array\n\n        # Convert to float32 (this is the key fix)\n        image = torch.tensor(dicom_image, dtype=torch.float32)\n\n        # If transform is specified, apply the transformation pipeline\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:38.593165Z","iopub.execute_input":"2025-05-10T10:04:38.593599Z","iopub.status.idle":"2025-05-10T10:04:38.599167Z","shell.execute_reply.started":"2025-05-10T10:04:38.593582Z","shell.execute_reply":"2025-05-10T10:04:38.598319Z"}},"outputs":[],"execution_count":11},{"cell_type":"code","source":"# Define transformations (with normalization)\nfrom torchvision import transforms\n\ntransform = transforms.Compose([\n    transforms.ToPILImage(),  # Convert numpy array to PIL Image if needed\n    transforms.Resize((224, 224)),  # Resize to the input size required by ViT and Swin\n    transforms.Grayscale(num_output_channels=3),  # Ensure the image is grayscale\n    transforms.ToTensor(),  # Convert to Tensor and normalize to [0, 1]\n    transforms.Normalize(mean=[0.5], std=[0.5])  # Normalize the image to the range expected by the models\n])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:38.599972Z","iopub.execute_input":"2025-05-10T10:04:38.60025Z","iopub.status.idle":"2025-05-10T10:04:38.617132Z","shell.execute_reply.started":"2025-05-10T10:04:38.600228Z","shell.execute_reply":"2025-05-10T10:04:38.616387Z"}},"outputs":[],"execution_count":12},{"cell_type":"code","source":"\n\n# Filter out rare classes\ncounts = metadata['class_id'].value_counts()\nvalid_classes = counts[counts > 1].index\nmetadata = metadata[metadata['class_id'].isin(valid_classes)]\n\n# Create a mapping from class_id to contiguous indices (0 to N-1)\nclass_id_to_contiguous = {cls_id: idx for idx, cls_id in enumerate(valid_classes)}\n\n# Compute class weights for valid classes\ncomputed_weights = compute_class_weight(\n    class_weight='balanced', \n    classes=valid_classes, \n    y=metadata['class_id']\n)\n\n# Initialize full weight tensor for 15 classes (the original number of classes)\nfull_class_weights = torch.zeros(15, dtype=torch.float)\n\n# Populate weights only for valid classes\nfor class_id, weight_idx in class_id_to_contiguous.items():\n    full_class_weights[class_id] = computed_weights[weight_idx]\n\n# Send weights to device\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nclass_weights = full_class_weights.to(device)\ncriterion = torch.nn.CrossEntropyLoss(weight=class_weights)\n\n# Now split into training and validation sets\ntrain_metadata, val_metadata = train_test_split(\n    metadata, test_size=0.2, stratify=metadata['class_id'], random_state=42\n)\n\n# Use weights for sampler (train set)\nsample_weights = torch.tensor(train_metadata['class_id'].map(lambda class_id: class_weights[class_id].item()).tolist(), dtype=torch.float)\nsampler = WeightedRandomSampler(weights=sample_weights, num_samples=len(train_metadata), replacement=True)\n\n# Now you can proceed to create your DataLoaders with the train and validation datasets","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:38.617992Z","iopub.execute_input":"2025-05-10T10:04:38.618365Z","iopub.status.idle":"2025-05-10T10:04:38.846881Z","shell.execute_reply.started":"2025-05-10T10:04:38.618317Z","shell.execute_reply":"2025-05-10T10:04:38.846255Z"}},"outputs":[],"execution_count":13},{"cell_type":"code","source":"\n# Create datasets\ntrain_dataset = DICOMDataset(metadata_df=train_metadata,  transform=transform)\nval_dataset = DICOMDataset(metadata_df=val_metadata, transform=transform)\n\n# Create DataLoaders for training and validation\ntrain_loader = DataLoader(train_dataset, batch_size=32, sampler=sampler)  # Balanced sampling\nval_loader = DataLoader(val_dataset, batch_size=32, shuffle=True)  # Regular shuffle for validation","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:38.84761Z","iopub.execute_input":"2025-05-10T10:04:38.84789Z","iopub.status.idle":"2025-05-10T10:04:38.852387Z","shell.execute_reply.started":"2025-05-10T10:04:38.847865Z","shell.execute_reply":"2025-05-10T10:04:38.85177Z"}},"outputs":[],"execution_count":14},{"cell_type":"code","source":"from transformers import ViTForImageClassification\nfrom torch import nn\n\n# Load pre-trained ViT model\nvit_model = ViTForImageClassification.from_pretrained('google/vit-base-patch16-224-in21k')\n\n# Replace the classification head with a new one (15 classes)\nvit_model.classifier = nn.Linear(vit_model.classifier.in_features, 15)  # 15 is the number of classes\n\n# Ensure model is on the same device (GPU/CPU)\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nvit_model = vit_model.to(device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:38.853182Z","iopub.execute_input":"2025-05-10T10:04:38.853469Z","iopub.status.idle":"2025-05-10T10:04:41.474016Z","shell.execute_reply.started":"2025-05-10T10:04:38.853454Z","shell.execute_reply":"2025-05-10T10:04:41.473478Z"}},"outputs":[{"output_type":"display_data","data":{"text/plain":"config.json:   0%|          | 0.00/502 [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"644a79374833431abaf319acae33d3f0"}},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"model.safetensors:   0%|          | 0.00/346M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"8720bad9441847c7a9fcba50304f0fac"}},"metadata":{}},{"name":"stderr","text":"Some weights of ViTForImageClassification were not initialized from the model checkpoint at google/vit-base-patch16-224-in21k and are newly initialized: ['classifier.bias', 'classifier.weight']\nYou should probably TRAIN this model on a down-stream task to be able to use it for predictions and inference.\n","output_type":"stream"}],"execution_count":15},{"cell_type":"code","source":"from transformers import SwinForImageClassification\nfrom torch import nn\n\n# Load pre-trained Swin model\nswin_model = SwinForImageClassification.from_pretrained('microsoft/swin-base-patch4-window7-224')\n\n# Replace the classification head with a new one (15 classes)\nswin_model.classifier = nn.Linear(swin_model.classifier.in_features, 15)  # 15 is the number of classes\n\n# Ensure model is on the same device (GPU/CPU)\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nswin_model = swin_model.to(device)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:41.474925Z","iopub.execute_input":"2025-05-10T10:04:41.475114Z","iopub.status.idle":"2025-05-10T10:04:43.564185Z","shell.execute_reply.started":"2025-05-10T10:04:41.475099Z","shell.execute_reply":"2025-05-10T10:04:43.563452Z"}},"outputs":[{"output_type":"display_data","data":{"text/plain":"config.json:   0%|          | 0.00/71.8k [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"b9ee208b3a0d4dadb7a1c89e540d73ee"}},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"model.safetensors:   0%|          | 0.00/352M [00:00<?, ?B/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"acc04277a25841a888ad0ccbdf41f2d5"}},"metadata":{}}],"execution_count":16},{"cell_type":"code","source":"# Initialize optimizers for both models\noptimizer_vit = optim.AdamW(vit_model.parameters(), lr=1e-4)\noptimizer_swin = optim.AdamW(swin_model.parameters(), lr=1e-4)\n\n\n# Training loop\nnum_epochs = 10\nfor epoch in range(num_epochs):\n    vit_model.train()\n    swin_model.train()\n\n    running_loss_vit = 0.0\n    running_loss_swin = 0.0\n\n    for images, labels in train_loader:\n        images, labels = images.to(device), labels.to(device)\n\n        # Train ViT\n        optimizer_vit.zero_grad()\n        outputs_vit = vit_model(images).logits\n        loss_vit = criterion(outputs_vit, labels)\n        loss_vit.backward() \n        optimizer_vit.step()\n        running_loss_vit += loss_vit.item()\n\n        # Train Swin Transformer\n        optimizer_swin.zero_grad()\n        outputs_swin = swin_model(images).logits\n        loss_swin = criterion(outputs_swin, labels)\n        loss_swin.backward()\n        optimizer_swin.step()\n        running_loss_swin += loss_swin.item()\n\n    print(f\"Epoch [{epoch+1}/{num_epochs}] - ViT Loss: {running_loss_vit:.4f}, Swin Loss: {running_loss_swin:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T10:04:43.565057Z","iopub.execute_input":"2025-05-10T10:04:43.565331Z","iopub.status.idle":"2025-05-10T12:30:04.615099Z","shell.execute_reply.started":"2025-05-10T10:04:43.565308Z","shell.execute_reply":"2025-05-10T12:30:04.614228Z"}},"outputs":[{"name":"stdout","text":"Epoch [1/10] - ViT Loss: 48.8314, Swin Loss: 30.8264\nEpoch [2/10] - ViT Loss: 26.4380, Swin Loss: 7.0080\nEpoch [3/10] - ViT Loss: 15.7371, Swin Loss: 2.3593\nEpoch [4/10] - ViT Loss: 11.2003, Swin Loss: 1.1100\nEpoch [5/10] - ViT Loss: 7.4889, Swin Loss: 0.5120\nEpoch [6/10] - ViT Loss: 5.3536, Swin Loss: 0.3202\nEpoch [7/10] - ViT Loss: 4.0327, Swin Loss: 0.1572\nEpoch [8/10] - ViT Loss: 3.0042, Swin Loss: 0.1097\nEpoch [9/10] - ViT Loss: 2.4728, Swin Loss: 0.1278\nEpoch [10/10] - ViT Loss: 2.0999, Swin Loss: 0.1013\n","output_type":"stream"}],"execution_count":17},{"cell_type":"code","source":"# Save ViT model\ntorch.save(vit_model.state_dict(), \"/kaggle/working/vit_model.pth\")\n\n# Save Swin Transformer model\ntorch.save(swin_model.state_dict(), \"/kaggle/working/swin_model.pth\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T12:32:10.690231Z","iopub.execute_input":"2025-05-10T12:32:10.690972Z","iopub.status.idle":"2025-05-10T12:32:11.605694Z","shell.execute_reply.started":"2025-05-10T12:32:10.69095Z","shell.execute_reply":"2025-05-10T12:32:11.605023Z"}},"outputs":[],"execution_count":18},{"cell_type":"code","source":"import torch\nfrom torchvision import transforms\nimport pydicom\nimport matplotlib.pyplot as plt\n\n# Load and preprocess the image\ndef preprocess_image(dicom_path):\n    dicom_image = pydicom.dcmread(dicom_path).pixel_array\n    image = torch.tensor(dicom_image, dtype=torch.float32)\n    \n    transform = transforms.Compose([\n        transforms.ToPILImage(),\n        transforms.Resize((224, 224)),\n        transforms.Grayscale(num_output_channels=3),\n        transforms.ToTensor(),\n        transforms.Normalize(mean=[0.5], std=[0.5])\n    ])\n    \n    image = transform(image)\n    image = image.unsqueeze(0)  # Add batch dimension\n    return image\n\n# Inference function\ndef predict_class(model, dicom_path, class_names):\n    model.eval()\n    image = preprocess_image(dicom_path).to(device)\n    \n    with torch.no_grad():\n        logits = model(image).logits  # For ViT and Swin models\n        probs = torch.softmax(logits, dim=1)\n        predicted_class = torch.argmax(probs, dim=1).item()\n    \n    print(f\"Predicted class: {class_names[predicted_class]}\")\n    return predicted_class, probs.cpu().numpy()\n\n# Example usage\nnew_dicom_path = \"/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train/2a848e44e179f5e9b7d708835cfa5109.dicom\"\nclass_names = [\n    \"Aortic enlargement\", \"Atelectasis\", \"Calcification\", \"Cardiomegaly\",\n    \"Consolidation\", \"ILD\", \"Infiltration\", \"Lung Opacity\", \"Nodule/Mass\",\n    \"Other lesion\", \"Pleural effusion\", \"Pleural thickening\", \"Pneumothorax\",\n    \"Pulmonary fibrosis\", \"No finding\"\n]\n\n# Choose one of the trained models\npredict_class(vit_model, new_dicom_path, class_names)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T12:47:17.43018Z","iopub.execute_input":"2025-05-10T12:47:17.430917Z","iopub.status.idle":"2025-05-10T12:47:19.80165Z","shell.execute_reply.started":"2025-05-10T12:47:17.430887Z","shell.execute_reply":"2025-05-10T12:47:19.800827Z"}},"outputs":[{"name":"stdout","text":"Predicted class: Aortic enlargement\n","output_type":"stream"},{"execution_count":24,"output_type":"execute_result","data":{"text/plain":"(0,\n array([[0.67088604, 0.01157351, 0.01311743, 0.11564867, 0.01327584,\n         0.01077338, 0.01051545, 0.0150988 , 0.01406651, 0.01886388,\n         0.01079162, 0.03021857, 0.01799492, 0.02055695, 0.02661847]],\n       dtype=float32))"},"metadata":{}}],"execution_count":24},{"cell_type":"code","source":"predict_class(swin_model,new_dicom_path,class_names)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T12:47:31.14543Z","iopub.execute_input":"2025-05-10T12:47:31.145735Z","iopub.status.idle":"2025-05-10T12:47:33.533507Z","shell.execute_reply.started":"2025-05-10T12:47:31.145716Z","shell.execute_reply":"2025-05-10T12:47:33.532904Z"}},"outputs":[{"name":"stdout","text":"Predicted class: Lung Opacity\n","output_type":"stream"},{"execution_count":25,"output_type":"execute_result","data":{"text/plain":"(7,\n array([[3.3181933e-01, 1.8496761e-05, 2.3031657e-04, 2.4356321e-02,\n         5.2644254e-04, 8.0999243e-06, 3.9093429e-03, 5.4631841e-01,\n         1.7522972e-05, 4.1905302e-04, 2.9774781e-06, 1.7737367e-04,\n         9.8294644e-05, 1.4266321e-03, 9.0671360e-02]], dtype=float32))"},"metadata":{}}],"execution_count":25},{"cell_type":"code","source":"image_paths = list(input_dir.glob(\"*.dicom\"))\nimage_1301_path = image_paths[1300]  # Index 1300 for the 1301st image\nprint(image_1301_path)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T12:43:57.406785Z","iopub.execute_input":"2025-05-10T12:43:57.407553Z","iopub.status.idle":"2025-05-10T12:43:57.443225Z","shell.execute_reply.started":"2025-05-10T12:43:57.40753Z","shell.execute_reply":"2025-05-10T12:43:57.442471Z"}},"outputs":[{"name":"stdout","text":"/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train/2a848e44e179f5e9b7d708835cfa5109.dicom\n","output_type":"stream"}],"execution_count":22},{"cell_type":"code","source":"import pandas as pd\n\n# Load training data CSV\ndf = pd.read_csv(\"/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train.csv\")\n\n# Filter for the image you're testing\nimage_id = \"2a848e44e179f5e9b7d708835cfa5109\"\ndf_image = df[df[\"image_id\"] == image_id]\n\n# Get unique labels for this image\ntrue_labels = df_image[\"class_id\"].unique()\nprint(\"Ground truth labels:\", true_labels)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T12:44:23.682159Z","iopub.execute_input":"2025-05-10T12:44:23.682445Z","iopub.status.idle":"2025-05-10T12:44:23.765785Z","shell.execute_reply.started":"2025-05-10T12:44:23.682427Z","shell.execute_reply":"2025-05-10T12:44:23.765106Z"}},"outputs":[{"name":"stdout","text":"Ground truth labels: [ 3 10 11  0]\n","output_type":"stream"}],"execution_count":23},{"cell_type":"code","source":"def evaluate_model(model, data_loader, device, threshold=0.5):\n    model.eval()\n    all_targets = []\n    all_outputs = []\n\n    with torch.no_grad():\n        for images, labels in tqdm(data_loader):\n            images = images.to(device)\n            labels = labels.to(device)\n\n            outputs = model(images).logits  # Shape: [B, num_classes]\n            probs = torch.sigmoid(outputs)  # Multi-label sigmoid\n            all_outputs.append(probs.cpu())\n            all_targets.append(labels.cpu())\n\n    y_true = torch.cat(all_targets).numpy()         # Shape: [N, num_classes]\n    y_probs = torch.cat(all_outputs).numpy()        # Shape: [N, num_classes]\n    y_pred = (y_probs >= threshold).astype(int)     # Binary predictions\n\n    # Compute evaluation metrics\n    accuracy = accuracy_score(y_true, y_pred)\n    macro_f1 = f1_score(y_true, y_pred, average='macro')\n    macro_precision = precision_score(y_true, y_pred, average='macro')\n    macro_recall = recall_score(y_true, y_pred, average='macro')\n    try:\n        macro_auc = roc_auc_score(y_true, y_probs, average='macro')  # probs not pred\n    except ValueError:\n        macro_auc = None  # Some classes might be missing\n\n    # Class-wise AUCs\n    class_aucs = []\n    for i in range(y_true.shape[1]):\n        try:\n            auc = roc_auc_score(y_true[:, i], y_probs[:, i])\n        except ValueError:\n            auc = None\n        class_aucs.append(auc)\n\n    results = {\n        \"accuracy\": accuracy,\n        \"macro_f1\": macro_f1,\n        \"macro_precision\": macro_precision,\n        \"macro_recall\": macro_recall,\n        \"macro_auc\": macro_auc,\n        \"class_wise_auc\": class_aucs\n    }\n\n    return results\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T12:53:09.684405Z","iopub.execute_input":"2025-05-10T12:53:09.684914Z","iopub.status.idle":"2025-05-10T12:53:09.692153Z","shell.execute_reply.started":"2025-05-10T12:53:09.684892Z","shell.execute_reply":"2025-05-10T12:53:09.691343Z"}},"outputs":[],"execution_count":27},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nfrom sklearn.metrics import accuracy_score, f1_score, roc_auc_score, precision_score, recall_score\nfrom tqdm import tqdm\nimport numpy as np\n\n\nvit_results = evaluate_model(vit_model, val_loader, device)\nswin_results = evaluate_model(swin_model, val_loader, device)\n\nprint(\"ViT Evaluation:\")\nprint(vit_results)\n\nprint(\"\\nSwin Transformer Evaluation:\")\nprint(swin_results)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-10T12:59:47.086482Z","iopub.execute_input":"2025-05-10T12:59:47.086791Z","iopub.status.idle":"2025-05-10T13:02:44.024697Z","shell.execute_reply.started":"2025-05-10T12:59:47.086772Z","shell.execute_reply":"2025-05-10T13:02:44.023792Z"}},"outputs":[{"name":"stderr","text":"100%|██████████| 7/7 [02:56<00:00, 25.27s/it]\n","output_type":"stream"},{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mValueError\u001b[0m                                Traceback (most recent call last)","\u001b[0;32m/tmp/ipykernel_31/985417795.py\u001b[0m in \u001b[0;36m<cell line: 0>\u001b[0;34m()\u001b[0m\n\u001b[1;32m      6\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m      7\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m----> 8\u001b[0;31m \u001b[0mvit_results\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mevaluate_model\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mvit_model\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mval_loader\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mdevice\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m      9\u001b[0m \u001b[0mswin_results\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mevaluate_model\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mswin_model\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mval_loader\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mdevice\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     10\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/tmp/ipykernel_31/3051755691.py\u001b[0m in \u001b[0;36mevaluate_model\u001b[0;34m(model, data_loader, device, threshold)\u001b[0m\n\u001b[1;32m     19\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     20\u001b[0m     \u001b[0;31m# Compute evaluation metrics\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 21\u001b[0;31m     \u001b[0maccuracy\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0maccuracy_score\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0my_true\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0my_pred\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     22\u001b[0m     \u001b[0mmacro_f1\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mf1_score\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0my_true\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0my_pred\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0maverage\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;34m'macro'\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     23\u001b[0m     \u001b[0mmacro_precision\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mprecision_score\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0my_true\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0my_pred\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0maverage\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;34m'macro'\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/sklearn/utils/_param_validation.py\u001b[0m in \u001b[0;36mwrapper\u001b[0;34m(*args, **kwargs)\u001b[0m\n\u001b[1;32m    190\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    191\u001b[0m             \u001b[0;32mtry\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 192\u001b[0;31m                 \u001b[0;32mreturn\u001b[0m \u001b[0mfunc\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m*\u001b[0m\u001b[0margs\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    193\u001b[0m             \u001b[0;32mexcept\u001b[0m \u001b[0mInvalidParameterError\u001b[0m \u001b[0;32mas\u001b[0m \u001b[0me\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    194\u001b[0m                 \u001b[0;31m# When the function is just a wrapper around an estimator, we allow\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/sklearn/metrics/_classification.py\u001b[0m in \u001b[0;36maccuracy_score\u001b[0;34m(y_true, y_pred, normalize, sample_weight)\u001b[0m\n\u001b[1;32m    219\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    220\u001b[0m     \u001b[0;31m# Compute accuracy for each possible representation\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 221\u001b[0;31m     \u001b[0my_type\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0my_true\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0my_pred\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0m_check_targets\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0my_true\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0my_pred\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    222\u001b[0m     \u001b[0mcheck_consistent_length\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0my_true\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0my_pred\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0msample_weight\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    223\u001b[0m     \u001b[0;32mif\u001b[0m \u001b[0my_type\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mstartswith\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"multilabel\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/sklearn/metrics/_classification.py\u001b[0m in \u001b[0;36m_check_targets\u001b[0;34m(y_true, y_pred)\u001b[0m\n\u001b[1;32m     93\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     94\u001b[0m     \u001b[0;32mif\u001b[0m \u001b[0mlen\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0my_type\u001b[0m\u001b[0;34m)\u001b[0m \u001b[0;34m>\u001b[0m \u001b[0;36m1\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 95\u001b[0;31m         raise ValueError(\n\u001b[0m\u001b[1;32m     96\u001b[0m             \"Classification metrics can't handle a mix of {0} and {1} targets\".format(\n\u001b[1;32m     97\u001b[0m                 \u001b[0mtype_true\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtype_pred\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;31mValueError\u001b[0m: Classification metrics can't handle a mix of multiclass and multilabel-indicator targets"],"ename":"ValueError","evalue":"Classification metrics can't handle a mix of multiclass and multilabel-indicator targets","output_type":"error"}],"execution_count":31},{"cell_type":"code","source":"import pandas as pd\nfrom pathlib import Path\ntest_images_dir = Path(\"/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train\")\n\nfull_metadata = pd.read_csv(\"/kaggle/input/vinbigdata-chest-xray-abnormalities-detection/train.csv\")\n\n# Extract DICOM paths 2001–3000\ntest_image_paths = list(test_images_dir.glob(\"*.dicom\"))[2000:3000]\ntest_image_ids = [p.stem for p in test_image_paths]\n\n# Filter metadata\ntest_metadata = full_metadata[full_metadata[\"image_id\"].isin(test_image_ids)]\n\n# Multi-label: group labels per image\ntest_metadata_grouped = test_metadata.groupby(\"image_id\")[\"class_id\"].apply(list).reset_index()\n\n# Map image_id to path\ntest_path_map = {p.stem: str(p) for p in test_image_paths}\ntest_metadata_grouped[\"image_path\"] = test_metadata_grouped[\"image_id\"].map(test_path_map)\n\n# Reset index\ntest_metadata_grouped = test_metadata_grouped.reset_index(drop=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T07:27:13.146116Z","iopub.execute_input":"2025-05-11T07:27:13.146403Z","iopub.status.idle":"2025-05-11T07:27:13.395928Z","shell.execute_reply.started":"2025-05-11T07:27:13.146381Z","shell.execute_reply":"2025-05-11T07:27:13.394919Z"}},"outputs":[],"execution_count":5},{"cell_type":"code","source":"test_metadata_grouped","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T07:27:32.12339Z","iopub.execute_input":"2025-05-11T07:27:32.123729Z","iopub.status.idle":"2025-05-11T07:27:32.149659Z","shell.execute_reply.started":"2025-05-11T07:27:32.123697Z","shell.execute_reply":"2025-05-11T07:27:32.148879Z"}},"outputs":[{"execution_count":6,"output_type":"execute_result","data":{"text/plain":"                             image_id  \\\n0    000d68e42b71d3eac10ccc077aba07c1   \n1    009d4c31ebf87e51c5c8c160a4bd8006   \n2    010018c93ed33ae56ed048ee54867e46   \n3    01546d3e6175ceaabd7d92f0c566579d   \n4    015bf89fc34cde9fafe7c79366fecee7   \n..                                ...   \n995  ff4cd5b2a61258ac551673d1311fa2a8   \n996  ff7a3edad07fc3a153f865ce329307a4   \n997  ff87bf702fb77cbf48bdf732d2f2defa   \n998  ffbecda150d808687714dd54bd3cb2a6   \n999  ffeffc54594debf3716d6fcd2402a99f   \n\n                                              class_id  \\\n0    [9, 9, 9, 9, 11, 13, 11, 0, 9, 9, 7, 9, 9, 9, ...   \n1                      [13, 7, 10, 4, 0, 7, 10, 7, 10]   \n2    [13, 13, 13, 11, 11, 13, 8, 0, 0, 3, 11, 11, 8...   \n3                                   [3, 3, 0, 9, 8, 8]   \n4                                         [14, 14, 14]   \n..                                                 ...   \n995                                       [14, 14, 14]   \n996                                       [14, 14, 14]   \n997                                       [14, 14, 14]   \n998                                       [14, 14, 14]   \n999                                          [0, 0, 0]   \n\n                                            image_path  \n0    /kaggle/input/vinbigdata-chest-xray-abnormalit...  \n1    /kaggle/input/vinbigdata-chest-xray-abnormalit...  \n2    /kaggle/input/vinbigdata-chest-xray-abnormalit...  \n3    /kaggle/input/vinbigdata-chest-xray-abnormalit...  \n4    /kaggle/input/vinbigdata-chest-xray-abnormalit...  \n..                                                 ...  \n995  /kaggle/input/vinbigdata-chest-xray-abnormalit...  \n996  /kaggle/input/vinbigdata-chest-xray-abnormalit...  \n997  /kaggle/input/vinbigdata-chest-xray-abnormalit...  \n998  /kaggle/input/vinbigdata-chest-xray-abnormalit...  \n999  /kaggle/input/vinbigdata-chest-xray-abnormalit...  \n\n[1000 rows x 3 columns]","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>image_id</th>\n      <th>class_id</th>\n      <th>image_path</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>0</th>\n      <td>000d68e42b71d3eac10ccc077aba07c1</td>\n      <td>[9, 9, 9, 9, 11, 13, 11, 0, 9, 9, 7, 9, 9, 9, ...</td>\n      <td>/kaggle/input/vinbigdata-chest-xray-abnormalit...</td>\n    </tr>\n    <tr>\n      <th>1</th>\n      <td>009d4c31ebf87e51c5c8c160a4bd8006</td>\n      <td>[13, 7, 10, 4, 0, 7, 10, 7, 10]</td>\n      <td>/kaggle/input/vinbigdata-chest-xray-abnormalit...</td>\n    </tr>\n    <tr>\n      <th>2</th>\n      <td>010018c93ed33ae56ed048ee54867e46</td>\n      <td>[13, 13, 13, 11, 11, 13, 8, 0, 0, 3, 11, 11, 8...</td>\n      <td>/kaggle/input/vinbigdata-chest-xray-abnormalit...</td>\n    </tr>\n    <tr>\n      <th>3</th>\n      <td>01546d3e6175ceaabd7d92f0c566579d</td>\n      <td>[3, 3, 0, 9, 8, 8]</td>\n      <td>/kaggle/input/vinbigdata-chest-xray-abnormalit...</td>\n    </tr>\n    <tr>\n      <th>4</th>\n      <td>015bf89fc34cde9fafe7c79366fecee7</td>\n      <td>[14, 14, 14]</td>\n      <td>/kaggle/input/vinbigdata-chest-xray-abnormalit...</td>\n    </tr>\n    <tr>\n      <th>...</th>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n    </tr>\n    <tr>\n      <th>995</th>\n      <td>ff4cd5b2a61258ac551673d1311fa2a8</td>\n      <td>[14, 14, 14]</td>\n      <td>/kaggle/input/vinbigdata-chest-xray-abnormalit...</td>\n    </tr>\n    <tr>\n      <th>996</th>\n      <td>ff7a3edad07fc3a153f865ce329307a4</td>\n      <td>[14, 14, 14]</td>\n      <td>/kaggle/input/vinbigdata-chest-xray-abnormalit...</td>\n    </tr>\n    <tr>\n      <th>997</th>\n      <td>ff87bf702fb77cbf48bdf732d2f2defa</td>\n      <td>[14, 14, 14]</td>\n      <td>/kaggle/input/vinbigdata-chest-xray-abnormalit...</td>\n    </tr>\n    <tr>\n      <th>998</th>\n      <td>ffbecda150d808687714dd54bd3cb2a6</td>\n      <td>[14, 14, 14]</td>\n      <td>/kaggle/input/vinbigdata-chest-xray-abnormalit...</td>\n    </tr>\n    <tr>\n      <th>999</th>\n      <td>ffeffc54594debf3716d6fcd2402a99f</td>\n      <td>[0, 0, 0]</td>\n      <td>/kaggle/input/vinbigdata-chest-xray-abnormalit...</td>\n    </tr>\n  </tbody>\n</table>\n<p>1000 rows × 3 columns</p>\n</div>"},"metadata":{}}],"execution_count":6},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\n\n\nclass MultiLabelDICOMDataset(Dataset):\n    def __init__(self, metadata_df, num_classes=15, transform=None):\n        self.metadata_df = metadata_df\n        self.transform = transform\n        self.num_classes = num_classes\n\n    def __len__(self):\n        return len(self.metadata_df)\n\n    def __getitem__(self, idx):\n        row = self.metadata_df.iloc[idx]\n        image_path = row['image_path']\n        labels = row['class_id']  # This is a list\n\n        dicom_image = pydicom.dcmread(image_path).pixel_array\n        image = torch.tensor(dicom_image, dtype=torch.float32)\n\n        if self.transform:\n            image = self.transform(image)\n\n        # Multi-hot encode labels\n        target = torch.zeros(self.num_classes, dtype=torch.float32)\n        for label in labels:\n            target[label] = 1.0\n\n        return image, target\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T07:29:45.637449Z","iopub.execute_input":"2025-05-11T07:29:45.63775Z","iopub.status.idle":"2025-05-11T07:29:50.719804Z","shell.execute_reply.started":"2025-05-11T07:29:45.637725Z","shell.execute_reply":"2025-05-11T07:29:50.718877Z"}},"outputs":[],"execution_count":8},{"cell_type":"code","source":"from torchvision import transforms\n\ntransform = transforms.Compose([\n    transforms.ToPILImage(),  # Convert numpy array to PIL Image\n    transforms.Resize((224, 224)),  # Match input size of ViT/Swin\n    transforms.Grayscale(num_output_channels=3),  # Convert single channel to 3-channel\n    transforms.ToTensor(),  # Convert to Tensor\n    transforms.Normalize(mean=[0.5], std=[0.5])  # Normalize (you can adjust mean/std if needed)\n])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T07:30:32.964116Z","iopub.execute_input":"2025-05-11T07:30:32.96443Z","iopub.status.idle":"2025-05-11T07:30:37.043011Z","shell.execute_reply.started":"2025-05-11T07:30:32.964407Z","shell.execute_reply":"2025-05-11T07:30:37.041954Z"}},"outputs":[],"execution_count":10},{"cell_type":"code","source":"test_dataset = MultiLabelDICOMDataset(metadata_df=test_metadata_grouped, transform=transform)\ntest_loader = DataLoader(test_dataset, batch_size=32, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T07:30:45.786511Z","iopub.execute_input":"2025-05-11T07:30:45.786825Z","iopub.status.idle":"2025-05-11T07:30:45.792708Z","shell.execute_reply.started":"2025-05-11T07:30:45.786801Z","shell.execute_reply":"2025-05-11T07:30:45.791365Z"}},"outputs":[],"execution_count":12},{"cell_type":"code","source":"test_dataset","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T07:30:53.96423Z","iopub.execute_input":"2025-05-11T07:30:53.964546Z","iopub.status.idle":"2025-05-11T07:30:53.970941Z","shell.execute_reply.started":"2025-05-11T07:30:53.964523Z","shell.execute_reply":"2025-05-11T07:30:53.97004Z"}},"outputs":[{"execution_count":13,"output_type":"execute_result","data":{"text/plain":"<__main__.MultiLabelDICOMDataset at 0x7d4e87756210>"},"metadata":{}}],"execution_count":13},{"cell_type":"code","source":"from sklearn.metrics import f1_score, roc_auc_score, accuracy_score, classification_report\n\ndef evaluate_model(model, test_loader, threshold=0.5):\n    model.eval()\n    all_labels = []\n    all_preds = []\n\n    with torch.no_grad():\n        for images, labels in test_loader:\n            images = images.to(device)\n            labels = labels.to(device)\n            outputs = torch.sigmoid(model(images).logits)  # Multi-label\n\n            preds = (outputs > threshold).float()\n\n            all_labels.append(labels.cpu())\n            all_preds.append(preds.cpu())\n\n    y_true = torch.cat(all_labels).numpy()\n    y_pred = torch.cat(all_preds).numpy()\n\n    # Compute per-metric\n    f1 = f1_score(y_true, y_pred, average='macro', zero_division=0)\n    acc = accuracy_score(y_true, y_pred)\n    roc_auc = roc_auc_score(y_true, y_pred, average='macro')\n\n    print(\"Classification Report:\\n\", classification_report(y_true, y_pred, zero_division=0))\n\n    return {\n        'F1 Score': f1,\n        'Accuracy': acc,\n        'ROC AUC': roc_auc\n    }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T07:31:54.034416Z","iopub.execute_input":"2025-05-11T07:31:54.034736Z","iopub.status.idle":"2025-05-11T07:31:54.883759Z","shell.execute_reply.started":"2025-05-11T07:31:54.034712Z","shell.execute_reply":"2025-05-11T07:31:54.88292Z"}},"outputs":[],"execution_count":14},{"cell_type":"code","source":"from transformers import ViTForImageClassification, SwinForImageClassification\nimport torch.nn as nn\nimport torch\n\n# Step 1: Define model architectures (15 classes)\nvit_model = ViTForImageClassification.from_pretrained('google/vit-base-patch16-224-in21k')\nvit_model.classifier = nn.Linear(vit_model.classifier.in_features, 15)\n\nswin_model = SwinForImageClassification.from_pretrained('microsoft/swin-base-patch4-window7-224')\nswin_model.classifier = nn.Linear(swin_model.classifier.in_features, 15)\n\n# Step 2: Load weights\nvit_model.load_state_dict(torch.load(\"/kaggle/input/vit-and-swin/tensorflow2/default/1/vit_model.pth\", map_location='cpu'))\nswin_model.load_state_dict(torch.load(\"/kaggle/input/vit-and-swin/tensorflow2/default/1/swin_model.pth\", map_location='cpu'))\n\n# Step 3: Move to device and set to eval mode\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nvit_model = vit_model.to(device)\nswin_model = swin_model.to(device)\n\nvit_model.eval()\nswin_model.eval()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T07:46:29.450072Z","iopub.execute_input":"2025-05-11T07:46:29.450421Z","iopub.status.idle":"2025-05-11T07:46:35.042473Z","shell.execute_reply.started":"2025-05-11T07:46:29.450395Z","shell.execute_reply":"2025-05-11T07:46:35.041305Z"}},"outputs":[{"name":"stderr","text":"Some weights of ViTForImageClassification were not initialized from the model checkpoint at google/vit-base-patch16-224-in21k and are newly initialized: ['classifier.bias', 'classifier.weight']\nYou should probably TRAIN this model on a down-stream task to be able to use it for predictions and inference.\n/tmp/ipykernel_31/801021637.py:13: FutureWarning: You are using `torch.load` with `weights_only=False` (the current default value), which uses the default pickle module implicitly. It is possible to construct malicious pickle data which will execute arbitrary code during unpickling (See https://github.com/pytorch/pytorch/blob/main/SECURITY.md#untrusted-models for more details). In a future release, the default value for `weights_only` will be flipped to `True`. This limits the functions that could be executed during unpickling. Arbitrary objects will no longer be allowed to be loaded via this mode unless they are explicitly allowlisted by the user via `torch.serialization.add_safe_globals`. We recommend you start setting `weights_only=True` for any use case where you don't have full control of the loaded file. Please open an issue on GitHub for any issues related to this experimental feature.\n  vit_model.load_state_dict(torch.load(\"/kaggle/input/vit-and-swin/tensorflow2/default/1/vit_model.pth\", map_location='cpu'))\n/tmp/ipykernel_31/801021637.py:14: FutureWarning: You are using `torch.load` with `weights_only=False` (the current default value), which uses the default pickle module implicitly. It is possible to construct malicious pickle data which will execute arbitrary code during unpickling (See https://github.com/pytorch/pytorch/blob/main/SECURITY.md#untrusted-models for more details). In a future release, the default value for `weights_only` will be flipped to `True`. This limits the functions that could be executed during unpickling. Arbitrary objects will no longer be allowed to be loaded via this mode unless they are explicitly allowlisted by the user via `torch.serialization.add_safe_globals`. We recommend you start setting `weights_only=True` for any use case where you don't have full control of the loaded file. Please open an issue on GitHub for any issues related to this experimental feature.\n  swin_model.load_state_dict(torch.load(\"/kaggle/input/vit-and-swin/tensorflow2/default/1/swin_model.pth\", map_location='cpu'))\n","output_type":"stream"},{"execution_count":17,"output_type":"execute_result","data":{"text/plain":"SwinForImageClassification(\n  (swin): SwinModel(\n    (embeddings): SwinEmbeddings(\n      (patch_embeddings): SwinPatchEmbeddings(\n        (projection): Conv2d(3, 128, kernel_size=(4, 4), stride=(4, 4))\n      )\n      (norm): LayerNorm((128,), eps=1e-05, elementwise_affine=True)\n      (dropout): Dropout(p=0.0, inplace=False)\n    )\n    (encoder): SwinEncoder(\n      (layers): ModuleList(\n        (0): SwinStage(\n          (blocks): ModuleList(\n            (0): SwinLayer(\n              (layernorm_before): LayerNorm((128,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=128, out_features=128, bias=True)\n                  (key): Linear(in_features=128, out_features=128, bias=True)\n                  (value): Linear(in_features=128, out_features=128, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=128, out_features=128, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): Identity()\n              (layernorm_after): LayerNorm((128,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=128, out_features=512, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=512, out_features=128, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (1): SwinLayer(\n              (layernorm_before): LayerNorm((128,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=128, out_features=128, bias=True)\n                  (key): Linear(in_features=128, out_features=128, bias=True)\n                  (value): Linear(in_features=128, out_features=128, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=128, out_features=128, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.004347826354205608)\n              (layernorm_after): LayerNorm((128,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=128, out_features=512, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=512, out_features=128, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n          )\n          (downsample): SwinPatchMerging(\n            (reduction): Linear(in_features=512, out_features=256, bias=False)\n            (norm): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n          )\n        )\n        (1): SwinStage(\n          (blocks): ModuleList(\n            (0): SwinLayer(\n              (layernorm_before): LayerNorm((256,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=256, out_features=256, bias=True)\n                  (key): Linear(in_features=256, out_features=256, bias=True)\n                  (value): Linear(in_features=256, out_features=256, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=256, out_features=256, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.008695652708411217)\n              (layernorm_after): LayerNorm((256,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=256, out_features=1024, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=1024, out_features=256, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (1): SwinLayer(\n              (layernorm_before): LayerNorm((256,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=256, out_features=256, bias=True)\n                  (key): Linear(in_features=256, out_features=256, bias=True)\n                  (value): Linear(in_features=256, out_features=256, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=256, out_features=256, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.013043479062616825)\n              (layernorm_after): LayerNorm((256,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=256, out_features=1024, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=1024, out_features=256, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n          )\n          (downsample): SwinPatchMerging(\n            (reduction): Linear(in_features=1024, out_features=512, bias=False)\n            (norm): LayerNorm((1024,), eps=1e-05, elementwise_affine=True)\n          )\n        )\n        (2): SwinStage(\n          (blocks): ModuleList(\n            (0): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.017391305416822433)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (1): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.021739132702350616)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (2): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.02608695812523365)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (3): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.030434783548116684)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (4): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.03478261083364487)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (5): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.03913043811917305)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (6): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.04347826540470123)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (7): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.04782608896493912)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (8): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.052173912525177)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (9): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.056521736085414886)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (10): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.06086956337094307)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (11): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.06521739065647125)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (12): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.06956521421670914)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (13): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.07391304522752762)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (14): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.0782608687877655)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (15): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.08260869979858398)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (16): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.08695652335882187)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (17): SwinLayer(\n              (layernorm_before): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=512, out_features=512, bias=True)\n                  (key): Linear(in_features=512, out_features=512, bias=True)\n                  (value): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=512, out_features=512, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.09130434691905975)\n              (layernorm_after): LayerNorm((512,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=512, out_features=2048, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=2048, out_features=512, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n          )\n          (downsample): SwinPatchMerging(\n            (reduction): Linear(in_features=2048, out_features=1024, bias=False)\n            (norm): LayerNorm((2048,), eps=1e-05, elementwise_affine=True)\n          )\n        )\n        (3): SwinStage(\n          (blocks): ModuleList(\n            (0): SwinLayer(\n              (layernorm_before): LayerNorm((1024,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=1024, out_features=1024, bias=True)\n                  (key): Linear(in_features=1024, out_features=1024, bias=True)\n                  (value): Linear(in_features=1024, out_features=1024, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=1024, out_features=1024, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.09565217792987823)\n              (layernorm_after): LayerNorm((1024,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=1024, out_features=4096, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=4096, out_features=1024, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n            (1): SwinLayer(\n              (layernorm_before): LayerNorm((1024,), eps=1e-05, elementwise_affine=True)\n              (attention): SwinAttention(\n                (self): SwinSelfAttention(\n                  (query): Linear(in_features=1024, out_features=1024, bias=True)\n                  (key): Linear(in_features=1024, out_features=1024, bias=True)\n                  (value): Linear(in_features=1024, out_features=1024, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n                (output): SwinSelfOutput(\n                  (dense): Linear(in_features=1024, out_features=1024, bias=True)\n                  (dropout): Dropout(p=0.0, inplace=False)\n                )\n              )\n              (drop_path): SwinDropPath(p=0.10000000149011612)\n              (layernorm_after): LayerNorm((1024,), eps=1e-05, elementwise_affine=True)\n              (intermediate): SwinIntermediate(\n                (dense): Linear(in_features=1024, out_features=4096, bias=True)\n                (intermediate_act_fn): GELUActivation()\n              )\n              (output): SwinOutput(\n                (dense): Linear(in_features=4096, out_features=1024, bias=True)\n                (dropout): Dropout(p=0.0, inplace=False)\n              )\n            )\n          )\n        )\n      )\n    )\n    (layernorm): LayerNorm((1024,), eps=1e-05, elementwise_affine=True)\n    (pooler): AdaptiveAvgPool1d(output_size=1)\n  )\n  (classifier): Linear(in_features=1024, out_features=15, bias=True)\n)"},"metadata":{}}],"execution_count":17},{"cell_type":"code","source":"import pydicom\nmetrics_vit = evaluate_model(vit_model, test_loader)\nmetrics_swin = evaluate_model(swin_model, test_loader)\n\nimport pandas as pd\ncomparison_df = pd.DataFrame([metrics_vit, metrics_swin], index=[\"ViT\", \"Swin\"])\nprint(comparison_df)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T07:47:59.506056Z","iopub.execute_input":"2025-05-11T07:47:59.506377Z","iopub.status.idle":"2025-05-11T08:25:14.635847Z","shell.execute_reply.started":"2025-05-11T07:47:59.506354Z","shell.execute_reply":"2025-05-11T08:25:14.633776Z"}},"outputs":[{"name":"stdout","text":"Classification Report:\n               precision    recall  f1-score   support\n\n           0       0.22      0.98      0.36       222\n           1       0.00      0.00      0.00        13\n           2       0.04      0.10      0.05        41\n           3       0.16      0.97      0.28       158\n           4       0.00      0.00      0.00        26\n           5       0.07      0.12      0.09        32\n           6       0.07      0.14      0.10        36\n           7       0.13      0.75      0.22        79\n           8       0.08      0.12      0.09        58\n           9       0.06      0.37      0.10        84\n          10       0.12      0.16      0.14        81\n          11       0.15      0.78      0.25       148\n          12       0.00      0.00      0.00         6\n          13       0.11      0.53      0.18       124\n          14       0.69      0.84      0.76       695\n\n   micro avg       0.22      0.70      0.34      1803\n   macro avg       0.13      0.39      0.17      1803\nweighted avg       0.35      0.70      0.42      1803\n samples avg       0.23      0.78      0.33      1803\n\n","output_type":"stream"},{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mKeyboardInterrupt\u001b[0m                         Traceback (most recent call last)","\u001b[0;32m/tmp/ipykernel_31/3051162046.py\u001b[0m in \u001b[0;36m<cell line: 0>\u001b[0;34m()\u001b[0m\n\u001b[1;32m      1\u001b[0m \u001b[0;32mimport\u001b[0m \u001b[0mpydicom\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m      2\u001b[0m \u001b[0mmetrics_vit\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mevaluate_model\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mvit_model\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtest_loader\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m----> 3\u001b[0;31m \u001b[0mmetrics_swin\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mevaluate_model\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mswin_model\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mtest_loader\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m      4\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m      5\u001b[0m \u001b[0;32mimport\u001b[0m \u001b[0mpandas\u001b[0m \u001b[0;32mas\u001b[0m \u001b[0mpd\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/tmp/ipykernel_31/2924289302.py\u001b[0m in \u001b[0;36mevaluate_model\u001b[0;34m(model, test_loader, threshold)\u001b[0m\n\u001b[1;32m      7\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m      8\u001b[0m     \u001b[0;32mwith\u001b[0m \u001b[0mtorch\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mno_grad\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m----> 9\u001b[0;31m         \u001b[0;32mfor\u001b[0m \u001b[0mimages\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mlabels\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mtest_loader\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     10\u001b[0m             \u001b[0mimages\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mimages\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mto\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mdevice\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     11\u001b[0m             \u001b[0mlabels\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mlabels\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mto\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mdevice\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/torch/utils/data/dataloader.py\u001b[0m in \u001b[0;36m__next__\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m    699\u001b[0m                 \u001b[0;31m# TODO(https://github.com/pytorch/pytorch/issues/76750)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    700\u001b[0m                 \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_reset\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m  \u001b[0;31m# type: ignore[call-arg]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 701\u001b[0;31m             \u001b[0mdata\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_next_data\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    702\u001b[0m             \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_num_yielded\u001b[0m \u001b[0;34m+=\u001b[0m \u001b[0;36m1\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    703\u001b[0m             if (\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/torch/utils/data/dataloader.py\u001b[0m in \u001b[0;36m_next_data\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m    755\u001b[0m     \u001b[0;32mdef\u001b[0m \u001b[0m_next_data\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    756\u001b[0m         \u001b[0mindex\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_next_index\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m  \u001b[0;31m# may raise StopIteration\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 757\u001b[0;31m         \u001b[0mdata\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_dataset_fetcher\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mfetch\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mindex\u001b[0m\u001b[0;34m)\u001b[0m  \u001b[0;31m# may raise StopIteration\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    758\u001b[0m         \u001b[0;32mif\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_pin_memory\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    759\u001b[0m             \u001b[0mdata\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0m_utils\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpin_memory\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpin_memory\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mdata\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_pin_memory_device\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/torch/utils/data/_utils/fetch.py\u001b[0m in \u001b[0;36mfetch\u001b[0;34m(self, possibly_batched_index)\u001b[0m\n\u001b[1;32m     50\u001b[0m                 \u001b[0mdata\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdataset\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m__getitems__\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mpossibly_batched_index\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     51\u001b[0m             \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 52\u001b[0;31m                 \u001b[0mdata\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;34m[\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdataset\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0midx\u001b[0m\u001b[0;34m]\u001b[0m \u001b[0;32mfor\u001b[0m \u001b[0midx\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mpossibly_batched_index\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     53\u001b[0m         \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     54\u001b[0m             \u001b[0mdata\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdataset\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mpossibly_batched_index\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/torch/utils/data/_utils/fetch.py\u001b[0m in \u001b[0;36m<listcomp>\u001b[0;34m(.0)\u001b[0m\n\u001b[1;32m     50\u001b[0m                 \u001b[0mdata\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdataset\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m__getitems__\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mpossibly_batched_index\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     51\u001b[0m             \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 52\u001b[0;31m                 \u001b[0mdata\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;34m[\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdataset\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0midx\u001b[0m\u001b[0;34m]\u001b[0m \u001b[0;32mfor\u001b[0m \u001b[0midx\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mpossibly_batched_index\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     53\u001b[0m         \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     54\u001b[0m             \u001b[0mdata\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdataset\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mpossibly_batched_index\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/tmp/ipykernel_31/3321203798.py\u001b[0m in \u001b[0;36m__getitem__\u001b[0;34m(self, idx)\u001b[0m\n\u001b[1;32m     16\u001b[0m         \u001b[0mlabels\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mrow\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;34m'class_id'\u001b[0m\u001b[0;34m]\u001b[0m  \u001b[0;31m# This is a list\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     17\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 18\u001b[0;31m         \u001b[0mdicom_image\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mpydicom\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdcmread\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mimage_path\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpixel_array\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m     19\u001b[0m         \u001b[0mimage\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mtorch\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtensor\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mdicom_image\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mdtype\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mtorch\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mfloat32\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m     20\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/pydicom/dataset.py\u001b[0m in \u001b[0;36mpixel_array\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m   2191\u001b[0m             \u001b[0mthat\u001b[0m \u001b[0miterates\u001b[0m \u001b[0mthrough\u001b[0m \u001b[0mthe\u001b[0m \u001b[0mimage\u001b[0m \u001b[0mframes\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   2192\u001b[0m         \"\"\"\n\u001b[0;32m-> 2193\u001b[0;31m         \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mconvert_pixel_data\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   2194\u001b[0m         \u001b[0;32mreturn\u001b[0m \u001b[0mcast\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"numpy.ndarray\"\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_pixel_array\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   2195\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/pydicom/dataset.py\u001b[0m in \u001b[0;36mconvert_pixel_data\u001b[0;34m(self, handler_name)\u001b[0m\n\u001b[1;32m   1724\u001b[0m             \u001b[0;31m# Use 'pydicom.pixels' backend\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1725\u001b[0m             \u001b[0mopts\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;34m\"decoding_plugin\"\u001b[0m\u001b[0;34m]\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mname\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1726\u001b[0;31m             \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_pixel_array\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mpixel_array\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mopts\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1727\u001b[0m             \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_pixel_id\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mget_image_pixel_ids\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1728\u001b[0m         \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/pydicom/pixels/utils.py\u001b[0m in \u001b[0;36mpixel_array\u001b[0;34m(src, ds_out, specific_tags, index, raw, decoding_plugin, **kwargs)\u001b[0m\n\u001b[1;32m   1428\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1429\u001b[0m         \u001b[0mopts\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mas_pixel_options\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mds\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m**\u001b[0m\u001b[0mkwargs\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1430\u001b[0;31m         return decoder.as_array(\n\u001b[0m\u001b[1;32m   1431\u001b[0m             \u001b[0mds\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1432\u001b[0m             \u001b[0mindex\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mindex\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/pydicom/pixels/decoders/base.py\u001b[0m in \u001b[0;36mas_array\u001b[0;34m(self, src, index, validate, raw, decoding_plugin, **kwargs)\u001b[0m\n\u001b[1;32m    998\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    999\u001b[0m         \u001b[0mas_frame\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mindex\u001b[0m \u001b[0;32mis\u001b[0m \u001b[0;32mnot\u001b[0m \u001b[0;32mNone\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1000\u001b[0;31m         \u001b[0marr\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mrunner\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mreshape\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mfunc\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mrunner\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mindex\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mas_frame\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mas_frame\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1001\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1002\u001b[0m         \u001b[0;32mif\u001b[0m \u001b[0mrunner\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_test_for\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m\"sign_correction\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/pydicom/pixels/decoders/base.py\u001b[0m in \u001b[0;36m_as_array_encapsulated\u001b[0;34m(runner, index)\u001b[0m\n\u001b[1;32m   1065\u001b[0m         \u001b[0mframe_generator\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mrunner\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0miter_decode\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1066\u001b[0m         \u001b[0;32mfor\u001b[0m \u001b[0midx\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mrange\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mrunner\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mnumber_of_frames\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m-> 1067\u001b[0;31m             \u001b[0mframe\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mnext\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mframe_generator\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m   1068\u001b[0m             \u001b[0mstart\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0midx\u001b[0m \u001b[0;34m*\u001b[0m \u001b[0mpixels_per_frame\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m   1069\u001b[0m             arr[start : start + pixels_per_frame] = np.frombuffer(\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/pydicom/pixels/decoders/base.py\u001b[0m in \u001b[0;36miter_decode\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m    456\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    457\u001b[0m             \u001b[0;31m# Otherwise try all decoders\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 458\u001b[0;31m             \u001b[0;32myield\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_decode_frame\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0msrc\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    459\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    460\u001b[0m         \u001b[0;32mif\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mis_binary\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/pydicom/pixels/decoders/base.py\u001b[0m in \u001b[0;36m_decode_frame\u001b[0;34m(self, src)\u001b[0m\n\u001b[1;32m    353\u001b[0m             \u001b[0;32mtry\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    354\u001b[0m                 \u001b[0;31m# Attempt to decode the frame\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 355\u001b[0;31m                 \u001b[0mframe\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mfunc\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0msrc\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    356\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    357\u001b[0m                 \u001b[0;31m# Decode success, if we were previously successful then\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/pydicom/pixels/decoders/pillow.py\u001b[0m in \u001b[0;36m_decode_frame\u001b[0;34m(src, runner)\u001b[0m\n\u001b[1;32m    100\u001b[0m     \u001b[0;31m# Pillow converts N-bit signed/unsigned data to 8- or 16-bit unsigned data\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    101\u001b[0m     \u001b[0;31m#   See Pillow src/libImaging/Jpeg2KDecode.c::j2ku_gray_i\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 102\u001b[0;31m     \u001b[0mbuffer\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mbytearray\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mimage\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtobytes\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m)\u001b[0m  \u001b[0;31m# so the array is writeable\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    103\u001b[0m     \u001b[0;32mdel\u001b[0m \u001b[0mimage\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    104\u001b[0m     \u001b[0mdtype\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mrunner\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpixel_dtype\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/PIL/Image.py\u001b[0m in \u001b[0;36mtobytes\u001b[0;34m(self, encoder_name, *args)\u001b[0m\n\u001b[1;32m    794\u001b[0m             \u001b[0mencoder_args\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mmode\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    795\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 796\u001b[0;31m         \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mload\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    797\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    798\u001b[0m         \u001b[0;32mif\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mwidth\u001b[0m \u001b[0;34m==\u001b[0m \u001b[0;36m0\u001b[0m \u001b[0;32mor\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mheight\u001b[0m \u001b[0;34m==\u001b[0m \u001b[0;36m0\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/PIL/Jpeg2KImagePlugin.py\u001b[0m in \u001b[0;36mload\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m    349\u001b[0m             \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mtile\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;34m[\u001b[0m\u001b[0mImageFile\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0m_Tile\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mt\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;36m0\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;34m(\u001b[0m\u001b[0;36m0\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0;36m0\u001b[0m\u001b[0;34m)\u001b[0m \u001b[0;34m+\u001b[0m \u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0msize\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mt\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;36m2\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m,\u001b[0m \u001b[0mt3\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    350\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 351\u001b[0;31m         \u001b[0;32mreturn\u001b[0m \u001b[0mImageFile\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mImageFile\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mload\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    352\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    353\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;32m/usr/local/lib/python3.11/dist-packages/PIL/ImageFile.py\u001b[0m in \u001b[0;36mload\u001b[0;34m(self)\u001b[0m\n\u001b[1;32m    273\u001b[0m                     \u001b[0;32mif\u001b[0m \u001b[0mdecoder\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mpulls_fd\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    274\u001b[0m                         \u001b[0mdecoder\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0msetfd\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mself\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mfp\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m--> 275\u001b[0;31m                         \u001b[0merr_code\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mdecoder\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mdecode\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34mb\"\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0;36m1\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m    276\u001b[0m                     \u001b[0;32melse\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m    277\u001b[0m                         \u001b[0mb\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0mprefix\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n","\u001b[0;31mKeyboardInterrupt\u001b[0m: "],"ename":"KeyboardInterrupt","evalue":"","output_type":"error"}],"execution_count":19},{"cell_type":"code","source":"test_metadata_grouped = test_metadata.groupby(\"image_id\")[\"class_id\"].apply(lambda x: list(set(x))).reset_index()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T08:25:24.019094Z","iopub.execute_input":"2025-05-11T08:25:24.01947Z","iopub.status.idle":"2025-05-11T08:25:24.057675Z","shell.execute_reply.started":"2025-05-11T08:25:24.01944Z","shell.execute_reply":"2025-05-11T08:25:24.056477Z"}},"outputs":[],"execution_count":20},{"cell_type":"code","source":"test_metadata_grouped","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T08:25:32.243626Z","iopub.execute_input":"2025-05-11T08:25:32.244037Z","iopub.status.idle":"2025-05-11T08:25:32.258383Z","shell.execute_reply.started":"2025-05-11T08:25:32.244008Z","shell.execute_reply":"2025-05-11T08:25:32.2576Z"}},"outputs":[{"execution_count":21,"output_type":"execute_result","data":{"text/plain":"                             image_id           class_id\n0    000d68e42b71d3eac10ccc077aba07c1  [0, 7, 9, 11, 13]\n1    009d4c31ebf87e51c5c8c160a4bd8006  [0, 4, 7, 10, 13]\n2    010018c93ed33ae56ed048ee54867e46  [0, 3, 8, 11, 13]\n3    01546d3e6175ceaabd7d92f0c566579d       [0, 9, 3, 8]\n4    015bf89fc34cde9fafe7c79366fecee7               [14]\n..                                ...                ...\n995  ff4cd5b2a61258ac551673d1311fa2a8               [14]\n996  ff7a3edad07fc3a153f865ce329307a4               [14]\n997  ff87bf702fb77cbf48bdf732d2f2defa               [14]\n998  ffbecda150d808687714dd54bd3cb2a6               [14]\n999  ffeffc54594debf3716d6fcd2402a99f                [0]\n\n[1000 rows x 2 columns]","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>image_id</th>\n      <th>class_id</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>0</th>\n      <td>000d68e42b71d3eac10ccc077aba07c1</td>\n      <td>[0, 7, 9, 11, 13]</td>\n    </tr>\n    <tr>\n      <th>1</th>\n      <td>009d4c31ebf87e51c5c8c160a4bd8006</td>\n      <td>[0, 4, 7, 10, 13]</td>\n    </tr>\n    <tr>\n      <th>2</th>\n      <td>010018c93ed33ae56ed048ee54867e46</td>\n      <td>[0, 3, 8, 11, 13]</td>\n    </tr>\n    <tr>\n      <th>3</th>\n      <td>01546d3e6175ceaabd7d92f0c566579d</td>\n      <td>[0, 9, 3, 8]</td>\n    </tr>\n    <tr>\n      <th>4</th>\n      <td>015bf89fc34cde9fafe7c79366fecee7</td>\n      <td>[14]</td>\n    </tr>\n    <tr>\n      <th>...</th>\n      <td>...</td>\n      <td>...</td>\n    </tr>\n    <tr>\n      <th>995</th>\n      <td>ff4cd5b2a61258ac551673d1311fa2a8</td>\n      <td>[14]</td>\n    </tr>\n    <tr>\n      <th>996</th>\n      <td>ff7a3edad07fc3a153f865ce329307a4</td>\n      <td>[14]</td>\n    </tr>\n    <tr>\n      <th>997</th>\n      <td>ff87bf702fb77cbf48bdf732d2f2defa</td>\n      <td>[14]</td>\n    </tr>\n    <tr>\n      <th>998</th>\n      <td>ffbecda150d808687714dd54bd3cb2a6</td>\n      <td>[14]</td>\n    </tr>\n    <tr>\n      <th>999</th>\n      <td>ffeffc54594debf3716d6fcd2402a99f</td>\n      <td>[0]</td>\n    </tr>\n  </tbody>\n</table>\n<p>1000 rows × 2 columns</p>\n</div>"},"metadata":{}}],"execution_count":21},{"cell_type":"code","source":"def evaluate_any_correct(model, dataloader, threshold=0.5, device='cpu'):\n    model.eval()\n    correct = 0\n    total = 0\n\n    with torch.no_grad():\n        for images, targets in dataloader:\n            images = images.to(device)\n            outputs = model(images).logits  # Assuming your model returns logits\n            probs = torch.sigmoid(outputs)  # Convert to probabilities\n\n            # Predicted class is the one with highest probability\n            preds = torch.argmax(probs, dim=1)\n\n            for pred, true_labels in zip(preds.cpu(), targets):\n                # If the predicted class is in the list of true labels → correct\n                if pred.item() in true_labels:\n                    correct += 1\n                total += 1\n\n    accuracy = correct / total if total > 0 else 0\n    print(f\"Any-Correct Accuracy: {accuracy:.4f}\")\n    return accuracy\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T08:29:01.108487Z","iopub.execute_input":"2025-05-11T08:29:01.10883Z","iopub.status.idle":"2025-05-11T08:29:01.117565Z","shell.execute_reply.started":"2025-05-11T08:29:01.108807Z","shell.execute_reply":"2025-05-11T08:29:01.116709Z"}},"outputs":[],"execution_count":22},{"cell_type":"code","source":"accuracy_vit = evaluate_any_correct(vit_model, test_loader, device=device)\naccuracy_swin = evaluate_any_correct(swin_model, test_loader, device=device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T08:29:14.963703Z","iopub.execute_input":"2025-05-11T08:29:14.964063Z","iopub.status.idle":"2025-05-11T09:10:33.642688Z","shell.execute_reply.started":"2025-05-11T08:29:14.964038Z","shell.execute_reply":"2025-05-11T09:10:33.641034Z"}},"outputs":[{"name":"stdout","text":"Any-Correct Accuracy: 0.6220\nAny-Correct Accuracy: 0.1810\n","output_type":"stream"}],"execution_count":23}]}