{"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":"markdown","source":"<a id=\"table\"></a>\n<h1 style=\"background-color:pink;font-family:newtimeroman;font-size:350%;text-align:center;border-radius: 15px 50px;\">Table of Content</h1>\n\n* [1. IMPORTING LIBRARIES](#1)\n\n* [2. CONFIG](#2)    \n\n* [3. MONITORING SETTINGS](#3)\n\n* [4. DATA LOADING](#4)\n  \n* [5. TRAIN/VALIDATION](#5)\n\n* [6. TRANSFORMS & DATA GENERATOR](#6)\n\n* [7. EXAMPLE IMAGE SAMPLES](#7)\n\n* [8. MODEL](#8)\n\n* [9. CLASS TRAINER](#9)\n\n* [10. FIT](#10)\n\n* [11. PLOT PROBABILISTICS F1/LOSS](#11)","metadata":{}},{"cell_type":"markdown","source":"<a id=\"1\"></a>\n# <p style=\"padding:10px;background-color:lightpink;margin:0;color:black;font-family:newtimeroman;font-size:100%;text-align:center;border-radius: 15px 50px;overflow:hidden;font-weight:500\">Importing Libraries</p>","metadata":{}},{"cell_type":"code","source":"import sys\n\nsys.path.append('../input/timm-pytorch-image-models/pytorch-image-models-master')\n! pip install ../input/einops-030/einops-0.3.0-py2.py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:54:57.739211Z","iopub.execute_input":"2023-01-04T15:54:57.739586Z","iopub.status.idle":"2023-01-04T15:55:08.977802Z","shell.execute_reply.started":"2023-01-04T15:54:57.739554Z","shell.execute_reply":"2023-01-04T15:55:08.976611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! pip install timm\n! pip install einops\n\n! pip install python-gdcm -q\n! pip install pylibjpeg -q\n\n! pip install efficientnet_pytorch -q","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:56:38.795142Z","iopub.status.idle":"2023-01-04T15:56:38.796022Z","shell.execute_reply.started":"2023-01-04T15:56:38.795757Z","shell.execute_reply":"2023-01-04T15:56:38.795781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\nimport os\nimport glob\nimport timm\nimport torch\nimport wandb\nimport pydicom\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport torch.nn as nn\nfrom tqdm import tqdm\nimport albumentations\nimport albumentations as A\nfrom einops import rearrange\nfrom torchvision import transforms\nfrom joblib import Parallel, delayed\nfrom matplotlib import pyplot as plt\nfrom kaggle_secrets import UserSecretsClient\nfrom mpl_toolkits.axes_grid1 import ImageGrid\nfrom albumentations.pytorch import ToTensorV2\nfrom efficientnet_pytorch import EfficientNet\nfrom torch.utils.data import Dataset, DataLoader\nfrom accelerate import Accelerator, notebook_launcher\n\nfrom albumentations.augmentations.dropout.coarse_dropout import CoarseDropout ","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:55:48.182476Z","iopub.execute_input":"2023-01-04T15:55:48.183733Z","iopub.status.idle":"2023-01-04T15:55:54.605958Z","shell.execute_reply.started":"2023-01-04T15:55:48.183686Z","shell.execute_reply":"2023-01-04T15:55:54.604896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"2\"></a>\n# <p style=\"padding:10px;background-color:lightpink;margin:0;color:black;font-family:newtimeroman;font-size:100%;text-align:center;border-radius: 15px 50px;overflow:hidden;font-weight:500\">Config</p>","metadata":{}},{"cell_type":"code","source":"class CFG:\n    class data:\n        fold=0\n        batch_size=32\n        image_size=(224, 224)\n        path_to_train=\"../input/split-folds-rsna/5_folds_data.csv\"\n        path_to_dcms=\"../input/rsna-breast-cancer-detection/train_images\"\n        path_to_train_images=\"/kaggle/input/rsna-bcd-roi-1024x-png-dataset/train_images\"\n        path_to_dcm_images = \"/kaggle/input/rsna-breast-cancer-detection/train_images/*/*.dcm\"\n        \n    class monitoring:\n        accelerator=None\n        \n    class model:\n        pretrained_name='vit_base_patch16_224'\n        learning_rate=1e-4\n        auto_break_n=3\n        n_epochs=6\n        n_labels=1","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:55:54.608971Z","iopub.execute_input":"2023-01-04T15:55:54.60978Z","iopub.status.idle":"2023-01-04T15:55:54.616316Z","shell.execute_reply.started":"2023-01-04T15:55:54.609738Z","shell.execute_reply":"2023-01-04T15:55:54.614968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"3\"></a>\n# <p style=\"padding:10px;background-color:lightpink;margin:0;color:black;font-family:newtimeroman;font-size:100%;text-align:center;border-radius: 15px 50px;overflow:hidden;font-weight:500\">Monitoring Settings</p>","metadata":{}},{"cell_type":"code","source":"user_secrets = UserSecretsClient()\n\nwb_key = user_secrets.get_secret(\"WANDB_API_KEY\")\nwandb.login(key=wb_key)\n\nCFG.monitoring.accelerator = Accelerator(mixed_precision='fp16', log_with='wandb')\nCFG.monitoring.accelerator = Accelerator()\n\nCFG.monitoring.accelerator.init_trackers(\"rsna_mammography_pytorch\")","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:55:54.617879Z","iopub.execute_input":"2023-01-04T15:55:54.618604Z","iopub.status.idle":"2023-01-04T15:55:56.055927Z","shell.execute_reply.started":"2023-01-04T15:55:54.618565Z","shell.execute_reply":"2023-01-04T15:55:56.054847Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"4\"></a>\n# <p style=\"padding:10px;background-color:lightpink;margin:0;color:black;font-family:newtimeroman;font-size:100%;text-align:center;border-radius: 15px 50px;overflow:hidden;font-weight:500\">Data Loading</p>","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(CFG.data.path_to_train)\ntrain_df['image_path'] = f'{CFG.data.path_to_train_images}'\\\n                    + '/' + train_df.patient_id.astype(str)\\\n                    + '/' + train_df.image_id.astype(str)\\\n                    + '.png'\n\nprint(f\"train.shape = {train_df.shape}\")\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:55:56.05803Z","iopub.execute_input":"2023-01-04T15:55:56.058658Z","iopub.status.idle":"2023-01-04T15:55:56.302006Z","shell.execute_reply.started":"2023-01-04T15:55:56.058617Z","shell.execute_reply":"2023-01-04T15:55:56.300839Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Label Preprocessing (Encoding) \n","metadata":{}},{"cell_type":"code","source":"from sklearn.preprocessing import LabelEncoder, normalize\n\n# Keep only columns in test + target variable\n\ntrain_df = train_df[[\"patient_id\", \"image_id\", \"laterality\", \"view\", \"age\", \"implant\", \"image_path\", \"cancer\", \"fold\"]]\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:55:56.304317Z","iopub.execute_input":"2023-01-04T15:55:56.305513Z","iopub.status.idle":"2023-01-04T15:55:56.331216Z","shell.execute_reply.started":"2023-01-04T15:55:56.305461Z","shell.execute_reply":"2023-01-04T15:55:56.330068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Encode categorical variables\n\nlabel_laterality = LabelEncoder()\nlabel_view = LabelEncoder()\n\ntrain_df.isnull().sum()[train_df.isnull().sum()  > 0]","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:55:56.332605Z","iopub.execute_input":"2023-01-04T15:55:56.332999Z","iopub.status.idle":"2023-01-04T15:55:56.360382Z","shell.execute_reply.started":"2023-01-04T15:55:56.332962Z","shell.execute_reply":"2023-01-04T15:55:56.359198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['age'] = train_df['age'].fillna(58)\n\ntrain_df[\"laterality\"] = label_laterality.fit_transform(train_df.laterality)\ntrain_df[\"view\"] = label_view.fit_transform(train_df.view)\n\ntrain_df.to_csv(\"train_path.csv\", index=False)\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:55:56.36288Z","iopub.execute_input":"2023-01-04T15:55:56.36406Z","iopub.status.idle":"2023-01-04T15:55:56.601331Z","shell.execute_reply.started":"2023-01-04T15:55:56.364018Z","shell.execute_reply":"2023-01-04T15:55:56.600182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"5\"></a>\n# <p style=\"padding:10px;background-color:lightpink;margin:0;color:black;font-family:newtimeroman;font-size:100%;text-align:center;border-radius: 15px 50px;overflow:hidden;font-weight:500\">Train/Validatio</p>","metadata":{}},{"cell_type":"code","source":"train = train_df.query(f'fold != {CFG.data.fold}').reset_index(drop=True)\nvalid = train_df.query(f'fold == {CFG.data.fold}').reset_index(drop=True)\n\ntrain.shape, valid.shape","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:55:56.607103Z","iopub.execute_input":"2023-01-04T15:55:56.608069Z","iopub.status.idle":"2023-01-04T15:55:56.634892Z","shell.execute_reply.started":"2023-01-04T15:55:56.608023Z","shell.execute_reply":"2023-01-04T15:55:56.633739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"6\"></a>\n# <p style=\"padding:10px;background-color:lightpink;margin:0;color:black;font-family:newtimeroman;font-size:100%;text-align:center;border-radius: 15px 50px;overflow:hidden;font-weight:500\">Transforms & Data Generator</p>","metadata":{}},{"cell_type":"code","source":"class RSNAData(Dataset):\n    def __init__(self, df, img_folder, transform=None, is_test=False):\n        self.df = df\n        self.is_test = is_test\n        self.transform = transform\n        self.img_folder = img_folder\n        \n        self.csv_columns = ['laterality', 'view', 'age', 'implant']\n\n    def __getitem__(self, idx):\n        img_path = self.df['image_path'][idx]\n        img = cv2.imread(img_path)\n\n        if self.transform:\n            img = self.transform(image=img)[\"image\"]\n        \n#         img = rearrange(img, 'h w c -> c h w')\n        img = torch.tensor(img, dtype=torch.float)\n        \n        meta_data = np.array(\n            self.df.iloc[idx][self.csv_columns].values, \n            dtype=np.float32\n        )\n            \n        if not self.is_test:\n            target = self.df['cancer'][idx]\n            target = torch.tensor(target, dtype=torch.float)\n            \n            return {\n                \"image\": img,\n                \"meta\": meta_data,\n                \"target\": target,\n            }\n        \n        return {\n            \"image\": img,\n            \"meta\": meta_data,\n        }\n    \n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:55:56.636659Z","iopub.execute_input":"2023-01-04T15:55:56.637267Z","iopub.status.idle":"2023-01-04T15:55:56.649472Z","shell.execute_reply.started":"2023-01-04T15:55:56.637229Z","shell.execute_reply":"2023-01-04T15:55:56.648293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def augmantation(image_size, is_train=True):\n    vertical_flip = 0.5\n    horizontal_flip = 0.5\n        \n    if is_train:\n        transforms = [\n#             A.RandomResizedCrop(height=224, width=224),\n            A.Resize(height=224, width=224),  \n            A.ShiftScaleRotate(rotate_limit=90, scale_limit = [0.8, 1.2]),\n            A.HorizontalFlip(p = horizontal_flip),\n            A.VerticalFlip(p = vertical_flip)          \n        ]\n        \n    else:\n        transforms = [\n            A.Resize(height=224, width=224),  \n        ]\n        \n    transforms.extend([\n        ToTensorV2()\n    ])\n    \n    manipulation =  A.Compose(transforms, p=1)\n    \n    return manipulation","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:55:56.651104Z","iopub.execute_input":"2023-01-04T15:55:56.65157Z","iopub.status.idle":"2023-01-04T15:55:56.662673Z","shell.execute_reply.started":"2023-01-04T15:55:56.65153Z","shell.execute_reply":"2023-01-04T15:55:56.661663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = RSNAData(df=train, img_folder=CFG.data.path_to_train_images, transform=augmantation(CFG.data.image_size, True))\ntrain_loader = DataLoader(train_dataset, batch_size=CFG.data.batch_size, shuffle=True)\n\nvalid_dataset = RSNAData(df=valid, img_folder=CFG.data.path_to_train_images, transform=augmantation(CFG.data.image_size, False))\nvalid_loader = DataLoader(valid_dataset, batch_size=CFG.data.batch_size, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:55:56.664075Z","iopub.execute_input":"2023-01-04T15:55:56.664659Z","iopub.status.idle":"2023-01-04T15:55:56.674207Z","shell.execute_reply.started":"2023-01-04T15:55:56.664622Z","shell.execute_reply":"2023-01-04T15:55:56.673266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"next(iter(train_loader)).keys()","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:55:56.67548Z","iopub.execute_input":"2023-01-04T15:55:56.676045Z","iopub.status.idle":"2023-01-04T15:55:57.425859Z","shell.execute_reply.started":"2023-01-04T15:55:56.676007Z","shell.execute_reply":"2023-01-04T15:55:57.424743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"7\"></a>\n# <p style=\"padding:10px;background-color:lightpink;margin:0;color:black;font-family:newtimeroman;font-size:100%;text-align:center;border-radius: 15px 50px;overflow:hidden;font-weight:500\">Example image samples</p>","metadata":{}},{"cell_type":"code","source":"batch_sample = next(iter(train_loader))\nimg = batch_sample[\"image\"][0]\nimg.size()","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:55:57.428333Z","iopub.execute_input":"2023-01-04T15:55:57.429069Z","iopub.status.idle":"2023-01-04T15:55:58.183315Z","shell.execute_reply.started":"2023-01-04T15:55:57.429027Z","shell.execute_reply":"2023-01-04T15:55:58.182103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure(figsize=(30, 20))\ngrid = ImageGrid(fig, 111,\n                 nrows_ncols=(4, 8),\n                 axes_pad=0.25\n)\n\nfor ax, img in zip(grid, batch_sample[\"image\"]):\n    ax.imshow(img.type(torch.int).permute(1, 2, 0))\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:55:58.18501Z","iopub.execute_input":"2023-01-04T15:55:58.185467Z","iopub.status.idle":"2023-01-04T15:56:02.49879Z","shell.execute_reply.started":"2023-01-04T15:55:58.185423Z","shell.execute_reply":"2023-01-04T15:56:02.497345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"8\"></a>\n# <p style=\"padding:10px;background-color:lightpink;margin:0;color:black;font-family:newtimeroman;font-size:100%;text-align:center;border-radius: 15px 50px;overflow:hidden;font-weight:500\">Model</p>","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(\n            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        \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 = nn.functional.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":{"iopub.status.busy":"2023-01-04T15:56:38.797489Z","iopub.status.idle":"2023-01-04T15:56:38.798366Z","shell.execute_reply.started":"2023-01-04T15:56:38.798078Z","shell.execute_reply":"2023-01-04T15:56:38.798102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"output_size = 1\ncsv_columns = ['laterality', 'view', 'age', 'implant']\nno_columns = len(csv_columns)","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:57:43.37225Z","iopub.execute_input":"2023-01-04T15:57:43.372678Z","iopub.status.idle":"2023-01-04T15:57:43.377958Z","shell.execute_reply.started":"2023-01-04T15:57:43.372644Z","shell.execute_reply":"2023-01-04T15:57:43.376884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with CFG.monitoring.accelerator.main_process_first():\n    model = EffNetNetwork(output_size=output_size, no_columns=no_columns)","metadata":{"execution":{"iopub.status.busy":"2023-01-04T16:00:47.173152Z","iopub.execute_input":"2023-01-04T16:00:47.173763Z","iopub.status.idle":"2023-01-04T16:00:53.130082Z","shell.execute_reply.started":"2023-01-04T16:00:47.173725Z","shell.execute_reply":"2023-01-04T16:00:53.128785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"out = model(batch_sample[\"image\"], batch_sample[\"meta\"], prints=True)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = torch.optim.Adam(params=model.parameters(), lr=CFG.model.learning_rate)\ncriterion = nn.BCEWithLogitsLoss()\n\nmodel, optimizer, train_loader, valid_loader = CFG.monitoring.accelerator.prepare(\n    model, optimizer, train_loader, valid_loader\n)","metadata":{"execution":{"iopub.status.busy":"2023-01-04T16:04:03.111453Z","iopub.execute_input":"2023-01-04T16:04:03.111856Z","iopub.status.idle":"2023-01-04T16:04:03.334875Z","shell.execute_reply.started":"2023-01-04T16:04:03.111824Z","shell.execute_reply":"2023-01-04T16:04:03.333532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"9\"></a>\n# <p style=\"padding:10px;background-color:lightpink;margin:0;color:black;font-family:newtimeroman;font-size:100%;text-align:center;border-radius: 15px 50px;overflow:hidden;font-weight:500\">Class Trainer</p>","metadata":{}},{"cell_type":"code","source":"class Trainer:\n    def __init__(self, model, accelerator, criterion, optimizer, epochs, auto_break_n, monitoring=True):\n        self.model = model\n        self.accelerator = accelerator\n        self.criterion = criterion\n        self.optimizer = optimizer\n        self.epochs = epochs\n        \n        self.history_train_loss = []\n        self.history_train_score = []\n        \n        self.history_val_loss = []\n        self.history_val_score = []\n        \n        self.auto_break_n = auto_break_n\n        self.monitoring = monitoring\n        \n        self.best_val_score = 0        \n        self.no_improvement_epoch = 0\n\n\n    def fit(self, train_loader, valid_loader):\n        \n        for epoch in range(self.epochs):\n            self.accelerator.print(f\"{'=' * 20} Epoch: {epoch + 1} {'=' * 20}\")\n\n            self.model.train()\n            \n            avg_loss = 0\n            all_outputs, all_targets = [], []\n            \n            pbar_train = tqdm(train_loader)\n            for index, batch in enumerate(pbar_train, start=1):\n\n                images, meta, targets = batch[\"image\"], batch[\"meta\"], batch[\"target\"]\n                \n                # forward()\n                outputs = self.model(images, meta).view(-1)\n                loss = self.criterion(outputs, targets)\n                    \n                avg_loss += loss.item()\n                all_outputs.extend(torch.sigmoid(outputs).cpu().detach().tolist())\n                all_targets.extend(targets.cpu().detach().tolist())\n                \n                pbar_train.set_postfix({\n                    'batch': index,\n                    'train_loss': loss.item(),\n                })\n                \n                self.accelerator.backward(loss)\n                self.optimizer.step()\n                self.optimizer.zero_grad(set_to_none=True)\n\n                \n            avg_loss /= len(train_loader)\n            prob_f1_score = self.probabilistic_f1(all_targets, all_outputs, beta=0.5)\n            \n            self.history_train_loss.append(avg_loss)\n            self.history_train_score.append(prob_f1_score)\n\n            self.accelerator.print(f\"\\nEpoch: {epoch + 1} / {self.epochs}  |  Training Loss: {avg_loss:.4f}  |  F1 Score: {prob_f1_score:.4f}\\n\")\n            \n            if self.monitoring:\n                self.accelerator.log({\n                    'train_loss': avg_loss,\n                    'train_prob_f1': prob_f1_score,\n                }\n            )\n                \n            if valid_loader:\n                prob_f1_score, avg_loss = self.valid_score(valid_loader)\n                \n                self.history_val_loss.append(avg_loss)\n                self.history_val_score.append(prob_f1_score)\n                \n                self.accelerator.print(f\"\\nEpoch: {epoch + 1} / {self.epochs}  |  Validation Loss: {avg_loss:.4f}  |  F1 Score: {prob_f1_score:.4f}\\n\")\n                \n                if self.monitoring:\n                    self.accelerator.log({\n                        'valid_loss': avg_loss,\n                        'valid_prob_f1': prob_f1_score,\n                    }\n                )\n                \n                \n                if prob_f1_score > self.best_val_score:\n                    self.no_improvement_epoch = 0\n                    self.best_val_score = prob_f1_score\n                    \n                    # Save the model\n                    self.accelerator.wait_for_everyone() \n                    model = self.accelerator.unwrap_model(self.model)\n                    self.accelerator.save({\n                        \"model\": self.model.state_dict(),\n                        \"optimizer\": self.optimizer.optimizer.state_dict()\n                    }, f\"epoch_{epoch + 1}_model.pth\")\n                    \n                else:  \n                    self.no_improvement_epoch += 1\n                  \n                self.accelerator.print(f\"no improvement iter = {self.no_improvement_epoch}\")\n\n                if self.no_improvement_epoch == self.auto_break_n:\n                    self.accelerator.print('Auto_break !!!')\n\n                    if self.monitoring:\n                        wandb.finish()\n                    break\n            \n    def valid_score(self, valid_loader):\n        self.model.eval()  # switch for some specific layers/parts\n        \n        avg_loss = 0\n        all_outputs, all_targets = [], []\n        \n        with torch.no_grad():\n            \n            pbar_valid = tqdm(valid_loader)\n            for index, batch in enumerate(pbar_valid, start=1):\n\n                images, meta, targets = batch[\"image\"], batch[\"meta\"], batch[\"target\"]\n                \n                outputs = self.model(images, meta).view(-1)\n                loss = self.criterion(outputs, targets)\n                    \n                avg_loss += loss.item()\n                outputs, targets = self.accelerator.gather_for_metrics((\n                    outputs, targets\n                ))\n                \n                all_outputs.extend(torch.sigmoid(outputs).cpu().detach().tolist())\n                all_targets.extend(targets.cpu().detach().tolist())\n                \n                pbar_valid.set_postfix({\n                    'batch': index,\n                    'valid_loss': loss.item(),\n                })\n                \n        prob_f1_score = self.probabilistic_f1(all_targets, all_outputs, beta=0.5)\n        avg_loss /= len(valid_loader)\n        \n        return prob_f1_score, avg_loss\n    \n    \n    @staticmethod\n    def probabilistic_f1(labels, predictions, beta=0.5):\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                cfp += 1 - prediction\n            else:\n                cfp += prediction\n\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","metadata":{"execution":{"iopub.status.busy":"2023-01-04T16:04:06.50651Z","iopub.execute_input":"2023-01-04T16:04:06.506916Z","iopub.status.idle":"2023-01-04T16:04:06.533289Z","shell.execute_reply.started":"2023-01-04T16:04:06.506883Z","shell.execute_reply":"2023-01-04T16:04:06.531729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"10\"></a>\n# <p style=\"padding:10px;background-color:lightpink;margin:0;color:black;font-family:newtimeroman;font-size:100%;text-align:center;border-radius: 15px 50px;overflow:hidden;font-weight:500\">Fit</p>","metadata":{}},{"cell_type":"code","source":"my_model = Trainer(model, CFG.monitoring.accelerator, criterion, optimizer, CFG.model.n_epochs, CFG.model.auto_break_n)\nnotebook_launcher(my_model.fit, args=(train_loader, valid_loader), num_processes=2)","metadata":{"execution":{"iopub.status.busy":"2023-01-04T16:04:07.980908Z","iopub.execute_input":"2023-01-04T16:04:07.981321Z","iopub.status.idle":"2023-01-04T16:05:24.52558Z","shell.execute_reply.started":"2023-01-04T16:04:07.981285Z","shell.execute_reply":"2023-01-04T16:05:24.523868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"10\"></a>\n# <p style=\"padding:10px;background-color:lightpink;margin:0;color:black;font-family:newtimeroman;font-size:100%;text-align:center;border-radius: 15px 50px;overflow:hidden;font-weight:500\">Plot probabilistic F1/Loss</p>","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(1, 2, figsize=(30, 10))\n\nax[0].plot(my_model.history_train_score, '-o')\nax[0].plot(my_model.history_val_score, '-o')\n\nax[1].plot(my_model.history_train_loss, '-o')\nax[1].plot(my_model.history_val_loss, '-o')\n\nax[0].set_xlabel(\"epoch\")\nax[1].set_xlabel(\"epoch\")\n\nax[0].set_ylabel(\"probabilistic f1\")\nax[1].set_ylabel(\"loss\")\n\nax[0].legend(['Train','Valid'])\nax[0].set_title('Train vs Valid')\nax[1].legend(['Train','Valid'])\nax[1].set_title('Train vs Valid')\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-01-04T15:56:38.792778Z","iopub.status.idle":"2023-01-04T15:56:38.793142Z","shell.execute_reply.started":"2023-01-04T15:56:38.792959Z","shell.execute_reply":"2023-01-04T15:56:38.792976Z"},"trusted":true},"execution_count":null,"outputs":[]}]}