{"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":"!pip install -qU python-gdcm pydicom pylibjpeg","metadata":{"_uuid":"d4878292-b81e-43bb-a4ad-9a8f256e1b77","_cell_guid":"90519183-cb4d-43b8-bc8d-61c8ff5dc08b","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:05:05.603076Z","iopub.execute_input":"2023-02-27T06:05:05.603531Z","iopub.status.idle":"2023-02-27T06:05:19.449356Z","shell.execute_reply.started":"2023-02-27T06:05:05.60343Z","shell.execute_reply":"2023-02-27T06:05:19.448155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Cancer Image Classification\n\nReferences\n[Pytorch:aux target](https://www.kaggle.com/code/vslaykovsky/train-pytorch-aux-targets-weighted-loss-thres?scriptVersionId=114303403)","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append('../input/timm-pytorch-image-models/pytorch-image-models-master')\nimport timm\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport os \nimport random\nimport gdcm\nimport pydicom\nfrom pydicom.pixel_data_handlers import apply_windowing\nfrom tqdm.auto import tqdm\nfrom skimage.io import imsave\nfrom skimage.io import imsave\nimport cv2\nimport PIL\nfrom PIL import Image\nimport glob\nfrom pathlib import Path\nfrom joblib import Parallel,delayed\n\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom sklearn.preprocessing import LabelEncoder,normalize\n\nimport torch\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn.utils import clip_grad_norm_\nimport torchvision\nfrom torchvision import models\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport torchvision.transforms as transforms\nfrom torch.optim import Adam\nfrom torch.utils.data import WeightedRandomSampler\nfrom torchvision.io import read_image\n\nfrom logging import getLogger,INFO,FileHandler,Formatter,StreamHandler","metadata":{"_uuid":"ec9ae3d7-858b-4bb5-96b7-cb0f0f7691ce","_cell_guid":"cb105830-0f83-4871-aa4d-7a6517587d5c","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:05:19.452518Z","iopub.execute_input":"2023-02-27T06:05:19.45331Z","iopub.status.idle":"2023-02-27T06:05:24.553587Z","shell.execute_reply.started":"2023-02-27T06:05:19.453267Z","shell.execute_reply":"2023-02-27T06:05:24.552547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"timm.__version__","metadata":{"_uuid":"531ebc87-e367-46f8-99c1-b4524ec77c75","_cell_guid":"13ac05e2-1a81-4b9a-a568-2a8568cd4c86","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:05:24.554979Z","iopub.execute_input":"2023-02-27T06:05:24.555333Z","iopub.status.idle":"2023-02-27T06:05:24.566883Z","shell.execute_reply.started":"2023-02-27T06:05:24.555299Z","shell.execute_reply":"2023-02-27T06:05:24.565978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TARGET_FEATURES = ['site_id', 'patient_id', 'laterality', 'view', 'age',\n        'biopsy', 'invasive', 'BIRADS', 'implant', 'density',\n       'machine_id', 'difficult_negative_case', ]\nTARGET_LABEL = 'cancer'\nFEATURES = [TARGET_LABEL] + TARGET_FEATURES","metadata":{"_uuid":"bd44d9c3-2e16-40c7-8ef5-78b237ca9e53","_cell_guid":"ea8adda2-38ab-4717-b3fc-48275e2db7c7","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:05:24.570758Z","iopub.execute_input":"2023-02-27T06:05:24.571339Z","iopub.status.idle":"2023-02-27T06:05:24.578434Z","shell.execute_reply.started":"2023-02-27T06:05:24.571311Z","shell.execute_reply":"2023-02-27T06:05:24.577519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#project configurations\nclass Config:\n    def __init__(self):\n        self.lr = 3e-4\n        self.batch_size = 32\n        self.num_epochs = 10\n        self.model_name=\"efficientnet_b4\"\n        self.folds = 4\n        self.num_classes = 4\n        self.SEED = 42\n        self.N_SAMPLES = 500\n        self.wd = 1e-6\n        self.num_workers = 4\n        \ncfg = Config()","metadata":{"_uuid":"1fd68453-5e94-48ce-9510-dc686a9c3efd","_cell_guid":"3722b108-00f1-42fc-ab37-89a1783ab83c","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:05:24.580455Z","iopub.execute_input":"2023-02-27T06:05:24.580741Z","iopub.status.idle":"2023-02-27T06:05:24.587943Z","shell.execute_reply.started":"2023-02-27T06:05:24.580715Z","shell.execute_reply":"2023-02-27T06:05:24.586897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{"_uuid":"f25935b0-0802-4a49-abe6-9e9fc4ec0775","_cell_guid":"a250db13-6f7d-4308-b8b2-48f1dc4932d4","trusted":true}},{"cell_type":"code","source":"def set_seed(seed):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    \n    \ndef dcm_png_conv(dcm_path,train=True,img_size: int = 1024):\n    subfolder = 'train' if train else 'test'\n    #patient = dcm_path.split('/')[-2]\n    image_number = dcm_path.split('/')[-1][:-4]\n    dicom = pydicom.dcmread(dcm_path)\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        data = np.amax(dicom.pixel_array) - dicom.pixel_array\n    else:\n        data = dicom.pixel_array\n    img = apply_windowing(data,dicom)\n    #img = (img.astype(float) - img.min() /(img.max() - img.min()))\n    #img = Image.fromarray((img*255).astype(np.uint8))\n    img = (img - img.min())/(img.max() - img.min())\n    img = cv2.resize(img,(img_size,img_size))\n    img = (img*255).astype(np.uint8)\n    \n    cv2.imwrite(f'output/{subfolder}/{image_number}.png',img)\n    return img\n\ndef get_aug(data):\n    def transforms(img):\n        if data == 'train':\n            train_transform = [\n                torchvision.transforms.RandomHorizontalFlip(0.5),\n                torchvision.transforms.RandomRotation(degrees=(-5, 5)), \n                torchvision.transforms.RandomResizedCrop((1024, 512), scale=(0.8, 1), ratio=(0.45, 0.55)) \n            ]\n            img = torchvision.transforms.Compose(train_transform + [            \n                torchvision.transforms.ToTensor(),\n                torchvision.transforms.Normalize(mean=0.2179, std=0.0529),\n            ])(img)\n            return img\n        elif data == 'valid':\n            test_transform = [\n                torchvision.transforms.RandomHorizontalFlip(0.5),\n                torchvision.transforms.Resize((1024, 512))\n            ]\n            img = torchvision.transforms.Compose(test_transform + [            \n                torchvision.transforms.ToTensor(),\n                torchvision.transforms.Normalize(mean=0.2179, std=0.0529),\n            ])(img)\n            return img\n        return lambda img: transforms(img)","metadata":{"_uuid":"314729a7-fae5-4e3c-87ea-e60eba970a7d","_cell_guid":"35a49822-78d3-4ff3-975e-48d89fa7e84e","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:05:24.58965Z","iopub.execute_input":"2023-02-27T06:05:24.59006Z","iopub.status.idle":"2023-02-27T06:05:24.605905Z","shell.execute_reply.started":"2023-02-27T06:05:24.590027Z","shell.execute_reply":"2023-02-27T06:05:24.604035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#images dataset\ntrain_dir = glob.glob(\"/kaggle/input/rsna-breast-cancer-detection/train_images/*/*.dcm\")\ntest_dir =  glob.glob(\"/kaggle/input/rsna-breast-cancer-detection/test_images/*/*.dcm\")\n\n#csv dataset \ntrain_df = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/train.csv')\ntest_df = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/test.csv')","metadata":{"_uuid":"f1044039-76cb-467e-9c1a-075f161ad301","_cell_guid":"22fe969c-90a9-46d8-9cc3-7d6dfdeb9236","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:05:24.607396Z","iopub.execute_input":"2023-02-27T06:05:24.607835Z","iopub.status.idle":"2023-02-27T06:06:08.099019Z","shell.execute_reply.started":"2023-02-27T06:05:24.607801Z","shell.execute_reply":"2023-02-27T06:06:08.098057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = train_df['cancer'].values\nprint(labels)","metadata":{"_uuid":"3d5de48b-7bbc-4176-aa92-4c23bebf34bd","_cell_guid":"2b6d0525-3919-4158-b360-c7ff1c4a288b","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:06:08.100339Z","iopub.execute_input":"2023-02-27T06:06:08.100931Z","iopub.status.idle":"2023-02-27T06:06:08.111351Z","shell.execute_reply.started":"2023-02-27T06:06:08.100894Z","shell.execute_reply":"2023-02-27T06:06:08.110449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_path = '/kaggle/input/rsna-breast-cancer-detection/train_images'\n# train_path + patient_id + image_id \ndef make_paths(patientId , imageId):\n    return train_path + str(patientId) + \"/\" + str(imageId) +\".dcm\"\n\ntrain_df['dcm_path'] = train_df.apply(lambda x: make_paths(x.patient_id , x.image_id), axis = 1)","metadata":{"execution":{"iopub.status.busy":"2023-02-27T06:09:25.121918Z","iopub.execute_input":"2023-02-27T06:09:25.122291Z","iopub.status.idle":"2023-02-27T06:09:26.222611Z","shell.execute_reply.started":"2023-02-27T06:09:25.122245Z","shell.execute_reply":"2023-02-27T06:09:26.221628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#create new directories to store the converted images\nos.makedirs('output/train/',exist_ok=True)\nos.makedirs('output/test/',exist_ok=True)\nTRAIN_IMAGES_PATH = '/kaggle/input/rsna-1024-images/output/train'\nTEST_IMAGES_PATH = '/kaggle/input/rsna-1024-images/output/test'\ntrain_dir.sort()\ntrain_dir = train_dir[:cfg.N_SAMPLES]\n#create a new column img_path for the conveted png images\ntrain_df['img_path'] = TRAIN_IMAGES_PATH+'/'+train_df.image_id.astype('string')+'.png'\ntest_df['img_path'] = TEST_IMAGES_PATH+'/'+test_df.image_id.astype('string')+'.png'\ntest_df[['patient_id','image_id','img_path']].head()","metadata":{"_uuid":"12a28d96-f483-4c9e-88fb-21da06c095da","_cell_guid":"231e622b-98ec-4966-8d42-6c19810cfbbe","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:06:08.372134Z","iopub.status.idle":"2023-02-27T06:06:08.372467Z","shell.execute_reply.started":"2023-02-27T06:06:08.372307Z","shell.execute_reply":"2023-02-27T06:06:08.372324Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = train_df[:cfg.N_SAMPLES]\nprint([sample[56:-4].split('/') for sample in train_dir][-5:])\ntrain_df.tail(5)","metadata":{"_uuid":"0e44cddf-7a4d-452d-8e60-23876a72faa2","_cell_guid":"dd91b62b-c196-47fa-be0e-71b3714267b5","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:06:08.374311Z","iopub.status.idle":"2023-02-27T06:06:08.375291Z","shell.execute_reply.started":"2023-02-27T06:06:08.37502Z","shell.execute_reply":"2023-02-27T06:06:08.375048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data cleaning","metadata":{}},{"cell_type":"code","source":"train_df.isnull()","metadata":{"_uuid":"92a33a6a-fd94-4cd0-97af-f7051ec8d0b9","_cell_guid":"33725387-92c6-4239-8869-3e9e55ec9b1f","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:06:08.376522Z","iopub.status.idle":"2023-02-27T06:06:08.377493Z","shell.execute_reply.started":"2023-02-27T06:06:08.377231Z","shell.execute_reply":"2023-02-27T06:06:08.377259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.age.fillna(train_df.age.mean(),inplace=True)\ntrain_df['age'] = pd.qcut(train_df.age,10,labels=range(10),retbins=False).astype(int)\ntrain_df[TARGET_FEATURES] = train_df[TARGET_FEATURES].apply(LabelEncoder().fit_transform)\ntrain_df[FEATURES]","metadata":{"_uuid":"9da2148b-ee3a-4d91-845c-6d0713ca02af","_cell_guid":"43acafdf-575f-433f-a323-d96fdf2099bd","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:06:08.378744Z","iopub.status.idle":"2023-02-27T06:06:08.379719Z","shell.execute_reply.started":"2023-02-27T06:06:08.379441Z","shell.execute_reply":"2023-02-27T06:06:08.37947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.groupby(\"cancer\")['cancer'].count()","metadata":{"_uuid":"fc0cce49-693c-4c47-bf79-20b323cad64e","_cell_guid":"313eeb4a-dc71-4148-8aea-ea1e081b7f2f","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:06:08.380937Z","iopub.status.idle":"2023-02-27T06:06:08.381805Z","shell.execute_reply.started":"2023-02-27T06:06:08.381516Z","shell.execute_reply":"2023-02-27T06:06:08.381542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#use GPU for fast convrsion\n_ = Parallel(n_jobs=5)(delayed(dcm_png_conv)(x,train=True)\n                      for x in tqdm(train_dir)                       )\n_ = Parallel(n_jobs=4)(delayed(dcm_png_conv)(x,train=False)\n                       for x in tqdm(test_dir)\n                       )","metadata":{"_uuid":"042c2fc4-6740-4aa5-ad85-aeb9346caf74","_cell_guid":"04881703-a758-4da4-ae17-a25eaf22d8f6","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:06:08.383289Z","iopub.status.idle":"2023-02-27T06:06:08.384113Z","shell.execute_reply.started":"2023-02-27T06:06:08.383866Z","shell.execute_reply":"2023-02-27T06:06:08.383889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#download the converted images\n#!zip -r rsna__png_1024.zip /kaggle/working/\n#!zip -r rsna_test_png_1024.zip ","metadata":{"execution":{"iopub.status.busy":"2023-02-27T06:06:08.385457Z","iopub.status.idle":"2023-02-27T06:06:08.386276Z","shell.execute_reply.started":"2023-02-27T06:06:08.38603Z","shell.execute_reply":"2023-02-27T06:06:08.386053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#from IPython.display import FileLink\n#FileLink(r'rsna__png_1024.zip')","metadata":{"execution":{"iopub.status.busy":"2023-02-27T06:06:08.38764Z","iopub.status.idle":"2023-02-27T06:06:08.38845Z","shell.execute_reply.started":"2023-02-27T06:06:08.388196Z","shell.execute_reply":"2023-02-27T06:06:08.38822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_cols,ds_spls = 5,10\nnp.random.shuffle(train_dir)\nfig,ax_arr = plt.subplots(ncols=ds_cols,nrows=ds_spls // ds_cols,\n                        figsize=(4* ds_cols,5*ds_spls/ds_cols))\nfor i, dicom_path in enumerate(train_dir[:ds_spls]):\n    dicom = pydicom.dcmread(dicom_path)\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        data = np.amax(dicom.pixel_array) - dicom.pixel_array\n    else:\n        data = dicom.pixel_array\n    img = apply_windowing(data,dicom)\n        \n    print(img.shape,img.min(),img.max())\n    img = (img.astype(float) - img.min()) / (img.max() - img.min())\n    ax_arr[i // ds_cols,i % ds_cols].imshow(img,cmap=\"gray\")\n    ax_arr[i // ds_cols,i % ds_cols].set_axis_off()\n    fig.tight_layout()","metadata":{"_uuid":"eb170d7c-c2af-4fe8-bb6f-4518cc19aa53","_cell_guid":"3514ced4-f854-4d2e-8fca-25250193323d","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:06:08.38979Z","iopub.status.idle":"2023-02-27T06:06:08.390615Z","shell.execute_reply.started":"2023-02-27T06:06:08.390352Z","shell.execute_reply":"2023-02-27T06:06:08.390376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class RSNAMammographyDataset(Dataset):\n    \n    def __init__(self,df,path,transform=None,is_train=True):\n        self.df = df\n        self.path = path\n        self.transform = transform\n        self.is_test = is_train\n        \n    def __len__(self):\n        return len(self.df)\n        \n        \n    def __getitem__(self, idx):\n        img_path = self.df['dcm_path'][idx]\n        img = pydicom.dcmread(img_path).pixel_array.astype(np.float32)\n        \n        csv_data = np.array(self.df.ilox[idx][TARGET_FEATURES].values,dtype=np.float32)\n        \n        if self.is_train:\n            return {\"image\"\n                    \"meta\":csv_data,\n                    \"target\":self.df['cancer'][idx]}","metadata":{"execution":{"iopub.status.busy":"2023-02-27T06:06:08.391949Z","iopub.status.idle":"2023-02-27T06:06:08.392759Z","shell.execute_reply.started":"2023-02-27T06:06:08.392504Z","shell.execute_reply":"2023-02-27T06:06:08.392527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def data_to_device(data):\n    \n    image, metadata, targets = data.values()\n    return image.to(DEVICE), metadata.to(DEVICE), targets.to(DEVICE)","metadata":{"execution":{"iopub.status.busy":"2023-02-27T06:06:08.394077Z","iopub.status.idle":"2023-02-27T06:06:08.395051Z","shell.execute_reply.started":"2023-02-27T06:06:08.394804Z","shell.execute_reply":"2023-02-27T06:06:08.394828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Sample data\nsample_df = train_df.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":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pfbeta(labels, predictions, beta=1):\n    y_true_count = 0\n    ctp = 0\n    cfp = 0\n\n    for idx in range(len(labels)):\n        prediction = min(max(predictions[idx], 0), 1)\n        if (labels[idx]):\n            y_true_count += 1\n            ctp += prediction\n        else:\n            cfp += prediction\n\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp)\n    c_recall = ctp / max(y_true_count,1)\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\n    \ndef optimal_f1(labels, predictions):\n    thres = np.linspace(0, 1, 101)\n    f1s = [pfbeta(labels, predictions > thr) for thr in thres]\n    idx = np.argmax(f1s)\n    return f1s[idx], thres[idx]","metadata":{"_uuid":"1b5c55b9-4d93-4a08-b554-2fe84bd6728d","_cell_guid":"39ab004a-5643-49e6-8b06-54af5c1b4743","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:06:08.396174Z","iopub.status.idle":"2023-02-27T06:06:08.397946Z","shell.execute_reply.started":"2023-02-27T06:06:08.39768Z","shell.execute_reply":"2023-02-27T06:06:08.397705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model building","metadata":{}},{"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","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"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":{"_uuid":"e00bb7a9-e5d4-426d-847e-a03de339731b","_cell_guid":"d787c661-d64c-4c45-afda-5ab258e347da","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-02-27T06:06:08.403725Z","iopub.status.idle":"2023-02-27T06:06:08.404493Z","shell.execute_reply.started":"2023-02-27T06:06:08.40424Z","shell.execute_reply":"2023-02-27T06:06:08.404263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_folds(model, train_original):\n    # Creates a .txt file that will contain the logs\n    # logs == what we also print to console\n    f = open(f\"logs_{VERSION}.txt\", \"w+\")\n    \n    # Split in folds\n    group_fold = GroupKFold(n_splits = FOLDS)\n\n    # Generate indices to split data into training and test set.\n    k_folds = group_fold.split(X = np.zeros(len(train_original)), \n                               y = train_original['cancer'], \n                               groups = train_original['patient_id'].tolist())\n    \n    # For each fold\n    for i, (train_index, valid_index) in enumerate(k_folds):\n        \n        print(clr.S+f\"---------- Fold: {i+1} ----------\"+clr.E)\n        add_in_file(f\"---------- Fold: {i+1} ----------\", f)\n        \n        # 🐝 W&B Tracking\n        RUN_CONFIG = CONFIG.copy()\n        params = dict(model=MODEL, \n                      version=VERSION,\n                      fold=i,\n                      epochs=EPOCHS, \n                      batch=BATCH_SIZE1,\n                      lr=LR,\n                      weight_decay=WD)\n        RUN_CONFIG.update(params)\n        #run = wandb.init(project='RSNA_Breast_Cancer', config=RUN_CONFIG)\n\n        w#andb.watch(model, log_freq=100) # 🐝\n\n        # --- Create Instances ---\n        # Best ROC score in this fold\n        best_roc = None\n        # Reset patience before every fold\n        patience_f = PATIENCE\n\n        # Optimizer/ Scheduler/ Criterion\n        optimizer = torch.optim.Adam(model.parameters(), lr = LR, \n                                     weight_decay=WD)\n        scheduler = ReduceLROnPlateau(optimizer=optimizer, mode='max', \n                                      patience=LR_PATIENCE, verbose=True, factor=LR_FACTOR)\n        criterion = nn.BCEWithLogitsLoss()\n\n\n        # --- Read in Data ---\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        # Create Data instances\n        train = RSNADataset(train_data, vertical_flip, horizontal_flip, \n                            is_train=True)\n        valid = RSNADataset(valid_data, vertical_flip, horizontal_flip,\n                            is_train=True)\n\n        # Dataloaders\n        train_loader = DataLoader(train, batch_size=BATCH_SIZE1, \n                                  shuffle=True, num_workers=WORKERS)\n        valid_loader = DataLoader(valid, batch_size=BATCH_SIZE2, \n                                  shuffle=False, num_workers=WORKERS)\n\n\n        # === EPOCHS ===\n        for epoch in range(EPOCHS):\n            start_time = time()\n            correct = 0\n            train_losses = 0\n\n            # === TRAIN ===\n            # Sets the module in training mode.\n            model.train()\n\n            # For each batch\n            for k, data in tqdm(enumerate(train_loader)):\n                # Save them to device\n                image, meta, targets = data_to_device(data)\n\n                # Clear gradients first; very important\n                # usually done BEFORE prediction\n                optimizer.zero_grad()\n\n                # Log Probabilities & Backpropagation\n                out = model(image, meta)\n                loss = criterion(out, targets.unsqueeze(1).float())\n                loss.backward()\n                optimizer.step()\n\n                # --- Save information after this batch ---\n                # Save loss\n                train_losses += loss.item()\n                #wandb.log({\"train_loss\": loss.item()}, step=epoch) # 🐝\n                log.info({\"train_loss\": loss.item()}, step=epoch) # 🐝\n                # From log probabilities to actual probabilities\n                train_preds = torch.round(torch.sigmoid(out)) # 0 and 1\n                # Number of correct predictions\n                correct += (train_preds.cpu() == targets.cpu().unsqueeze(1)).sum().item()\n\n            # Compute Train Accuracy\n            train_acc = correct / len(train_index)\n            #wandb.log({\"train_acc\": train_acc}) # 🐝\n            log.info({\"train_acc\": train_acc}) # 🐝\n\n\n            # === EVAL ===\n            # Sets the model in evaluation mode.\n            model.eval()\n\n            # Create matrix to store evaluation predictions (for accuracy)\n            valid_preds = torch.zeros(size = (len(valid_index), 1), \n                                      device=DEVICE, dtype=torch.float32)\n\n\n            # Disables gradients (we need to be sure no optimization happens)\n            with torch.no_grad():\n                for k, data in tqdm(enumerate(valid_loader)):\n                    # Save them to device\n                    image, meta, targets = data_to_device(data)\n\n                    out = model(image, meta)\n                    pred = torch.sigmoid(out)\n                    valid_preds[k*image.shape[0] : k*image.shape[0] + image.shape[0]] = pred\n\n                # Calculate accuracy\n                valid_acc = accuracy_score(valid_data['cancer'].values, \n                                           torch.round(valid_preds.cpu()))\n                #wandb.log({\"valid_acc\": valid_acc}) # 🐝\n                log.info({\"valid_acc\": valid_acc}) # 🐝\n                # Calculate ROC\n                valid_roc = roc_auc_score(valid_data['cancer'].values, \n                                          valid_preds.cpu())\n                #wandb.log({\"valid_roc\": valid_roc}) # 🐝\n                log.info({\"valid_roc\": valid_roc}) # 🐝\n\n                # Calculate time on Train + Eval\n                duration = str(dtime.timedelta(seconds=time() - start_time))[:7]\n\n\n                # PRINT INFO\n                final_logs = '{} | Epoch: {}/{} | Loss: {:.4} | Acc_tr: {:.3} | Acc_vd: {:.3} | ROC: {:.3}'.\\\n                                format(duration, epoch+1, EPOCHS, \n                                       train_losses, train_acc, valid_acc, valid_roc)\n                add_in_file(final_logs,f)\n                print(final_logs)\n\n\n                # === SAVE MODEL ===\n\n                # Update scheduler (for learning_rate)\n                scheduler.step(valid_roc)\n                # Name the model\n                model_name = f\"Fold{i+1}_Epoch{epoch+1}_ValidAcc{valid_acc:.3f}_ROC{valid_roc:.3f}.pth\"\n\n                # Update best_roc\n                if not best_roc: # If best_roc = None\n                    best_roc = valid_roc\n                    torch.save(model.state_dict(), model_name)\n                    continue\n\n                if valid_roc > best_roc:\n                    best_roc = valid_roc\n                    # Reset patience (because we have improvement)\n                    patience_f = PATIENCE\n                    torch.save(model.state_dict(), model_name)\n                else:\n                    # Decrease patience (no improvement in ROC)\n                    patience_f = patience_f - 1\n                    if patience_f == 0:\n                        stop_logs = 'Early stopping (no improvement since 3 models) | Best ROC: {}'.\\\n                                    format(best_roc)\n                        add_in_file(stop_logs, f)\n                        print(stop_logs)\n                        break\n\n\n        # === CLEANING ===\n        # Clear memory\n        del train, valid, train_loader, valid_loader, image, targets\n        gc.collect()\n        \n        # 🐝 Experiment End for this fold\n        #wandb.finish()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}