{"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":"# About this notebook\n- PyTorch tf_efficientnet_b7_ns starter code\n- StratifiedKFold 5 folds\n- Inference notebook is [here](https://www.kaggle.com/yasufuminakama/g2net-efficientnet-b7-baseline-inference)\n- Spectrogram generation code\n    - https://www.kaggle.com/yasufuminakama/g2net-n-mels-128-train-images is generated by https://www.kaggle.com/yasufuminakama/g2net-spectrogram-generation-train\n    - https://www.kaggle.com/yasufuminakama/g2net-n-mels-128-test-images is generated by https://www.kaggle.com/yasufuminakama/g2net-spectrogram-generation-test\n- version 2: melspectrogram approach using above dataset\n- version 3: nnAudio Q-transform approach\n    - Here is nnAudio Constant Q-transform Demonstration  \n        - https://www.kaggle.com/atamazian/nnaudio-constant-q-transform-demonstration\n        - https://www.kaggle.com/c/g2net-gravitational-wave-detection/discussion/250621\n    - Thanks for sharing @atamazian  \n- version 4: tf_efficientnet_b0_ns -> tf_efficientnet_b7_ns\n- version 6: W&B and Grad-CAM\n    - Pytorch W&B Usage Examples from https://docs.wandb.ai/guides/integrations/pytorch\n\nIf this notebook is helpful, feel free to upvote :)","metadata":{"papermill":{"duration":0.018345,"end_time":"2021-07-01T14:31:32.640858","exception":false,"start_time":"2021-07-01T14:31:32.622513","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import IPython.display\nIPython.display.YouTubeVideo('hhbMpe17fzA', width=800, height=500)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2021-08-03T19:06:36.11495Z","iopub.execute_input":"2021-08-03T19:06:36.115391Z","iopub.status.idle":"2021-08-03T19:06:36.252468Z","shell.execute_reply.started":"2021-08-03T19:06:36.115299Z","shell.execute_reply":"2021-08-03T19:06:36.251487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q nnAudio\n!pip install -q --upgrade wandb\n!pip install -q grad-cam\n!pip install -q ttach","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2021-08-03T19:06:36.253868Z","iopub.execute_input":"2021-08-03T19:06:36.254259Z","iopub.status.idle":"2021-08-03T19:07:17.525863Z","shell.execute_reply.started":"2021-08-03T19:06:36.2542Z","shell.execute_reply":"2021-08-03T19:07:17.524801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{"papermill":{"duration":0.016769,"end_time":"2021-07-01T14:31:32.675036","exception":false,"start_time":"2021-07-01T14:31:32.658267","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom matplotlib import pyplot as plt\nimport seaborn as sns","metadata":{"papermill":{"duration":0.717732,"end_time":"2021-07-01T14:31:33.409839","exception":false,"start_time":"2021-07-01T14:31:32.692107","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:07:17.528103Z","iopub.execute_input":"2021-08-03T19:07:17.528487Z","iopub.status.idle":"2021-08-03T19:07:18.275622Z","shell.execute_reply.started":"2021-08-03T19:07:17.528447Z","shell.execute_reply":"2021-08-03T19:07:18.274783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('../input/g2net-gravitational-wave-detection/training_labels.csv')\ntest = pd.read_csv('../input/g2net-gravitational-wave-detection/sample_submission.csv')\n\ndef get_train_file_path(image_id):\n    return \"../input/g2net-gravitational-wave-detection/train/{}/{}/{}/{}.npy\".format(\n        image_id[0], image_id[1], image_id[2], image_id)\n\ndef get_test_file_path(image_id):\n    return \"../input/g2net-gravitational-wave-detection/test/{}/{}/{}/{}.npy\".format(\n        image_id[0], image_id[1], image_id[2], image_id)\n\ntrain['file_path'] = train['id'].apply(get_train_file_path)\ntest['file_path'] = test['id'].apply(get_test_file_path)\n\ndisplay(train.head())\ndisplay(test.head())","metadata":{"papermill":{"duration":1.251215,"end_time":"2021-07-01T14:31:34.678928","exception":false,"start_time":"2021-07-01T14:31:33.427713","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:07:18.277247Z","iopub.execute_input":"2021-08-03T19:07:18.277573Z","iopub.status.idle":"2021-08-03T19:07:19.427608Z","shell.execute_reply.started":"2021-08-03T19:07:18.277539Z","shell.execute_reply":"2021-08-03T19:07:19.426853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Quick EDA","metadata":{"papermill":{"duration":0.018004,"end_time":"2021-07-01T14:31:34.717088","exception":false,"start_time":"2021-07-01T14:31:34.699084","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import torch\nfrom nnAudio.Spectrogram import CQT1992v2\n\ndef apply_qtransform(waves, transform=CQT1992v2(sr=2048, fmin=20, fmax=1024, hop_length=64)):\n    waves = np.hstack(waves)\n    waves = waves / np.max(waves)\n    waves = torch.from_numpy(waves).float()\n    image = transform(waves)\n    return image\n\nfor i in range(5):\n    waves = np.load(train.loc[i, 'file_path'])\n    image = apply_qtransform(waves)\n    target = train.loc[i, 'target']\n    plt.imshow(image[0])\n    plt.title(f\"target: {target}\")\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2021-08-03T19:07:19.428871Z","iopub.execute_input":"2021-08-03T19:07:19.429238Z","iopub.status.idle":"2021-08-03T19:07:21.629061Z","shell.execute_reply.started":"2021-08-03T19:07:19.429184Z","shell.execute_reply":"2021-08-03T19:07:21.628289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train['target'].hist()","metadata":{"papermill":{"duration":0.193214,"end_time":"2021-07-01T14:31:37.746016","exception":false,"start_time":"2021-07-01T14:31:37.552802","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:07:21.630375Z","iopub.execute_input":"2021-08-03T19:07:21.630728Z","iopub.status.idle":"2021-08-03T19:07:21.786216Z","shell.execute_reply.started":"2021-08-03T19:07:21.63069Z","shell.execute_reply":"2021-08-03T19:07:21.785441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Directory settings","metadata":{"papermill":{"duration":0.030055,"end_time":"2021-07-01T14:31:37.805204","exception":false,"start_time":"2021-07-01T14:31:37.775149","status":"completed"},"tags":[]}},{"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)","metadata":{"papermill":{"duration":0.036354,"end_time":"2021-07-01T14:31:37.870775","exception":false,"start_time":"2021-07-01T14:31:37.834421","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:07:21.787466Z","iopub.execute_input":"2021-08-03T19:07:21.787791Z","iopub.status.idle":"2021-08-03T19:07:21.792346Z","shell.execute_reply.started":"2021-08-03T19:07:21.787755Z","shell.execute_reply":"2021-08-03T19:07:21.791347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CFG","metadata":{"papermill":{"duration":0.028932,"end_time":"2021-07-01T14:31:37.928007","exception":false,"start_time":"2021-07-01T14:31:37.899075","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\nclass CFG:\n    apex=False\n    debug=False\n    print_freq=100\n    num_workers=4\n    model_name='tf_efficientnet_b7_ns'\n    scheduler='CosineAnnealingLR' # ['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts']\n    epochs=3\n    #factor=0.2 # ReduceLROnPlateau\n    #patience=4 # ReduceLROnPlateau\n    #eps=1e-6 # ReduceLROnPlateau\n    T_max=3 # CosineAnnealingLR\n    #T_0=3 # CosineAnnealingWarmRestarts\n    lr=1e-4\n    min_lr=1e-6\n    batch_size=48\n    weight_decay=1e-6\n    gradient_accumulation_steps=1\n    max_grad_norm=1000\n    qtransform_params={\"sr\": 2048, \"fmin\": 20, \"fmax\": 1024, \"hop_length\": 32, \"bins_per_octave\": 8}\n    seed=42\n    target_size=1\n    target_col='target'\n    n_fold=5\n    trn_fold=[0] # [0, 1, 2, 3, 4]\n    train=True\n    grad_cam=True\n    \nif CFG.debug:\n    CFG.epochs = 1\n    train = train.sample(n=10000, random_state=CFG.seed).reset_index(drop=True)","metadata":{"papermill":{"duration":0.181532,"end_time":"2021-07-01T14:31:38.138409","exception":false,"start_time":"2021-07-01T14:31:37.956877","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:07:21.795498Z","iopub.execute_input":"2021-08-03T19:07:21.795878Z","iopub.status.idle":"2021-08-03T19:07:21.909672Z","shell.execute_reply.started":"2021-08-03T19:07:21.79584Z","shell.execute_reply":"2021-08-03T19:07:21.908687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Library","metadata":{"papermill":{"duration":0.028374,"end_time":"2021-07-01T14:31:38.19586","exception":false,"start_time":"2021-07-01T14:31:38.167486","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Library\n# ====================================================\nimport sys\nsys.path.append('../input/pytorch-image-models/pytorch-image-models-master')\n\nimport os\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.model_selection import StratifiedKFold, GroupKFold, KFold\n\nfrom tqdm.auto 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\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\n\nfrom pytorch_grad_cam.utils.image import show_cam_on_image\nfrom pytorch_grad_cam import GradCAM, ScoreCAM, GradCAMPlusPlus, AblationCAM, XGradCAM, EigenCAM\n\nimport timm\n\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"papermill":{"duration":3.270545,"end_time":"2021-07-01T14:31:41.495669","exception":false,"start_time":"2021-07-01T14:31:38.225124","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:07:21.911692Z","iopub.execute_input":"2021-08-03T19:07:21.91206Z","iopub.status.idle":"2021-08-03T19:07:23.871052Z","shell.execute_reply.started":"2021-08-03T19:07:21.912021Z","shell.execute_reply":"2021-08-03T19:07:23.870214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nwandb_api = user_secrets.get_secret(\"wandb_api\")","metadata":{"execution":{"iopub.status.busy":"2021-08-03T19:07:23.874278Z","iopub.execute_input":"2021-08-03T19:07:23.874545Z","iopub.status.idle":"2021-08-03T19:07:24.166646Z","shell.execute_reply.started":"2021-08-03T19:07:23.874519Z","shell.execute_reply":"2021-08-03T19:07:24.16574Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import wandb\nwandb.login(key=wandb_api)\n\ndef class2dict(f):\n    return dict((name, getattr(f, name)) for name in dir(f) if not name.startswith('__'))\n\nrun = wandb.init(project=\"G2Net-Public-experiments\", \n                 name=\"exp1\",\n                 config=class2dict(CFG),\n                 group=CFG.model_name,\n                 job_type=\"train\")","metadata":{"execution":{"iopub.status.busy":"2021-08-03T19:07:24.167892Z","iopub.execute_input":"2021-08-03T19:07:24.168243Z","iopub.status.idle":"2021-08-03T19:07:32.487383Z","shell.execute_reply.started":"2021-08-03T19:07:24.168189Z","shell.execute_reply":"2021-08-03T19:07:32.486344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{"papermill":{"duration":0.028157,"end_time":"2021-07-01T14:31:41.552625","exception":false,"start_time":"2021-07-01T14:31:41.524468","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Utils\n# ====================================================\ndef get_score(y_true, y_pred):\n    score = roc_auc_score(y_true, y_pred)\n    return score\n\n\ndef init_logger(log_file=OUTPUT_DIR+'train.log'):\n    from logging import getLogger, INFO, FileHandler,  Formatter,  StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = init_logger()\n\n\ndef seed_torch(seed=42):\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=CFG.seed)","metadata":{"papermill":{"duration":0.042029,"end_time":"2021-07-01T14:31:41.623242","exception":false,"start_time":"2021-07-01T14:31:41.581213","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:07:32.489087Z","iopub.execute_input":"2021-08-03T19:07:32.489701Z","iopub.status.idle":"2021-08-03T19:07:32.505376Z","shell.execute_reply.started":"2021-08-03T19:07:32.489657Z","shell.execute_reply":"2021-08-03T19:07:32.503953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV split","metadata":{"papermill":{"duration":0.02877,"end_time":"2021-07-01T14:31:41.680818","exception":false,"start_time":"2021-07-01T14:31:41.652048","status":"completed"},"tags":[]}},{"cell_type":"code","source":"Fold = StratifiedKFold(n_splits=CFG.n_fold, shuffle=True, random_state=CFG.seed)\nfor n, (train_index, val_index) in enumerate(Fold.split(train, train[CFG.target_col])):\n    train.loc[val_index, 'fold'] = int(n)\ntrain['fold'] = train['fold'].astype(int)\ndisplay(train.groupby(['fold', 'target']).size())","metadata":{"papermill":{"duration":0.060375,"end_time":"2021-07-01T14:31:41.769944","exception":false,"start_time":"2021-07-01T14:31:41.709569","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:07:32.506958Z","iopub.execute_input":"2021-08-03T19:07:32.507543Z","iopub.status.idle":"2021-08-03T19:07:32.538317Z","shell.execute_reply.started":"2021-08-03T19:07:32.507506Z","shell.execute_reply":"2021-08-03T19:07:32.535291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{"papermill":{"duration":0.028894,"end_time":"2021-07-01T14:31:41.827575","exception":false,"start_time":"2021-07-01T14:31:41.798681","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Dataset\n# ====================================================\nclass TrainDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df['file_path'].values\n        self.labels = df[CFG.target_col].values\n        self.wave_transform = CQT1992v2(**CFG.qtransform_params)\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def apply_qtransform(self, waves, transform):\n        waves = np.hstack(waves)\n        waves = waves / np.max(waves)\n        waves = torch.from_numpy(waves).float()\n        image = transform(waves)\n        return image\n\n    def __getitem__(self, idx):\n        file_path = self.file_names[idx]\n        waves = np.load(file_path)\n        image = self.apply_qtransform(waves, self.wave_transform)\n        image = image.squeeze().numpy()\n        if self.transform:\n            image = self.transform(image=image)['image']\n        label = torch.tensor(self.labels[idx]).float()\n        return image, label","metadata":{"papermill":{"duration":0.040385,"end_time":"2021-07-01T14:31:41.897587","exception":false,"start_time":"2021-07-01T14:31:41.857202","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:07:32.542Z","iopub.execute_input":"2021-08-03T19:07:32.542335Z","iopub.status.idle":"2021-08-03T19:07:32.555272Z","shell.execute_reply.started":"2021-08-03T19:07:32.542306Z","shell.execute_reply":"2021-08-03T19:07:32.554369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GradCAMDataset(Dataset):\n    def __init__(self, df):\n        self.df = df\n        self.image_ids = df['id'].values\n        self.file_names = df['file_path'].values\n        self.labels = df[CFG.target_col].values\n        self.wave_transform = CQT1992v2(**CFG.qtransform_params)\n        self.transform = get_transforms(data='valid')\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def apply_qtransform(self, waves, transform):\n        waves = np.hstack(waves)\n        waves = waves / np.max(waves)\n        waves = torch.from_numpy(waves).float()\n        image = transform(waves)\n        return image\n\n    def __getitem__(self, idx):\n        image_id = self.image_ids[idx]\n        file_path = self.file_names[idx]\n        waves = np.load(file_path)\n        image = self.apply_qtransform(waves, self.wave_transform)\n        image = image.squeeze().numpy()\n        vis_image = image.copy()\n        if self.transform:\n            image = self.transform(image=image)['image']\n        label = torch.tensor(self.labels[idx]).float()\n        return image_id, image, vis_image, label","metadata":{"execution":{"iopub.status.busy":"2021-08-03T19:07:32.55677Z","iopub.execute_input":"2021-08-03T19:07:32.557536Z","iopub.status.idle":"2021-08-03T19:07:32.569568Z","shell.execute_reply.started":"2021-08-03T19:07:32.557497Z","shell.execute_reply":"2021-08-03T19:07:32.568538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transforms","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Transforms\n# ====================================================\ndef get_transforms(*, data):\n    \n    if data == 'train':\n        return A.Compose([\n            ToTensorV2(),\n        ])\n\n    elif data == 'valid':\n        return A.Compose([\n            ToTensorV2(),\n        ])","metadata":{"execution":{"iopub.status.busy":"2021-08-03T19:07:32.570957Z","iopub.execute_input":"2021-08-03T19:07:32.571633Z","iopub.status.idle":"2021-08-03T19:07:32.580583Z","shell.execute_reply.started":"2021-08-03T19:07:32.571594Z","shell.execute_reply":"2021-08-03T19:07:32.579728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = TrainDataset(train, transform=get_transforms(data='train'))\n\nfor i in range(5):\n    plt.figure(figsize=(16,12))\n    image, label = train_dataset[i]\n    plt.imshow(image[0])\n    plt.title(f'label: {label}')\n    plt.show() ","metadata":{"papermill":{"duration":1.037231,"end_time":"2021-07-01T14:31:42.96351","exception":false,"start_time":"2021-07-01T14:31:41.926279","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:07:32.582216Z","iopub.execute_input":"2021-08-03T19:07:32.583011Z","iopub.status.idle":"2021-08-03T19:07:33.585466Z","shell.execute_reply.started":"2021-08-03T19:07:32.582943Z","shell.execute_reply":"2021-08-03T19:07:33.584518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# MODEL","metadata":{"papermill":{"duration":0.03649,"end_time":"2021-07-01T14:31:43.035743","exception":false,"start_time":"2021-07-01T14:31:42.999253","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# MODEL\n# ====================================================\nclass CustomModel(nn.Module):\n    def __init__(self, cfg, pretrained=False):\n        super().__init__()\n        self.cfg = cfg\n        self.model = timm.create_model(self.cfg.model_name, pretrained=pretrained, in_chans=1)\n        self.n_features = self.model.classifier.in_features\n        self.model.classifier = nn.Linear(self.n_features, self.cfg.target_size)\n\n    def forward(self, x):\n        output = self.model(x)\n        return output","metadata":{"papermill":{"duration":0.044023,"end_time":"2021-07-01T14:31:43.114443","exception":false,"start_time":"2021-07-01T14:31:43.07042","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:07:33.586851Z","iopub.execute_input":"2021-08-03T19:07:33.587193Z","iopub.status.idle":"2021-08-03T19:07:33.594205Z","shell.execute_reply.started":"2021-08-03T19:07:33.587158Z","shell.execute_reply":"2021-08-03T19:07:33.593263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{"papermill":{"duration":0.034387,"end_time":"2021-07-01T14:31:43.183231","exception":false,"start_time":"2021-07-01T14:31:43.148844","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Helper functions\n# ====================================================\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 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\n\ndef asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\n\ndef timeSince(since, percent):\n    now = time.time()\n    s = now - since\n    es = s / (percent)\n    rs = es - s\n    return '%s (remain %s)' % (asMinutes(s), asMinutes(rs))\n\n\ndef train_fn(fold, train_loader, model, criterion, optimizer, epoch, scheduler, device):\n    if CFG.apex:\n        scaler = GradScaler()\n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    # switch to train mode\n    model.train()\n    start = end = time.time()\n    global_step = 0\n    for step, (images, labels) in enumerate(train_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        images = images.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        if CFG.apex:\n            with autocast():\n                y_preds = model(images)\n                loss = criterion(y_preds.view(-1), labels)\n        else:\n            y_preds = model(images)\n            loss = criterion(y_preds.view(-1), labels)\n        # record loss\n        losses.update(loss.item(), batch_size)\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        if CFG.apex:\n            scaler.scale(loss).backward()\n        else:\n            loss.backward()\n        grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.max_grad_norm)\n        if (step + 1) % CFG.gradient_accumulation_steps == 0:\n            if CFG.apex:\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                optimizer.step()\n            optimizer.zero_grad()\n            global_step += 1\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(train_loader)-1):\n            print('Epoch: [{0}][{1}/{2}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  'Grad: {grad_norm:.4f}  '\n                  'LR: {lr:.6f}  '\n                  .format(epoch+1, step, len(train_loader), \n                          remain=timeSince(start, float(step+1)/len(train_loader)),\n                          loss=losses,\n                          grad_norm=grad_norm,\n                          lr=scheduler.get_lr()[0]))\n        wandb.log({f\"[fold{fold}] loss\": losses.val,\n                   f\"[fold{fold}] lr\": scheduler.get_lr()[0]})\n    return losses.avg\n\n\ndef valid_fn(valid_loader, model, criterion, device):\n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    # switch to evaluation mode\n    model.eval()\n    preds = []\n    start = end = time.time()\n    for step, (images, labels) in enumerate(valid_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        images = images.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        # compute loss\n        with torch.no_grad():\n            y_preds = model(images)\n        loss = criterion(y_preds.view(-1), labels)\n        losses.update(loss.item(), batch_size)\n        # record accuracy\n        preds.append(y_preds.sigmoid().to('cpu').numpy())\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(valid_loader)-1):\n            print('EVAL: [{0}/{1}] '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  .format(step, len(valid_loader),\n                          loss=losses,\n                          remain=timeSince(start, float(step+1)/len(valid_loader))))\n    predictions = np.concatenate(preds)\n    return losses.avg, predictions","metadata":{"papermill":{"duration":0.190966,"end_time":"2021-07-01T14:31:43.408934","exception":false,"start_time":"2021-07-01T14:31:43.217968","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:07:33.595792Z","iopub.execute_input":"2021-08-03T19:07:33.596167Z","iopub.status.idle":"2021-08-03T19:07:33.622055Z","shell.execute_reply.started":"2021-08-03T19:07:33.596112Z","shell.execute_reply":"2021-08-03T19:07:33.621133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_grad_cam(model, device, x_tensor, img, label, plot=False):\n    \n    result = {\"vis\": None, \"img\": None, \"prob\": None, \"label\": None}\n    \n    # model prob\n    with torch.no_grad():\n        prob = model(x_tensor.unsqueeze(0).to(device))\n    prob = np.concatenate(prob.sigmoid().to('cpu').numpy())[0]\n    \n    # grad-cam\n    target_layer = model.model.conv_head\n    cam = GradCAM(model=model, target_layer=target_layer, use_cuda=True)\n    output = cam(input_tensor=x_tensor.unsqueeze(0))\n    try:\n        vis = show_cam_on_image(x_tensor.numpy().transpose((1, 2, 0)), output[0])\n    except:\n        return result\n\n    # plot result\n    if plot:\n        fig, axes = plt.subplots(figsize=(16, 12), ncols=2)\n        axes[0].imshow(vis)\n        axes[0].set_title(f\"prob={prob:.4f}\")\n        axes[1].imshow(img)\n        axes[1].set_title(f\"target={label}\")\n        plt.show()\n        \n    result = {\"vis\": vis, \"img\": img, \"prob\": prob, \"label\": label}\n\n    return result","metadata":{"execution":{"iopub.status.busy":"2021-08-03T19:07:33.623942Z","iopub.execute_input":"2021-08-03T19:07:33.624506Z","iopub.status.idle":"2021-08-03T19:07:33.636124Z","shell.execute_reply.started":"2021-08-03T19:07:33.624467Z","shell.execute_reply":"2021-08-03T19:07:33.635364Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train loop","metadata":{"papermill":{"duration":0.034375,"end_time":"2021-07-01T14:31:43.478039","exception":false,"start_time":"2021-07-01T14:31:43.443664","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# Train loop\n# ====================================================\ndef train_loop(folds, fold):\n    \n    LOGGER.info(f\"========== fold: {fold} training ==========\")\n\n    # ====================================================\n    # loader\n    # ====================================================\n    trn_idx = folds[folds['fold'] != fold].index\n    val_idx = folds[folds['fold'] == 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    valid_labels = valid_folds[CFG.target_col].values\n\n    train_dataset = TrainDataset(train_folds, transform=get_transforms(data='train'))\n    valid_dataset = TrainDataset(valid_folds, transform=get_transforms(data='train'))\n\n    train_loader = DataLoader(train_dataset,\n                              batch_size=CFG.batch_size, \n                              shuffle=True, \n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=True)\n    valid_loader = DataLoader(valid_dataset, \n                              batch_size=CFG.batch_size * 2, \n                              shuffle=False, \n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n    \n    # ====================================================\n    # scheduler \n    # ====================================================\n    def get_scheduler(optimizer):\n        if CFG.scheduler=='ReduceLROnPlateau':\n            scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=CFG.factor, patience=CFG.patience, verbose=True, eps=CFG.eps)\n        elif CFG.scheduler=='CosineAnnealingLR':\n            scheduler = CosineAnnealingLR(optimizer, T_max=CFG.T_max, eta_min=CFG.min_lr, last_epoch=-1)\n        elif CFG.scheduler=='CosineAnnealingWarmRestarts':\n            scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=CFG.T_0, T_mult=1, eta_min=CFG.min_lr, last_epoch=-1)\n        return scheduler\n\n    # ====================================================\n    # model & optimizer\n    # ====================================================\n    model = CustomModel(CFG, pretrained=True)\n    model.to(device)\n\n    optimizer = Adam(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay, amsgrad=False)\n    scheduler = get_scheduler(optimizer)\n\n    # ====================================================\n    # loop\n    # ====================================================\n    criterion = nn.BCEWithLogitsLoss()\n\n    best_score = 0.\n    best_loss = np.inf\n    \n    for epoch in range(CFG.epochs):\n        \n        start_time = time.time()\n        \n        # train\n        avg_loss = train_fn(fold, train_loader, model, criterion, optimizer, epoch, scheduler, device)\n\n        # eval\n        avg_val_loss, preds = valid_fn(valid_loader, model, criterion, device)\n        \n        if isinstance(scheduler, ReduceLROnPlateau):\n            scheduler.step(avg_val_loss)\n        elif isinstance(scheduler, CosineAnnealingLR):\n            scheduler.step()\n        elif isinstance(scheduler, CosineAnnealingWarmRestarts):\n            scheduler.step()\n\n        # scoring\n        score = get_score(valid_labels, preds)\n\n        elapsed = time.time() - start_time\n\n        LOGGER.info(f'Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  avg_val_loss: {avg_val_loss:.4f}  time: {elapsed:.0f}s')\n        LOGGER.info(f'Epoch {epoch+1} - Score: {score:.4f}')\n        wandb.log({f\"[fold{fold}] epoch\": epoch+1, \n                   f\"[fold{fold}] avg_train_loss\": avg_loss, \n                   f\"[fold{fold}] avg_val_loss\": avg_val_loss,\n                   f\"[fold{fold}] score\": score})\n\n        if score > best_score:\n            best_score = score\n            LOGGER.info(f'Epoch {epoch+1} - Save Best Score: {best_score:.4f} Model')\n            torch.save({'model': model.state_dict(), \n                        'preds': preds},\n                        OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best_score.pth')\n        \n        if avg_val_loss < best_loss:\n            best_loss = avg_val_loss\n            LOGGER.info(f'Epoch {epoch+1} - Save Best Loss: {best_loss:.4f} Model')\n            torch.save({'model': model.state_dict(), \n                        'preds': preds},\n                        OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best_loss.pth')\n    \n    valid_folds['preds'] = torch.load(OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best_score.pth', \n                                      map_location=torch.device('cpu'))['preds']\n\n    return valid_folds","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":0.053758,"end_time":"2021-07-01T14:31:43.566555","exception":false,"start_time":"2021-07-01T14:31:43.512797","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:07:33.638002Z","iopub.execute_input":"2021-08-03T19:07:33.63836Z","iopub.status.idle":"2021-08-03T19:07:33.658158Z","shell.execute_reply.started":"2021-08-03T19:07:33.638326Z","shell.execute_reply":"2021-08-03T19:07:33.657358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# main\n# ====================================================\ndef main():\n\n    \"\"\"\n    Prepare: 1.train \n    \"\"\"\n\n    def get_result(result_df):\n        preds = result_df['preds'].values\n        labels = result_df[CFG.target_col].values\n        score = get_score(labels, preds)\n        LOGGER.info(f'Score: {score:<.4f}')\n    \n    if CFG.train:\n        # train \n        oof_df = pd.DataFrame()\n        for fold in range(CFG.n_fold):\n            if fold in CFG.trn_fold:\n                _oof_df = train_loop(train, fold)\n                oof_df = pd.concat([oof_df, _oof_df])\n                LOGGER.info(f\"========== fold: {fold} result ==========\")\n                get_result(_oof_df)\n        # CV result\n        LOGGER.info(f\"========== CV ==========\")\n        get_result(oof_df)\n        # save result\n        oof_df.to_csv(OUTPUT_DIR+'oof_df.csv', index=False)\n    \n    if CFG.grad_cam:\n        N = 5\n        wandb_table = wandb.Table(columns=[\"id\", \"target\", \"prob\", \"image\", \"grad_cam_image\"])\n        for fold in range(CFG.n_fold):\n            if fold in CFG.trn_fold:\n                # load model\n                model = CustomModel(CFG, pretrained=False)\n                state = torch.load(OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best_loss.pth', \n                                   map_location=torch.device('cpu'))['model']\n                model.load_state_dict(state)\n                model.to(device)\n                model.eval()\n                # load oof\n                oof = pd.read_csv(OUTPUT_DIR+'oof_df.csv')\n                oof = oof[oof['fold'] == fold].reset_index(drop=True)\n                # grad-cam (oof ascending=False)\n                count = 0\n                oof = oof.sort_values('preds', ascending=False)\n                valid_dataset = GradCAMDataset(oof)\n                for i in range(len(valid_dataset)):\n                    image_id, x_tensor, img, label = valid_dataset[i]\n                    result = get_grad_cam(model, device, x_tensor, img, label, plot=True)\n                    if result[\"vis\"] is not None:\n                        count += 1\n                        wandb_table.add_data(image_id, \n                                             result[\"label\"], \n                                             result[\"prob\"], \n                                             wandb.Image(result[\"img\"]), \n                                             wandb.Image(result[\"vis\"]))\n                    if count >= N:\n                        break\n                # grad-cam (oof ascending=True)\n                count = 0\n                oof = oof.sort_values('preds', ascending=True)\n                valid_dataset = GradCAMDataset(oof)\n                for i in range(len(valid_dataset)):\n                    image_id, x_tensor, img, label = valid_dataset[i]\n                    result = get_grad_cam(model, device, x_tensor, img, label, plot=True)\n                    if result[\"vis\"] is not None:\n                        count += 1\n                        wandb_table.add_data(image_id, \n                                             result[\"label\"], \n                                             result[\"prob\"], \n                                             wandb.Image(result[\"img\"]), \n                                             wandb.Image(result[\"vis\"]))\n                    if count >= N:\n                        break\n        wandb.log({'grad_cam': wandb_table})\n    \n    wandb.finish()","metadata":{"papermill":{"duration":0.047199,"end_time":"2021-07-01T14:31:43.648478","exception":false,"start_time":"2021-07-01T14:31:43.601279","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:08:22.072387Z","iopub.execute_input":"2021-08-03T19:08:22.072769Z","iopub.status.idle":"2021-08-03T19:08:22.090442Z","shell.execute_reply.started":"2021-08-03T19:08:22.072737Z","shell.execute_reply":"2021-08-03T19:08:22.089451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == '__main__':\n    main()","metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","papermill":{"duration":1704.25542,"end_time":"2021-07-01T15:00:07.939127","exception":false,"start_time":"2021-07-01T14:31:43.683707","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-08-03T19:08:22.501971Z","iopub.execute_input":"2021-08-03T19:08:22.502349Z","iopub.status.idle":"2021-08-03T19:12:12.62414Z","shell.execute_reply.started":"2021-08-03T19:08:22.502314Z","shell.execute_reply":"2021-08-03T19:12:12.623151Z"},"trusted":true},"execution_count":null,"outputs":[]}]}