{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":11848,"databundleVersionId":862157,"sourceType":"competition"},{"sourceId":11163745,"sourceType":"datasetVersion","datasetId":6966331}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Problem Overview","metadata":{}},{"cell_type":"markdown","source":"The goal of this playground is to analyze metastatic cancer given image patches from pathology scans. Essentially, for each `id` in the set, we need to predict a probability whether there is a pixel of tumor tissue or not at a center 32x32px region.\n\nFor this playground. I will be using a 3-layer CNN alongside a fully connected layer to predict the tumor pixel. While this is a smaller NN for cancer detection, it can start as a benchmark for larger models.\n\nCreated with PyTorch.","metadata":{}},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport torch\nimport matplotlib.pyplot as plt\nimport os\nimport tifffile as tiff\nfrom torch.utils.data import Dataset, DataLoader, random_split, Subset\nfrom sklearn.model_selection import train_test_split\nfrom pathlib import Path\nfrom typing import Literal, Callable\nfrom PIL import Image\nimport torchvision.transforms as transforms\nimport torch.nn as nn\nfrom collections import Counter\nfrom tqdm.notebook import tqdm\nfrom torchinfo import summary\nfrom sklearn.metrics import precision_score, recall_score, f1_score\nimport torchvision.models as models","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:35:54.917310Z","iopub.execute_input":"2025-03-25T21:35:54.917600Z","iopub.status.idle":"2025-03-25T21:36:01.865976Z","shell.execute_reply.started":"2025-03-25T21:35:54.917577Z","shell.execute_reply":"2025-03-25T21:36:01.865306Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else 'cpu')\nprint(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:36:01.867024Z","iopub.execute_input":"2025-03-25T21:36:01.867668Z","iopub.status.idle":"2025-03-25T21:36:01.922375Z","shell.execute_reply.started":"2025-03-25T21:36:01.867635Z","shell.execute_reply":"2025-03-25T21:36:01.921433Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data EDA","metadata":{}},{"cell_type":"markdown","source":"We will take a look at the data first to understand what we are working with, then creat our dataset and showcase some images from the dataset.\n\nThere should be a healthy split of non-cancerous and cancerous tumor images. We should not have duplicates as well.","metadata":{}},{"cell_type":"code","source":"dirname = '/kaggle/input/histopathologic-cancer-detection'\n\ntrain_df = pd.read_csv(f'{dirname}/train_labels.csv')\n\ntrain_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:36:01.924242Z","iopub.execute_input":"2025-03-25T21:36:01.924520Z","iopub.status.idle":"2025-03-25T21:36:02.358753Z","shell.execute_reply.started":"2025-03-25T21:36:01.924494Z","shell.execute_reply":"2025-03-25T21:36:02.357851Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"No Cancer: {len(train_df[train_df['label'] == 0])}, Cancer: {len(train_df[train_df['label'] == 1])}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:36:02.359976Z","iopub.execute_input":"2025-03-25T21:36:02.360284Z","iopub.status.idle":"2025-03-25T21:36:02.382978Z","shell.execute_reply.started":"2025-03-25T21:36:02.360260Z","shell.execute_reply":"2025-03-25T21:36:02.382153Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(8,5))\ntrain_df['label'].map({1: 'Cancer', 0: 'No Cancer'}).value_counts().plot(kind='bar', color=['gray', 'black'])\nplt.xlabel('Label')\nplt.ylabel('Count')\nplt.xticks(rotation=45)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:36:02.383817Z","iopub.execute_input":"2025-03-25T21:36:02.384084Z","iopub.status.idle":"2025-03-25T21:36:02.669870Z","shell.execute_reply.started":"2025-03-25T21:36:02.384032Z","shell.execute_reply":"2025-03-25T21:36:02.669005Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_imgs = os.listdir(f'{dirname}/train')\ntest_imgs = os.listdir(f'{dirname}/test')\nprint(f'Example Files: {train_imgs[:5]}')\nprint(f\"Shape: {Image.open(os.path.join(dirname, 'train', train_imgs[0])).convert('RGB').size}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:36:02.670528Z","iopub.execute_input":"2025-03-25T21:36:02.670767Z","iopub.status.idle":"2025-03-25T21:36:05.653085Z","shell.execute_reply.started":"2025-03-25T21:36:02.670746Z","shell.execute_reply":"2025-03-25T21:36:05.652188Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"From this we know a couple facts about the dataset:\n- About 60% of the dataset has no cancer, and 40% does.\n- We are dealing with 96x96 RGB `tif` images \n\nNow, we need to create our dataset, inspect the images for any patterns for our model, and then create our model","metadata":{}},{"cell_type":"markdown","source":"## Dataset Creation","metadata":{}},{"cell_type":"code","source":"class CancerDS(Dataset):\n    '''\n        Custom dataset for histopathologic cancer detection images\n\n        Args:\n            data_dir: root directory of image data\n            transform: transformation function to apply to the images\n            imgs: list of the image filenames \n            labels: list of the labels (0 (no cancer) or 1 (cancer))\n    '''\n    def __init__(self, data_dir: str = \"\", transform: Callable = None, d_type: Literal['train', 'test'] = 'train'):\n        self.data_dir = Path(data_dir)\n        self.transform = transform\n\n        # find images\n        image_dir = self.data_dir / d_type\n        if not image_dir.exists():\n            raise FileNotFoundError(f'Directory {image_dir} not found')\n    \n        self.imgs = list(image_dir.glob('*.tif'))\n\n        # find labels\n        labels_dir = self.data_dir / 'train_labels.csv'\n        if not labels_dir.exists():\n            raise FileNotFoundError(f'Directory {labels_dir} not found')\n\n        df = pd.read_csv(labels_dir).set_index('id')\n        self.labels = [df.loc[img.stem].values[0] for img in self.imgs]\n    \n    def __len__(self):\n        return len(self.imgs)\n    def __getitem__(self, idx):\n        # open image\n        img_path = self.imgs[idx]\n        img = Image.open(img_path).convert('RGB')\n        label = torch.tensor(self.labels[idx], dtype=torch.float32)\n\n        # apply transform if it exists\n        if self.transform:\n            img = self.transform(img)\n\n        # return image, label, and the image name (the ID)\n        img_id = img_path.stem\n        return img, label, img_id","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:36:05.654136Z","iopub.execute_input":"2025-03-25T21:36:05.654457Z","iopub.status.idle":"2025-03-25T21:36:05.661818Z","shell.execute_reply.started":"2025-03-25T21:36:05.654423Z","shell.execute_reply":"2025-03-25T21:36:05.660820Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Transform for training\ntransform_train = transforms.Compose([\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n# Transform for testing\ntransform_test = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\n# create dataset and split into train, testing\nds = CancerDS(dirname, transform_train, 'train')\n\nindices = np.arange(len(ds))\ntrain_indices, temp_indices = train_test_split(indices, test_size=0.3)\ntest_indices, val_indices = train_test_split(temp_indices, test_size=0.5)\n\ntrain_ds = Subset(ds, train_indices)\ntest_ds = Subset(ds, test_indices)\nval_ds = Subset(ds, val_indices)\ntest_ds.dataset.transform = transform_test\nval_ds.dataset.transform = transform_test","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:36:05.663951Z","iopub.execute_input":"2025-03-25T21:36:05.664244Z","iopub.status.idle":"2025-03-25T21:36:13.609898Z","shell.execute_reply.started":"2025-03-25T21:36:05.664218Z","shell.execute_reply":"2025-03-25T21:36:13.609231Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f'Training: {len(train_ds)}, Testing: {len(test_ds)}, Validation: {len(val_ds)}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:36:13.611299Z","iopub.execute_input":"2025-03-25T21:36:13.611617Z","iopub.status.idle":"2025-03-25T21:36:13.616011Z","shell.execute_reply.started":"2025-03-25T21:36:13.611587Z","shell.execute_reply":"2025-03-25T21:36:13.615151Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dl = DataLoader(train_ds, batch_size=32, shuffle=True, pin_memory=True, num_workers=4, prefetch_factor=4, persistent_workers=True)\ntest_dl = DataLoader(test_ds, batch_size=32, shuffle=False, pin_memory=True, num_workers=4, prefetch_factor=4, persistent_workers=True)\nval_dl = DataLoader(val_ds, batch_size=32, shuffle=False, pin_memory=True, num_workers=4, prefetch_factor=4, persistent_workers=True)\n\nfor indices in train_dl:\n    print(indices[0].shape)\n    break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:36:13.616702Z","iopub.execute_input":"2025-03-25T21:36:13.616927Z","iopub.status.idle":"2025-03-25T21:36:14.320697Z","shell.execute_reply.started":"2025-03-25T21:36:13.616906Z","shell.execute_reply":"2025-03-25T21:36:14.318842Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Example Images","metadata":{}},{"cell_type":"code","source":"# showcasing a couple of the images from the dataset, without the normalization\nexample = next(iter(train_dl))\n\nimages = example[0]\nlabels = example[1]\n\nfig, axes = plt.subplots(1, 5, sharex=True, sharey=True)\nmean = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)\nstd = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)\n\nrandom_indices = torch.randint(low=0, high=32, size=(5,))\nfor i, num in enumerate(random_indices):\n    # take out the normalization and reorganize so matplotlib can show the images\n    img = images[num]\n    img = img * std + mean\n    img = img.permute(1, 2, 0)\n    label = 'Cancer' if labels[num] == 1 else 'No Cancer'\n    \n    axes[i].set_title(label)\n    axes[i].imshow(img)\n    axes[i].axis('off')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:36:14.322679Z","iopub.execute_input":"2025-03-25T21:36:14.323017Z","iopub.status.idle":"2025-03-25T21:36:16.093837Z","shell.execute_reply.started":"2025-03-25T21:36:14.322977Z","shell.execute_reply":"2025-03-25T21:36:16.092798Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Creation","metadata":{}},{"cell_type":"markdown","source":"For this model, I chose to use 3 CNN layers, alongside fully connected layers for our classification purposes. Considering the limited GPU and CPU space, this will work well for our case.\n\nWe also have training and validation datasets to further ensure this model can attain a decent accuracy, and to analyze whether the mode is overfitting / underfitting. The model uses:\n- Dropout layers for improving generalization / reducing overfitting\n- Early stopping for prevention of overfitting (only save the model that does better to also prevent)\n- Learning Rate Scheduler for avoiding local minimas (and better training results)","metadata":{}},{"cell_type":"code","source":"class CancerModel(nn.Module):\n    def __init__(self):\n        super(CancerModel, self).__init__()\n        self.convs = nn.Sequential(\n            nn.Conv2d(3, 32, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(2, 2), # 32, 48 x 48\n\n            nn.Conv2d(32, 64, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(2, 2), # 64, 24 x 24\n\n            nn.Conv2d(64, 128, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.MaxPool2d(2, 2) # 128, 12 x 12\n        )\n\n        self.fc = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(128 * 12 * 12, 256),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(256, 1),\n            nn.Sigmoid()\n        )\n\n    def forward(self, x):\n        x = self.convs(x)\n        x = self.fc(x)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:36:16.094919Z","iopub.execute_input":"2025-03-25T21:36:16.095214Z","iopub.status.idle":"2025-03-25T21:36:16.103205Z","shell.execute_reply.started":"2025-03-25T21:36:16.095190Z","shell.execute_reply":"2025-03-25T21:36:16.102355Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = CancerModel().to(device)\noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)\nlr = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=5)\n\nparams = {\n    'epochs': 25,\n    'optimizer': optimizer,\n    'lr_scheduler': lr,\n    'weight_path': 'cnn1_weights.pt',\n    'loss_fn': nn.BCELoss(),\n    'patience': 7,\n}\n\nsummary(model, input_size=(1, 3, 96, 96))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:36:16.104329Z","iopub.execute_input":"2025-03-25T21:36:16.104633Z","iopub.status.idle":"2025-03-25T21:36:16.856765Z","shell.execute_reply.started":"2025-03-25T21:36:16.104601Z","shell.execute_reply":"2025-03-25T21:36:16.855997Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training / Testing of the Model","metadata":{}},{"cell_type":"code","source":"def train(model, dl, loss_fn, opt, device):\n    model.train()\n    for batch, (X, y, _) in enumerate(dl):\n        X, y = X.to(device), y.to(device).float().unsqueeze(1)\n        pred = model(X)\n        loss = loss_fn(pred, y)\n\n        loss.backward()\n        opt.step()\n        opt.zero_grad()\n\n    return loss.item()\n\ndef val(model, dl, loss_fn, device):\n    model.eval()\n    total_loss, corr = 0, 0\n    with torch.no_grad():\n        for X, y, _ in dl:\n            X, y = X.to(device), y.to(device).float().unsqueeze(1)\n            pred = model(X)\n            total_loss += loss_fn(pred, y).item() * X.size(0)\n            pred = (pred > 0.5).float() # binary output\n            corr += (pred == y).sum().item()\n    avg_loss = total_loss / len(dl.dataset)\n    accuracy = corr / len(dl.dataset)\n    return avg_loss, accuracy\n\ndef train_model(model, train_dl, val_dl, params, device):\n    epochs = params['epochs']\n    optimizer = params['optimizer']\n    scheduler = params['lr_scheduler']\n    weight_path = params['weight_path']\n    loss_func = params['loss_fn']\n    patience = params['patience']\n    \n    best_loss = float('inf')\n    no_improve = 0\n    total_train_loss = []\n    total_val_loss = []\n    \n    for epoch in tqdm(range(epochs), desc='Training'):\n        # training \n        print(f'Epoch {epoch+1}/{epochs}')\n        train_loss = train(model, train_dl, loss_func, optimizer, device)\n        val_loss, val_acc = val(model, val_dl, loss_func, device)\n        print(f'Training Loss: {train_loss:.4f}, Validation Loss: {val_loss:.4f}, Accuracy: {val_acc:.4f}')\n        total_train_loss.append(train_loss)\n        total_val_loss.append(val_loss)\n\n        scheduler.step(val_loss)\n\n        if val_loss < best_loss:\n            best_loss = val_loss\n            no_improve = 0\n            torch.save(model.state_dict(), weight_path)\n            print('New best model.')\n        else:\n            no_improve += 1\n            if no_improve >= patience:\n                print('Model might be overfitting. Stopping early...')\n                break\n    return total_train_loss, total_val_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:36:16.857668Z","iopub.execute_input":"2025-03-25T21:36:16.857973Z","iopub.status.idle":"2025-03-25T21:36:16.867486Z","shell.execute_reply.started":"2025-03-25T21:36:16.857933Z","shell.execute_reply":"2025-03-25T21:36:16.866563Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_losses, val_losses = train_model(model, train_dl, val_dl, params, device) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T16:42:11.650727Z","iopub.execute_input":"2025-03-25T16:42:11.650936Z","iopub.status.idle":"2025-03-25T17:13:06.600028Z","shell.execute_reply.started":"2025-03-25T16:42:11.650916Z","shell.execute_reply":"2025-03-25T17:13:06.598905Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Outcomes","metadata":{}},{"cell_type":"markdown","source":"We can check the accuracy and how precise the model is with the loss and accuracy. However considering our dataset consists of about 40% cancer tumors, and 60% non-cancerous (a little imbalance will give a higher accuracy than wanted), we should use another metric to be more precise.\n- Recall for missing cancer cases\n- Precision for understanding false positives","metadata":{}},{"cell_type":"code","source":"best_model = CancerModel().to(device)\n# best_model.load_state_dict(torch.load(params['weight_path'])) # -> if just trained data\nbest_model.load_state_dict(torch.load('/kaggle/input/cnn-cancerdetection-weights/cnn1_weights.pt', weights_only=True), strict=True) # when uploading weights previously trained\nbest_model.eval()\nbest_loss, best_acc = val(best_model, val_dl, params['loss_fn'], device)\nprint(f'Loss: {best_loss:.4}, Accuracy: {best_acc:.4}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:36:16.869467Z","iopub.execute_input":"2025-03-25T21:36:16.869761Z","iopub.status.idle":"2025-03-25T21:37:31.294385Z","shell.execute_reply.started":"2025-03-25T21:36:16.869726Z","shell.execute_reply":"2025-03-25T21:37:31.293178Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_predictions = []\nall_labels = []\n\nwith torch.no_grad():\n    for X, y, _ in val_dl:\n        X = X.to(device)\n        y = y.to(device)\n        preds = (best_model(X) > 0.5).float()\n\n        all_predictions.extend(preds.cpu().numpy())\n        all_labels.extend(y.cpu().numpy())\n\nprec = precision_score(all_labels, all_predictions)\nrecall = recall_score(all_labels, all_predictions)\nf1 = f1_score(all_labels, all_predictions)\n\nprint(f'Precision: {prec:.4f}, Recall: {recall:.4f}, F1: {f1:.4f}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:37:31.296142Z","iopub.execute_input":"2025-03-25T21:37:31.296443Z","iopub.status.idle":"2025-03-25T21:37:54.936408Z","shell.execute_reply.started":"2025-03-25T21:37:31.296415Z","shell.execute_reply":"2025-03-25T21:37:54.935427Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.plot(train_losses, label='Train Loss')\nplt.plot(val_losses, label='Val Loss')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.title('Train vs Val Loss')\nplt.legend()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T17:13:44.705220Z","iopub.execute_input":"2025-03-25T17:13:44.705531Z","iopub.status.idle":"2025-03-25T17:13:44.925041Z","shell.execute_reply.started":"2025-03-25T17:13:44.705508Z","shell.execute_reply":"2025-03-25T17:13:44.924240Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"From this, we can see a couple things\n- When we detect cancer, we are 93% correct (7% false positive rate)\n- Catch 90% of all cancer cases (10% false negative rate)\n- 92% f1 score (well balanced)\n- After Epoch 6, we started to overfit our model\n\nHonestly this was not as good as what we were hoping for, as an AI in the medical field should not give out false positives 7% of the time, and failing to notice cancerous tumors 10% of the time.\n\nThere is room for improvement though, such as fine-tuning an existing model like **Resnet50** on the dataset, creating a deeper and more intricate CancerModel, or creating an ensemble method between the two.  ","metadata":{}},{"cell_type":"markdown","source":"# Submission / Conclusion","metadata":{}},{"cell_type":"markdown","source":"The final result of this dataset was a score of 0.81, or 81% overall. Some improvements could be:\n- Fine-tuning of more intricate image classification models\n- More Pre-processing of images\n- Modifying the CNN with more layers for feature extraction and perhaps more dropout layers for more generalization\n\nI attempted to fine-tune a Resnet50 model on the data as well but it took far longer than wanted and did not show any promising results.","metadata":{}},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, test_dir, transform=None):\n        self.test_dir = Path(test_dir)\n        self.image_ids = sorted(os.listdir(self.test_dir))\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.image_ids)\n\n    def __getitem__(self, idx):\n        image_id = self.image_ids[idx]\n        img_path = self.test_dir / image_id\n        img = Image.open(img_path).convert(\"RGB\")\n        if self.transform:\n            img = self.transform(img)\n        return img, image_id","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:57:01.456592Z","iopub.execute_input":"2025-03-25T21:57:01.456911Z","iopub.status.idle":"2025-03-25T21:57:01.462100Z","shell.execute_reply.started":"2025-03-25T21:57:01.456888Z","shell.execute_reply":"2025-03-25T21:57:01.461131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dir = f\"{dirname}/test\"\nbatch_size = 64\n\ntest_dataset = TestDataset(test_dir, transform=transform_test)\ntest_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)\n\nbest_model.eval()\nfinal_preds = []\nfinal_ids = []\n\nwith torch.no_grad():\n    for images, ids in test_loader:\n        images = images.to(device)\n        outputs = best_model(images)\n        preds = (outputs > 0.5).float().squeeze().cpu().numpy()\n\n        if preds.ndim == 0:\n            preds = [preds.item()]\n        else:\n            preds = preds.tolist()\n\n        final_preds.extend(preds)\n        final_ids.extend([i.split('.')[0] for i in ids])\n\nresults_df = pd.DataFrame({\n    'id': final_ids,\n    'label': final_preds\n})\n\nresults_df.to_csv('submission.csv', index=False)\nresults_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-25T21:57:04.023620Z","iopub.execute_input":"2025-03-25T21:57:04.023905Z","iopub.status.idle":"2025-03-25T21:59:49.244283Z","shell.execute_reply.started":"2025-03-25T21:57:04.023885Z","shell.execute_reply":"2025-03-25T21:59:49.243515Z"}},"outputs":[],"execution_count":null}]}