{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pydicom\nimport numpy as np\nimport pandas as pd\nimport scipy\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision\nimport torchvision.transforms as transforms\nfrom torchvision.models import ResNet50_Weights\nfrom torch.utils.data import Dataset, DataLoader, Subset, WeightedRandomSampler\nfrom tqdm import tqdm\nimport sklearn.metrics\nfrom sklearn.model_selection import train_test_split\nfrom skimage import io\nimport os\nimport matplotlib.pyplot as plt\nimport cv2\nimport gc\nimport glob\nfrom tabulate import tabulate\nimport random\nimport seaborn as sns\nsns.set_theme(style=\"whitegrid\", palette=\"Set2\")\nimport warnings\nwarnings.filterwarnings(\"ignore\")\nimport wandb\nfrom joblib import Parallel, delayed\n\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nwandb_api = user_secrets.get_secret(\"wandb_api\") \nwandb.login(key=wandb_api)\n\ndevice = torch.device(\"cuda\" if (torch.cuda.is_available()) else \"cpu\")","metadata":{"_uuid":"95dd93b0-712b-4006-b488-78604dc7d1de","_cell_guid":"36d12cdb-7133-48a9-a111-4477479a6469","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-09-30T14:57:58.988601Z","iopub.execute_input":"2023-09-30T14:57:58.988947Z","iopub.status.idle":"2023-09-30T14:57:59.713061Z","shell.execute_reply.started":"2023-09-30T14:57:58.98892Z","shell.execute_reply":"2023-09-30T14:57:59.712045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Motivation and Goals\nThis notebook shows my attempts to learn Pytorch and apply it towards a medical imaging classification problem. Using data from the RSNA Mammography challenge, the aim is to be able to identify breast cancer from mammogram images. It involves some challenge in processing the associated tabular data as well as the DICOM format mammogram images, but is overall a good introduction to working with a relatively large dataset (large enough that Kaggle doesn't have enough RAM to avoid having to read images from disk!).","metadata":{"_uuid":"1232df41-22cd-43a2-ac61-a096eb6c0feb","_cell_guid":"aa2459ee-278a-4712-8db7-d61d794254ee","trusted":true}},{"cell_type":"markdown","source":"## Pre-Processing\n### Dataframe Creation\nThe first step I took was to drop some columns that I didn't intent to use. `site_id` and `machine_id` seem to have little useful relation to presence of cancer, while `biopsy` and `BIRADS` kind of give the answer away. `density` and `difficult_negative_case` had many missing values, so I decided to drop them for convenience.\n\nIn general, handling of missing values could be worked on. For now, I am just dropping rows with missing ages, as that is a feature that will be used. The remaining features seem to give important information about the image, so I encoded them as binary numbers to be able to pass in to a neural network. Note that CC (craniocaudal) and MLO (mediolateral oblique) are by far the most common imaging views, and so the few unusual views were dropped as well. \n\nFinally, image data needs to be somehow placed into the dataframe for later retrieval. The biggest challenge here is the size of the images; the raw DICOM images are very high resolution and thus reading and resizing is expensive. My original plan was to store resized images in a dataframe column for quick and easy retrieval. This dataframe with images (stored as a .feather file) can be recreated with the following code, which both processes non-numerical columns and processes images into flattened arrays (required for .feather format). Load the rsna-mammography-256x256.feather file for the final result. However, this uses a humongous amount of RAM (>30 GB).\n\n[theoviel's dataset](https://www.kaggle.com/code/theoviel/dicom-resized-png-jpg/) takes a much smarter approach in just resizing and saving each image to disk. This means that images can be read one-by-one and do not have to all be loaded into RAM like with a singular table. I used the above dataset to access the resized images on disk. This probably slows down training due to the overhead in retrieving each image, but I did not have the compute resources for my original approach.","metadata":{"_uuid":"debed861-34d6-43be-bd6c-243713c45aa1","_cell_guid":"4a0b0591-02a2-4f40-a91c-352c069d22dc","trusted":true}},{"cell_type":"code","source":"# df_train = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/train.csv')\n# img_dir = '/kaggle/input/rsna-breast-cancer-detection/train_images/'\n\n# # drop unused features\n# drop = [\"site_id\", \"biopsy\", \"BIRADS\", \"density\", \"machine_id\", \"difficult_negative_case\"]\n# df_train = df_train.drop(columns=drop)\n# # drop rows with missing ages\n# df_train = df_train.dropna(subset=[\"age\"])\n# # drop uncommon views\n# df_train = df_train[(df_train.view == \"MLO\") | (df_train.view == \"CC\")]\n\n# # convert laterality and view to binary variables\n# # L=0, R=1\n# # MLO=0, CC=1\n# df_train.loc[df_train[\"laterality\"]==\"L\", \"laterality\"] = 0\n# df_train.loc[df_train[\"laterality\"]==\"R\", \"laterality\"] = 1\n# df_train.loc[df_train[\"view\"]==\"MLO\", \"view\"] = 0\n# df_train.loc[df_train[\"view\"]==\"CC\", \"view\"] = 1\n\n# # add image paths to df\n# df_train[\"path\"] = img_dir + df_train['patient_id'].astype(str) + \"/\" + df_train['image_id'].astype(str) + \".dcm\"\n# df_train.head()","metadata":{"_uuid":"a78a6f6f-8ce3-4be3-90d3-bc67afec402c","_cell_guid":"8dd1b060-d950-48ff-ae5b-3476334b5271","_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-09-30T14:57:59.715061Z","iopub.execute_input":"2023-09-30T14:57:59.715381Z","iopub.status.idle":"2023-09-30T14:57:59.720227Z","shell.execute_reply.started":"2023-09-30T14:57:59.71535Z","shell.execute_reply":"2023-09-30T14:57:59.718944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def readDICOMRow(row, size=256):\n#     dicom = pydicom.dcmread(row)\n#     img = dicom.pixel_array\n#     img = (img - img.min()) / (img.max() - img.min())\n\n#     if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n#         img = 1 - img\n#     img = cv2.resize(img, (size, size))\n#     return img.reshape(-1)\n\n# start = 40960+4096+4096+4096\n# end = start+4096\n# fname = f\"df_train_{start}-{end-1}.feather\"\n\n# df_train = df_train.iloc[start:end]\n# df_train = df_train.reset_index()\n# df_train[\"img\"] = df_train[\"path\"].swifter.apply(readDICOMRow)\n# df_train.head()","metadata":{"_uuid":"bfc121cc-ebb5-4b8e-a87a-c96c9689bbb1","_cell_guid":"f421a8b1-4da7-4033-9978-38db81d63953","_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-09-30T14:57:59.721754Z","iopub.execute_input":"2023-09-30T14:57:59.722417Z","iopub.status.idle":"2023-09-30T14:57:59.732441Z","shell.execute_reply.started":"2023-09-30T14:57:59.722387Z","shell.execute_reply":"2023-09-30T14:57:59.731643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# df_train.laterality = df_train.laterality.astype(\"int64\")\n# df_train.view = df_train.view.astype(\"int64\")\n# df_train.path = df_train.path.astype(\"string\")\n# df_train.to_feather(\"/kaggle/working/\" + fname)\n# from IPython.display import FileLink\n# %cd /kaggle/working\n# FileLink(fname)","metadata":{"_uuid":"b652c9ed-54b3-44ac-a7c7-ebd0a89ef65c","_cell_guid":"51509b62-d1b2-421c-b9ba-f69e97146a42","_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-09-30T14:57:59.735023Z","iopub.execute_input":"2023-09-30T14:57:59.735322Z","iopub.status.idle":"2023-09-30T14:57:59.746665Z","shell.execute_reply.started":"2023-09-30T14:57:59.735294Z","shell.execute_reply":"2023-09-30T14:57:59.74575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# df_train = pd.read_feather(\"/kaggle/input/rsna-mammography-256x256/rsna_mammography_256x256.feather\")\n# df_train.head()","metadata":{"_uuid":"29f48d68-1298-4137-8f1c-5b889560188f","_cell_guid":"10590d49-53e0-4d86-945c-c13f1a1210d5","_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-09-30T14:57:59.747966Z","iopub.execute_input":"2023-09-30T14:57:59.748665Z","iopub.status.idle":"2023-09-30T14:57:59.757606Z","shell.execute_reply.started":"2023-09-30T14:57:59.748637Z","shell.execute_reply":"2023-09-30T14:57:59.756649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataframe Creation v2: Using Downloaded Images:\nImages filenames are formatted as \"{patient\\_id}\\_{image\\_id}.png\", so the path can be generated by just concatenating. Processing of the remaining columns is the same.","metadata":{"_uuid":"50aa6d46-f623-42b2-ad82-251d4b56ab20","_cell_guid":"a9b82da9-aabb-46fc-9811-8bb901588849","trusted":true}},{"cell_type":"code","source":"df_train = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/train.csv')\nimg_dir = '/kaggle/input/rsna-breast-cancer-512-pngs/'\n\n# drop unused features\ndrop = [\"site_id\", \"biopsy\", \"BIRADS\", \"density\", \"machine_id\", \"difficult_negative_case\"]\ndf_train = df_train.drop(columns=drop)\n# drop rows with missing ages\ndf_train = df_train.dropna(subset=[\"age\"])\n# drop uncommon views\ndf_train = df_train[(df_train.view == \"MLO\") | (df_train.view == \"CC\")]\n\n# convert laterality and view to binary variables\n# L=0, R=1\n# MLO=0, CC=1\ndf_train.loc[df_train[\"laterality\"]==\"L\", \"laterality\"] = 0\ndf_train.loc[df_train[\"laterality\"]==\"R\", \"laterality\"] = 1\ndf_train.loc[df_train[\"view\"]==\"MLO\", \"view\"] = 0\ndf_train.loc[df_train[\"view\"]==\"CC\", \"view\"] = 1\n\n# add image paths to df\ndf_train[\"path\"] = img_dir + df_train['patient_id'].astype(str) + \"_\" + df_train['image_id'].astype(str) + \".png\"\ndf_train.head()","metadata":{"_uuid":"fbdbf122-25c7-45bb-ae5c-bebd3f7df9c2","_cell_guid":"cba3351b-e874-47f4-bf59-01598d50f97a","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-09-30T14:57:59.758924Z","iopub.execute_input":"2023-09-30T14:57:59.759677Z","iopub.status.idle":"2023-09-30T14:58:00.020867Z","shell.execute_reply.started":"2023-09-30T14:57:59.75964Z","shell.execute_reply":"2023-09-30T14:58:00.019952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### EDA\nSome basic plotting practice with Seaborn. The main idea is to visualize the proportions of cancerous vs non-cancerous, as the distribution of ages and cancer throughout. The obvious takeaway is the imbalanced nature of the dataset, which I later address with a weighted loss. Plots also reveal the somewhat higher age of cancer patients.","metadata":{"_uuid":"462191f8-f06b-4dd9-8cf5-5083b7b22dbe","_cell_guid":"eb0948dd-3c8f-4c1a-94e3-bfef3c1b8de5","trusted":true}},{"cell_type":"code","source":"h = 2.5\na = 3\n\nsns.catplot(data=df_train, y=\"cancer\", hue=\"invasive\", kind=\"count\", height=h, aspect=a)    \nsns.displot(data=df_train, x=\"age\", hue=df_train[[\"cancer\",\"invasive\"]].apply(tuple, axis=1), binwidth=1, stat=\"density\", common_norm=False, multiple=\"stack\", height=h, aspect=a)\nsns.catplot(data=df_train, y=\"cancer\", x=\"age\", kind=\"box\", orient=\"h\", height=h, aspect=a, width=0.5)","metadata":{"_uuid":"1c5427d6-d282-413f-8707-d7c30b017ee7","_cell_guid":"5e5a5958-2878-4e4e-affe-c5192bc633b4","execution":{"iopub.status.busy":"2023-09-30T14:58:00.022214Z","iopub.execute_input":"2023-09-30T14:58:00.023141Z","iopub.status.idle":"2023-09-30T14:58:02.258696Z","shell.execute_reply.started":"2023-09-30T14:58:00.023107Z","shell.execute_reply":"2023-09-30T14:58:02.257363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# standardize age after plotting\ndf_train[\"age\"] = (df_train[\"age\"] - df_train[\"age\"].mean()) / df_train[\"age\"].std()","metadata":{"_uuid":"cdaa3342-4b3f-4745-bb1b-621c5c0d6727","_cell_guid":"63842b3e-88ba-4919-bb93-abe80632612f","execution":{"iopub.status.busy":"2023-09-30T14:58:02.260009Z","iopub.execute_input":"2023-09-30T14:58:02.260803Z","iopub.status.idle":"2023-09-30T14:58:02.268696Z","shell.execute_reply.started":"2023-09-30T14:58:02.260769Z","shell.execute_reply":"2023-09-30T14:58:02.267728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def readDICOM(path, size=256):\n    dicom = pydicom.dcmread(path)\n    img = dicom.pixel_array\n    img = (img - img.min()) / (img.max() - img.min())\n\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        img = 1 - img\n    img = cv2.resize(img, (size, size))\n    return img\n\ndef readArray(arr):\n    img = np.array(arr.reshape(256,256))\n    # add in a channel dimension\n    img = img[:,:,np.newaxis]\n    return img\n\ndef readPath(path):\n    # img = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n    img = io.imread(path)\n    # add in a channel dimension\n    img = img[:,:,np.newaxis]\n    return img\n\ndef plotDICOM(path, size=256):\n    dicom = pydicom.dcmread(path)\n    img = dicom.pixel_array\n    img = (img - img.min()) / (img.max() - img.min())\n\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        img = 1 - img\n        \n    img = cv2.resize(img, (size, size))\n    plt.figure()\n    plt.axis('off')\n    plt.grid(visible=None)\n    plt.imshow(img, cmap=\"twilight\")\n    plt.show()\n\ndef plotImage(img):\n    plt.figure()\n    plt.axis('off')\n    plt.grid(visible=None)\n    plt.imshow(img, cmap=\"twilight\")\n    plt.show()","metadata":{"_uuid":"7638d279-30a3-4a19-b9a1-3d663d6b4c71","_cell_guid":"5cf2d31e-78be-455f-a04e-34f5ce21cb27","execution":{"iopub.status.busy":"2023-09-30T14:58:02.270006Z","iopub.execute_input":"2023-09-30T14:58:02.270357Z","iopub.status.idle":"2023-09-30T14:58:02.280724Z","shell.execute_reply.started":"2023-09-30T14:58:02.270324Z","shell.execute_reply":"2023-09-30T14:58:02.279868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_invasive = df_train[df_train.invasive==1]\n\nplt.figure()\nfig, axs = plt.subplots(4,4, figsize=(10,10))\naxs = axs.flatten()\nfor i, ax in enumerate(axs):\n    path = df_invasive.iloc[i].path\n    ax.axis('off')\n    ax.imshow(readPath(path), cmap=\"twilight\")\n\nplt.figtext(0.5, 0.90, \"Invasive Cancer\", ha=\"center\", fontsize=16)\nfig.subplots_adjust(wspace=0.05, hspace=0.05)\nplt.show()\n\n#--------------------------------------------------------------#\n\ndf_cancer = df_train[(df_train.cancer==1) & (df_train.invasive==0)]\n\nplt.figure()\nfig, axs = plt.subplots(4,4, figsize=(10,10))\naxs = axs.flatten()\nfor i, ax in enumerate(axs):\n    path = df_cancer.iloc[i].path\n    ax.axis('off')\n    ax.imshow(readPath(path), cmap=\"twilight\")\n\nplt.figtext(0.5, 0.90, \"Non-Invasive Cancer\", ha=\"center\", fontsize=16)\nfig.subplots_adjust(wspace=0.05, hspace=0.05)\n\n#--------------------------------------------------------------#\n\ndf_none = df_train[df_train.cancer==0]\n\nplt.figure()\nfig, axs = plt.subplots(4,4, figsize=(10,10))\naxs = axs.flatten()\nfor i, ax in enumerate(axs):\n    path = df_none.iloc[i].path\n    ax.axis('off')\n    ax.imshow(readPath(path), cmap=\"twilight\")\n\nplt.figtext(0.5, 0.90, \"No Cancer\", ha=\"center\", fontsize=16)\nfig.subplots_adjust(wspace=0.05, hspace=0.05)\nplt.show()","metadata":{"_uuid":"cf764859-0023-46d6-89be-ec2a273241bf","_cell_guid":"7407e17f-9679-43e3-9936-79b392d3f7e7","execution":{"iopub.status.busy":"2023-09-30T14:58:02.28417Z","iopub.execute_input":"2023-09-30T14:58:02.284405Z","iopub.status.idle":"2023-09-30T14:58:47.815943Z","shell.execute_reply.started":"2023-09-30T14:58:02.284385Z","shell.execute_reply":"2023-09-30T14:58:47.814865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MammogramDataset(Dataset):\n    def __init__(self, df, transform=None, test=False):\n        self.df = df\n        self.transform = transform\n        self.test = test\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        if self.test:\n            label = self.df[\"prediction_id\"].iloc[idx]\n        else:\n            label = self.df.cancer.iloc[idx]\n        \n        img = readPath(self.df.iloc[idx].path)\n        if self.transform:\n            img = self.transform(img)\n        \n        features = [\"laterality\", \"view\", \"age\", \"implant\"]\n        features = self.df[features].iloc[idx].to_numpy().astype(float)\n        return img, features, label","metadata":{"_uuid":"eb1c0e44-6c66-453b-bcec-02b25b214a0c","_cell_guid":"169c4831-ef40-4faa-9905-b765d5e4dfae","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-09-30T14:58:47.817477Z","iopub.execute_input":"2023-09-30T14:58:47.818303Z","iopub.status.idle":"2023-09-30T14:58:47.826789Z","shell.execute_reply.started":"2023-09-30T14:58:47.818271Z","shell.execute_reply":"2023-09-30T14:58:47.825982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def buildDataframe(path, img_dir, test=False, calc=False, tmean=1, tstd=0):\n    df_train = pd.read_csv(path)\n\n    # drop unused features\n    if test:\n        drop = [\"site_id\", \"machine_id\"]\n    else:\n        drop = [\"site_id\", \"biopsy\", \"BIRADS\", \"density\", \"machine_id\", \"difficult_negative_case\"]\n    df_train = df_train.drop(columns=drop)\n    # drop rows with missing ages\n    df_train = df_train.dropna(subset=[\"age\"])\n    # drop uncommon views\n    df_train = df_train[(df_train.view == \"MLO\") | (df_train.view == \"CC\")]\n\n    # convert laterality and view to binary variables\n    # L=0, R=1\n    # MLO=0, CC=1\n    df_train.loc[df_train[\"laterality\"]==\"L\", \"laterality\"] = 0\n    df_train.loc[df_train[\"laterality\"]==\"R\", \"laterality\"] = 1\n    df_train.loc[df_train[\"view\"]==\"MLO\", \"view\"] = 0\n    df_train.loc[df_train[\"view\"]==\"CC\", \"view\"] = 1\n\n    # add image paths to df\n    df_train[\"path\"] = img_dir + df_train[\"patient_id\"].astype(str) + \"_\" + df_train[\"image_id\"].astype(str) + \".png\"\n    \n    if not calc:\n        if test:\n            df_train[\"age\"] = (df_train[\"age\"] - tmean) / tstd\n        else:\n            # standardize age\n            df_train[\"age\"] = (df_train[\"age\"] - df_train[\"age\"].mean()) / df_train[\"age\"].std()\n\n    return df_train\n\ndef buildDatasets(df):\n    labels = df.cancer.to_numpy()\n    train_idx, val_idx = train_test_split(np.arange(len(df)), test_size=0.2, shuffle=True, stratify=labels)\n    transform = transforms.Compose([transforms.ToTensor(),\n#                                     transforms.RandomAffine(degrees=0, translate=(0.01, 0.01), fill=0),\n#                                     transforms.ColorJitter(brightness=0.5),\n                                    transforms.Normalize((0.5), (0.5))])\n    dataset = MammogramDataset(df=df, transform=transform)\n\n    train_dataset = Subset(dataset, train_idx)\n    validation_dataset = Subset(dataset, val_idx)\n\n    return train_idx, train_dataset, validation_dataset\n\ndef buildDataloaders(df, train_idx, train_dataset, validation_dataset, batch_size, num_workers, imbalance_strategy):\n    if (imbalance_strategy == \"oversample\") or (imbalance_strategy == \"oversample + weighted-loss\"):\n        # weight by proportion of cancer to non cancer\n        counts = np.bincount(df.iloc[train_idx].cancer.to_numpy())\n        label_weights = 1. / counts\n        weights = label_weights[df.iloc[train_idx].cancer.to_numpy()]\n        sampler = WeightedRandomSampler(weights, len(weights))\n        train_loader = DataLoader(train_dataset, batch_size=batch_size, num_workers=num_workers, sampler=sampler, pin_memory=True, persistent_workers=True)\n    else:\n        train_loader = DataLoader(train_dataset, batch_size=batch_size, num_workers=num_workers, shuffle=True, pin_memory=True, persistent_workers=True)\n    val_loader = DataLoader(validation_dataset, batch_size=batch_size, num_workers=num_workers, pin_memory=True, persistent_workers=True)\n\n    return train_loader, val_loader\n\ndef calculateMetrics(ep_labels, ep_predictions, running_loss, num_batches, verbose=True):\n    tn,fp,fn,tp = sklearn.metrics.confusion_matrix(ep_labels, ep_predictions, labels=[0, 1]).ravel()\n    sensitivity = tp / (tp+fn)\n    specificity = tn / (fp+tn)\n    precision = tp / (tp+fp)\n    accuracy = (tp+tn) / len(ep_labels)\n    f1 = 2*(precision*sensitivity) / (precision+sensitivity)\n    avg_loss = running_loss/num_batches\n\n    if (verbose):\n        matrix = [[f\"tn: {tn}\", f\"fp: {fp}\"], [f\"fn: {fn}\", f\"tp: {tp}\"]]\n        print(tabulate(matrix, tablefmt=\"grid\"))\n        print(\"# predicted ones: \", sum(ep_predictions))\n        print(\"# labeled ones: \", sum(ep_labels))\n        print(\"avg loss: \", avg_loss)\n        print(\"sensitivity: \", sensitivity)\n        print(\"\")\n\n    return sensitivity, specificity, precision, accuracy, f1, avg_loss\n\ndef init_weights(m):\n    if isinstance(m, nn.Conv2d):\n        nn.init.kaiming_normal_(m.weight, mode=\"fan_out\", nonlinearity=\"relu\")\n    elif isinstance(m, nn.BatchNorm2d):\n        nn.init.constant_(m.weight, 1)\n        nn.init.constant_(m.bias, 0)","metadata":{"_uuid":"423a1b83-9e49-42c2-9748-9aefec9b6925","_cell_guid":"28156414-9d36-47dd-bea9-e06e573f2b43","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-09-30T14:58:47.828362Z","iopub.execute_input":"2023-09-30T14:58:47.828699Z","iopub.status.idle":"2023-09-30T14:58:47.847295Z","shell.execute_reply.started":"2023-09-30T14:58:47.82867Z","shell.execute_reply":"2023-09-30T14:58:47.846371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make(config=None):\n    print(device)\n    print(torch.cuda.get_device_name(0))\n    if config.resolution == 256:\n        img_dir = \"/kaggle/input/rsna-breast-cancer-256-pngs/\"\n    elif config.resolution == 512:\n        img_dir = \"/kaggle/input/rsna-breast-cancer-512-pngs/\"\n    elif config.resolution == 1024:\n        img_dir = \"/kaggle/input/rsna-breast-cancer-1024-pngs/output/\"\n    df = buildDataframe(\"/kaggle/input/rsna-breast-cancer-detection/train.csv\", \n                        img_dir)\n    train_idx, train_dataset, validation_dataset = buildDatasets(df)\n    train_loader, val_loader = buildDataloaders(df, train_idx, train_dataset, validation_dataset, config.batch_size, config.num_workers, config.imbalance_strategy)\n\n    if config.architecture == \"VGGNet\":\n        net = VGGNet().to(device)\n    elif config.architecture == \"Net\":\n        net = Net().to(device)\n    elif config.architecture == \"ResNet\":\n        net = ResNet().to(device)\n    elif config.architecture == \"ResNet34\":\n        net = ResNet34().to(device)\n    elif config.architecture == \"ResNet34_NoTab\":\n        net = ResNet34_NoTab().to(device)\n    elif config.architecture == \"ResNet50\":\n        net = ResNet50().to(device)\n    elif config.architecture == \"MNet_Linear\":\n        net = MNet_Linear(config.resolution).to(device)\n    elif config.architecture == \"MNet_Linear2\":\n        net = MNet_Linear2(config.resolution).to(device)\n    elif config.architecture == \"MNet_Linear2BN\":\n        net = MNet_Linear2BN(config.resolution).to(device)\n    elif config.architecture == \"MNet_Linear2_Pretrained\":\n        net = MNet_Linear2_Pretrained(config.resolution).to(device)\n    elif config.architecture == \"MNet_ViewSplit\":\n        net = MNet_ViewSplit().to(device)\n    net.apply(init_weights)\n\n    if (config.imbalance_strategy == \"weighted-loss\") or (config.imbalance_strategy == \"oversample + weighted-loss\"):\n        # weight by proportion of cancer to non cancer\n        weight = torch.tensor(len(df[df.cancer==0]) / len(df[df.cancer==1])).to(device)\n        loss_fn = nn.BCEWithLogitsLoss(pos_weight=weight)\n    else:\n        loss_fn = nn.BCEWithLogitsLoss()\n    opt = torch.optim.AdamW(net.parameters(), lr=config.learning_rate, weight_decay=config.weight_decay, betas=(config.beta1, config.beta2), fused=True)\n    sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=config.max_learning_rate, \n                                                steps_per_epoch=len(train_loader), epochs=config.epochs)\n    scaler = torch.cuda.amp.GradScaler(enabled=use_amp)\n\n    print(\"MODEL: \", config.architecture)\n    print(\"OPTIMIZER: \", config.optimizer)\n    progress = tqdm(range(config.epochs))\n    wandb.watch(net, loss_fn, log=\"all\", log_freq=8)\n\n    return progress, net, train_loader, val_loader, opt, scaler, loss_fn, sched\n\ndef train(config, progress, net, train_loader, val_loader, opt, scaler, loss_fn, sched):\n    examples = 0\n    for epoch in progress:\n        net.train()\n        running_loss = 0.0\n        ep_predictions = []\n        ep_labels = []\n        for i, (imgs,features,labels) in enumerate(train_loader):\n            with torch.autocast(device_type=\"cuda\", dtype=torch.float16, enabled=use_amp):\n                imgs = imgs.float().to(device, non_blocking=True)\n                features = features.float().to(device, non_blocking=True)\n                labels = labels.float().to(device, non_blocking=True)\n                opt.zero_grad(set_to_none=True)\n                out = net(imgs, features).squeeze()\n\n                loss = loss_fn(out, labels)\n            scaler.scale(loss).backward()\n            scaler.step(opt)\n            scaler.update()\n            sched.step()\n            examples += config.batch_size\n            \n            out = F.sigmoid(out)\n            running_loss += loss.item()\n            ep_predictions.extend(out.detach().cpu().to(torch.float).round().numpy())\n            ep_labels.extend(labels.detach().cpu().numpy())\n            # wandb.log({\"batch loss\": loss.item()})\n            del loss, out\n\n        print(\"\\nEPOCH: \", epoch)\n        print(\"---TRAIN---\")\n\n        sensitivity, specificity, precision, accuracy, f1, avg_loss = calculateMetrics(\n            ep_labels, ep_predictions, running_loss, len(train_loader))\n        wandb.log({\"train\": {\"epoch\": epoch, \"avg_loss\": avg_loss, \n                                \"sensitivity\": sensitivity, \"specificity\": specificity,\n                                \"precision\": precision,\"accuracy\": accuracy,\"f1\": f1}}, step=examples)\n        del sensitivity, specificity, precision, accuracy, f1, avg_loss, ep_predictions, ep_labels, running_loss\n        # validate\n        if (epoch % 2 == 0) or (epoch == config.epochs - 1):\n            net.eval()\n            torch.save(net.state_dict(), os.path.join(wandb.run.dir, f\"net.pt\"))\n            wandb.save(f\"net.pt\")\n            ep_predictions = []\n            ep_probs = []\n            ep_labels = []\n            running_loss = 0.0\n            with torch.no_grad():\n                for i, (imgs,features,labels) in enumerate(val_loader):\n                    with torch.autocast(device_type=\"cuda\", dtype=torch.float16, enabled=use_amp):\n                        imgs = imgs.float().to(device, non_blocking=True)\n                        features = features.float().to(device, non_blocking=True)\n                        labels = labels.float().to(device, non_blocking=True)\n\n                        out = net(imgs, features).squeeze()\n                        loss = loss_fn(out, labels)\n\n                    out = F.sigmoid(out)\n                    running_loss += loss.item()\n                    ep_probs.extend(out.detach().cpu().to(torch.float).numpy())\n                    ep_labels.extend(labels.detach().cpu().numpy())\n                    del loss, out\n                \n                print(\"---VALIDATION---\")\n\n                ep_predictions = np.round(ep_probs)\n                sensitivity, specificity, precision, accuracy, f1, avg_loss = calculateMetrics(\n                ep_labels, ep_predictions, running_loss, len(train_loader))\n\n                if np.isnan(avg_loss):\n                    raise Exception(f'val avg_loss is NaN')\n\n                wandb.log({\"val\": {\"epoch\": epoch, \"avg_loss\": avg_loss, \n                                    \"sensitivity\": sensitivity, \"specificity\": specificity,\n                                    \"precision\": precision,\"accuracy\": accuracy,\"f1\": f1}}, step=examples)\n\n                if (epoch == config.epochs - 1):\n                    cm = wandb.plot.confusion_matrix(None, ep_labels, ep_predictions, class_names=[\"no cancer\", \"cancer\"])\n                    # transform to 2 class format\n                    ep_probs = [[1-prob, prob] for prob in ep_probs]\n                    pr = wandb.plot.pr_curve(ep_labels, ep_probs, labels=[\"no cancer\", \"cancer\"])\n\n                    wandb.log({\"cm\": cm, \"pr\": pr})\n                    del cm, pr\n                del sensitivity, specificity, precision, accuracy, f1, avg_loss, ep_labels, running_loss, ep_probs, ep_predictions\n    return net\n\ndef run(config=None):\n    with wandb.init(project=\"rsna-mammogram\", config=config, mode=\"online\"):\n        config = wandb.config\n        params = make(config)\n        gc.collect()\n        net = train(config, *params)\n        wandb.finish()\n        return net","metadata":{"_uuid":"a0e4a2f0-7699-4f43-a205-6a2a279bee30","_cell_guid":"264a25d6-c6b0-4294-b17b-cabc48d8d9ff","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-09-30T14:58:47.848434Z","iopub.execute_input":"2023-09-30T14:58:47.848763Z","iopub.status.idle":"2023-09-30T14:58:47.872268Z","shell.execute_reply.started":"2023-09-30T14:58:47.848734Z","shell.execute_reply":"2023-09-30T14:58:47.871169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Network Architecture\nI begin with a simple deep CNN. Later on, the tabular features as well as more advanced features such as attention and skip connections will be implemented.","metadata":{"_uuid":"29f1c23e-5e59-41bd-8a6c-3c48b40919c9","_cell_guid":"0de2005a-9e57-4215-9997-2650d3302a48","trusted":true}},{"cell_type":"code","source":"class VGGNet(nn.Module):\n    def __init__(self, inc=1):\n        super().__init__()\n        self.block1 = nn.Sequential(\n            nn.Conv2d(inc, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Conv2d(64, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.MaxPool2d(2, 2)\n        )\n        self.block2 = nn.Sequential(\n            nn.Conv2d(64, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.MaxPool2d(2, 2)\n        )\n        self.block3 = nn.Sequential(\n            nn.Conv2d(128, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.MaxPool2d(2, 2)\n        )\n        self.block4 = nn.Sequential(\n            nn.Conv2d(256, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.MaxPool2d(2, 2)\n        )\n        self.block5 = nn.Sequential(\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.MaxPool2d(2, 2)\n        )\n        self.fc1 = nn.Linear(512*8*8 + 4, 512)\n        self.fc2 = nn.Linear(512, 256)\n        self.fc3 = nn.Linear(256, 256)\n        self.fc4 = nn.Linear(256, 1)\n    \n    def forward(self, x, features):\n        x = self.block1(x)\n        x = self.block2(x)\n        x = self.block3(x)\n        x = self.block4(x)\n        x = self.block5(x)\n\n        x  = torch.flatten(x, 1)\n        x = torch.cat((x, features.squeeze(1)), dim=1)\n\n        x = F.dropout(x, p=0.5, training=self.training)\n        x = F.relu(self.fc1(x))\n        x = F.dropout(x, p=0.5, training=self.training)\n        x = F.relu(self.fc2(x))\n        x = F.dropout(x, p=0.5, training=self.training)\n        x = F.relu(self.fc3(x))\n        x = self.fc4(x)\n\n        return x","metadata":{"_uuid":"41bbc214-f12d-4f05-b7d9-5f64040af7b4","_cell_guid":"cbdad002-50b1-4d0c-88db-f8c08cc20e3f","execution":{"iopub.status.busy":"2023-09-30T14:58:47.87368Z","iopub.execute_input":"2023-09-30T14:58:47.87398Z","iopub.status.idle":"2023-09-30T14:58:47.888745Z","shell.execute_reply.started":"2023-09-30T14:58:47.873949Z","shell.execute_reply":"2023-09-30T14:58:47.887811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ResNet(nn.Module):\n    def __init__(self, inc=1):\n        super().__init__()\n        self.down1 = nn.Sequential(\n            nn.Conv2d(1, 64, 1, stride=1, bias=False),\n            nn.BatchNorm2d(64)\n        )\n        self.down2 = nn.Sequential(\n            nn.Conv2d(64, 128, 1, stride=1, bias=False),\n            nn.BatchNorm2d(128)\n        )\n        self.down3 = nn.Sequential(\n            nn.Conv2d(128, 256, 1, stride=1, bias=False),\n            nn.BatchNorm2d(256)\n        )\n        self.down4 = nn.Sequential(\n            nn.Conv2d(256, 512, 1, stride=1, bias=False),\n            nn.BatchNorm2d(512)\n        )\n        self.down5 = nn.Sequential(\n            nn.Conv2d(512, 512, 1, stride=1, bias=False),\n            nn.BatchNorm2d(512)\n        )\n        self.block1 = nn.Sequential(\n            nn.Conv2d(inc, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Conv2d(64, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n        )\n        self.block2 = nn.Sequential(\n            nn.Conv2d(64, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n        )\n        self.block3 = nn.Sequential(\n            nn.Conv2d(128, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.block4 = nn.Sequential(\n            nn.Conv2d(256, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n        )\n        self.block5 = nn.Sequential(\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n        )\n\n        self.avgpool = nn.AdaptiveAvgPool2d((1,1))\n        self.fc1 = nn.Linear(512 + 4, 64)\n        self.fc2 = nn.Linear(64, 1)\n    \n    def forward(self, x, features):\n        res = self.down1(x)\n        x = self.block1(x)\n        x = x + res\n        x = F.max_pool2d(F.relu(x), 2, 2)\n        \n        res = self.down2(x)\n        x = self.block2(x)\n        x = x + res\n        x = F.max_pool2d(F.relu(x), 2, 2)\n\n        res = self.down3(x)\n        x = self.block3(x)\n        x = x + res\n        x = F.max_pool2d(F.relu(x), 2, 2)\n\n        res = self.down4(x)\n        x = self.block4(x)\n        x = x + res\n        x = F.max_pool2d(F.relu(x), 2, 2)\n\n        res = x\n        x = self.block5(x)\n        x = x + res\n        x = F.max_pool2d(F.relu(x), 2, 2)\n\n        x = self.avgpool(x)\n        x  = torch.flatten(x, 1)\n        x = torch.cat((x, features.squeeze(1)), dim=1)\n        \n        x = F.relu(self.fc1(x))\n        x = self.fc2(x)\n        return x","metadata":{"_uuid":"e147cb2c-f836-4849-ad98-f154b6c6d5c6","_cell_guid":"f237ded5-345e-4372-abb5-0ebc839fbb68","execution":{"iopub.status.busy":"2023-09-30T14:58:47.890184Z","iopub.execute_input":"2023-09-30T14:58:47.890784Z","iopub.status.idle":"2023-09-30T14:58:47.907001Z","shell.execute_reply.started":"2023-09-30T14:58:47.890755Z","shell.execute_reply":"2023-09-30T14:58:47.905988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ResNet34(nn.Module):\n    def __init__(self, inc=1):\n        super().__init__()\n        self.down2_3 = nn.Sequential(\n            nn.Conv2d(64, 128, 1, stride=2, bias=False),\n            nn.BatchNorm2d(128)\n        )\n        self.down3_4 = nn.Sequential(\n            nn.Conv2d(128, 256, 1, stride=2, bias=False),\n            nn.BatchNorm2d(256)\n        )\n        self.down4_5 = nn.Sequential(\n            nn.Conv2d(256, 512, 1, stride=2, bias=False),\n            nn.BatchNorm2d(512)\n        )\n   \n        self.conv1 = nn.Sequential(\n#             nn.Conv2d(inc, 64, 7, stride=2, bias=False),\n            nn.Conv2d(1, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False),\n            nn.MaxPool2d(3, 2)\n        )\n        self.conv2_1 = nn.Sequential(\n            nn.Conv2d(64, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Conv2d(64, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n        )\n        self.conv2_2 = nn.Sequential(\n            nn.Conv2d(64, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Conv2d(64, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n        )\n        self.conv2_3 = nn.Sequential(\n            nn.Conv2d(64, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Conv2d(64, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n        )\n        self.conv3_1 = nn.Sequential(\n            nn.Conv2d(64, 128, 3, padding=1, stride=2, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n        )\n        self.conv3_2 = nn.Sequential(\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n        )\n        self.conv3_3 = nn.Sequential(\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n        )\n        self.conv3_4 = nn.Sequential(\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n        )\n        self.conv4_1 = nn.Sequential(\n            nn.Conv2d(128, 256, 3, padding=1, stride=2, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_2 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_3 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_4 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_5 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_6 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv5_1 = nn.Sequential(\n            nn.Conv2d(256, 512, 3, padding=1, stride=2, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n        )\n        self.conv5_2 = nn.Sequential(\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n        )\n        self.conv5_3 = nn.Sequential(\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n        )\n\n        self.avgpool = nn.AdaptiveAvgPool2d((1,1))\n        self.fc1 = nn.Linear(512 + 4, 64)\n        self.fc2 = nn.Linear(64, 1)\n    \n    def forward(self, x, features):\n        x = self.conv1(x)\n\n        res = x\n        x = self.conv2_1(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv2_2(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv2_3(x)\n        x += res\n        x = F.relu(x)\n\n        res = self.down2_3(x)\n        x = self.conv3_1(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv3_2(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv3_3(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv3_4(x)\n        x += res\n        x = F.relu(x)\n\n        res = self.down3_4(x)\n        x = self.conv4_1(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_2(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_3(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_4(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_5(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_6(x)\n        x += res\n        x = F.relu(x)\n\n        res = self.down4_5(x)\n        x = self.conv5_1(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv5_2(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv5_3(x)\n        x += res\n        x = F.relu(x)\n\n        x = self.avgpool(x)\n        x  = torch.flatten(x, 1)\n        x = torch.cat((x, features.squeeze(1)), dim=1)\n\n        x = F.relu(self.fc1(x))\n        x = self.fc2(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-09-30T14:58:47.908331Z","iopub.execute_input":"2023-09-30T14:58:47.908903Z","iopub.status.idle":"2023-09-30T14:58:47.933852Z","shell.execute_reply.started":"2023-09-30T14:58:47.908873Z","shell.execute_reply":"2023-09-30T14:58:47.93295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ResNet50(nn.Module):\n    def __init__(self, inc=1):\n        super().__init__()\n        self.down2_3 = nn.Sequential(\n            nn.Conv2d(64, 128, 1, stride=2, bias=False),\n            nn.BatchNorm2d(128)\n        )\n        self.down3_4 = nn.Sequential(\n            nn.Conv2d(128, 256, 1, stride=2, bias=False),\n            nn.BatchNorm2d(256)\n        )\n        self.down4_5 = nn.Sequential(\n            nn.Conv2d(256, 512, 1, stride=2, bias=False),\n            nn.BatchNorm2d(512)\n        )\n   \n        self.conv1 = nn.Sequential(\n            nn.Conv2d(inc, 64, 7, stride=2, bias=False),\n            nn.BatchNorm2d(64),\n            nn.MaxPool2d(2, 2)\n        )\n        self.conv2_1 = nn.Sequential(\n            nn.Conv2d(64, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Conv2d(64, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n        )\n        self.conv2_2 = nn.Sequential(\n            nn.Conv2d(64, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Conv2d(64, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n        )\n        self.conv2_3 = nn.Sequential(\n            nn.Conv2d(64, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Conv2d(64, 64, 3, padding=1, bias=False),\n            nn.BatchNorm2d(64),\n        )\n        self.conv3_1 = nn.Sequential(\n            nn.Conv2d(64, 128, 3, padding=1, stride=2, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n        )\n        self.conv3_2 = nn.Sequential(\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n        )\n        self.conv3_3 = nn.Sequential(\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n        )\n        self.conv3_4 = nn.Sequential(\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Conv2d(128, 128, 3, padding=1, bias=False),\n            nn.BatchNorm2d(128),\n        )\n        self.conv4_1 = nn.Sequential(\n            nn.Conv2d(128, 256, 3, padding=1, stride=2, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_2 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_3 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_4 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_5 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_6 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_7 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_8 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_9 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_10 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_11 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_12 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_13 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_14 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_15 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_16 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_17 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_18 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_19 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_20 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_21 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv4_22 = nn.Sequential(\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, 3, padding=1, bias=False),\n            nn.BatchNorm2d(256),\n        )\n        self.conv5_1 = nn.Sequential(\n            nn.Conv2d(256, 512, 3, padding=1, stride=2, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n        )\n        self.conv5_2 = nn.Sequential(\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n        )\n        self.conv5_3 = nn.Sequential(\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, 3, padding=1, bias=False),\n            nn.BatchNorm2d(512),\n        )\n\n        self.avgpool = nn.AdaptiveAvgPool2d((1,1))\n        self.fc1 = nn.Linear(512 + 4, 64)\n        self.fc2 = nn.Linear(64, 1)\n    \n    def forward(self, x, features):\n        x = self.conv1(x)\n\n        res = x\n        x = self.conv2_1(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv2_2(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv2_3(x)\n        x += res\n        x = F.relu(x)\n\n        res = self.down2_3(x)\n        x = self.conv3_1(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv3_2(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv3_3(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv3_4(x)\n        x += res\n        x = F.relu(x)\n\n        res = self.down3_4(x)\n        x = self.conv4_1(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_2(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_3(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_4(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_5(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_6(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_7(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_8(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_9(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_10(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_11(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_12(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_13(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_14(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_15(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_16(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_17(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_18(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_19(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_20(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_21(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv4_22(x)\n        x += res\n        x = F.relu(x)\n\n        res = self.down4_5(x)\n        x = self.conv5_1(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv5_2(x)\n        x += res\n        x = F.relu(x)\n        res = x\n        x = self.conv5_3(x)\n        x += res\n        x = F.relu(x)\n\n        x = self.avgpool(x)\n        x  = torch.flatten(x, 1)\n        x = torch.cat((x, features.squeeze(1)), dim=1)\n\n        x = F.relu(self.fc1(x))\n        x = self.fc2(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-09-30T14:58:47.935242Z","iopub.execute_input":"2023-09-30T14:58:47.935774Z","iopub.status.idle":"2023-09-30T14:58:47.974286Z","shell.execute_reply.started":"2023-09-30T14:58:47.935739Z","shell.execute_reply":"2023-09-30T14:58:47.973516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MNet_Linear(nn.Module):\n    def __init__(self, res, inc=1):\n        super().__init__()\n        self.res = res\n        self.inc = inc\n        self.linear = nn.Sequential(\n            nn.Linear(4, res*inc),\n            nn.ReLU(),\n            nn.Linear(res*inc, res*inc),\n            nn.ReLU(),\n            nn.Linear(res*inc, res*res*inc),\n            nn.Sigmoid()\n        )\n        self.rnet = ResNet34(inc=inc)\n    \n    def forward(self, x, features):\n        f = self.linear(features)\n        f = torch.reshape(f, (-1, self.inc, self.res, self.res))\n        x = f + x\n\n        x = self.rnet(x, features)\n\n        return x\n\nclass MNet_Linear2(nn.Module):\n    def __init__(self, res, inc=1):\n        super().__init__()\n        self.res = res\n        self.inc = inc\n        self.linear = nn.Sequential(\n            nn.Linear(2, res*inc),\n            nn.ReLU(),\n            nn.Linear(res*inc, res*inc),\n            nn.ReLU(),\n            nn.Linear(res*inc, res*res*inc),\n            nn.Sigmoid()\n        )\n        self.rnet = ResNet34(inc=inc)\n    \n    def forward(self, x, features):\n        f = self.linear(features[:, 0:2])\n        f = torch.reshape(f, (-1, self.inc, self.res, self.res))\n        x = f + x\n\n        x = self.rnet(x, features)\n\n        return x\n\nclass MNet_Linear2BN(nn.Module):\n    def __init__(self, res, inc=1):\n        super().__init__()\n        self.res = res\n        self.inc = inc\n        self.linear = nn.Sequential(\n            nn.Linear(2, res*inc, bias=False),\n            nn.BatchNorm1d(res*inc),\n            nn.ReLU(),\n            nn.Linear(res*inc, res*inc, bias=False),\n            nn.BatchNorm1d(res*inc),\n            nn.ReLU(),\n            nn.Linear(res*inc, res*res*inc, bias=False),\n            nn.BatchNorm1d(res*res*inc),\n            nn.Sigmoid()\n        )\n        self.rnet = ResNet34(inc=inc)\n    \n    def forward(self, x, features):\n        f = self.linear(features[:, 0:2])\n        f = torch.reshape(f, (-1, self.inc, self.res, self.res))\n        x = f + x\n\n        x = self.rnet(x, features)\n\n        return x    \n\nclass MNet_Linear2_Pretrained(nn.Module):\n    def __init__(self, res, inc=1):\n        super().__init__()\n        self.res = res\n        self.inc = inc\n        self.linear = nn.Sequential(\n            nn.Linear(2, res*inc),\n            nn.ReLU(),\n            nn.Linear(res*inc, res*inc),\n            nn.ReLU(),\n            nn.Linear(res*inc, res*res*inc),\n            nn.Sigmoid()\n        )\n        self.rnet = torchvision.models.resnet50(weights=ResNet50_Weights.DEFAULT)\n        self.rnet.conv1 = nn.Conv2d(1, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n        self.rnet.fc = nn.Linear(in_features=2048, out_features=508)\n        self.fc1 = nn.Linear(512, 1)\n    \n    def forward(self, x, features):\n        f = self.linear(features[:, 0:2])\n        f = torch.reshape(f, (-1, self.inc, self.res, self.res))\n        x = f + x\n\n        x = F.relu(self.rnet(x))\n        \n        x = torch.cat((x, features.squeeze(1)), dim=1)\n        x = self.fc1(x)\n\n        return x\n\nclass MNet_ViewSplit(nn.Module):\n    def __init__(self, inc=1):\n        super().__init__()\n        self.inc = inc\n        self.MLOrnet = ResNet34(inc=inc)\n        self.CCrnet = ResNet34(inc=inc)\n\n    def forward(self, x, features):\n        bsize = features.shape[0]\n        mlo_idxs = (features[1] == 0).squeeze().nonzero().squeeze()\n        cc_idxs = (features[1] == 1).squeeze().nonzero().squeeze()\n\n        mlo_x = x.index_select(0, mlo_idxs)\n        mlo_features = features.index_select(0, mlo_idxs)\n\n        cc_x = x.index_select(0, cc_idxs)\n        cc_features = features.index_select(0, cc_idxs)\n\n        mlo_x = self.MLOrnet(mlo_x, mlo_features)\n        cc_x = self.CCrnet(cc_x, cc_features)\n\n        out = torch.zeros(bsize, 1, dtype=torch.float16, device=device)\n        out.index_add_(0, mlo_idxs, mlo_x)\n        out.index_add_(0, cc_idxs, cc_x)\n        return out","metadata":{"execution":{"iopub.status.busy":"2023-09-30T14:58:47.975531Z","iopub.execute_input":"2023-09-30T14:58:47.976048Z","iopub.status.idle":"2023-09-30T14:58:47.992627Z","shell.execute_reply.started":"2023-09-30T14:58:47.976019Z","shell.execute_reply":"2023-09-30T14:58:47.991729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process(f, size=512, save_folder=\"\", extension=\"png\"):\n    patient = f.split('/')[-2]\n    image = f.split('/')[-1][:-4]\n\n    dicom = pydicom.dcmread(f)\n    img = dicom.pixel_array\n\n    img = (img - img.min()) / (img.max() - img.min())\n\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        img = 1 - img\n\n    img = cv2.resize(img, (size, size))\n\n    cv2.imwrite(save_folder + f\"{patient}_{image}.{extension}\", (img * 255).astype(np.uint8))","metadata":{"execution":{"iopub.status.busy":"2023-09-30T14:58:47.993876Z","iopub.execute_input":"2023-09-30T14:58:47.994179Z","iopub.status.idle":"2023-09-30T14:58:48.01027Z","shell.execute_reply.started":"2023-09-30T14:58:47.994151Z","shell.execute_reply":"2023-09-30T14:58:48.009317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Ensure deterministic behavior\ntorch.backends.cudnn.deterministic = True\nrandom.seed(hash(\"setting random seeds\") % 2**32 - 1)\nnp.random.seed(hash(\"improves reproducibility\") % 2**32 - 1)\ntorch.manual_seed(hash(\"by removing stochasticity\") % 2**32 - 1)\ntorch.cuda.manual_seed_all(hash(\"so runs are repeatable\") % 2**32 - 1)\n\n# optimize\ntorch.autograd.set_detect_anomaly(False, check_nan=False)\ntorch.set_float32_matmul_precision(\"medium\")\ntorch.backends.cuda.matmul.allow_tf32 = True\ntorch.backends.cudnn.allow_tf32 = True\nuse_amp = True\ntorch.backends.cudnn.benchmark = True\n\nconfig = dict(\n    epochs = 32,\n    batch_size = 16,\n    architecture = \"MNet_ViewSplit\",\n    optimizer = \"AdamW\",\n    learning_rate = 1e-4,\n    max_learning_rate = 1e-3,\n    num_workers = 8,\n    imbalance_strategy = \"weighted-loss\",\n    weight_decay = 1e-2,\n    resolution = 256,\n    beta1 = 0.9,\n    beta2 = 0.999\n)\n\nsweep=False\nwandb.login()\nif sweep:\n    sweep_id = wandb.sweep(sweep_config, project=\"rsna-mammogram\")\n    # sweep_id = \"pyroknight/rsna-mammogram/xp0rvcvl\"\n    wandb.agent(sweep_id, run)\nelse:\n    net = run(config)\n#     net.eval()\n    \n#     if config[\"resolution\"] == 256:\n#         img_dir = \"/kaggle/input/rsna-breast-cancer-256-pngs/\"\n#     elif config[\"resolution\"] == 512:\n#         img_dir = \"/kaggle/input/rsna-breast-cancer-512-pngs/\"\n#     elif config[\"resolution\"] == 1024:\n#         img_dir = \"/kaggle/input/rsna-breast-cancer-1024-pngs/\"\n#     df = buildDataframe(\"/kaggle/input/rsna-breast-cancer-detection/train.csv\", \n#                         img_dir, calc=True)\n#     tmean = df[\"age\"].mean()\n#     tstd = df[\"age\"].std()\n    \n#     test_images = glob.glob(\"/kaggle/input/rsna-breast-cancer-detection/test_images/*/*.dcm\")\n#     SAVE_FOLDER = \"test-pngs/\"\n#     SIZE = config[\"resolution\"]\n#     EXTENSION = \"png\"\n#     os.makedirs(SAVE_FOLDER, exist_ok=True)\n#     _ = Parallel(n_jobs=4)(\n#     delayed(process)(uid, size=SIZE, save_folder=SAVE_FOLDER, extension=EXTENSION)\n#     for uid in tqdm(test_images)\n#     )\n    \n#     img_dir = \"/kaggle/working/test-pngs/\"\n#     df_test = buildDataframe(\"/kaggle/input/rsna-breast-cancer-detection/test.csv\", \n#                         img_dir, test=True, tmean=tmean, tstd=tstd)\n#     transform = transforms.Compose([transforms.ToTensor(),\n#                                     transforms.Normalize((0.5), (0.5))])\n#     test_dataset = MammogramDataset(df=df_test, transform=transform, test=True)\n#     test_loader = DataLoader(test_dataset, batch_size=4)\n    \n#     ep_predictions = []\n#     ep_probs = []\n#     ep_labels = []\n#     ep_pids = []\n#     running_loss = 0.0\n#     with torch.no_grad():\n#         for imgs,features,labels in test_loader:\n#             with torch.autocast(device_type=\"cuda\", dtype=torch.float16, enabled=use_amp):\n#                 imgs = imgs.float().to(device, non_blocking=True)\n#                 features = features.float().to(device, non_blocking=True)\n#                 out = net(imgs, features).squeeze()\n\n#             out = F.sigmoid(out)\n#             ep_probs.extend(out.detach().cpu().to(torch.float).numpy())\n#             ep_pids.extend(labels)\n\n#         out = pd.DataFrame(data={\"prediction_id\":ep_pids, \"cancer\":ep_probs})\n#         out = out.groupby(\"prediction_id\").mean()\n#         display(out)\n#         out.to_csv(\"submission.csv\", index=False)","metadata":{"_uuid":"29bc3ac6-7f7d-4e06-b41a-187ff1bff46a","_cell_guid":"9fe08583-2766-4607-918d-1a715c3fd4fc","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-09-30T14:58:48.01168Z","iopub.execute_input":"2023-09-30T14:58:48.011972Z","iopub.status.idle":"2023-09-30T15:02:07.450192Z","shell.execute_reply.started":"2023-09-30T14:58:48.011944Z","shell.execute_reply":"2023-09-30T15:02:07.448799Z"},"trusted":true},"execution_count":null,"outputs":[]}]}