{"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- Training notebook is [here](https://www.kaggle.com/yasufuminakama/g2net-efficientnet-b7-baseline-training)\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 3: melspectrogram approach using above dataset\n- version 4: 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 6: tf_efficientnet_b0_ns -> tf_efficientnet_b7_ns\n- version 7: update model\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-04T08:38:26.724096Z","iopub.execute_input":"2021-08-04T08:38:26.724432Z","iopub.status.idle":"2021-08-04T08:38:26.828165Z","shell.execute_reply.started":"2021-08-04T08:38:26.724358Z","shell.execute_reply":"2021-08-04T08:38:26.827399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q nnAudio","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2021-08-04T08:38:26.829459Z","iopub.execute_input":"2021-08-04T08:38:26.829838Z","iopub.status.idle":"2021-08-04T08:38:34.854525Z","shell.execute_reply.started":"2021-08-04T08:38:26.829801Z","shell.execute_reply":"2021-08-04T08:38:34.853498Z"},"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    debug=False\n    num_workers=4\n    model_name='tf_efficientnet_b7_ns'\n    model_dir='../input/g2net-efficientnet-b7-baseline-training/'\n    batch_size=512\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]","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-04T08:40:14.587270Z","iopub.execute_input":"2021-08-04T08:40:14.587603Z","iopub.status.idle":"2021-08-04T08:40:14.598166Z","shell.execute_reply.started":"2021-08-04T08:40:14.587570Z","shell.execute_reply":"2021-08-04T08:40:14.597072Z"},"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\nimport timm\n\nfrom torch.cuda.amp import autocast, GradScaler\n\nfrom nnAudio.Spectrogram import CQT1992v2\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-04T08:38:43.390603Z","iopub.execute_input":"2021-08-04T08:38:43.390967Z","iopub.status.idle":"2021-08-04T08:38:47.472674Z","shell.execute_reply.started":"2021-08-04T08:38:43.390937Z","shell.execute_reply":"2021-08-04T08:38:47.471751Z"},"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 get_result(result_df):\n    preds = result_df['preds'].values\n    labels = result_df['target'].values\n    score = get_score(labels, preds)\n    return score\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-04T08:38:47.475907Z","iopub.execute_input":"2021-08-04T08:38:47.476172Z","iopub.status.idle":"2021-08-04T08:38:47.486728Z","shell.execute_reply.started":"2021-08-04T08:38:47.476145Z","shell.execute_reply":"2021-08-04T08:38:47.485912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{}},{"cell_type":"code","source":"test = pd.read_csv('../input/g2net-gravitational-wave-detection/sample_submission.csv')\n\nif CFG.debug:\n    test = test.sample(n=1000, random_state=CFG.seed).reset_index(drop=True)\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\ntest['file_path'] = test['id'].apply(get_test_file_path)\n\ndisplay(test.head())","metadata":{"execution":{"iopub.status.busy":"2021-08-04T08:38:47.490120Z","iopub.execute_input":"2021-08-04T08:38:47.490430Z","iopub.status.idle":"2021-08-04T08:38:47.834479Z","shell.execute_reply.started":"2021-08-04T08:38:47.490393Z","shell.execute_reply":"2021-08-04T08:38:47.833582Z"},"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 TestDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df['file_path'].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        return image","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-04T08:40:17.625170Z","iopub.execute_input":"2021-08-04T08:40:17.625522Z","iopub.status.idle":"2021-08-04T08:40:17.633350Z","shell.execute_reply.started":"2021-08-04T08:40:17.625487Z","shell.execute_reply":"2021-08-04T08:40:17.632296Z"},"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-04T08:40:17.957289Z","iopub.execute_input":"2021-08-04T08:40:17.957623Z","iopub.status.idle":"2021-08-04T08:40:17.962792Z","shell.execute_reply.started":"2021-08-04T08:40:17.957591Z","shell.execute_reply":"2021-08-04T08:40:17.961742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from matplotlib import pyplot as plt\n\ntest_dataset = TestDataset(test, transform=get_transforms(data='valid'))\n\nfor i in range(5):\n    plt.figure(figsize=(16,12))\n    image = test_dataset[i]\n    plt.imshow(image[0])\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-04T08:40:18.131322Z","iopub.execute_input":"2021-08-04T08:40:18.131622Z","iopub.status.idle":"2021-08-04T08:40:18.880804Z","shell.execute_reply.started":"2021-08-04T08:40:18.131593Z","shell.execute_reply":"2021-08-04T08:40:18.879866Z"},"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-04T08:40:18.882381Z","iopub.execute_input":"2021-08-04T08:40:18.882731Z","iopub.status.idle":"2021-08-04T08:40:18.889336Z","shell.execute_reply.started":"2021-08-04T08:40:18.882695Z","shell.execute_reply":"2021-08-04T08:40:18.888350Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# inference","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# inference\n# ====================================================\ndef inference(model, states, test_loader, device):\n    model.to(device)\n    tk0 = tqdm(enumerate(test_loader), total=len(test_loader))\n    probs = []\n    for i, (images) in tk0:\n        images = images.to(device)\n        avg_preds = []\n        for state in states:\n            model.load_state_dict(state['model'])\n            model.eval()\n            with torch.no_grad():\n                y_preds = model(images)\n            avg_preds.append(y_preds.sigmoid().to('cpu').numpy())\n        avg_preds = np.mean(avg_preds, axis=0)\n        probs.append(avg_preds)\n    probs = np.concatenate(probs)\n    return probs","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-04T08:40:19.146720Z","iopub.execute_input":"2021-08-04T08:40:19.147011Z","iopub.status.idle":"2021-08-04T08:40:19.153137Z","shell.execute_reply.started":"2021-08-04T08:40:19.146983Z","shell.execute_reply":"2021-08-04T08:40:19.152331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CustomModel(CFG, pretrained=False)\nstates = [torch.load(CFG.model_dir+f'{CFG.model_name}_fold{fold}_best_score.pth') for fold in CFG.trn_fold]\ntest_dataset = TestDataset(test, transform=get_transforms(data='valid'))\ntest_loader = DataLoader(test_dataset, batch_size=CFG.batch_size, shuffle=False, \n                         num_workers=CFG.num_workers, pin_memory=True)\npredictions = inference(model, states, test_loader, device)","metadata":{"execution":{"iopub.status.busy":"2021-08-04T08:40:19.304816Z","iopub.execute_input":"2021-08-04T08:40:19.305073Z","iopub.status.idle":"2021-08-04T08:41:06.714620Z","shell.execute_reply.started":"2021-08-04T08:40:19.305048Z","shell.execute_reply":"2021-08-04T08:41:06.713314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"test['target'] = predictions\ntest[['id', 'target']].to_csv('submission.csv', index=False)\ntest.head()","metadata":{"execution":{"iopub.status.busy":"2021-07-01T19:27:01.839039Z","iopub.execute_input":"2021-07-01T19:27:01.83936Z","iopub.status.idle":"2021-07-01T19:27:02.084101Z","shell.execute_reply.started":"2021-07-01T19:27:01.839322Z","shell.execute_reply":"2021-07-01T19:27:02.083342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test['target'].hist()","metadata":{"execution":{"iopub.status.busy":"2021-07-01T19:27:11.201719Z","iopub.execute_input":"2021-07-01T19:27:11.202044Z","iopub.status.idle":"2021-07-01T19:27:11.349503Z","shell.execute_reply.started":"2021-07-01T19:27:11.202014Z","shell.execute_reply":"2021-07-01T19:27:11.348558Z"},"trusted":true},"execution_count":null,"outputs":[]}]}