{"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":"## reference\n- https://www.kaggle.com/code/myso1987/g2net-basic-audio-data-augmentation","metadata":{}},{"cell_type":"code","source":" COLAB = False\n\nif COLAB == True:\n    from google.colab import drive\n    drive.mount('/content/drive')\n    %cd '/content/drive/MyDrive/Colab Notebooks/kaggle/G2Net2022/code'","metadata":{"id":"w5UprQs8ga3c","outputId":"5374b854-2ac0-42d5-9c0c-dede89cdaf8a","execution":{"iopub.status.busy":"2022-11-14T08:12:31.03651Z","iopub.execute_input":"2022-11-14T08:12:31.036881Z","iopub.status.idle":"2022-11-14T08:12:31.058875Z","shell.execute_reply.started":"2022-11-14T08:12:31.036805Z","shell.execute_reply":"2022-11-14T08:12:31.057912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! pip3 install timm -q","metadata":{"id":"uBZokyT1gYWD","outputId":"aed15e9e-e33e-4780-9434-cf85f12c2229","execution":{"iopub.status.busy":"2022-11-14T08:12:31.0617Z","iopub.execute_input":"2022-11-14T08:12:31.062061Z","iopub.status.idle":"2022-11-14T08:12:46.788464Z","shell.execute_reply.started":"2022-11-14T08:12:31.062017Z","shell.execute_reply":"2022-11-14T08:12:46.787311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport time\nimport h5py\nimport timm\nimport torch\nimport torch.nn as nn\nimport torchaudio\nimport torchvision.transforms as TF\n\n\nfrom tqdm.auto import tqdm\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import roc_auc_score\nfrom timm.scheduler import CosineLRScheduler\n\ndevice = torch.device('cuda')\ncriterion = nn.BCEWithLogitsLoss()\n\n# Train metadata\ndi = '../input/g2net-detecting-continuous-gravitational-waves'\ndf = pd.read_csv(di + '/train_labels.csv')\ndf = df[df.target >= 0]  # Remove 3 unknowns (target = -1)","metadata":{"id":"3Xgw6q0ugYWH","execution":{"iopub.status.busy":"2022-11-14T08:12:46.790799Z","iopub.execute_input":"2022-11-14T08:12:46.791512Z","iopub.status.idle":"2022-11-14T08:12:49.829664Z","shell.execute_reply.started":"2022-11-14T08:12:46.791462Z","shell.execute_reply":"2022-11-14T08:12:49.828534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{"id":"tbekTVh-gYWI"}},{"cell_type":"code","source":"transforms_time_mask = nn.Sequential(\n                torchaudio.transforms.TimeMasking(time_mask_param=10),\n            )\n\ntransforms_freq_mask = nn.Sequential(\n                torchaudio.transforms.FrequencyMasking(freq_mask_param=10),\n            )\n\nflip_rate = 0.0 # probability of applying the horizontal flip and vertical flip \nfre_shift_rate = 0.0 # probability of applying the vertical shift\n\ntime_mask_num = 0 # number of time masking\nfreq_mask_num = 0 # number of frequency masking","metadata":{"id":"7Z3rNynhq1gr","execution":{"iopub.status.busy":"2022-11-14T08:12:49.831417Z","iopub.execute_input":"2022-11-14T08:12:49.831776Z","iopub.status.idle":"2022-11-14T08:12:49.84156Z","shell.execute_reply.started":"2022-11-14T08:12:49.831742Z","shell.execute_reply":"2022-11-14T08:12:49.840569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Dataset(torch.utils.data.Dataset):\n    \"\"\"\n    dataset = Dataset(data_type, df)\n\n    img, y = dataset[i]\n      img (np.float32): 2 x 360 x 128\n      y (np.float32): label 0 or 1\n    \"\"\"\n    def __init__(self, data_type, df, tfms=False):\n        self.data_type = data_type\n        self.df = df\n        self.tfms = tfms\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, i):\n        \"\"\"\n        i (int): get ith data\n        \"\"\"\n        r = self.df.iloc[i]\n        y = np.float32(r.target)\n        file_id = r.id\n\n        img = np.empty((2, 360, 128), dtype=np.float32)\n\n        filename = '%s/%s/%s.hdf5' % (di, self.data_type, file_id)\n        with h5py.File(filename, 'r') as f:\n            g = f[file_id]\n\n            for ch, s in enumerate(['H1', 'L1']):\n                a = g[s]['SFTs'][:, :4096] * 1e22  # Fourier coefficient complex64\n\n                p = a.real**2 + a.imag**2  # power\n                p /= np.mean(p)  # normalize\n                p = np.mean(p.reshape(360, 128, 32), axis=2)  # compress 4096 -> 128\n                img[ch] = p\n\n        if self.tfms:\n            if np.random.rand() <= flip_rate: # horizontal flip\n                img = np.flip(img, axis=1).copy()\n            if np.random.rand() <= flip_rate: # vertical flip\n                img = np.flip(img, axis=2).copy()\n            if np.random.rand() <= fre_shift_rate: # vertical shift\n                img = np.roll(img, np.random.randint(low=0, high=img.shape[1]), axis=1)\n            \n            img = torch.from_numpy(img)\n\n            for _ in range(time_mask_num): # tima masking\n                img = transforms_time_mask(img)\n            for _ in range(freq_mask_num): # frequency masking\n                img = transforms_freq_mask(img)\n        \n        else:\n            img = torch.from_numpy(img)\n                \n        return img, y","metadata":{"id":"3kOMMsyagYWN","execution":{"iopub.status.busy":"2022-11-14T08:12:49.845574Z","iopub.execute_input":"2022-11-14T08:12:49.846537Z","iopub.status.idle":"2022-11-14T08:12:49.859446Z","shell.execute_reply.started":"2022-11-14T08:12:49.846474Z","shell.execute_reply":"2022-11-14T08:12:49.858159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Audio Data Augmentation","metadata":{}},{"cell_type":"markdown","source":"* horizontal flip\n* vertical flip\n* vertical shift\n* time masking*\n* frequency masking*\n\n*Reference  \nSpecAugment  \nhttps://arxiv.org/abs/1904.08779","metadata":{}},{"cell_type":"markdown","source":"## Horizontal flip and Vertical flip ","metadata":{}},{"cell_type":"code","source":"dataset = Dataset('train', df, tfms=False)\nimg, y = dataset[10]\n\n\nplt.figure(figsize=(8, 3))\nplt.title('Spectrogram')\nplt.xlabel('time')\nplt.ylabel('frequency')\nplt.imshow(img[0, 0:360])\nplt.colorbar()\nplt.show()\n\n\nflip_rate = 1.0 # probability of applying the horizontal flip and vertical flip \n\ndataset = Dataset('train', df, tfms=True)\nimg, y = dataset[10]\n\nplt.figure(figsize=(8, 3))\nplt.title('Spectrogram')\nplt.xlabel('time')\nplt.ylabel('frequency')\nplt.imshow(img[0, 0:360])\nplt.colorbar()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-14T08:12:49.86241Z","iopub.execute_input":"2022-11-14T08:12:49.862779Z","iopub.status.idle":"2022-11-14T08:12:51.21868Z","shell.execute_reply.started":"2022-11-14T08:12:49.862752Z","shell.execute_reply":"2022-11-14T08:12:51.217557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Vertical shift","metadata":{}},{"cell_type":"code","source":"dataset = Dataset('train', df, tfms=False)\nimg, y = dataset[10]\n\n\nplt.figure(figsize=(8, 3))\nplt.title('Spectrogram')\nplt.xlabel('time')\nplt.ylabel('frequency')\nplt.imshow(img[0, 0:360])\nplt.colorbar()\nplt.show()\n\n\nflip_rate = 0.0 # probability of applying the horizontal flip and vertical flip \nfre_shift_rate = 1.0 # probability of applying the vertical shift\n\ndataset = Dataset('train', df, tfms=True)\nimg, y = dataset[10]\n\nplt.figure(figsize=(8, 3))\nplt.title('Spectrogram')\nplt.xlabel('time')\nplt.ylabel('frequency')\nplt.imshow(img[0, 0:360])\nplt.colorbar()\nplt.show()","metadata":{"id":"KKX7AmjTU8qI","outputId":"448f6ad0-bae2-4e83-8028-c86fdd2fa17b","execution":{"iopub.status.busy":"2022-11-14T08:12:51.220948Z","iopub.execute_input":"2022-11-14T08:12:51.221604Z","iopub.status.idle":"2022-11-14T08:12:52.165672Z","shell.execute_reply.started":"2022-11-14T08:12:51.221566Z","shell.execute_reply":"2022-11-14T08:12:52.164597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Time masking","metadata":{}},{"cell_type":"code","source":"dataset = Dataset('train', df, tfms=False)\nimg, y = dataset[10]\n\n\nplt.figure(figsize=(8, 3))\nplt.title('Spectrogram')\nplt.xlabel('time')\nplt.ylabel('frequency')\nplt.imshow(img[0, 0:360])\nplt.colorbar()\nplt.show()\n\n\nflip_rate = 0.0 # probability of applying the horizontal flip and vertical flip \nfre_shift_rate = 0.0 # probability of applying the vertical shift\ntime_mask_num = 3 # number of time masking\n\ndataset = Dataset('train', df, tfms=True)\nimg, y = dataset[10]\n\nplt.figure(figsize=(8, 3))\nplt.title('Spectrogram')\nplt.xlabel('time')\nplt.ylabel('frequency')\nplt.imshow(img[0, 0:360])\nplt.colorbar()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-14T08:12:52.167167Z","iopub.execute_input":"2022-11-14T08:12:52.167511Z","iopub.status.idle":"2022-11-14T08:12:53.137008Z","shell.execute_reply.started":"2022-11-14T08:12:52.167476Z","shell.execute_reply":"2022-11-14T08:12:53.136018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Frequency masking","metadata":{}},{"cell_type":"code","source":"dataset = Dataset('train', df, tfms=False)\nimg, y = dataset[10]\n\n\nplt.figure(figsize=(8, 3))\nplt.title('Spectrogram')\nplt.xlabel('time')\nplt.ylabel('frequency')\nplt.imshow(img[0, 0:360])\nplt.colorbar()\nplt.show()\n\n\nflip_rate = 0.0 # probability of applying the horizontal flip and vertical flip \nfre_shift_rate = 0.0 # probability of applying the vertical shift\ntime_mask_num = 0 # number of time masking\nfreq_mask_num = 3 # number of frequency masking\n\ndataset = Dataset('train', df, tfms=True)\nimg, y = dataset[10]\n\nplt.figure(figsize=(8, 3))\nplt.title('Spectrogram')\nplt.xlabel('time')\nplt.ylabel('frequency')\nplt.imshow(img[0, 0:360])\nplt.colorbar()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-11-14T08:12:53.138806Z","iopub.execute_input":"2022-11-14T08:12:53.139543Z","iopub.status.idle":"2022-11-14T08:12:54.265638Z","shell.execute_reply.started":"2022-11-14T08:12:53.139503Z","shell.execute_reply":"2022-11-14T08:12:54.264593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"id":"ejqEYZTxgYWP"}},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self, name, *, pretrained=False):\n        \"\"\"\n        name (str): timm model name, e.g. tf_efficientnet_b2_ns\n        \"\"\"\n        super().__init__()\n\n        # Use timm\n        model = timm.create_model(name, pretrained=pretrained, in_chans=2)\n\n        clsf = model.default_cfg['classifier']\n        n_features = model._modules[clsf].in_features\n        model._modules[clsf] = nn.Identity()\n\n        self.fc = nn.Linear(n_features, 1)\n        self.model = model\n\n    def forward(self, x):\n        x = self.model(x)\n        x = self.fc(x)\n        return x","metadata":{"id":"W3JQE_UbgYWP","execution":{"iopub.status.busy":"2022-11-14T08:12:54.267718Z","iopub.execute_input":"2022-11-14T08:12:54.268338Z","iopub.status.idle":"2022-11-14T08:12:54.275925Z","shell.execute_reply.started":"2022-11-14T08:12:54.268297Z","shell.execute_reply":"2022-11-14T08:12:54.274947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict and evaluate","metadata":{"id":"eO6YqnT0gYWQ"}},{"cell_type":"code","source":"def evaluate(model, loader_val, *, compute_score=True, pbar=None):\n    \"\"\"\n    Predict and compute loss and score\n    \"\"\"\n    tb = time.time()\n    was_training = model.training\n    model.eval()\n\n    loss_sum = 0.0\n    n_sum = 0\n    y_all = []\n    y_pred_all = []\n\n    if pbar is not None:\n        pbar = tqdm(desc='Predict', nrows=78, total=pbar)\n\n    for img, y in loader_val:\n        n = y.size(0)\n        img = img.to(device)\n        y = y.to(device)\n\n        with torch.no_grad():\n                y_pred = model(img.to(device))\n\n        loss = criterion(y_pred.view(-1), y)\n\n        n_sum += n\n        loss_sum += n * loss.item()\n\n        y_all.append(y.cpu().detach().numpy())\n        y_pred_all.append(y_pred.sigmoid().squeeze().cpu().detach().numpy())\n\n        if pbar is not None:\n            pbar.update(len(img))\n        \n        del loss, y_pred, img, y\n\n    loss_val = loss_sum / n_sum\n\n    y = np.concatenate(y_all)\n    y_pred = np.concatenate(y_pred_all)\n\n    score = roc_auc_score(y, y_pred) if compute_score else None\n\n    ret = {'loss': loss_val,\n           'score': score,\n           'y': y,\n           'y_pred': y_pred,\n           'time': time.time() - tb}\n    \n    model.train(was_training)  # back to train from eval if necessary\n\n    return ret","metadata":{"id":"B-xWfjWYgYWR","execution":{"iopub.status.busy":"2022-11-14T08:12:54.277491Z","iopub.execute_input":"2022-11-14T08:12:54.27798Z","iopub.status.idle":"2022-11-14T08:12:54.290488Z","shell.execute_reply.started":"2022-11-14T08:12:54.27794Z","shell.execute_reply":"2022-11-14T08:12:54.28948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{"id":"3arVeLkRgYWS"}},{"cell_type":"code","source":"model_name = 'tf_efficientnet_b7_ns'\nnfold = 5\nkfold = KFold(n_splits=nfold, random_state=42, shuffle=True)\n\nepochs = 25\nbatch_size = 16\nnum_workers = 2\nweight_decay = 1e-6\nmax_grad_norm = 1000\n\nlr_max = 4e-4\nepochs_warmup = 1.0\n\n\n## setting of audio data augmentation \nflip_rate = 0.5 # probability of applying the horizontal flip and vertical flip \nfre_shift_rate = 1.0 # probability of applying the vertical shift\ntime_mask_num = 1 # number of time masking\nfreq_mask_num = 2 # number of frequency masking\n","metadata":{"id":"tkWJ1eXpgYWS","outputId":"ed0dd2e0-114f-4dd5-c5a6-2d8511e78359","execution":{"iopub.status.busy":"2022-11-14T08:12:54.292041Z","iopub.execute_input":"2022-11-14T08:12:54.292718Z","iopub.status.idle":"2022-11-14T08:12:54.304322Z","shell.execute_reply.started":"2022-11-14T08:12:54.292683Z","shell.execute_reply":"2022-11-14T08:12:54.303282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict and submit","metadata":{"id":"xRXxBCwRgYWU"}},{"cell_type":"code","source":"# Load model (if necessary)\nsubmit = pd.read_csv(di + '/sample_submission.csv')\nsubmit['target']=0\nif COLAB == False:\n    # Load model (if necessary)\n    for i in range(5):\n        model = Model(model_name, pretrained=False)\n        filename = f'../input/g2net-b7-aug/model{i}.pytorch'\n        model.to(device)\n        model.load_state_dict(torch.load(filename, map_location=device))\n        model.eval()\n\n        # Predict\n        dataset_test = Dataset('test', submit)\n        loader_test = torch.utils.data.DataLoader(dataset_test, batch_size=64,\n                                                num_workers=num_workers, pin_memory=True)\n\n        test = evaluate(model, loader_test, compute_score=False, pbar=len(submit))\n\n        # Write prediction\n        submit['target'] += test['y_pred'] /5\nsubmit.to_csv('submission.csv', index=False)\nprint('target range [%.2f, %.2f]' % (submit['target'].min(), submit['target'].max()))","metadata":{"execution":{"iopub.status.busy":"2022-11-14T08:17:25.508402Z","iopub.execute_input":"2022-11-14T08:17:25.509179Z"},"trusted":true},"execution_count":null,"outputs":[]}]}