{"cells":[{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:38.186639Z","iopub.status.busy":"2021-02-28T08:48:38.185758Z","iopub.status.idle":"2021-02-28T08:48:38.188037Z","shell.execute_reply":"2021-02-28T08:48:38.187379Z"},"papermill":{"duration":0.039097,"end_time":"2021-02-28T08:48:38.188208","exception":false,"start_time":"2021-02-28T08:48:38.149111","status":"completed"},"scrolled":true,"tags":[],"trusted":true},"cell_type":"code","source":"import sys\nsys.path.append('../input/pytorch-image-models/pytorch-image-models-master')\n# sys.setrecursionlimit(10**6)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:38.251626Z","iopub.status.busy":"2021-02-28T08:48:38.25077Z","iopub.status.idle":"2021-02-28T08:48:38.252982Z","shell.execute_reply":"2021-02-28T08:48:38.252346Z"},"papermill":{"duration":0.036951,"end_time":"2021-02-28T08:48:38.253141","exception":false,"start_time":"2021-02-28T08:48:38.21619","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)\n\nTRAIN_PATH = '../input/ranzcr-clip-catheter-line-classification/train'","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:38.319065Z","iopub.status.busy":"2021-02-28T08:48:38.318294Z","iopub.status.idle":"2021-02-28T08:48:42.322258Z","shell.execute_reply":"2021-02-28T08:48:42.320969Z"},"papermill":{"duration":4.041383,"end_time":"2021-02-28T08:48:42.322416","exception":false,"start_time":"2021-02-28T08:48:38.281033","status":"completed"},"scrolled":true,"tags":[],"trusted":true},"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom datetime import datetime\nfrom tqdm.notebook import tqdm\nfrom pprint import pprint\nimport cv2, glob, time, random, os, ast, glob\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport timm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms as T\nfrom torch.utils.data import Dataset, DataLoader\n# https://nvlabs.github.io/iccv2019-mixed-precision-tutorial/files/dusan_stosic_intro_to_mixed_precision_training.pdf\n# https://analyticsindiamag.com/pytorch-mixed-precision-training/\n# https://pytorch.org/docs/stable/notes/amp_examples.html\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau, CosineAnnealingWarmRestarts\nfrom torch.optim import Adam, AdamW, SGD\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom sklearn.metrics import accuracy_score, roc_auc_score\nfrom sklearn.model_selection import StratifiedKFold","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:42.43489Z","iopub.status.busy":"2021-02-28T08:48:42.434212Z","iopub.status.idle":"2021-02-28T08:48:42.437126Z","shell.execute_reply":"2021-02-28T08:48:42.437576Z"},"papermill":{"duration":0.035484,"end_time":"2021-02-28T08:48:42.437726","exception":false,"start_time":"2021-02-28T08:48:42.402242","status":"completed"},"tags":[],"trusted":true},"cell_type":"raw","source":"# ====================================================\n# Library\n# ====================================================\nimport sys\nsys.path.append('../input/pytorch-image-models/pytorch-image-models-master')\n\nimport os\nimport ast\nimport copy\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn import preprocessing\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.utils import check_random_state\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold\n\nfrom tqdm.autonotebook import tqdm\nfrom functools import partial\n\nimport cv2\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD\nimport torchvision.models as models\nfrom torch.nn.parameter import Parameter\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\nfrom torch.cuda.amp import autocast, GradScaler\n\nfrom albumentations import (\n    Compose, OneOf, Normalize, Resize, RandomResizedCrop, RandomCrop, HorizontalFlip, VerticalFlip, \n    RandomBrightness, RandomContrast, RandomBrightnessContrast, Rotate, ShiftScaleRotate, Cutout, \n    IAAAdditiveGaussianNoise, Transpose, HueSaturationValue, CoarseDropout\n    )\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\n\nimport timm\n\n# if CFG.device == 'TPU':\n#     import ignite.distributed as idist\n# elif CFG.device == 'GPU':\n# #     from torch.cuda.amp import autocast, GradScaler\n#     from torch.cuda.amp import autocast, GradScaler\nimport warnings \nwarnings.filterwarnings('ignore')"},{"metadata":{"_kg_hide-output":true,"execution":{"iopub.execute_input":"2021-02-28T08:48:42.497768Z","iopub.status.busy":"2021-02-28T08:48:42.494121Z","iopub.status.idle":"2021-02-28T08:48:42.50813Z","shell.execute_reply":"2021-02-28T08:48:42.507082Z"},"papermill":{"duration":0.048935,"end_time":"2021-02-28T08:48:42.508305","exception":false,"start_time":"2021-02-28T08:48:42.45937","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# get the list of pretrained models\nmodel_names = timm.list_models()\npprint(model_names)","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.023119,"end_time":"2021-02-28T08:48:42.564433","exception":false,"start_time":"2021-02-28T08:48:42.541314","status":"completed"},"tags":[]},"cell_type":"markdown","source":"<a id = \"cont\"></a>\n## CFG"},{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:43.017532Z","iopub.status.busy":"2021-02-28T08:48:43.016433Z","iopub.status.idle":"2021-02-28T08:48:43.024991Z","shell.execute_reply":"2021-02-28T08:48:43.024524Z"},"papermill":{"duration":0.394971,"end_time":"2021-02-28T08:48:43.025183","exception":false,"start_time":"2021-02-28T08:48:42.630212","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"BATCH_SIZE = 8 # 8 for bigger architectures\nVAL_BATCH_SIZE = 16\nEPOCHS = 15 # train upto 10 epochs\nIMG_SIZE = 640 # 384 for bigger architectures\nif BATCH_SIZE == 8:\n    ITER_FREQ = 400\nelse:\n    ITER_FREQ = 200\nNUM_WORKERS = 8\nMEAN = [0.485, 0.456, 0.406]\nSTD = [0.229, 0.224, 0.225]\nSEED = 999\nN_FOLDS = 5\nTR_FOLDS = [0,1,2,3,4]\nSTART_FOLD = 0\n\ntarget_cols=['ETT - Abnormal', 'ETT - Borderline', 'ETT - Normal',\n                 'NGT - Abnormal', 'NGT - Borderline', 'NGT - Incompletely Imaged', 'NGT - Normal', \n                 'CVC - Abnormal', 'CVC - Borderline', 'CVC - Normal',\n                 'Swan Ganz Catheter Present']\n\nMODEL_PATH = None\nMODEL_ARCH = 'resnet200d_320' # tf_efficientnet_b4_ns, tf_efficientnet_b5_ns, resnext50_32x4d\nITERS_TO_ACCUMULATE = 1\n\nLR = 5e-4\nMIN_LR = 1e-6 # SAM, CosineAnnealingWarmRestarts\nWEIGHT_DECAY = 1e-6\nMOMENTUM = 0.9\nT_0 = EPOCHS # SAM, CosineAnnealingWarmRestarts\nMAX_NORM = 1000\nT_MAX = 5 # CosineAnnealingLR\n\nBASE_OPTIMIZER = SGD #for SAM, Ranger\nOPTIMIZER = 'Adam' # Ranger, Adam, AdamP, SGD, SAM\n\nSCHEDULER = 'CosineAnnealingWarmRestarts' # ReduceLROnPlateau, CosineAnnealingLR, CosineAnnealingWarmRestarts, OneCycleLR\nSCHEDULER_UPDATE = 'epoch' # batch\n\nCRITERION = 'BCE' # CrossEntropyLoss, TaylorSmoothedLoss, LabelSmoothedLoss\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:43.115862Z","iopub.status.busy":"2021-02-28T08:48:43.11489Z","iopub.status.idle":"2021-02-28T08:48:43.127067Z","shell.execute_reply":"2021-02-28T08:48:43.128112Z"},"papermill":{"duration":0.055518,"end_time":"2021-02-28T08:48:43.128327","exception":false,"start_time":"2021-02-28T08:48:43.072809","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"class AverageMeter(object):    \n    def __init__(self):\n        self.reset()\n        \n    def reset(self):\n        self.val = 0\n        self.sum = 0\n        self.avg = 0\n        self.count = 0\n        \n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val*n\n        self.count += n\n        self.avg = self.sum / self.count\n\ndef seed_torch(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    \nseed_torch(SEED)\n\ndef macro_multilabel_auc(label, pred):\n    aucs = []\n    for i in range(len(target_cols)):\n        aucs.append(roc_auc_score(label[:, i], pred[:, i]))\n#     print(np.round(aucs, 4))\n    return np.mean(aucs)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:43.205046Z","iopub.status.busy":"2021-02-28T08:48:43.204236Z","iopub.status.idle":"2021-02-28T08:48:43.602453Z","shell.execute_reply":"2021-02-28T08:48:43.60131Z"},"papermill":{"duration":0.436319,"end_time":"2021-02-28T08:48:43.602605","exception":false,"start_time":"2021-02-28T08:48:43.166286","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"TRAIN_DIR = '../input/ranzcr-clip-catheter-line-classification/train/'\ntrain_df = pd.read_csv('../input/ranzcr-clip-catheter-line-classification/train.csv')\nfolds = pd.read_csv('../input/ranzcr-folds/train_folds.csv')\ntrain_annotations = pd.read_csv('../input/ranzcr-clip-catheter-line-classification/train_annotations.csv')","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:43.751363Z","iopub.status.busy":"2021-02-28T08:48:43.750647Z","iopub.status.idle":"2021-02-28T08:48:43.754698Z","shell.execute_reply":"2021-02-28T08:48:43.754203Z"},"papermill":{"duration":0.040033,"end_time":"2021-02-28T08:48:43.754812","exception":false,"start_time":"2021-02-28T08:48:43.714779","status":"completed"},"scrolled":true,"tags":[],"trusted":true},"cell_type":"code","source":"COLOR_MAP = {'ETT - Abnormal': (255, 0, 0),\n             'ETT - Borderline': (0, 255, 0),\n             'ETT - Normal': (0, 0, 255),\n             'NGT - Abnormal': (255, 255, 0),\n             'NGT - Borderline': (255, 0, 255),\n             'NGT - Incompletely Imaged': (0, 255, 255),\n             'NGT - Normal': (128, 0, 0),\n             'CVC - Abnormal': (0, 128, 0),\n             'CVC - Borderline': (0, 0, 128),\n             'CVC - Normal': (128, 128, 0),\n             'Swan Ganz Catheter Present': (128, 0, 128),\n            }\n\n\nclass RanzcrDataset(Dataset):\n    def __init__(self, df, df_annotations, annot_size=50, transform=None):\n        self.df = df\n        self.df_annotations = df_annotations\n        self.annot_size = annot_size\n        self.image_id = df['StudyInstanceUID'].values\n        self.labels = df[target_cols].values\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_name = self.image_id[idx]\n        file_path = f'{TRAIN_DIR}{file_name}.jpg'\n        image = cv2.imread(file_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        query_string = f\"StudyInstanceUID == '{file_name}'\"\n        df = self.df_annotations.query(query_string)\n        for i, row in df.iterrows():\n            label = row[\"label\"]\n            data = np.array(ast.literal_eval(row[\"data\"]))\n            for d in data:\n                image[d[1]-self.annot_size//2:d[1]+self.annot_size//2,\n                      d[0]-self.annot_size//2:d[0]+self.annot_size//2,\n                      :] = COLOR_MAP[label]\n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented['image']\n        label = torch.tensor(self.labels[idx]).float()\n        return image, label","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:43.810044Z","iopub.status.busy":"2021-02-28T08:48:43.808327Z","iopub.status.idle":"2021-02-28T08:48:43.810891Z","shell.execute_reply":"2021-02-28T08:48:43.811367Z"},"papermill":{"duration":0.034126,"end_time":"2021-02-28T08:48:43.811498","exception":false,"start_time":"2021-02-28T08:48:43.777372","status":"completed"},"scrolled":true,"tags":[],"trusted":true},"cell_type":"code","source":"def get_transform(*, train=True):\n    \n    if train:\n        return A.Compose([\n            A.RandomResizedCrop(IMG_SIZE, IMG_SIZE, scale=(0.85, 1.0)),\n            A.HorizontalFlip(p=0.5),\n            A.RandomBrightnessContrast(p=0.2, brightness_limit=(-0.2, 0.2), contrast_limit=(-0.2, 0.2)),\n            A.HueSaturationValue(p=0.2, hue_shift_limit=0.2, sat_shift_limit=0.2, val_shift_limit=0.2),\n            A.ShiftScaleRotate(p=0.2, shift_limit=0.0625, scale_limit=0.2, rotate_limit=20),\n            A.CoarseDropout(p=0.2),\n            A.Cutout(p=0.2, max_h_size=16, max_w_size=16, fill_value=(0., 0., 0.), num_holes=16),\n            A.Normalize(mean=MEAN, std=STD),\n            ToTensorV2(),\n        ])\n    else:\n        return A.Compose([\n#             A.CenterCrop(IMG_SIZE, IMG_SIZE),\n            A.Resize(IMG_SIZE, IMG_SIZE),\n            A.Normalize(mean=MEAN, std=STD, max_pixel_value=255.0, p=1.0),\n            ToTensorV2(),\n        ])","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.0224,"end_time":"2021-02-28T08:48:43.856345","exception":false,"start_time":"2021-02-28T08:48:43.833945","status":"completed"},"tags":[]},"cell_type":"markdown","source":"## Model"},{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:43.914334Z","iopub.status.busy":"2021-02-28T08:48:43.909097Z","iopub.status.idle":"2021-02-28T08:48:43.927184Z","shell.execute_reply":"2021-02-28T08:48:43.926717Z"},"papermill":{"duration":0.048342,"end_time":"2021-02-28T08:48:43.927298","exception":false,"start_time":"2021-02-28T08:48:43.878956","status":"completed"},"scrolled":true,"tags":[],"trusted":true},"cell_type":"code","source":"class ResNet200D(nn.Module):\n    def __init__(self, model_arch, out_dim, pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_arch, pretrained=False)\n        if pretrained:\n            pretrained_path = '../input/startingpointschestx/resnet200d_320_chestx.pth'\n            self.model.load_state_dict(torch.load(pretrained_path)['model'])\n        n_features = self.model.fc.in_features\n        self.model.global_pool = nn.Identity()\n        self.model.fc = nn.Identity()\n        self.pooling = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Linear(n_features, out_dim)\n\n    def forward(self, x):\n        bs = x.size(0)\n        features = self.model(x)\n        pooled_features = self.pooling(features).view(bs, -1)\n        output = self.fc(pooled_features)\n        return output\n    \nclass CustomResNet200D(nn.Module):\n    def __init__(self, model_arch, n_classes, pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_arch, pretrained=False)\n        n_features = self.model.fc.in_features\n        self.model.fc = nn.Linear(n_features, n_classes)\n        if pretrained:\n#             pretrained_path = '../input/startingpointschestx/resnet200d_320_chestx.pth'\n#             state_dict = dict()\n#             for k, v in torch.load(pretrained_path, map_location='cpu')[\"model\"].items():\n#                 if k[:6] == \"model.\":\n#                     k = k.replace(\"model.\", \"\")\n#                 state_dict[k] = v\n# #             base_model.load_state_dict(state_dict)\n#             self.model.load_state_dict(state_dict)\n#             self.model.reset_classifier(0, '')\n#             print(f'load {model_name} pretrained model')\n            pretrained_path = '../input/startingpointschestx/resnet200d_320_chestx.pth'\n            checkpoint = torch.load(pretrained_path)['model']\n            for key in list(checkpoint.keys()):\n                if 'model.' in key:\n                    checkpoint[key.replace('model.', '')] = checkpoint[key]\n                    del checkpoint[key]\n            self.model.load_state_dict(checkpoint) \n            print(f'load {model_arch} pretrained model')\n        n_features = self.model.fc.in_features\n        self.model.global_pool = nn.Identity()\n        self.model.fc = nn.Identity()\n        self.pooling = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Linear(n_features, n_classes)\n\n    def forward(self, x):\n        bs = x.size(0)\n        features = self.model(x)\n        pooled_features = self.pooling(features).view(bs, -1)\n        output = self.fc(pooled_features)\n        return features, pooled_features, output\n\nclass SeResnet152D(nn.Module): \n    def __init__(self, model_arch, n_classes, pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_arch, pretrained=pretrained)\n        n_features = self.model.fc.in_features\n        self.model.global_pool = nn.Identity()\n        self.model.fc = nn.Identity()\n        self.pooling = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Linear(n_features, n_classes)\n\n    def forward(self, x):\n        bs = x.size(0)\n        features = self.model(x)\n        pooled_features = self.pooling(features).view(bs, -1)\n        output = self.fc(pooled_features)\n        return output\n            \nclass CustomEffNet(nn.Module):\n    def __init__(self, model_arch, n_classes, pretrained=True):\n        super().__init__()\n        self.model = timm.create_model(model_arch, pretrained, n_class)\n        n_features = self.model.classifier.in_features\n        self.model.global_pool = nn.Identity()\n        self.model.classifier = nn.Identity()\n        self.pooling = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Linear(n_features,n_classes)\n        \n    def forward(self,x): \n        bs = x.size(0)\n        features = self.model(x)\n        pooled_features = self.pooling(features).view(bs,-1)\n        output = self.fc(pooled_features)\n        return output \n    \nclass CustomResNext(nn.Module):\n    def __init__(self, model_arch, n_classes, pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_arch, pretrained=pretrained)\n        n_features = self.model.fc.in_features\n        self.model.global_pool = nn.Identity()\n        self.model.fc = nn.Identity()\n        self.pooling = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Linear(n_features, n_classes)\n\n    def forward(self, x):\n        bs = x.size(0)\n        features = self.model(x)\n        pooled_features = self.pooling(features).view(bs, -1)\n        output = self.fc(pooled_features)\n        return output","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.023264,"end_time":"2021-02-28T08:48:43.97332","exception":false,"start_time":"2021-02-28T08:48:43.950056","status":"completed"},"tags":[]},"cell_type":"markdown","source":"[Back to CFG(Click here)](#cont)"},{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:44.034419Z","iopub.status.busy":"2021-02-28T08:48:44.033678Z","iopub.status.idle":"2021-02-28T08:48:44.037124Z","shell.execute_reply":"2021-02-28T08:48:44.03755Z"},"papermill":{"duration":0.040894,"end_time":"2021-02-28T08:48:44.037681","exception":false,"start_time":"2021-02-28T08:48:43.996787","status":"completed"},"scrolled":true,"tags":[],"trusted":true},"cell_type":"code","source":"def GetCriterion(criterion_name, criterion=None):\n#     if criterion_name == 'BiTemperedLoss':\n#         criterion = BiTemperedLogistic()\n#     elif criterion_name == 'SymmetricCrossEntropyLoss':\n#         criterion = SymmetricCrossEntropy()\n    if criterion_name == 'CrossEntropyLoss':\n        criterion = nn.CrossEntropyLoss()\n    elif criterion_name == 'LabelSmoothingLoss':\n        criterion = LabelSmoothingLoss()\n#     elif criterion_name == 'FocalLoss':\n#         criterion = FocalLoss()\n#     elif criterion_name == 'FocalCosineLoss':\n#         criterion = FocalCosineLoss()\n    elif criterion_name == 'TaylorCrossEntropyLoss':\n        criterion = TaylorCrossEntropyLoss()\n    elif criterion_name == 'TaylorSmoothedLoss':\n        criterion = TaylorSmoothedLoss()\n    elif criterion_name == 'CutMix':\n        criterion = CutMixCriterion(criterion)\n    elif criterion_name == 'SnapMix':\n        criterion = SnapMixLoss()\n    elif criterion_name == 'CustomLoss':\n        criterion = CustomLoss(WEIGHTS)\n    elif criterion_name == 'BCE':\n        criterion = nn.BCEWithLogitsLoss()\n    return criterion\n    \n    \ndef GetScheduler(scheduler_name, optimizer, batches=None):\n    #['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts', 'OneCycleLR', 'GradualWarmupSchedulerV2']\n    if scheduler_name == 'OneCycleLR':\n        return torch.optim.lr_scheduler.OneCycleLR(optimizer,max_lr = 1e-2,epochs = CFG.EPOCHS,\n                                                   steps_per_epoch = batches+1,pct_start = 0.1)\n    if scheduler_name == 'CosineAnnealingWarmRestarts':\n        return torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0 = T_0, T_mult=1,\n                                                                    eta_min=MIN_LR, last_epoch=-1)\n    elif scheduler_name == 'CosineAnnealingLR':\n        return torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=T_MAX, eta_min=0, last_epoch=-1)\n    elif scheduler_name == 'ReduceLROnPlateau':\n        return torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, factor=0.1, patience=1, threshold=0.0001,\n                                                          cooldown=0, min_lr=MIN_LR)\n#     elif scheduler_name == 'GradualWarmupSchedulerV2':\n#         return GradualWarmupSchedulerV2(optimizer=optimizer)\n    \ndef GetOptimizer(optimizer_name,parameters):\n    #['Adam','Ranger']\n    if optimizer_name == 'Adam':\n#         if CFG.scheduler_name == 'GradualWarmupSchedulerV2':\n#             return torch.optim.Adam(parameters, lr=CFG.LR_START, weight_decay=CFG.weight_decay, amsgrad=False)\n#         else:\n        return torch.optim.Adam(parameters, lr=LR, weight_decay=WEIGHT_DECAY, amsgrad=False)\n    elif optimizer_name == 'AdamW':\n#         if CFG.scheduler_name == 'GradualWarmupSchedulerV2':\n#             return torch.optim.AdamW(parameters, lr=CFG.LR_START, weight_decay=CFG.weight_decay, amsgrad=False)\n#         else:\n        return torch.optim.Adam(parameters, lr=LR, weight_decay=WEIGHT_DECAY, amsgrad=False)\n    elif optimizer_name == 'AdamP':\n#         if CFG.scheduler_name == 'GradualWarmupSchedulerV2':\n#             return AdamP(parameters, lr=CFG.LR_START, weight_decay=CFG.weight_decay)\n#         else:\n        return AdamP(parameters, lr=LR, weight_decay=WEIGHT_DECAY)\n    elif optimizer_name == 'Ranger':\n        return Ranger(parameters, lr = LR, alpha = 0.5, k = 6, N_sma_threshhold = 5, \n                      betas = (0.95,0.999), weight_decay=WEIGHT_DECAY)\n    elif optimizer_name == 'SAM':\n        return SAM(parameters, BASE_OPTIMIZER, lr=0.1, momentum=0.9,weight_decay=0.0005)\n    \n    elif optimizer_name == 'AdamP':\n        return AdamP(parameters, lr=LR, weight_decay=WEIGHT_DECAY)","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.022577,"end_time":"2021-02-28T08:48:44.083479","exception":false,"start_time":"2021-02-28T08:48:44.060902","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Train and validation functions"},{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:44.18668Z","iopub.status.busy":"2021-02-28T08:48:44.185933Z","iopub.status.idle":"2021-02-28T08:48:44.189726Z","shell.execute_reply":"2021-02-28T08:48:44.189207Z"},"papermill":{"duration":0.038071,"end_time":"2021-02-28T08:48:44.189839","exception":false,"start_time":"2021-02-28T08:48:44.151768","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"def train_fn(model, dataloader, device, epoch, optimizer, criterion, scheduler):\n    \n    data_time = AverageMeter()\n    batch_time = AverageMeter()\n    losses = AverageMeter()\n    accuracies = AverageMeter()\n    model.train()\n    scaler = GradScaler()\n    start_time = time.time()\n    loader = tqdm(dataloader, total=len(dataloader))\n    for step, (images, labels) in enumerate(loader):\n        \n        images = images.to(device).float()\n        labels = labels.to(device)\n        data_time.update(time.time() - start_time)\n\n        with autocast():\n\n            _, _, output = model(images)\n            loss = criterion(output, labels)\n            losses.update(loss.item(), BATCH_SIZE)\n            scaler.scale(loss).backward()\n            grad_norm = nn.utils.clip_grad_norm_(model.parameters(), max_norm = MAX_NORM)\n            if (step+1) % ITERS_TO_ACCUMULATE == 0:\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n        \n        if scheduler is not None and SCHEDULER_UPDATE == 'batch':\n            scheduler.step()\n\n        batch_time.update(time.time() - start_time)\n        start_time = time.time()\n        \n        if step % ITER_FREQ == 0:\n            \n            print('Epoch: [{0}][{1}/{2}]\\t'\n                  'Batch Time {batch_time.val:.3f}s ({batch_time.avg:.3f}s)\\t'\n                  'Data Time {data_time.val:.3f}s ({data_time.avg:.3f}s)\\t'\n                  'Loss: {loss.val:.4f} ({loss.avg:.4f})'.format((epoch+1),\n                                                                    step, len(dataloader),\n                                                                    batch_time=batch_time,\n                                                                    data_time=data_time,\n                                                                    loss=losses))\n                                                                             #accuracy=accuracies))\n        # To check the loss real-time while iterating over data.   'Accuracy {accuracy.val:.4f} ({accuracy.avg:.4f})'\n        loader.set_description(f'Training Epoch {epoch+1}/{EPOCHS}')\n        loader.set_postfix(loss=losses.avg) #accuracy=accuracies.avg)\n#         del images, labels\n    if scheduler is not None and SCHEDULER_UPDATE == 'epoch':\n        scheduler.step()\n        \n    return losses.avg#, accuracies.avg","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:44.24659Z","iopub.status.busy":"2021-02-28T08:48:44.244818Z","iopub.status.idle":"2021-02-28T08:48:44.247308Z","shell.execute_reply":"2021-02-28T08:48:44.247771Z"},"papermill":{"duration":0.034962,"end_time":"2021-02-28T08:48:44.247903","exception":false,"start_time":"2021-02-28T08:48:44.212941","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"def valid_fn(epoch, model, criterion, val_loader, device, scheduler):\n    \n    model.eval()\n    losses = AverageMeter()\n    accuracies = AverageMeter()\n    PREDS = []\n    TARGETS = []\n    loader = tqdm(val_loader, total=len(val_loader))\n    with torch.no_grad():  # without torch.no_grad() will make the CUDA run OOM.\n        for step, (images, labels) in enumerate(loader):\n\n            images = images.to(device)\n            labels = labels.to(device)\n\n            _, _, output = model(images)\n            loss = criterion(output, labels)\n            losses.update(loss.item(), BATCH_SIZE)\n            PREDS += [output.sigmoid()]\n            TARGETS += [labels.detach().cpu()]\n            loader.set_description(f'Validating Epoch {epoch+1}/{EPOCHS}')\n            loader.set_postfix(loss=losses.avg)#, accuracy=accuracies.avg)\n    PREDS = torch.cat(PREDS).cpu().numpy()\n    TARGETS = torch.cat(TARGETS).cpu().numpy()\n    roc_auc = macro_multilabel_auc(TARGETS, PREDS)\n    if scheduler is not None:\n        scheduler.step()\n        \n    return losses.avg, roc_auc# accuracies.avg","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.022832,"end_time":"2021-02-28T08:48:44.293646","exception":false,"start_time":"2021-02-28T08:48:44.270814","status":"completed"},"tags":[]},"cell_type":"markdown","source":"[Back to CFG(Click here)](#cont)"},{"metadata":{"papermill":{"duration":0.022864,"end_time":"2021-02-28T08:48:44.339644","exception":false,"start_time":"2021-02-28T08:48:44.31678","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# Main"},{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:44.401357Z","iopub.status.busy":"2021-02-28T08:48:44.400526Z","iopub.status.idle":"2021-02-28T08:48:44.403202Z","shell.execute_reply":"2021-02-28T08:48:44.402731Z"},"papermill":{"duration":0.040595,"end_time":"2021-02-28T08:48:44.403317","exception":false,"start_time":"2021-02-28T08:48:44.362722","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"def engine(device, folds, fold, model_path=None):\n    \n    trn_idx = folds[folds['kfold'] != fold].index\n    val_idx = folds[folds['kfold'] == fold].index\n\n    train_folds = folds.loc[trn_idx].reset_index(drop=True)\n    valid_folds = folds.loc[val_idx].reset_index(drop=True)\n\n    train_folds = train_folds[train_folds['StudyInstanceUID'].isin(train_annotations['StudyInstanceUID'].unique())].reset_index(drop=True)\n    valid_folds = valid_folds[valid_folds['StudyInstanceUID'].isin(train_annotations['StudyInstanceUID'].unique())].reset_index(drop=True)\n\n    train_data = RanzcrDataset(train_folds, train_annotations, transform=get_transform())\n    val_data = RanzcrDataset(valid_folds, train_annotations, transform=get_transform(train=False))        \n    \n    train_loader = DataLoader(train_data,\n                              batch_size=BATCH_SIZE, \n                              shuffle=True, \n                              num_workers=NUM_WORKERS,\n                              pin_memory=True, # enables faster data transfer to CUDA-enabled GPUs.\n                              drop_last=True)\n    val_loader = DataLoader(val_data,\n                            batch_size=VAL_BATCH_SIZE,\n                            num_workers=NUM_WORKERS,\n                            shuffle=False, \n                            pin_memory=True,\n                            drop_last=False)\n\n    if model_path is not None:\n        model = torch.load(model_path)\n        START_EPOCH = int(model_path.split('_')[-1])\n    else:\n        model = CustomResNet200D(MODEL_ARCH, 11, True)\n        START_EPOCH = 0\n    model.to(device)\n    \n    params = filter(lambda p: p.requires_grad, model.parameters())    \n    optimizer = GetOptimizer(OPTIMIZER, params)\n\n    criterion = GetCriterion(CRITERION).to(device)    \n    val_criterion = GetCriterion(CRITERION).to(device)\n\n    scheduler = GetScheduler(SCHEDULER, optimizer)\n    \n    loss = []\n    accuracy = []\n    for epoch in range(START_EPOCH, EPOCHS):\n        \n        epoch_start = time.time()        \n        avg_loss = train_fn(model, train_loader, device, epoch, optimizer, criterion, scheduler)\n\n        torch.cuda.empty_cache()\n        avg_val_loss, roc_auc_score = valid_fn(epoch, model, val_criterion, val_loader, device, scheduler)\n        epoch_end = time.time() - epoch_start\n        \n        print(f'Validation accuracy after epoch {epoch+1}: {roc_auc_score:.4f}')\n        loss.append(avg_loss)\n#         accuracy.append(avg_accuracy)\n        \n        content = f'Fold {fold} Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  avg_val_loss: {avg_val_loss:.4f} roc_auc_score: {roc_auc_score:.4f} time: {epoch_end:.0f}s'\n        with open(f'GPU_{MODEL_ARCH}_{OPTIMIZER}_{CRITERION}.txt', 'a') as appender:\n            appender.write(content + '\\n')                                         # avg_train_accuracy: {avg_accuracy:.4f}\n        \n        # Save the model to use it for inference.\n        torch.save(model.state_dict(), f'stage1_{MODEL_ARCH}_fold_{fold}_epoch_{(epoch+1)}.pth')\n#         torch.save(model, f'stage1_{MODEL_ARCH}_fold_{fold}_epoch_{(epoch+1)}')\n        torch.cuda.empty_cache()\n    \n    return loss","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-02-28T08:48:44.459221Z","iopub.status.busy":"2021-02-28T08:48:44.458467Z","iopub.status.idle":"2021-02-28T14:33:23.533769Z","shell.execute_reply":"2021-02-28T14:33:23.534614Z"},"papermill":{"duration":20679.10847,"end_time":"2021-02-28T14:33:23.53485","exception":false,"start_time":"2021-02-28T08:48:44.42638","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"if __name__ == '__main__':\n    \n    if MODEL_PATH is not None:\n        START_FOLD = int(MODEL_PATH.split('_')[-3])\n    \n    for fold in range(START_FOLD, N_FOLDS):\n        print(f'===== Fold {fold} Starting =====')\n        fold_start = time.time()\n        logs = engine(DEVICE, folds, fold, MODEL_PATH)\n        print(f'Time taken in fold {fold}: {time.time()-fold_start}')","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}