{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.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":39272,"databundleVersionId":4629629,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":4781642,"sourceType":"datasetVersion","datasetId":2767654},{"sourceId":11873665,"sourceType":"datasetVersion","datasetId":7461717}],"dockerImageVersionId":30302,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Downloading Libraries","metadata":{}},{"cell_type":"code","source":"!pip install opencv-python\n!pip install tqdm\n!pip install pandas\n!pip install scikit-image\n!pip install scikit-learn\n!pip install seaborn\n","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! pip install -qU \"python-gdcm\" pydicom pylibjpeg \"opencv-python-headless\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T12:16:04.967342Z","iopub.execute_input":"2025-05-21T12:16:04.96821Z","iopub.status.idle":"2025-05-21T12:16:25.378231Z","shell.execute_reply.started":"2025-05-21T12:16:04.968124Z","shell.execute_reply":"2025-05-21T12:16:25.376866Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ===== Minimal imports for local PNG dataset =====\n\nimport os\nimport re\nimport gc\nimport cv2\nimport random\nimport math\nfrom glob import glob\nfrom tqdm import tqdm\nfrom pprint import pprint\nfrom time import time\nimport datetime as dtime\nfrom datetime import datetime\nimport itertools\nimport warnings\nimport pandas as pd\nimport numpy as np\n\nfrom skimage.transform import resize\nfrom sklearn.preprocessing import LabelEncoder, normalize\n\n\n# Visualization\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom matplotlib.colors import ListedColormap\nplt.rcParams.update({'font.size': 16})\n\n# Suppress warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Custom colors and print helper (optional)\nclass clr:\n    S = '\\033[1m' + '\\033[91m'\n    E = '\\033[0m'\n    \nmy_colors = [\"#517664\", \"#73AA90\", \"#94DDBC\", \"#DAB06C\", \n             \"#DF928E\", \"#C97973\", \"#B25F57\"]\nCMAP1 = ListedColormap(my_colors)\n\nprint(clr.S+\"Notebook Color Schemes:\"+clr.E)\nsns.palplot(sns.color_palette(my_colors))\nplt.show()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T12:16:44.898526Z","iopub.execute_input":"2025-05-21T12:16:44.899498Z","iopub.status.idle":"2025-05-21T12:16:45.81053Z","shell.execute_reply.started":"2025-05-21T12:16:44.899451Z","shell.execute_reply":"2025-05-21T12:16:45.809276Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === General Functions ===\n\ndef set_seed(seed = 1234):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n\ndef show_values_on_bars(axs, h_v=\"v\", space=0.4):\n    '''Plots the value at the end of a seaborn barplot.\n    axs: the ax of the plot\n    h_v: whether the barplot is vertical/horizontal'''\n    \n    def _show_on_single_plot(ax):\n        if h_v == \"v\":\n            for p in ax.patches:\n                _x = p.get_x() + p.get_width() / 2\n                _y = p.get_y() + p.get_height()\n                value = int(p.get_height())\n                ax.text(_x, _y, format(value, ','), ha=\"center\") \n        elif h_v == \"h\":\n            for p in ax.patches:\n                _x = p.get_x() + p.get_width() + float(space)\n                _y = p.get_y() + p.get_height()\n                value = int(p.get_width())\n                ax.text(_x, _y, format(value, ','), ha=\"left\")\n\n    if isinstance(axs, np.ndarray):\n        for idx, ax in np.ndenumerate(axs):\n            _show_on_single_plot(ax)\n    else:\n        _show_on_single_plot(axs)\n        \n\n# === W&B functions removed ===\n\n# Replacing save_dataset_artifact with a simple save message\ndef save_dataset_artifact(run_name, artifact_name, path, data_type=\"dataset\"):\n    print(f\"[INFO] Dataset artifact '{artifact_name}' would be saved from {path} (wandb removed).\")\n    # You can implement local saving/versioning here if needed\n    \n\n# Skipping wandb plots functions since wandb is removed\ndef create_wandb_plot(*args, **kwargs):\n    print(\"[INFO] Wandb plot functions removed. Plotting skipped.\")\n\ndef create_wandb_hist(*args, **kwargs):\n    print(\"[INFO] Wandb histogram function removed. Plotting skipped.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T12:16:48.547948Z","iopub.execute_input":"2025-05-21T12:16:48.548297Z","iopub.status.idle":"2025-05-21T12:16:48.558673Z","shell.execute_reply.started":"2025-05-21T12:16:48.548268Z","shell.execute_reply":"2025-05-21T12:16:48.557861Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport pandas as pd\nfrom tqdm import tqdm\nfrom sklearn.preprocessing import LabelEncoder\n\n# Load the original CSV\ntrain = pd.read_csv(\"/kaggle/input/mammography-breast-cancer-detection/train.csv\")\n\nbase_path = \"/kaggle/input/mammography-breast-cancer-detection/train/\"\n\nall_paths = []\nfor idx, row in tqdm(train.iterrows(), total=len(train)):\n    label_folder = str(row['cancer'])  # 0 or 1\n    img_file = f\"{row['patient_id']}_{row['image_id']}.png\"\n    full_path = os.path.join(base_path, label_folder, img_file)\n    all_paths.append(full_path)\n\ntrain[\"path\"] = all_paths\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T12:16:50.754583Z","iopub.execute_input":"2025-05-21T12:16:50.755444Z","iopub.status.idle":"2025-05-21T12:16:53.457766Z","shell.execute_reply.started":"2025-05-21T12:16:50.755413Z","shell.execute_reply":"2025-05-21T12:16:53.456869Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(clr.S+\"Number of TOTAL images:\"+clr.E,\n      len(glob(\"/kaggle/input/rsna-breast-cancer-detection/train_images/*/*\")))\nprint(clr.S+\"Records gathered in Site 1:\"+clr.E, train[\"site_id\"].value_counts().values[0], \"\\n\"+\n      clr.S+\"Records gathered in Site 2:\"+clr.E, train[\"site_id\"].value_counts().values[1])\nprint(\"-------------------------------------------------\")\nprint(clr.S+\"Total unique patients:\"+clr.E, train[\"patient_id\"].nunique())\nprint(\"-------------------------------------------------\")\nprint(clr.S+\"Total unique images:\"+clr.E, train[\"image_id\"].nunique())\nprint(\"-------------------------------------------------\")\nprint(clr.S+\"Statistics: Images per Patient\"+clr.E)\nprint(train.groupby(\"patient_id\")[\"image_id\"].count().reset_index().describe()[\"image_id\"])\nprint(\"-------------------------------------------------\")\nprint(clr.S+\"Image records count per laterality (R):\"+clr.E, train[\"laterality\"].value_counts().values[0], \"\\n\"+\n      clr.S+\"Image records count per laterality (L):\"+clr.E, train[\"laterality\"].value_counts().values[1])\nprint(\"-------------------------------------------------\")\nprint(clr.S+\"Image records count per View:\"+clr.E)\nprint(train[\"view\"].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-18T15:45:22.897102Z","iopub.execute_input":"2025-05-18T15:45:22.897921Z","iopub.status.idle":"2025-05-18T15:47:25.589747Z","shell.execute_reply.started":"2025-05-18T15:45:22.897883Z","shell.execute_reply":"2025-05-18T15:47:25.588878Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Keep only columns in test + target variable\ntrain = train[[\"patient_id\", \"image_id\", \"laterality\", \"view\", \"age\", \"implant\", \"path\", \"cancer\"]]\n\n# Encode categorical variables\nle_laterality = LabelEncoder()\nle_view = LabelEncoder()\n\ntrain['laterality'] = le_laterality.fit_transform(train['laterality'])\ntrain['view'] = le_view.fit_transform(train['view'])\n\ntrain.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T12:37:01.199464Z","iopub.execute_input":"2025-05-21T12:37:01.199887Z","iopub.status.idle":"2025-05-21T12:37:01.225106Z","shell.execute_reply.started":"2025-05-21T12:37:01.199857Z","shell.execute_reply":"2025-05-21T12:37:01.224245Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(clr.S+\"Number of missing values in Age:\"+clr.E, train[\"age\"].isna().sum())\ntrain['age'] = train['age'].fillna(58)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T12:37:03.606307Z","iopub.execute_input":"2025-05-21T12:37:03.607131Z","iopub.status.idle":"2025-05-21T12:37:03.613866Z","shell.execute_reply.started":"2025-05-21T12:37:03.607099Z","shell.execute_reply":"2025-05-21T12:37:03.612907Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def save_dataset_artifact(run_name, artifact_name, path, data_type):\n    # wandb removed: just print info about saving\n    print(f\"[INFO] Dataset artifact '{artifact_name}' saved locally at: {path} (wandb logging removed).\")\n\n# Save new dataset locally\ntrain.to_csv(\"train_path.csv\", index=False)\n\n# Save \"artifact\" locally (just a print now)\nsave_dataset_artifact(run_name=\"save_train_prep\", \n                      artifact_name=\"train_prep\",\n                      path=\"train_path.csv\",\n                      data_type=\"dataset\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T12:37:08.193125Z","iopub.execute_input":"2025-05-21T12:37:08.193796Z","iopub.status.idle":"2025-05-21T12:37:08.408547Z","shell.execute_reply.started":"2025-05-21T12:37:08.193764Z","shell.execute_reply":"2025-05-21T12:37:08.407592Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q efficientnet_pytorch\n!pip install -q transformers\n!pip install -q albumentations albumentations.pytorch\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T12:37:11.499361Z","iopub.execute_input":"2025-05-21T12:37:11.499736Z","iopub.status.idle":"2025-05-21T12:37:31.717197Z","shell.execute_reply.started":"2025-05-21T12:37:11.499705Z","shell.execute_reply":"2025-05-21T12:37:31.716Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"    # PyTorch\n    import torch\n    import torchvision\n    import torch.nn as nn\n    import torch.nn.functional as F\n    from torch import FloatTensor, LongTensor\n    from torch.utils.data import Dataset, DataLoader, Subset\n    from torch.optim.lr_scheduler import ReduceLROnPlateau\n    from torch.optim import AdamW\n    \n    \n    # Data Augmentation for Image Preprocessing\n\n    from albumentations import (ToFloat, Normalize, VerticalFlip, HorizontalFlip, Compose, Resize,\n                            RandomBrightnessContrast, HueSaturationValue, Blur, GaussNoise,\n                            Rotate, RandomResizedCrop, ShiftScaleRotate, ToGray)\n    from albumentations.pytorch.transforms import ToTensorV2\n\n    \n    \n    from efficientnet_pytorch import EfficientNet\n    from torchvision.models import resnet34, resnet50\n    \n    # SKlearn\n    from sklearn.model_selection import StratifiedKFold, GroupKFold\n    from sklearn.metrics import accuracy_score, roc_auc_score, confusion_matrix","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T12:37:50.287834Z","iopub.execute_input":"2025-05-21T12:37:50.288912Z","iopub.status.idle":"2025-05-21T12:37:53.163701Z","shell.execute_reply.started":"2025-05-21T12:37:50.288869Z","shell.execute_reply":"2025-05-21T12:37:53.162977Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\n\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n\n# Seed\nset_seed()\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint('Device available now:', DEVICE)\n\n# Read in Data\ntrain = pd.read_csv(\"/kaggle/input/helper/train_path.csv\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T12:38:02.814759Z","iopub.execute_input":"2025-05-21T12:38:02.815679Z","iopub.status.idle":"2025-05-21T12:38:03.049709Z","shell.execute_reply.started":"2025-05-21T12:38:02.815639Z","shell.execute_reply":"2025-05-21T12:38:03.048751Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Shuffle full dataset first\n# full_df = train.sample(frac=1, random_state=13).reset_index(drop=True)\n\n# # Split into train (10k) and test (1k)\n# # train = full_df.iloc[0:3000].reset_index(drop=True)\n\n# # Filter positives and negatives\n# positive_df = train[train['cancer'] == 1].sample(n=200, random_state=13)\n# negatives = train[train['cancer'] == 0].iloc[5000:]\n# negative_df = negatives.sample(n=2800, random_state=13)\n\n# # Combine and shuffle\n# train = pd.concat([positive_df, negative_df]).sample(frac=1, random_state=13).reset_index(drop=True)\n\n\n\n\n# test = full_df.iloc[5000:6000].reset_index(drop=True)\n\n# # Check class balance\n# print(\"Train cancer distribution:\\n\", train[\"cancer\"].value_counts())\n# print(\"Test cancer distribution:\\n\", test[\"cancer\"].value_counts())\n\n\n\n\n\n# Assuming `train` is your original full dataset\n\n# Select training set\nnegatives_train = train[train['cancer'] == 0].head(2000)\npositives_train = train[train['cancer'] == 1].head(500)\ntrain_set = pd.concat([negatives_train, positives_train]).reset_index(drop=True)\n\n# Select test set from original dataset (not the reduced train_set)\ntest_set = train.iloc[-1000:].reset_index(drop=True)\n\nprint(\"Train cancer distribution:\\n\", train_set[\"cancer\"].value_counts())\nprint(f\"Train set size: {len(train_set)}\")\n\nprint(\"Test cancer distribution:\\n\", test_set[\"cancer\"].value_counts())\nprint(f\"Test set size: {len(test_set)}\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T12:38:13.471964Z","iopub.execute_input":"2025-05-21T12:38:13.472811Z","iopub.status.idle":"2025-05-21T12:38:13.49144Z","shell.execute_reply.started":"2025-05-21T12:38:13.472774Z","shell.execute_reply":"2025-05-21T12:38:13.490422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ----- GLOBAL PARAMS -----\nvertical_flip = 0.5\nhorizontal_flip = 0.5\n\ncsv_columns = ['laterality', 'view', 'age', 'implant']\nno_columns = len(csv_columns)\noutput_size = 1\n# -------------------------","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T12:38:20.743552Z","iopub.execute_input":"2025-05-21T12:38:20.744529Z","iopub.status.idle":"2025-05-21T12:38:20.749139Z","shell.execute_reply.started":"2025-05-21T12:38:20.744494Z","shell.execute_reply":"2025-05-21T12:38:20.748081Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from PIL import Image\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset\nfrom albumentations import (\n    Compose, RandomResizedCrop, HorizontalFlip, VerticalFlip, ShiftScaleRotate,\n    RandomBrightnessContrast, HueSaturationValue, Blur, GaussNoise, Resize\n)\nfrom albumentations.pytorch import ToTensorV2\n\nclass RSNADataset(Dataset):\n    def __init__(self, dataframe, vertical_flip, horizontal_flip, is_train=True):\n        self.dataframe = dataframe\n        self.is_train = is_train\n        self.vertical_flip = vertical_flip\n        self.horizontal_flip = horizontal_flip\n\n        self.positive_transform = Compose([\n            RandomResizedCrop(height=224, width=224, scale=(0.6, 1.0)),\n            HorizontalFlip(p=self.horizontal_flip),\n            VerticalFlip(p=self.vertical_flip),\n            ShiftScaleRotate(rotate_limit=45, scale_limit=0.3, shift_limit=0.2, p=0.7),\n            RandomBrightnessContrast(p=0.7),\n            HueSaturationValue(p=0.5),\n            Blur(blur_limit=3, p=0.3),\n            GaussNoise(var_limit=(10.0, 50.0), p=0.3),\n            ToTensorV2()\n        ])\n\n        self.general_transform = Compose([\n            RandomResizedCrop(height=224, width=224),\n            HorizontalFlip(p=self.horizontal_flip),\n            VerticalFlip(p=self.vertical_flip),\n            ToTensorV2()\n        ])\n\n        self.test_transform = Compose([\n            Resize(224, 224),\n            ToTensorV2()\n        ])\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, index):\n        row = self.dataframe.iloc[index]\n        # Load PNG image (PIL)\n        image = Image.open(row['path']).convert('RGB')  # ensures 3-channel\n        \n        # Convert to numpy array for albumentations\n        image_np = np.array(image)\n        \n        if self.is_train:\n            transform = self.positive_transform if row['cancer'] == 1 else self.general_transform\n        else:\n            transform = self.test_transform\n\n        transformed = transform(image=image_np)\n        image_tensor = transformed['image'].float()  # ensure float32 tensor\n\n        meta = torch.tensor(row[csv_columns].values.astype(np.float32))\n        target = torch.tensor(row['cancer'], dtype=torch.float32) if 'cancer' in row else None\n\n        if target is not None:\n            return {\"image\": image_tensor, \"meta\": meta, \"target\": target}\n        else:\n            return {\"image\": image_tensor, \"meta\": meta}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T13:02:40.080543Z","iopub.execute_input":"2025-05-21T13:02:40.080942Z","iopub.status.idle":"2025-05-21T13:02:40.092538Z","shell.execute_reply.started":"2025-05-21T13:02:40.080911Z","shell.execute_reply":"2025-05-21T13:02:40.091567Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def data_to_device(data):\n    image = data[\"image\"].to(DEVICE)\n    metadata = data[\"meta\"].to(DEVICE)\n    target = data.get(\"target\", None)\n    \n    if target is not None:\n        target = target.to(DEVICE)\n    \n    return image, metadata, target\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T13:02:42.328399Z","iopub.execute_input":"2025-05-21T13:02:42.328749Z","iopub.status.idle":"2025-05-21T13:02:42.333786Z","shell.execute_reply.started":"2025-05-21T13:02:42.328718Z","shell.execute_reply":"2025-05-21T13:02:42.332837Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Sample data\nsample_df = train.head(6)\n\n# Instantiate Dataset object\ndataset = RSNADataset(sample_df, vertical_flip, horizontal_flip,\n                      is_train=True)\n# The Dataloader\ndataloader = DataLoader(dataset, batch_size=3, shuffle=False)\n\n# Output of the Dataloader\nfor k, data in enumerate(dataloader):\n    image, meta, targets = data_to_device(data)\n    print(clr.S + f\"Batch: {k}\" + clr.E, \"\\n\" +\n          clr.S + \"Image:\" + clr.E, image.shape, \"\\n\" +\n          clr.S + \"Meta:\" + clr.E, meta, \"\\n\" +\n          clr.S + \"Targets:\" + clr.E, targets, \"\\n\" +\n          \"=\"*50)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T13:02:44.261253Z","iopub.execute_input":"2025-05-21T13:02:44.261602Z","iopub.status.idle":"2025-05-21T13:02:44.291122Z","shell.execute_reply.started":"2025-05-21T13:02:44.261573Z","shell.execute_reply":"2025-05-21T13:02:44.289991Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ResNet50Network(nn.Module):\n    def __init__(self, output_size, no_columns):\n        super().__init__()\n        self.no_columns, self.output_size = no_columns, output_size\n        \n        # Define Feature part (IMAGE)\n        self.features = resnet50(pretrained=True) # 1000 neurons out\n        # (metadata)\n        self.csv = nn.Sequential(nn.Linear(self.no_columns, 500),\n                                 nn.LayerNorm(500),\n                                 nn.ReLU(),\n                                 nn.Dropout(p=0.2))\n        \n        # Define Classification part\n        self.classification = nn.Linear(1000 + 500, output_size)\n        \n        \n    def forward(self, image, meta, prints=False):\n        if prints: print('Input Image shape:', image.shape, '\\n'+\n                         'Input metadata shape:', meta.shape)\n        \n        # Image CNN\n        image = self.features(image)\n        if prints: print('Features Image shape:', image.shape)\n        \n        # CSV FNN\n        meta = self.csv(meta)\n        if prints: print('Meta Data:', meta.shape)\n            \n        # Concatenate layers from image with layers from csv_data\n        image_meta_data = torch.cat((image, meta), dim=1)\n        if prints: print('Concatenated Data:', image_meta_data.shape)\n        \n        # CLASSIF\n        out = self.classification(image_meta_data)\n        if prints: print('Out shape:', out.shape)\n        \n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T13:03:13.352238Z","iopub.execute_input":"2025-05-21T13:03:13.352868Z","iopub.status.idle":"2025-05-21T13:03:13.361065Z","shell.execute_reply.started":"2025-05-21T13:03:13.352834Z","shell.execute_reply":"2025-05-21T13:03:13.360099Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load Model\nmodel_example = ResNet50Network(output_size=output_size, no_columns=no_columns).to(DEVICE)\n\n# Outputs\nout = model_example(image, meta, prints=True)\n\n# Criterion example\ncriterion_example = nn.BCEWithLogitsLoss()\n# Unsqueeze(1) from shape=[3] to shape=[3, 1]\nloss = criterion_example(out, targets.unsqueeze(1).float()) \nprint(\"=\"*50)\nprint(clr.S+'Loss:'+clr.E, loss.item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T13:03:16.209273Z","iopub.execute_input":"2025-05-21T13:03:16.20991Z","iopub.status.idle":"2025-05-21T13:03:16.853099Z","shell.execute_reply.started":"2025-05-21T13:03:16.209878Z","shell.execute_reply":"2025-05-21T13:03:16.85205Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EffNetNetwork(nn.Module):\n    def __init__(self, output_size, no_columns):\n        super().__init__()\n        self.no_columns, self.output_size = no_columns, output_size\n        \n        # Define Feature part (IMAGE)\n        self.features = EfficientNet.from_pretrained('efficientnet-b2')\n        \n        # (CSV)\n        self.csv = nn.Sequential(nn.Linear(self.no_columns, 250),\n                                 nn.BatchNorm1d(250),\n                                 nn.ReLU(),\n                                 nn.Dropout(p=0.2),\n                                 \n                                 nn.Linear(250, 250),\n                                 nn.BatchNorm1d(250),\n                                 nn.ReLU(),\n                                 nn.Dropout(p=0.2))\n        \n        # Define Classification part\n        self.classification = nn.Sequential(nn.Linear(1408 + 250, self.output_size))\n        \n        \n    def forward(self, image, meta, prints=False):   \n        \n        if prints: print('Input Image shape:', image.shape, '\\n'+\n                         'Input metadata shape:', meta.shape)\n        \n        # Image CNN\n        image = self.features.extract_features(image)\n        image = F.avg_pool2d(image, image.size()[2:]).reshape(-1, 1408)\n        if prints: print('Features Image shape:', image.shape)\n        \n        # CSV FNN\n        meta = self.csv(meta)\n        if prints: print('Meta Data:', meta.shape)\n            \n        # Concatenate layers from image with layers from csv_data\n        image_meta_data = torch.cat((image, meta), dim=1)\n        if prints: print('Concatenated Data:', image_meta_data.shape)\n        \n        # CLASSIF\n        out = self.classification(image_meta_data)\n        if prints: print('Out shape:', out.shape)\n        \n        return out\n\n\n\n\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T13:04:10.442127Z","iopub.execute_input":"2025-05-21T13:04:10.44249Z","iopub.status.idle":"2025-05-21T13:04:10.452308Z","shell.execute_reply.started":"2025-05-21T13:04:10.442458Z","shell.execute_reply":"2025-05-21T13:04:10.451311Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\n# Load Model\nmodel_example2 = EffNetNetwork(output_size=output_size, no_columns=no_columns).to(DEVICE)\n\n# Outputs\nout = model_example2(image, meta, prints=True)\n\n# Criterion example\ncriterion_example = nn.BCEWithLogitsLoss()\n# Unsqueeze(1) from shape=[3] to shape=[3, 1]\nloss = criterion_example(out, targets.unsqueeze(1).float()) \nprint(\"=\"*50)\nprint(clr.S+'Loss:'+clr.E, loss.item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T13:04:13.12593Z","iopub.execute_input":"2025-05-21T13:04:13.126272Z","iopub.status.idle":"2025-05-21T13:04:13.368533Z","shell.execute_reply.started":"2025-05-21T13:04:13.126243Z","shell.execute_reply":"2025-05-21T13:04:13.367499Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def add_in_file(text, f):\n    \n    with open(f'logs_{VERSION}.txt', 'a+') as f:\n        print(text, file=f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T13:04:24.498285Z","iopub.execute_input":"2025-05-21T13:04:24.498649Z","iopub.status.idle":"2025-05-21T13:04:24.503491Z","shell.execute_reply.started":"2025-05-21T13:04:24.498607Z","shell.execute_reply":"2025-05-21T13:04:24.502438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def mixup_data(x, meta, y, alpha=0.2):\n    '''Returns mixed inputs, pairs of targets, and lambda'''\n    if alpha > 0:\n        lam = np.random.beta(alpha, alpha)\n    else:\n        lam = 1\n\n    batch_size = x.size()[0]\n    index = torch.randperm(batch_size).to(x.device)\n\n    mixed_x = lam * x + (1 - lam) * x[index, :]\n    mixed_meta = lam * meta + (1 - lam) * meta[index, :]\n    y_a, y_b = y, y[index]\n    return mixed_x, mixed_meta, y_a, y_b, lam\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T12:40:45.504127Z","iopub.execute_input":"2025-05-21T12:40:45.504761Z","iopub.status.idle":"2025-05-21T12:40:45.510512Z","shell.execute_reply.started":"2025-05-21T12:40:45.504725Z","shell.execute_reply":"2025-05-21T12:40:45.509534Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, alpha=0.25, gamma=2.0, reduction='mean'):\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n\n    def forward(self, inputs, targets):\n        BCE_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction='none')\n        pt = torch.exp(-BCE_loss)\n        F_loss = self.alpha * (1 - pt) ** self.gamma * BCE_loss\n\n        if self.reduction == 'mean':\n            return F_loss.mean()\n        elif self.reduction == 'sum':\n            return F_loss.sum()\n        else:\n            return F_loss\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T12:40:53.746868Z","iopub.execute_input":"2025-05-21T12:40:53.747547Z","iopub.status.idle":"2025-05-21T12:40:53.754197Z","shell.execute_reply.started":"2025-05-21T12:40:53.747512Z","shell.execute_reply":"2025-05-21T12:40:53.753098Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix, accuracy_score, roc_auc_score, precision_score, recall_score, f1_score, log_loss\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom transformers import get_cosine_schedule_with_warmup\nimport gc\nimport datetime as dtime\nfrom time import time\nimport os\nfrom torch.optim import AdamW\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\nimport numpy as np\nimport torch\n\n\ndef plot_confusion_matrix(y_true, y_pred, epoch, fold):\n    cm = confusion_matrix(y_true, y_pred)\n    plt.figure(figsize=(5, 5))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n                xticklabels=['Negative', 'Positive'],\n                yticklabels=['Negative', 'Positive'])\n    plt.title(f'Fold {fold} - Epoch {epoch}\\nConfusion Matrix')\n    plt.ylabel('Actual')\n    plt.xlabel('Predicted')\n    plt.show()\n\n\ndef reset_weights(m):\n    # Recursively reset weights of a model (common layers)\n    for layer in m.children():\n        if hasattr(layer, 'reset_parameters'):\n            layer.reset_parameters()\n        else:\n            reset_weights(layer)\n\n\ndef train_folds(model, train_original):\n    f = open(f\"logs_{VERSION}.txt\", \"w+\")\n    os.makedirs(\"saved_models\", exist_ok=True)\n\n    group_fold = StratifiedGroupKFold(n_splits=FOLDS)\n    k_folds = group_fold.split(train_original, train_original['cancer'], groups=train_original['patient_id'])\n\n    for i, (train_index, valid_index) in enumerate(k_folds):\n        print(clr.S + f\"---------- Fold: {i+1} ----------\" + clr.E)\n        add_in_file(f\"---------- Fold: {i+1} ----------\", f)\n\n        # Reset weights to fresh state at start of fold\n        reset_weights(model)\n\n        best_roc = None\n        patience_f = PATIENCE\n\n        train_data = train_original.iloc[train_index].reset_index(drop=True)\n        valid_data = train_original.iloc[valid_index].reset_index(drop=True)\n\n        # Debug: print data distribution\n        print(f\"Fold {i+1}: Train size = {len(train_data)}, Positives = {train_data['cancer'].sum()}, \"\n              f\"Negatives = {len(train_data) - train_data['cancer'].sum()}\")\n        print(f\"Fold {i+1}: Valid size = {len(valid_data)}, Positives = {valid_data['cancer'].sum()}, \"\n              f\"Negatives = {len(valid_data) - valid_data['cancer'].sum()}\")\n\n        train_ds = RSNADataset(train_data, vertical_flip, horizontal_flip, is_train=True)\n        valid_ds = RSNADataset(valid_data, vertical_flip, horizontal_flip, is_train=True)\n\n        train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE1, shuffle=True, num_workers=WORKERS)\n        valid_loader = DataLoader(valid_ds, batch_size=BATCH_SIZE2, shuffle=False, num_workers=WORKERS)\n\n        total_steps = len(train_loader) * EPOCHS\n        warmup_steps = int(0.1 * total_steps)  # 10% warm-up\n\n        optimizer = AdamW(model.parameters(), lr=LR, weight_decay=WD)\n        scheduler = get_cosine_schedule_with_warmup(\n            optimizer,\n            num_warmup_steps=warmup_steps,\n            num_training_steps=total_steps\n        )\n        criterion = FocalLoss(alpha=0.25, gamma=2)\n\n        for epoch in range(EPOCHS):\n            start_time = time()\n            correct = 0\n            train_losses = 0\n\n            model.train()\n            for k, data in tqdm(enumerate(train_loader), total=len(train_loader)):\n                image, meta, targets = data_to_device(data)\n\n                # Sanity check on targets\n                assert not torch.isnan(targets).any(), \"NaN detected in targets!\"\n                assert not torch.isinf(targets).any(), \"Inf detected in targets!\"\n\n                optimizer.zero_grad()\n                mixed_x, mixed_meta, y_a, y_b, lam = mixup_data(image, meta, targets.unsqueeze(1).float())\n                out = model(mixed_x, mixed_meta)\n\n                # Sanity check on outputs\n                if torch.isnan(out).any() or torch.isinf(out).any():\n                    print(\"Warning: NaN or Inf detected in model output!\")\n                    continue  # Skip problematic batch\n\n                loss = lam * criterion(out, y_a) + (1 - lam) * criterion(out, y_b)\n\n                with torch.autograd.detect_anomaly():\n                    loss.backward()\n\n                torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n                optimizer.step()\n                scheduler.step()\n                train_losses += loss.item()\n                train_preds = torch.round(torch.sigmoid(out))\n                correct += (train_preds.cpu() == targets.cpu().unsqueeze(1)).sum().item()\n\n                if k % 10 == 0:\n                    print(f\"Batch {k}: GPU memory allocated = {torch.cuda.memory_allocated() / 1024 ** 2:.1f} MB\")\n\n            train_acc = correct / len(train_index)\n\n            model.eval()\n            valid_preds = torch.zeros(size=(len(valid_index), 1), device=DEVICE, dtype=torch.float32)\n\n            tta_steps = 5\n            with torch.no_grad():\n                for k, data in tqdm(enumerate(valid_loader), total=len(valid_loader)):\n                    image, meta, targets = data_to_device(data)\n                    batch_size = image.size(0)\n                    batch_preds = torch.zeros((batch_size, 1), device=DEVICE)\n\n                    for t in range(tta_steps):\n                        tta_images = []\n                        for b in range(batch_size):\n                            idx = k * batch_size + b\n                            if idx >= len(valid_data):\n                                continue\n                            sample = RSNADataset(valid_data.iloc[[idx]], vertical_flip, horizontal_flip, is_train=True)[0]\n                            tta_images.append(sample[\"image\"].unsqueeze(0))\n\n                        if len(tta_images) == 0:\n                            continue  # Skip empty TTA batch\n\n                        tta_batch = torch.cat(tta_images).to(DEVICE)\n                        meta_repeated = meta.repeat(tta_steps, 1)[:tta_batch.size(0)]\n                        out = model(tta_batch, meta_repeated)\n                        batch_preds += torch.sigmoid(out)\n\n                    avg_preds = batch_preds / tta_steps\n                    valid_preds[k * batch_size: k * batch_size + batch_size] = avg_preds\n\n            y_true = valid_data['cancer'].values\n            y_score = valid_preds.cpu().numpy()\n\n            # Sanitize predictions to remove NaNs/Infs\n            y_score = np.nan_to_num(y_score, nan=0.0, posinf=1.0, neginf=0.0)\n            val_preds_bin = np.round(y_score)\n\n            # Check for NaNs/Infs before plotting confusion matrix\n            if np.isnan(y_true).any() or np.isnan(val_preds_bin).any() or np.isinf(val_preds_bin).any():\n                print(\"Warning: NaN or Inf detected in true or predicted labels, skipping confusion matrix plot.\")\n            else:\n                plot_confusion_matrix(y_true, val_preds_bin, epoch + 1, i + 1)\n\n            valid_acc = accuracy_score(y_true, val_preds_bin)\n            valid_roc = 0.0 if len(np.unique(y_true)) < 2 else roc_auc_score(y_true, y_score)\n            val_precision = precision_score(y_true, val_preds_bin)\n            val_recall = recall_score(y_true, val_preds_bin)\n            val_f1 = f1_score(y_true, val_preds_bin)\n            try:\n                val_loss = log_loss(y_true, y_score)\n            except Exception:\n                val_loss = float('nan')\n\n            print(f\"loss: {train_losses:.4f} - accuracy: {train_acc:.3f} - val_loss: {val_loss:.4f} - val_accuracy: {valid_acc:.3f} - val_auc: {valid_roc:.3f} - precision: {val_precision:.3f} - recall: {val_recall:.3f} - f1score: {val_f1:.3f}\")\n\n            duration = str(dtime.timedelta(seconds=time() - start_time))[:7]\n            final_logs = '{} | Epoch: {}/{} | Loss: {:.4} | Acc_tr: {:.3} | Acc_vd: {:.3} | ROC: {:.3}'.format(\n                duration, epoch + 1, EPOCHS, train_losses, train_acc, valid_acc, valid_roc)\n            add_in_file(final_logs, f)\n            print(final_logs)\n\n            if not best_roc or valid_roc > best_roc:\n                best_roc = valid_roc\n                patience_f = PATIENCE\n                model_name = f\"BEST_Fold{i + 1}_Epoch{epoch + 1}_ROC{valid_roc:.3f}.pth\"\n                torch.save(model.state_dict(), os.path.join(\"saved_models\", model_name))\n                with open(\"saved_models/best_model_name.txt\", \"w\") as f_name:\n                    f_name.write(model_name)\n                print(f\"✅ Best model saved: {model_name}\")\n            else:\n                patience_f -= 1\n                if patience_f == 0:\n                    print(f\"⛔ Early stopping triggered — Best ROC: {best_roc:.4f}\")\n                    break\n\n        del train_ds, valid_ds, train_loader, valid_loader, image, targets\n        gc.collect()\n    f.close()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T14:48:18.739849Z","iopub.execute_input":"2025-05-21T14:48:18.74017Z","iopub.status.idle":"2025-05-21T14:48:18.770361Z","shell.execute_reply.started":"2025-05-21T14:48:18.740144Z","shell.execute_reply":"2025-05-21T14:48:18.769468Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix, accuracy_score, roc_auc_score, precision_score, recall_score, f1_score, log_loss\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom torch.optim import AdamW\nfrom time import time\nimport os\nimport gc\nimport numpy as np\nfrom tqdm import tqdm\n\ndef plot_confusion_matrix(y_true, y_pred, epoch, fold):\n    if len(y_true) == 0 or len(y_pred) == 0:\n        print(f\"⚠️ Skipping confusion matrix for Epoch {epoch}, Fold {fold} due to empty inputs.\")\n        return\n    cm = confusion_matrix(y_true, y_pred)\n    plt.figure(figsize=(5, 5))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n                xticklabels=['Negative', 'Positive'],\n                yticklabels=['Negative', 'Positive'])\n    plt.title(f'Fold {fold} - Epoch {epoch}\\nConfusion Matrix')\n    plt.ylabel('Actual')\n    plt.xlabel('Predicted')\n    plt.show()\n\ndef train_folds(model, train_original):\n    log_file = open(f\"logs_{VERSION}.txt\", \"w+\")\n    os.makedirs(\"saved_models\", exist_ok=True)\n    group_fold = StratifiedGroupKFold(n_splits=FOLDS)\n    k_folds = group_fold.split(train_original,\n                               train_original['cancer'],\n                               groups=train_original['patient_id'])\n\n    for fold_idx, (train_index, valid_index) in enumerate(k_folds, start=1):\n        print(f\"---------- Fold: {fold_idx} ----------\")\n        print(f\"---------- Fold: {fold_idx} ----------\", file=log_file)\n\n        best_roc = None\n        patience_f = PATIENCE\n\n        train_data = train_original.iloc[train_index].reset_index(drop=True)\n        valid_data = train_original.iloc[valid_index].reset_index(drop=True)\n\n        train_ds = RSNADataset(train_data, vertical_flip, horizontal_flip, is_train=True)\n        valid_ds = RSNADataset(valid_data, vertical_flip, horizontal_flip, is_train=False)\n\n        train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE1, shuffle=True, num_workers=WORKERS)\n        valid_loader = DataLoader(valid_ds, batch_size=BATCH_SIZE2, shuffle=False, num_workers=WORKERS)\n\n        optimizer = AdamW(model.parameters(), lr=LR, weight_decay=WD)\n        criterion = FocalLoss(alpha=0.25, gamma=2)\n\n        tta_steps = 5  # ✅ Defined here now\n\n        for epoch in range(EPOCHS):\n            start_time = time()\n            correct = 0\n            train_losses = 0\n\n            model.train()\n            for _, data in tqdm(enumerate(train_loader), total=len(train_loader)):\n                image, meta, targets = data_to_device(data)\n                optimizer.zero_grad()\n                mixed_x, mixed_meta, y_a, y_b, lam = mixup_data(image, meta, targets.unsqueeze(1).float())\n                out = model(mixed_x, mixed_meta)\n                loss = lam * criterion(out, y_a) + (1 - lam) * criterion(out, y_b)\n                loss.backward()\n                torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n                optimizer.step()\n                train_losses += loss.item()\n                train_preds = torch.round(torch.sigmoid(out))\n                correct += (train_preds.cpu() == targets.cpu().unsqueeze(1)).sum().item()\n\n            train_acc = correct / len(train_index)\n\n            # Validation with TTA\n            model.eval()\n            valid_preds = torch.zeros(size=(len(valid_data), 1), device=DEVICE, dtype=torch.float32)\n            tta_failed = True\n\n            with torch.no_grad():\n                for k, data in tqdm(enumerate(valid_loader), total=len(valid_loader)):\n                    image, meta, targets = data_to_device(data)\n                    batch_size = image.size(0)\n                    batch_preds = torch.zeros((batch_size, 1), device=DEVICE)\n\n                    for t in range(tta_steps):\n                        tta_images = []\n                        for i in range(batch_size):\n                            idx = k * BATCH_SIZE2 + i\n                            if idx >= len(valid_data):\n                                continue\n                            try:\n                                sample = RSNADataset(valid_data.iloc[[idx]],\n                                                     vertical_flip, horizontal_flip,\n                                                     is_train=True)[0]\n                                tta_images.append(sample[\"image\"].unsqueeze(0))\n                            except Exception as e:\n                                print(f\"[TTA ERROR] Skipping idx {idx}: {e}\")\n                                continue\n                        if len(tta_images) == 0:\n                            continue\n                        tta_batch = torch.cat(tta_images).to(DEVICE)\n                        out = model(tta_batch, meta[:len(tta_batch)])\n                        batch_preds[:len(tta_batch)] += torch.sigmoid(out)\n\n                    if batch_preds.isnan().any():\n                        print(f\"[WARN] TTA failed for batch {k}, skipping\")\n                        continue\n\n                    avg_preds = batch_preds / tta_steps\n                    start_idx = k * batch_size\n                    end_idx = start_idx + avg_preds.size(0)\n                    valid_preds[start_idx:end_idx] = avg_preds\n                    tta_failed = False\n\n            # Fallback if TTA failed completely\n            if tta_failed:\n                print(\"[FALLBACK] TTA failed — using direct inference instead.\")\n                with torch.no_grad():\n                    for k, data in tqdm(enumerate(valid_loader), total=len(valid_loader)):\n                        image, meta, targets = data_to_device(data)\n                        out = model(image, meta)\n                        valid_preds[k * image.size(0): k * image.size(0) + image.size(0)] = torch.sigmoid(out)\n\n            # Compute metrics\n            y_true = valid_data['cancer'].values\n            y_score = valid_preds[:len(valid_data)].cpu().numpy()\n            val_preds_bin = torch.round(valid_preds[:len(valid_data)].cpu()).numpy()\n\n            print(f\"[INFO] Fold {fold_idx}, Epoch {epoch+1} — Validation sample count: {len(y_true)}\")\n\n            val_preds_bin_1d = val_preds_bin.squeeze()\n            y_score_1d = y_score.squeeze()\n            valid_mask = (~np.isnan(y_true)) & (~np.isnan(val_preds_bin_1d)) & (~np.isnan(y_score_1d))\n\n            y_true = y_true[valid_mask]\n            val_preds_bin = val_preds_bin_1d[valid_mask]\n            y_score = y_score_1d[valid_mask]\n\n            if len(y_true) < 20:\n                print(f\"⚠️ Too few valid samples in Fold {fold_idx}, Epoch {epoch+1} — skipping confusion matrix.\")\n            else:\n                plot_confusion_matrix(y_true, val_preds_bin, epoch+1, fold_idx)\n\n            valid_acc = accuracy_score(y_true, val_preds_bin)\n            valid_roc = 0.0 if len(np.unique(y_true)) < 2 else roc_auc_score(y_true, y_score)\n            val_precision = precision_score(y_true, val_preds_bin, zero_division=0)\n            val_recall = recall_score(y_true, val_preds_bin, zero_division=0)\n            val_f1 = f1_score(y_true, val_preds_bin, zero_division=0)\n            try:\n                val_loss = log_loss(y_true, y_score)\n            except:\n                val_loss = float('nan')\n\n            print(f\"Epoch {epoch+1}/{EPOCHS} - loss: {train_losses:.4f} - accuracy: {train_acc:.3f} - \"\n                  f\"val_loss: {val_loss:.4f} - val_accuracy: {valid_acc:.3f} - val_auc: {valid_roc:.3f} - \"\n                  f\"precision: {val_precision:.3f} - recall: {val_recall:.3f} - f1score: {val_f1:.3f}\", flush=True)\n\n            print(f\"Epoch {epoch+1}/{EPOCHS} - loss: {train_losses:.4f} - accuracy: {train_acc:.3f} - \"\n                  f\"val_loss: {val_loss:.4f} - val_accuracy: {valid_acc:.3f} - val_auc: {valid_roc:.3f}\", file=log_file)\n\n            if not best_roc or valid_roc > best_roc:\n                best_roc = valid_roc\n                patience_f = PATIENCE\n                model_name = f\"BEST_Fold{fold_idx}_Epoch{epoch+1}_ROC{valid_roc:.3f}.pth\"\n                torch.save(model.state_dict(), os.path.join(\"saved_models\", model_name))\n                with open(\"saved_models/best_model_name.txt\", \"w\") as f:\n                    f.write(model_name)\n                print(f\"✅ Best model saved: {model_name}\")\n            else:\n                patience_f -= 1\n                if patience_f == 0:\n                    print(f\"⛔ Early stopping triggered — Best ROC: {best_roc:.4f}\")\n                    break\n\n        del train_ds, valid_ds, train_loader, valid_loader, image, targets\n        gc.collect()\n\n    log_file.close()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T15:21:10.273247Z","iopub.execute_input":"2025-05-21T15:21:10.273604Z","iopub.status.idle":"2025-05-21T15:21:10.304364Z","shell.execute_reply.started":"2025-05-21T15:21:10.273576Z","shell.execute_reply":"2025-05-21T15:21:10.303369Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"FOLDS = 3\nEPOCHS = 12\nPATIENCE = 12\nWORKERS = 8\n# iLR = 0.0005\n# WD = 0.0\nLR_PATIENCE = 1            # 1 model not improving until lr is decreasing\nLR_FACTOR = 0.4            # by how much the lr is decreasing\n\nLR = 2e-4       # Lower LR works better with AdamW\nWD = 1e-2       # Use non-zero weight decay for regularization\n\n\n\n# LR = 0.001\n# LR_PATIENCE = 2\n# LR_FACTOR = 0.5\n\n\n\n\nBATCH_SIZE1 = 32           # for train\nBATCH_SIZE2 = 16           # for valid\n\nVERSION = 'v1'\nMODEL = 'resnet50'\n\nmodel1 = ResNet50Network(output_size=output_size, no_columns=no_columns).to(DEVICE)\n\n# with open(\"saved_models/best_model_name.txt\", \"r\") as f:\n#     best_model_name = f.read().strip()\n\n# model_path = os.path.join(\"saved_models\", best_model_name)\n# model1.load_state_dict(torch.load(model_path, map_location=DEVICE))\n# print(f\"✅ Loaded model: {model_path}\")\n\n\n# ------------------\n\n# Run the cell below to train\n#Ran it locally on all data, see the results below\ntrain_folds(model=model1, train_original=train_set)\n\n# Print the logs during training\n# f = open('/kaggle/input/rsna-breast-cancer-helper-data/logs_v1.txt', \"r\")\n# contents = f.read()\n# print(contents)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T15:26:01.8834Z","iopub.execute_input":"2025-05-21T15:26:01.883778Z","iopub.status.idle":"2025-05-21T15:46:29.391034Z","shell.execute_reply.started":"2025-05-21T15:26:01.883749Z","shell.execute_reply":"2025-05-21T15:46:29.390229Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"FOLDS = 3\nEPOCHS = 3\nPATIENCE = 3\nWORKERS = 8\nLR = 0.0005\nWD = 0.0\nLR_PATIENCE = 1            # 1 model not improving until lr is decreasing\nLR_FACTOR = 0.4            # by how much the lr is decreasing\n\nBATCH_SIZE1 = 32           # for train\nBATCH_SIZE2 = 16           # for valid\n\nVERSION = 'v2'\nMODEL = 'effnet'\n\nmodel2 = EffNetNetwork(output_size=output_size, no_columns=no_columns).to(DEVICE)\n\n# ------------------\n\n# Run the cell below to train\n# Ran it locally on all data, see the results below\ntrain_folds(model=model2, train_original=train_set)\n\n# Print the logs during training\n# f = open('/kaggle/input/rsna-breast-cancer-helper-data/logs_v2.txt', \"r\")\n# contents = f.read()\n# print(contents)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T15:21:26.274722Z","iopub.execute_input":"2025-05-21T15:21:26.27559Z","iopub.status.idle":"2025-05-21T15:23:38.290519Z","shell.execute_reply.started":"2025-05-21T15:21:26.275556Z","shell.execute_reply":"2025-05-21T15:23:38.288976Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix, accuracy_score, roc_auc_score, precision_score, recall_score, f1_score, log_loss\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom transformers import get_cosine_schedule_with_warmup\nfrom torch.optim import AdamW\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\nimport numpy as np\nimport torch\nimport os\nimport gc\nfrom time import time\nimport datetime as dtime\n\n\ndef plot_confusion_matrix(y_true, y_pred, epoch, fold):\n    if len(y_true) == 0 or len(y_pred) == 0:\n        print(f\"⚠️ Skipping confusion matrix for Epoch {epoch}, Fold {fold} due to empty inputs.\")\n        return\n    cm = confusion_matrix(y_true, y_pred)\n    plt.figure(figsize=(5, 5))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n                xticklabels=['Negative', 'Positive'],\n                yticklabels=['Negative', 'Positive'])\n    plt.title(f'Fold {fold} - Epoch {epoch}\\nConfusion Matrix')\n    plt.ylabel('Actual')\n    plt.xlabel('Predicted')\n    plt.show()\n\n\ndef train_folds(model, train_original):\n    log_file = open(f\"logs_{VERSION}.txt\", \"w+\")\n    os.makedirs(\"saved_models\", exist_ok=True)\n\n    group_fold = StratifiedGroupKFold(n_splits=FOLDS)\n    k_folds = group_fold.split(train_original,\n                               train_original['cancer'],\n                               groups=train_original['patient_id'])\n\n    for fold_idx, (train_index, valid_index) in enumerate(k_folds, start=1):\n        print(f\"---------- Fold: {fold_idx} ----------\")\n        print(f\"---------- Fold: {fold_idx} ----------\", file=log_file)\n\n        best_roc = None\n        patience_f = PATIENCE\n\n        train_data = train_original.iloc[train_index].reset_index(drop=True)\n        valid_data = train_original.iloc[valid_index].reset_index(drop=True)\n\n        train_ds = RSNADataset(train_data, vertical_flip, horizontal_flip, is_train=True)\n        valid_ds = RSNADataset(valid_data, vertical_flip, horizontal_flip, is_train=False)\n\n        # Pre-create TTA dataset (with augmentation) for validation fold\n        valid_tta_ds = RSNADataset(valid_data, vertical_flip, horizontal_flip, is_train=True)\n\n        train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE1, shuffle=True, num_workers=WORKERS)\n        valid_loader = DataLoader(valid_ds, batch_size=BATCH_SIZE2, shuffle=False, num_workers=WORKERS)\n\n        total_steps = len(train_loader) * EPOCHS\n        warmup_steps = int(0.1 * total_steps)\n\n        optimizer = AdamW(model.parameters(), lr=LR, weight_decay=WD)\n        scheduler = get_cosine_schedule_with_warmup(optimizer, num_warmup_steps=warmup_steps, num_training_steps=total_steps)\n\n        criterion = FocalLoss(alpha=0.25, gamma=2)\n\n        tta_steps = 5  # Number of TTA augmentations\n\n        for epoch in range(EPOCHS):\n            start_time = time()\n            model.train()\n            correct = 0\n            train_losses = 0\n\n            for _, data in tqdm(enumerate(train_loader), total=len(train_loader)):\n                image, meta, targets = data_to_device(data)\n                optimizer.zero_grad()\n                mixed_x, mixed_meta, y_a, y_b, lam = mixup_data(image, meta, targets.unsqueeze(1).float())\n                out = model(mixed_x, mixed_meta)\n                loss = lam * criterion(out, y_a) + (1 - lam) * criterion(out, y_b)\n                loss.backward()\n                torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n                optimizer.step()\n                scheduler.step()\n                train_losses += loss.item()\n                train_preds = torch.round(torch.sigmoid(out))\n                correct += (train_preds.cpu() == targets.cpu().unsqueeze(1)).sum().item()\n\n            train_acc = correct / len(train_index)\n\n            model.eval()\n            valid_preds = torch.zeros(size=(len(valid_data), 1), device=DEVICE, dtype=torch.float32)\n\n            with torch.no_grad():\n                for batch_idx, data in tqdm(enumerate(valid_loader), total=len(valid_loader)):\n                    image, meta, targets = data_to_device(data)\n                    batch_size = image.size(0)\n                    batch_preds = torch.zeros((batch_size, 1), device=DEVICE)\n\n                    for t in range(tta_steps):\n                        tta_images = []\n                        # Collect augmented images for TTA\n                        for i in range(batch_size):\n                            idx = batch_idx * BATCH_SIZE2 + i\n                            # Fix: Make sure idx is within valid_tta_ds range\n                            if idx >= len(valid_tta_ds):\n                                continue\n                            sample = valid_tta_ds[idx]\n                            tta_images.append(sample[\"image\"].unsqueeze(0))\n\n                        if not tta_images:\n                            continue\n\n                        tta_batch = torch.cat(tta_images).to(DEVICE)\n                        meta_repeated = meta[:len(tta_batch)]  # Align meta size\n                        out = model(tta_batch, meta_repeated)\n                        batch_preds[:len(tta_batch)] += torch.sigmoid(out)\n\n                    avg_preds = batch_preds / tta_steps\n                    start_idx = batch_idx * batch_size\n                    end_idx = start_idx + avg_preds.size(0)\n                    valid_preds[start_idx:end_idx] = avg_preds\n\n            y_true = valid_data['cancer'].values\n            y_score = valid_preds[:len(valid_data)].cpu().numpy()\n            val_preds_bin = torch.round(valid_preds[:len(valid_data)].cpu()).numpy()\n\n            val_preds_bin_1d = val_preds_bin.squeeze()\n            y_score_1d = y_score.squeeze()\n            mask = (~np.isnan(y_true)) & (~np.isnan(val_preds_bin_1d)) & (~np.isnan(y_score_1d))\n\n            y_true = y_true[mask]\n            val_preds_bin = val_preds_bin_1d[mask]\n            y_score = y_score_1d[mask]\n\n            if len(y_true) >= 20:\n                plot_confusion_matrix(y_true, val_preds_bin, epoch + 1, fold_idx)\n\n            valid_acc = accuracy_score(y_true, val_preds_bin)\n            valid_roc = 0.0 if len(np.unique(y_true)) < 2 else roc_auc_score(y_true, y_score)\n            val_precision = precision_score(y_true, val_preds_bin, zero_division=0)\n            val_recall = recall_score(y_true, val_preds_bin, zero_division=0)\n            val_f1 = f1_score(y_true, val_preds_bin, zero_division=0)\n\n            try:\n                val_loss = log_loss(y_true, y_score)\n            except Exception:\n                val_loss = float('nan')\n\n            print(f\"Epoch {epoch+1}/{EPOCHS} - loss: {train_losses:.4f} - accuracy: {train_acc:.3f} - \"\n                  f\"val_loss: {val_loss:.4f} - val_accuracy: {valid_acc:.3f} - val_auc: {valid_roc:.3f} - \"\n                  f\"precision: {val_precision:.3f} - recall: {val_recall:.3f} - f1score: {val_f1:.3f}\")\n\n            print(f\"Epoch {epoch+1}/{EPOCHS} - loss: {train_losses:.4f} - accuracy: {train_acc:.3f} - \"\n                  f\"val_loss: {val_loss:.4f} - val_accuracy: {valid_acc:.3f} - val_auc: {valid_roc:.3f}\", file=log_file)\n\n            if best_roc is None or valid_roc > best_roc:\n                best_roc = valid_roc\n                patience_f = PATIENCE\n                model_name = f\"BEST_Fold{fold_idx}_Epoch{epoch+1}_ROC{valid_roc:.3f}.pth\"\n                torch.save(model.state_dict(), os.path.join(\"saved_models\", model_name))\n                with open(\"saved_models/best_model_name.txt\", \"w\") as f:\n                    f.write(model_name)\n                print(f\"✅ Best model saved: {model_name}\")\n            else:\n                patience_f -= 1\n                if patience_f == 0:\n                    print(f\"⛔ Early stopping triggered — Best ROC: {best_roc:.4f}\")\n                    break\n\n        del train_ds, valid_ds, train_loader, valid_loader, image, targets\n        gc.collect()\n\n    log_file.close()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T14:36:03.712582Z","iopub.execute_input":"2025-05-21T14:36:03.712982Z","iopub.status.idle":"2025-05-21T14:36:03.740401Z","shell.execute_reply.started":"2025-05-21T14:36:03.712951Z","shell.execute_reply":"2025-05-21T14:36:03.739377Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"2nd","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix, accuracy_score, roc_auc_score, precision_score, recall_score, f1_score, log_loss\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom transformers import get_cosine_schedule_with_warmup\nimport gc\nimport datetime as dtime\nfrom time import time\nimport os\nfrom torch.optim import AdamW\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\nimport numpy as np\nimport torch\n\n\ndef plot_confusion_matrix(y_true, y_pred, epoch, fold):\n    cm = confusion_matrix(y_true, y_pred)\n    plt.figure(figsize=(5, 5))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n                xticklabels=['Negative', 'Positive'],\n                yticklabels=['Negative', 'Positive'])\n    plt.title(f'Fold {fold} - Epoch {epoch}\\nConfusion Matrix')\n    plt.ylabel('Actual')\n    plt.xlabel('Predicted')\n    plt.show()\n\n\ndef reset_weights(m):\n    # Recursively reset weights of a model (common layers)\n    for layer in m.children():\n        if hasattr(layer, 'reset_parameters'):\n            layer.reset_parameters()\n        else:\n            reset_weights(layer)\n\n\ndef train_folds(model, train_original):\n    f = open(f\"logs_{VERSION}.txt\", \"w+\")\n    os.makedirs(\"saved_models\", exist_ok=True)\n\n    group_fold = StratifiedGroupKFold(n_splits=FOLDS)\n    k_folds = group_fold.split(train_original, train_original['cancer'], groups=train_original['patient_id'])\n\n    for i, (train_index, valid_index) in enumerate(k_folds):\n        print(clr.S + f\"---------- Fold: {i+1} ----------\" + clr.E)\n        add_in_file(f\"---------- Fold: {i+1} ----------\", f)\n\n        reset_weights(model)\n\n        best_roc = None\n        patience_f = PATIENCE\n\n        train_data = train_original.iloc[train_index].reset_index(drop=True)\n        valid_data = train_original.iloc[valid_index].reset_index(drop=True)\n\n        print(f\"Fold {i+1}: Train size = {len(train_data)}, Positives = {train_data['cancer'].sum()}, \"\n              f\"Negatives = {len(train_data) - train_data['cancer'].sum()}\")\n        print(f\"Fold {i+1}: Valid size = {len(valid_data)}, Positives = {valid_data['cancer'].sum()}, \"\n              f\"Negatives = {len(valid_data) - valid_data['cancer'].sum()}\")\n\n        train_ds = RSNADataset(train_data, vertical_flip, horizontal_flip, is_train=True)\n        valid_ds = RSNADataset(valid_data, vertical_flip, horizontal_flip, is_train=False)  # <--- No augmentation on validation\n\n        train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE1, shuffle=True, num_workers=WORKERS)\n        valid_loader = DataLoader(valid_ds, batch_size=BATCH_SIZE2, shuffle=False, num_workers=WORKERS)\n\n        total_steps = len(train_loader) * EPOCHS\n        warmup_steps = int(0.1 * total_steps)\n\n        optimizer = AdamW(model.parameters(), lr=LR, weight_decay=WD)\n        scheduler = get_cosine_schedule_with_warmup(optimizer,\n                                                    num_warmup_steps=warmup_steps,\n                                                    num_training_steps=total_steps)\n\n        criterion = FocalLoss(alpha=0.25, gamma=2)\n\n        tta_steps = 5\n\n        for epoch in range(EPOCHS):\n            start_time = time()\n            model.train()\n            correct = 0\n            train_losses = 0\n\n            for k, data in tqdm(enumerate(train_loader), total=len(train_loader)):\n                image, meta, targets = data_to_device(data)\n\n                optimizer.zero_grad()\n                mixed_x, mixed_meta, y_a, y_b, lam = mixup_data(image, meta, targets.unsqueeze(1).float())\n                out = model(mixed_x, mixed_meta)\n                loss = lam * criterion(out, y_a) + (1 - lam) * criterion(out, y_b)\n                loss.backward()\n\n                torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n                optimizer.step()\n                scheduler.step()\n\n                train_losses += loss.item()\n                train_preds = torch.round(torch.sigmoid(out))\n                correct += (train_preds.cpu() == targets.cpu().unsqueeze(1)).sum().item()\n\n                if k % 10 == 0:\n                    print(f\"Batch {k}: GPU memory allocated = {torch.cuda.memory_allocated() / 1024 ** 2:.1f} MB\")\n\n            train_acc = correct / len(train_index)\n\n            model.eval()\n            valid_preds = torch.zeros(size=(len(valid_data), 1), device=DEVICE, dtype=torch.float32)\n\n            with torch.no_grad():\n                for k, data in tqdm(enumerate(valid_loader), total=len(valid_loader)):\n                    image, meta, targets = data_to_device(data)\n                    batch_size = image.size(0)\n                    batch_preds = torch.zeros((batch_size, 1), device=DEVICE)\n\n                    # TTA: For each sample in batch, generate tta_steps augmented versions and average predictions\n                    for t in range(tta_steps):\n                        tta_images = []\n                        for b in range(batch_size):\n                            idx = k * BATCH_SIZE2 + b\n                            if idx >= len(valid_data):\n                                continue\n                            # Manually augment image for TTA\n                            sample = RSNADataset(valid_data.iloc[[idx]], vertical_flip, horizontal_flip, is_train=True)[0]\n                            tta_images.append(sample[\"image\"].unsqueeze(0))\n\n                        if not tta_images:\n                            continue\n\n                        tta_batch = torch.cat(tta_images).to(DEVICE)\n                        out = model(tta_batch, meta[:len(tta_batch)])\n                        batch_preds[:len(tta_batch)] += torch.sigmoid(out)\n\n                    avg_preds = batch_preds / tta_steps\n                    start_idx = k * batch_size\n                    valid_preds[start_idx:start_idx + avg_preds.size(0)] = avg_preds\n\n            y_true = valid_data['cancer'].values\n            y_score = valid_preds[:len(valid_data)].cpu().numpy()\n            val_preds_bin = np.round(y_score)\n\n            mask = (~np.isnan(y_true)) & (~np.isnan(val_preds_bin.squeeze())) & (~np.isnan(y_score.squeeze()))\n            y_true = y_true[mask]\n            val_preds_bin = val_preds_bin.squeeze()[mask]\n            y_score = y_score.squeeze()[mask]\n\n            if len(y_true) >= 20:\n                plot_confusion_matrix(y_true, val_preds_bin, epoch + 1, i + 1)\n\n            valid_acc = accuracy_score(y_true, val_preds_bin)\n            valid_roc = roc_auc_score(y_true, y_score) if len(np.unique(y_true)) > 1 else 0.0\n            val_precision = precision_score(y_true, val_preds_bin, zero_division=0)\n            val_recall = recall_score(y_true, val_preds_bin, zero_division=0)\n            val_f1 = f1_score(y_true, val_preds_bin, zero_division=0)\n\n            try:\n                val_loss = log_loss(y_true, y_score)\n            except Exception:\n                val_loss = float('nan')\n\n            print(f\"Epoch {epoch+1}/{EPOCHS} - loss: {train_losses:.4f} - accuracy: {train_acc:.3f} - \"\n                  f\"val_loss: {val_loss:.4f} - val_accuracy: {valid_acc:.3f} - val_auc: {valid_roc:.3f} - \"\n                  f\"precision: {val_precision:.3f} - recall: {val_recall:.3f} - f1score: {val_f1:.3f}\")\n\n            duration = str(dtime.timedelta(seconds=time() - start_time))[:7]\n            final_logs = '{} | Epoch: {}/{} | Loss: {:.4} | Acc_tr: {:.3} | Acc_vd: {:.3} | ROC: {:.3}'.format(\n                duration, epoch + 1, EPOCHS, train_losses, train_acc, valid_acc, valid_roc)\n            add_in_file(final_logs, f)\n            print(final_logs)\n\n            if best_roc is None or valid_roc > best_roc:\n                best_roc = valid_roc\n                patience_f = PATIENCE\n                model_name = f\"BEST_Fold{i + 1}_Epoch{epoch + 1}_ROC{valid_roc:.3f}.pth\"\n                torch.save(model.state_dict(), os.path.join(\"saved_models\", model_name))\n                with open(\"saved_models/best_model_name.txt\", \"w\") as f_name:\n                    f_name.write(model_name)\n                print(f\"✅ Best model saved: {model_name}\")\n            else:\n                patience_f -= 1\n                if patience_f == 0:\n                    print(f\"⛔ Early stopping triggered — Best ROC: {best_roc:.4f}\")\n                    break\n\n        del train_ds, valid_ds, train_loader, valid_loader, image, targets\n        gc.collect()\n    f.close()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T14:46:12.855219Z","iopub.execute_input":"2025-05-21T14:46:12.855546Z","iopub.status.idle":"2025-05-21T14:46:12.886428Z","shell.execute_reply.started":"2025-05-21T14:46:12.855519Z","shell.execute_reply":"2025-05-21T14:46:12.885466Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"3rd","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix, accuracy_score, roc_auc_score, precision_score, recall_score, f1_score, log_loss\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom transformers import get_cosine_schedule_with_warmup\nimport gc\nimport datetime as dtime\nfrom time import time\nimport os\nfrom torch.optim import AdamW\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\nimport numpy as np\nimport torch\n\ndef plot_confusion_matrix(y_true, y_pred, epoch, fold):\n    cm = confusion_matrix(y_true, y_pred)\n    plt.figure(figsize=(5, 5))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n                xticklabels=['Negative', 'Positive'],\n                yticklabels=['Negative', 'Positive'])\n    plt.title(f'Fold {fold} - Epoch {epoch}\\nConfusion Matrix')\n    plt.ylabel('Actual')\n    plt.xlabel('Predicted')\n    plt.show()\n\ndef reset_weights(m):\n    for layer in m.children():\n        if hasattr(layer, 'reset_parameters'):\n            layer.reset_parameters()\n        else:\n            reset_weights(layer)\n\ndef train_folds(model, train_original):\n    f = open(f\"logs_{VERSION}.txt\", \"w+\")\n    os.makedirs(\"saved_models\", exist_ok=True)\n\n    group_fold = StratifiedGroupKFold(n_splits=FOLDS)\n    k_folds = group_fold.split(train_original, train_original['cancer'], groups=train_original['patient_id'])\n\n    for i, (train_index, valid_index) in enumerate(k_folds):\n        print(clr.S + f\"---------- Fold: {i+1} ----------\" + clr.E)\n        add_in_file(f\"---------- Fold: {i+1} ----------\", f)\n\n        reset_weights(model)\n\n        best_roc = None\n        patience_f = PATIENCE\n\n        train_data = train_original.iloc[train_index].reset_index(drop=True)\n        valid_data = train_original.iloc[valid_index].reset_index(drop=True)\n\n        print(f\"Fold {i+1}: Train size = {len(train_data)}, Positives = {train_data['cancer'].sum()}, \"\n              f\"Negatives = {len(train_data) - train_data['cancer'].sum()}\")\n        print(f\"Fold {i+1}: Valid size = {len(valid_data)}, Positives = {valid_data['cancer'].sum()}, \"\n              f\"Negatives = {len(valid_data) - valid_data['cancer'].sum()}\")\n\n        train_ds = RSNADataset(train_data, vertical_flip, horizontal_flip, is_train=True)\n        valid_ds = RSNADataset(valid_data, vertical_flip, horizontal_flip, is_train=True)  # for TTA validation, use is_train=True\n\n        train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE1, shuffle=True, num_workers=WORKERS)\n        valid_loader = DataLoader(valid_ds, batch_size=BATCH_SIZE2, shuffle=False, num_workers=WORKERS)\n\n        total_steps = len(train_loader) * EPOCHS\n        warmup_steps = int(0.1 * total_steps)  # 10% warm-up\n\n        optimizer = AdamW(model.parameters(), lr=LR, weight_decay=WD)\n        scheduler = get_cosine_schedule_with_warmup(\n            optimizer,\n            num_warmup_steps=warmup_steps,\n            num_training_steps=total_steps\n        )\n        criterion = FocalLoss(alpha=0.25, gamma=2)\n\n        tta_steps = 5\n\n        for epoch in range(EPOCHS):\n            start_time = time()\n            correct = 0\n            train_losses = 0\n\n            model.train()\n            for k, data in tqdm(enumerate(train_loader), total=len(train_loader)):\n                image, meta, targets = data_to_device(data)\n\n                assert not torch.isnan(targets).any(), \"NaN detected in targets!\"\n                assert not torch.isinf(targets).any(), \"Inf detected in targets!\"\n\n                optimizer.zero_grad()\n                mixed_x, mixed_meta, y_a, y_b, lam = mixup_data(image, meta, targets.unsqueeze(1).float())\n                out = model(mixed_x, mixed_meta)\n\n                if torch.isnan(out).any() or torch.isinf(out).any():\n                    print(\"Warning: NaN or Inf detected in model output!\")\n                    continue\n\n                loss = lam * criterion(out, y_a) + (1 - lam) * criterion(out, y_b)\n\n                with torch.autograd.detect_anomaly():\n                    loss.backward()\n\n                torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n                optimizer.step()\n                scheduler.step()\n\n                train_losses += loss.item()\n                train_preds = torch.round(torch.sigmoid(out))\n                correct += (train_preds.cpu() == targets.cpu().unsqueeze(1)).sum().item()\n\n                if k % 10 == 0:\n                    print(f\"Batch {k}: GPU memory allocated = {torch.cuda.memory_allocated() / 1024 ** 2:.1f} MB\")\n\n            train_acc = correct / len(train_index)\n\n            model.eval()\n            valid_preds = torch.zeros(size=(len(valid_index), 1), device=DEVICE, dtype=torch.float32)\n\n            with torch.no_grad():\n                for k, data in tqdm(enumerate(valid_loader), total=len(valid_loader)):\n                    image, meta, targets = data_to_device(data)\n                    batch_size = image.size(0)\n                    batch_preds = torch.zeros((batch_size, 1), device=DEVICE)\n\n                    for t in range(tta_steps):\n                        tta_images = []\n                        for b in range(batch_size):\n                            idx = k * batch_size + b\n                            if idx >= len(valid_data):\n                                continue\n                            # Use valid_ds created above instead of recreating dataset sample by sample\n                            sample = valid_ds[idx]\n                            tta_images.append(sample[\"image\"].unsqueeze(0))\n\n                        if not tta_images:\n                            print(f\"[WARN] Empty TTA batch at batch {k}, skipping TTA for this batch.\")\n                            continue\n\n                        tta_batch = torch.cat(tta_images).to(DEVICE)\n                        meta_repeated = meta[:len(tta_batch)]\n                        out = model(tta_batch, meta_repeated)\n                        batch_preds[:len(tta_batch)] += torch.sigmoid(out)\n\n                    avg_preds = batch_preds / tta_steps\n                    valid_preds[k * batch_size: k * batch_size + batch_size] = avg_preds\n\n            y_true = valid_data['cancer'].values\n            y_score = valid_preds.cpu().numpy()\n\n            y_score = np.nan_to_num(y_score, nan=0.0, posinf=1.0, neginf=0.0)\n            val_preds_bin = np.round(y_score)\n\n            if np.isnan(y_true).any() or np.isnan(val_preds_bin).any() or np.isinf(val_preds_bin).any():\n                print(\"Warning: NaN or Inf detected in true or predicted labels, skipping confusion matrix plot.\")\n            else:\n                plot_confusion_matrix(y_true, val_preds_bin, epoch + 1, i + 1)\n\n            valid_acc = accuracy_score(y_true, val_preds_bin)\n            valid_roc = 0.0 if len(np.unique(y_true)) < 2 else roc_auc_score(y_true, y_score)\n            val_precision = precision_score(y_true, val_preds_bin)\n            val_recall = recall_score(y_true, val_preds_bin)\n            val_f1 = f1_score(y_true, val_preds_bin)\n            try:\n                val_loss = log_loss(y_true, y_score)\n            except Exception:\n                val_loss = float('nan')\n\n            print(f\"loss: {train_losses:.4f} - accuracy: {train_acc:.3f} - val_loss: {val_loss:.4f} - val_accuracy: {valid_acc:.3f} - val_auc: {valid_roc:.3f} - precision: {val_precision:.3f} - recall: {val_recall:.3f} - f1score: {val_f1:.3f}\")\n\n            duration = str(dtime.timedelta(seconds=time() - start_time))[:7]\n            final_logs = '{} | Epoch: {}/{} | Loss: {:.4} | Acc_tr: {:.3} | Acc_vd: {:.3} | ROC: {:.3}'.format(\n                duration, epoch + 1, EPOCHS, train_losses, train_acc, valid_acc, valid_roc)\n            add_in_file(final_logs, f)\n            print(final_logs)\n\n            if not best_roc or valid_roc > best_roc:\n                best_roc = valid_roc\n                patience_f = PATIENCE\n                model_name = f\"BEST_Fold{i + 1}_Epoch{epoch + 1}_ROC{valid_roc:.3f}.pth\"\n                torch.save(model.state_dict(), os.path.join(\"saved_models\", model_name))\n                with open(\"saved_models/best_model_name.txt\", \"w\") as f_name:\n                    f_name.write(model_name)\n                print(f\"✅ Best model saved: {model_name}\")\n            else:\n                patience_f -= 1\n                if patience_f == 0:\n                    print(f\"⛔ Early stopping triggered — Best ROC: {best_roc:.4f}\")\n                    break\n\n        del train_ds, valid_ds, train_loader, valid_loader, image, targets\n        gc.collect()\n    f.close()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-21T15:13:37.997541Z","iopub.execute_input":"2025-05-21T15:13:37.997964Z","iopub.status.idle":"2025-05-21T15:13:38.027529Z","shell.execute_reply.started":"2025-05-21T15:13:37.997934Z","shell.execute_reply":"2025-05-21T15:13:38.026597Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"__NOTE:__ Dataset used for training is from https://www.kaggle.com/datasets/anitho2910/rsna-mammography-breast-cancer-detection-png which was created from the given dicom image from my eda notebook https://www.kaggle.com/code/anitho2910/eda-notebook-breast-cancer ","metadata":{}},{"cell_type":"markdown","source":"# Importing Libraries","metadata":{}},{"cell_type":"code","source":"from fastai.vision.all import *\nfrom fastai.data.all import *\nfrom sklearn.model_selection import StratifiedShuffleSplit\nimport gc","metadata":{"execution":{"iopub.status.busy":"2022-12-17T08:07:19.970408Z","iopub.execute_input":"2022-12-17T08:07:19.971163Z","iopub.status.idle":"2022-12-17T08:07:21.808685Z","shell.execute_reply.started":"2022-12-17T08:07:19.971072Z","shell.execute_reply":"2022-12-17T08:07:21.807623Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Hyperparameters","metadata":{}},{"cell_type":"code","source":"seed = 42\nsave_path = '/kaggle/working'\ntrain_size = 0.8\nbatch_size = 32\nimage_resize = 256\nlr_unfreeze = slice(1e-7, 3e-6)\nn_epochs = 15","metadata":{"execution":{"iopub.status.busy":"2022-12-17T10:21:27.548513Z","iopub.execute_input":"2022-12-17T10:21:27.549168Z","iopub.status.idle":"2022-12-17T10:21:27.559512Z","shell.execute_reply.started":"2022-12-17T10:21:27.549116Z","shell.execute_reply":"2022-12-17T10:21:27.558464Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Helper Functions","metadata":{}},{"cell_type":"code","source":"#https://www.kaggle.com/competitions/rsna-breast-cancer-detection/discussion/369267  \ndef pfbeta_torch(preds, labels, beta=1):\n    softmax = torch.nn.Softmax(dim = -1)\n    preds = softmax(preds)\n    preds = preds[:, 1]\n    preds = preds.clip(0, 1)\n    y_true_count = labels.sum()\n    ctp = preds[labels==1].sum()\n    cfp = preds[labels==0].sum()\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp)\n    c_recall = ctp / y_true_count\n    if (c_precision > 0 and c_recall > 0):\n        result = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall)\n        return result\n    else:\n        return 0.0","metadata":{"execution":{"iopub.status.busy":"2022-12-17T08:07:21.827418Z","iopub.execute_input":"2022-12-17T08:07:21.830066Z","iopub.status.idle":"2022-12-17T08:07:21.840233Z","shell.execute_reply.started":"2022-12-17T08:07:21.830027Z","shell.execute_reply":"2022-12-17T08:07:21.839314Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Loading","metadata":{}},{"cell_type":"code","source":"base_path = Path('/kaggle/input')\nbase_images_path = base_path/'rsna-mammography-breast-cancer-detection-png'/'png_images'\nbase_data_path = base_path/'rsna-breast-cancer-detection'\ndf = pd.read_csv(base_data_path/'train.csv')\nprint(df.shape)\nprint(f\"Total Number of patient {len(df['patient_id'].unique())}\")\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-17T08:07:21.846333Z","iopub.execute_input":"2022-12-17T08:07:21.84881Z","iopub.status.idle":"2022-12-17T08:07:21.963706Z","shell.execute_reply.started":"2022-12-17T08:07:21.848772Z","shell.execute_reply":"2022-12-17T08:07:21.96266Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"only_cc_view_data = df[df['view'] == 'CC'].copy()\nprint(only_cc_view_data.shape)\nprint(f'''Number of patient: {len(only_cc_view_data['patient_id'].unique())} \nand patient with more than 2 scans {(only_cc_view_data['patient_id'].value_counts() > 2).sum()}''')\nonly_cc_view_data.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-17T08:07:21.967912Z","iopub.execute_input":"2022-12-17T08:07:21.970208Z","iopub.status.idle":"2022-12-17T08:07:22.010324Z","shell.execute_reply.started":"2022-12-17T08:07:21.970169Z","shell.execute_reply":"2022-12-17T08:07:22.009381Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"only_cc_view_data['path'] = only_cc_view_data.apply(lambda x: base_images_path/str(x['patient_id'])/(str(x['image_id'])+'.png'), axis = 1)\nfinal_subset = only_cc_view_data.drop_duplicates(subset = ['patient_id', 'laterality'], keep = 'last').copy()\nprint(final_subset.shape)\nprint(f'''Number of patient: {len(final_subset['patient_id'].unique())} \nand patient with more than 2 scans {(final_subset['patient_id'].value_counts() > 2).sum()}''')\nfinal_subset.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-17T08:07:22.014515Z","iopub.execute_input":"2022-12-17T08:07:22.017003Z","iopub.status.idle":"2022-12-17T08:07:23.155347Z","shell.execute_reply.started":"2022-12-17T08:07:22.016963Z","shell.execute_reply":"2022-12-17T08:07:23.154299Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_subset['is_valid'] = False\nstrata = StratifiedShuffleSplit(n_splits=2, train_size = train_size, random_state=seed)\nfor (train_idx, valid_idx) in strata.split(final_subset.index, final_subset['cancer']):\n    final_subset.iloc[train_idx, -1] = False\n    final_subset.iloc[valid_idx, -1] = True\n\nprint(final_subset['is_valid'].value_counts(normalize = True))\nfinal_subset.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-17T08:07:23.160066Z","iopub.execute_input":"2022-12-17T08:07:23.162675Z","iopub.status.idle":"2022-12-17T08:07:23.214124Z","shell.execute_reply.started":"2022-12-17T08:07:23.16263Z","shell.execute_reply":"2022-12-17T08:07:23.21316Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"number_class_0 = (final_subset['cancer'] == 0).sum()\nnumber_class_1 = (final_subset['cancer'] == 1).sum()\nweight_class_0 = 1\nweight_class_1 = number_class_0//number_class_1\nprint(f\"Cross Entropy weight for class 1: {weight_class_0}, and for class 0: {weight_class_1}\")\nweights = torch.tensor([weight_class_0, weight_class_1], dtype = torch.float32)\nprint(weights)","metadata":{"execution":{"iopub.status.busy":"2022-12-17T08:07:23.21865Z","iopub.execute_input":"2022-12-17T08:07:23.220912Z","iopub.status.idle":"2022-12-17T08:07:23.23374Z","shell.execute_reply.started":"2022-12-17T08:07:23.220873Z","shell.execute_reply":"2022-12-17T08:07:23.232328Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"datablock = DataBlock(blocks = (ImageBlock(), CategoryBlock),\n                     splitter = ColSplitter(),\n                     get_x = ColReader(-2),\n                     get_y = ColReader(6),\n                     item_tfms = Resize(image_resize, ResizeMethod.Pad, pad_mode = 'zeros'),)","metadata":{"execution":{"iopub.status.busy":"2022-12-17T08:07:23.238559Z","iopub.execute_input":"2022-12-17T08:07:23.240803Z","iopub.status.idle":"2022-12-17T08:07:23.250324Z","shell.execute_reply.started":"2022-12-17T08:07:23.240766Z","shell.execute_reply":"2022-12-17T08:07:23.249303Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataloaders = datablock.dataloaders(final_subset, bs = 2*batch_size)\ndataloaders.show_batch()","metadata":{"execution":{"iopub.status.busy":"2022-12-17T08:07:23.255391Z","iopub.execute_input":"2022-12-17T08:07:23.257849Z","iopub.status.idle":"2022-12-17T08:07:33.589403Z","shell.execute_reply.started":"2022-12-17T08:07:23.257811Z","shell.execute_reply":"2022-12-17T08:07:33.588359Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Procedure","metadata":{}},{"cell_type":"code","source":"learn = vision_learner(dataloaders, resnet50, loss_func = CrossEntropyLossFlat(weight = weights),\n                       metrics = [accuracy, pfbeta_torch]).to_fp16()\nlearn.fine_tune(3, cbs=[SaveModelCallback(monitor = 'pfbeta_torch', fname = 'resnet50')])\nlearn.recorder.plot_loss()","metadata":{"execution":{"iopub.status.busy":"2022-12-17T09:12:20.795167Z","iopub.execute_input":"2022-12-17T09:12:20.795707Z","iopub.status.idle":"2022-12-17T10:08:52.384761Z","shell.execute_reply.started":"2022-12-17T09:12:20.795663Z","shell.execute_reply":"2022-12-17T10:08:52.383528Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"interp = ClassificationInterpretation.from_learner(learn)\nlosses,idxs = interp.top_losses()\nlen(dataloaders.valid_ds)==len(losses)==len(idxs)\ninterp.plot_confusion_matrix(figsize=(7,7))","metadata":{"execution":{"iopub.status.busy":"2022-12-17T10:08:52.390139Z","iopub.execute_input":"2022-12-17T10:08:52.392722Z","iopub.status.idle":"2022-12-17T10:14:27.908361Z","shell.execute_reply.started":"2022-12-17T10:08:52.392673Z","shell.execute_reply":"2022-12-17T10:14:27.906906Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"interp.plot_top_losses(9, figsize=(15,10))","metadata":{"execution":{"iopub.status.busy":"2022-12-17T10:14:27.914693Z","iopub.execute_input":"2022-12-17T10:14:27.920599Z","iopub.status.idle":"2022-12-17T10:14:29.86001Z","shell.execute_reply.started":"2022-12-17T10:14:27.920535Z","shell.execute_reply":"2022-12-17T10:14:29.858969Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del learn\ntorch.cuda.empty_cache()\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-12-17T10:16:15.829995Z","iopub.execute_input":"2022-12-17T10:16:15.830539Z","iopub.status.idle":"2022-12-17T10:16:16.142969Z","shell.execute_reply.started":"2022-12-17T10:16:15.830492Z","shell.execute_reply":"2022-12-17T10:16:16.141929Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"learn = vision_learner(dataloaders, resnet50, loss_func = CrossEntropyLossFlat(weight = weights),\n                       metrics = [accuracy, pfbeta_torch]).to_fp16()\nlearn.load('/kaggle/working/models/resnet50')","metadata":{"execution":{"iopub.status.busy":"2022-12-17T10:16:17.560049Z","iopub.execute_input":"2022-12-17T10:16:17.560548Z","iopub.status.idle":"2022-12-17T10:16:18.669288Z","shell.execute_reply.started":"2022-12-17T10:16:17.560503Z","shell.execute_reply":"2022-12-17T10:16:18.668209Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"learn.unfreeze()\nlearn.lr_find()","metadata":{"execution":{"iopub.status.busy":"2022-12-17T10:16:27.354527Z","iopub.execute_input":"2022-12-17T10:16:27.354981Z","iopub.status.idle":"2022-12-17T10:19:46.399475Z","shell.execute_reply.started":"2022-12-17T10:16:27.354941Z","shell.execute_reply":"2022-12-17T10:19:46.39831Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"learn.fit_one_cycle(n_epochs, lr_unfreeze, wd = 0.1,\n                    cbs=[SaveModelCallback(monitor = 'pfbeta_torch', fname = 'resnet50_unfreeze'), \n                        EarlyStoppingCallback(monitor='pfbeta_torch', patience = 4)]) \nlearn.recorder.plot_loss()","metadata":{"execution":{"iopub.status.busy":"2022-12-17T10:21:36.57956Z","iopub.execute_input":"2022-12-17T10:21:36.580017Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"interp = ClassificationInterpretation.from_learner(learn)\nlosses,idxs = interp.top_losses()\nlen(dataloaders.valid_ds)==len(losses)==len(idxs)\ninterp.plot_confusion_matrix(figsize=(7,7))","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"interp.plot_top_losses(9, figsize=(15,10))","metadata":{},"outputs":[],"execution_count":null}]}