{"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":"code","source":"!pip install torch --upgrade --quiet","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:41:56.374670Z","iopub.execute_input":"2021-08-24T08:41:56.375345Z","iopub.status.idle":"2021-08-24T08:42:04.991826Z","shell.execute_reply.started":"2021-08-24T08:41:56.375299Z","shell.execute_reply":"2021-08-24T08:42:04.990016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfrom tqdm.auto import tqdm\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-08-24T08:42:04.994187Z","iopub.execute_input":"2021-08-24T08:42:04.994640Z","iopub.status.idle":"2021-08-24T08:42:05.009647Z","shell.execute_reply.started":"2021-08-24T08:42:04.994582Z","shell.execute_reply":"2021-08-24T08:42:05.008442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**There some difference in the `fmin` value**. Originally it was 20Hz and I slightly raise it to **21.83Hz** for this version. Lowering the value slightly will make bright part of the image dimmer (speaking in terms of image rather than frequency since easy to visualize) while raising the frequency slightly will brighten the strongest part, and some of the background noise on the RHS of the picture will also brighten into existence. \n\nIf you'd like to make your own dataset consider tuning this value to which you see fit. It might or might not fit better with brighter or dimmer value. \n\nSecond thing is `n_bins`. Tuning this too high will cause it to exceed the nyquist limit, while too low might have some bright image darkens. Consider tuning this as well. One changes it **from 55 to 63** to try out the difference. \n\nOf course, this is not a confirmation. Some of the bright image will dim out when increasing `fmin` and/or `n_bins`, hence this requires some experimentation. ","metadata":{}},{"cell_type":"code","source":"import fastai\nimport torch\nfastai.__version__","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:42:05.011719Z","iopub.execute_input":"2021-08-24T08:42:05.012394Z","iopub.status.idle":"2021-08-24T08:42:06.273565Z","shell.execute_reply.started":"2021-08-24T08:42:05.012352Z","shell.execute_reply":"2021-08-24T08:42:06.272397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\nimport pathlib\n\nhead = pathlib.Path(\"../input/g2net-gravitational-wave-detection\")\n\ntrain_files = sorted(glob.glob(\"../input/g2net-gravitational-wave-detection/train/*/*/*/*.npy\"))","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:42:06.275884Z","iopub.execute_input":"2021-08-24T08:42:06.276397Z","iopub.status.idle":"2021-08-24T08:45:31.740382Z","shell.execute_reply.started":"2021-08-24T08:42:06.276347Z","shell.execute_reply":"2021-08-24T08:45:31.739186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wave = np.load(train_files[0])","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:31.742114Z","iopub.execute_input":"2021-08-24T08:45:31.742472Z","iopub.status.idle":"2021-08-24T08:45:31.757173Z","shell.execute_reply.started":"2021-08-24T08:45:31.742436Z","shell.execute_reply":"2021-08-24T08:45:31.755962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import librosa\nimport librosa.display\nimport matplotlib.pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:31.758615Z","iopub.execute_input":"2021-08-24T08:45:31.758954Z","iopub.status.idle":"2021-08-24T08:45:34.277469Z","shell.execute_reply.started":"2021-08-24T08:45:31.758919Z","shell.execute_reply":"2021-08-24T08:45:34.276266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from numba import njit, jit, cuda, guvectorize\n\n@njit(nogil=True)\ndef min_max_scaler(wave):\n    for i in range(len(wave)):\n        wave[i] = (wave[i] - min(wave[i])) / (max(wave[i]) - min(wave[i]))\n        wave[i] = 2 * wave[i] - 1\n        \n    return wave","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:34.279077Z","iopub.execute_input":"2021-08-24T08:45:34.279460Z","iopub.status.idle":"2021-08-24T08:45:34.382584Z","shell.execute_reply.started":"2021-08-24T08:45:34.279426Z","shell.execute_reply":"2021-08-24T08:45:34.381662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wave1 = min_max_scaler(wave)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:34.383815Z","iopub.execute_input":"2021-08-24T08:45:34.384333Z","iopub.status.idle":"2021-08-24T08:45:35.698322Z","shell.execute_reply.started":"2021-08-24T08:45:34.384295Z","shell.execute_reply":"2021-08-24T08:45:35.697265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(dpi=120)\nfor i in range(len(wave)):\n    plt.plot(range(len(wave[i])), wave[i], label=f\"label_{i}\")\nplt.legend()","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:35.701162Z","iopub.execute_input":"2021-08-24T08:45:35.701720Z","iopub.status.idle":"2021-08-24T08:45:36.074848Z","shell.execute_reply.started":"2021-08-24T08:45:35.701682Z","shell.execute_reply":"2021-08-24T08:45:36.073393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Bandpass filter","metadata":{}},{"cell_type":"code","source":"from scipy.signal import butter, filtfilt, sosfiltfilt\n# from torchaudio.functional import bandpass_biquad\n\nT = 2 # sample period, s\nfs = 2048.0  # sample rate, Hz\ncutoff = 2.5  # desired cutoff frequency, slightly higher than actual 3 sine wave / 2 s = 1.5\n\nnyq = 0.5 * fs  # Nyquist frequency\n\norder = 3  # sine wave approx as quadratic\nn = int(T * fs)\nnormal_cutoff = cutoff / nyq","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:36.077666Z","iopub.execute_input":"2021-08-24T08:45:36.078148Z","iopub.status.idle":"2021-08-24T08:45:36.086248Z","shell.execute_reply.started":"2021-08-24T08:45:36.078095Z","shell.execute_reply":"2021-08-24T08:45:36.084965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def butter_bandpass_filter_torch(data, lowcut, highcut, fs):\n    return bandpass_biquad(data, fs, (highcut + lowcut) / 2, (highcut - lowcut) / (highcut + lowcut))","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:36.088171Z","iopub.execute_input":"2021-08-24T08:45:36.088688Z","iopub.status.idle":"2021-08-24T08:45:36.106181Z","shell.execute_reply.started":"2021-08-24T08:45:36.088626Z","shell.execute_reply":"2021-08-24T08:45:36.105024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# normal_cutoff = (21.83/fs, 500/fs)\n# def butter_bandpass_filter(data, normal_cutoff, fs, order=2):\n#     b, a = butter(order, normal_cutoff, btype=\"bandpass\", analog=False)\n#     y = filtfilt(b, a, data)\n#     return y","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:36.107929Z","iopub.execute_input":"2021-08-24T08:45:36.108564Z","iopub.status.idle":"2021-08-24T08:45:36.121381Z","shell.execute_reply.started":"2021-08-24T08:45:36.108495Z","shell.execute_reply":"2021-08-24T08:45:36.120351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def butter_bandpass_filter(data, low, high, fs, order):\n    sos = butter(order, [low, high], btype=\"bandpass\", output=\"sos\", fs=fs)\n    normalization = np.sqrt((high - low) / (fs / 2))\n    return sosfiltfilt(sos, data) / normalization","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:36.123077Z","iopub.execute_input":"2021-08-24T08:45:36.123669Z","iopub.status.idle":"2021-08-24T08:45:36.136090Z","shell.execute_reply.started":"2021-08-24T08:45:36.123615Z","shell.execute_reply":"2021-08-24T08:45:36.134835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def butter_lowpass_filter(data, normal_cutoff, fs, order):\n    \n    # Get filter coeff\n    b, a = butter(order, normal_cutoff, btype=\"lowpass\", analog=False)\n    y = filtfilt(b, a, data)\n    \n    return y","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:36.138078Z","iopub.execute_input":"2021-08-24T08:45:36.138457Z","iopub.status.idle":"2021-08-24T08:45:36.151125Z","shell.execute_reply.started":"2021-08-24T08:45:36.138422Z","shell.execute_reply":"2021-08-24T08:45:36.150044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# y = min_max_scaler(butter_bandpass_filter(wave, normal_cutoff, fs, 3))\ndata = torch.from_numpy(wave)\ny = butter_bandpass_filter(data, 21.83, 500, fs, 4)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:36.152695Z","iopub.execute_input":"2021-08-24T08:45:36.153352Z","iopub.status.idle":"2021-08-24T08:45:36.171302Z","shell.execute_reply.started":"2021-08-24T08:45:36.153312Z","shell.execute_reply":"2021-08-24T08:45:36.169903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(dpi=120)\nplt.plot(range(len(wave[0])), y[0])","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:36.173344Z","iopub.execute_input":"2021-08-24T08:45:36.173844Z","iopub.status.idle":"2021-08-24T08:45:36.375487Z","shell.execute_reply.started":"2021-08-24T08:45:36.173807Z","shell.execute_reply":"2021-08-24T08:45:36.374174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(dpi=120)\nplt.plot(range(len(wave[0])), y[0])","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:36.377426Z","iopub.execute_input":"2021-08-24T08:45:36.377911Z","iopub.status.idle":"2021-08-24T08:45:36.583773Z","shell.execute_reply.started":"2021-08-24T08:45:36.377854Z","shell.execute_reply":"2021-08-24T08:45:36.582583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Continuation","metadata":{}},{"cell_type":"code","source":"from scipy.signal import spectrogram\n\nplt.figure(dpi=120)\nfor i in range(len(wave)):\n    f, t, Sxx = spectrogram(wave1[i], fs=10)\n    plt.pcolormesh(t, f, Sxx, shading=\"gouraud\")","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:36.585109Z","iopub.execute_input":"2021-08-24T08:45:36.585422Z","iopub.status.idle":"2021-08-24T08:45:36.862613Z","shell.execute_reply.started":"2021-08-24T08:45:36.585393Z","shell.execute_reply":"2021-08-24T08:45:36.861499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plt.figure(dpi=120)\n# f, t, Sxx = spectrogram(wave1[0], fs=4096)\n# plt.pcolormesh(t, fftshift(f), fftshift(Sxx), shading=\"gouraud\")","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:36.863888Z","iopub.execute_input":"2021-08-24T08:45:36.864187Z","iopub.status.idle":"2021-08-24T08:45:36.868687Z","shell.execute_reply.started":"2021-08-24T08:45:36.864158Z","shell.execute_reply":"2021-08-24T08:45:36.867556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def wrapper_plot(m):\n    plt.figure(dpi=120)\n    m()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:36.869945Z","iopub.execute_input":"2021-08-24T08:45:36.870261Z","iopub.status.idle":"2021-08-24T08:45:36.885907Z","shell.execute_reply.started":"2021-08-24T08:45:36.870232Z","shell.execute_reply":"2021-08-24T08:45:36.884474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"stacked = []\nfor j in range(len(wave1)):\n    melspec = librosa.feature.melspectrogram(wave1[j], sr=4096, n_mels=128, fmin=21.83, fmax=2048)\n    melspec = librosa.power_to_db(melspec)\n    melspec = melspec.transpose((1, 0))\n    stacked.append(melspec)\nimage = np.vstack(stacked)\nwrapper_plot(lambda: plt.imshow(image))","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:36.887353Z","iopub.execute_input":"2021-08-24T08:45:36.887712Z","iopub.status.idle":"2021-08-24T08:45:37.164943Z","shell.execute_reply.started":"2021-08-24T08:45:36.887678Z","shell.execute_reply":"2021-08-24T08:45:37.163467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"t.min()","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:37.166820Z","iopub.execute_input":"2021-08-24T08:45:37.167321Z","iopub.status.idle":"2021-08-24T08:45:37.175825Z","shell.execute_reply.started":"2021-08-24T08:45:37.167266Z","shell.execute_reply":"2021-08-24T08:45:37.174422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Finish playing\nNow is time to use dataset created by Y. Nakama and continue. ","metadata":{}},{"cell_type":"code","source":"# X = np.load(\"../input/g2net-n-mels-128-train-images-aggregated/X.npy\")","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:37.177317Z","iopub.execute_input":"2021-08-24T08:45:37.177707Z","iopub.status.idle":"2021-08-24T08:45:37.192725Z","shell.execute_reply.started":"2021-08-24T08:45:37.177659Z","shell.execute_reply":"2021-08-24T08:45:37.191664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# y = np.load(\"../input/g2net-n-mels-128-train-images-aggregated/y.npy\")","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:37.193988Z","iopub.execute_input":"2021-08-24T08:45:37.194366Z","iopub.status.idle":"2021-08-24T08:45:37.210362Z","shell.execute_reply.started":"2021-08-24T08:45:37.194328Z","shell.execute_reply":"2021-08-24T08:45:37.208596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# X.shape","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:37.212911Z","iopub.execute_input":"2021-08-24T08:45:37.213574Z","iopub.status.idle":"2021-08-24T08:45:37.224931Z","shell.execute_reply.started":"2021-08-24T08:45:37.213501Z","shell.execute_reply":"2021-08-24T08:45:37.223711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q nnAudio","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:37.226402Z","iopub.execute_input":"2021-08-24T08:45:37.226791Z","iopub.status.idle":"2021-08-24T08:45:44.370156Z","shell.execute_reply.started":"2021-08-24T08:45:37.226757Z","shell.execute_reply":"2021-08-24T08:45:44.368418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from nnAudio.Spectrogram import *\nimport torch","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:44.375992Z","iopub.execute_input":"2021-08-24T08:45:44.376467Z","iopub.status.idle":"2021-08-24T08:45:44.395893Z","shell.execute_reply.started":"2021-08-24T08:45:44.376424Z","shell.execute_reply":"2021-08-24T08:45:44.394662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# @njit(nogil=True)\n# def min_max_scaler_hstack(wave):\n#     for i in range(len(wave)):\n#         wave[i] = (wave[i] - min(wave[i])) / (max(wave[i]) - min(wave[i]))\n        \n#     wave = np.hstack(wave)\n#     return wave","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:44.398220Z","iopub.execute_input":"2021-08-24T08:45:44.398564Z","iopub.status.idle":"2021-08-24T08:45:44.403380Z","shell.execute_reply.started":"2021-08-24T08:45:44.398515Z","shell.execute_reply":"2021-08-24T08:45:44.402072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ngc.collect()\n# import torch\n# torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:44.404742Z","iopub.execute_input":"2021-08-24T08:45:44.405187Z","iopub.status.idle":"2021-08-24T08:45:44.668498Z","shell.execute_reply.started":"2021-08-24T08:45:44.405154Z","shell.execute_reply":"2021-08-24T08:45:44.665845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"normal_cutoff = (20/nyq, 500/nyq)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:44.670099Z","iopub.execute_input":"2021-08-24T08:45:44.670436Z","iopub.status.idle":"2021-08-24T08:45:44.679235Z","shell.execute_reply.started":"2021-08-24T08:45:44.670404Z","shell.execute_reply":"2021-08-24T08:45:44.678018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from scipy.signal import cwt, ricker","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:44.680675Z","iopub.execute_input":"2021-08-24T08:45:44.681037Z","iopub.status.idle":"2021-08-24T08:45:44.694366Z","shell.execute_reply.started":"2021-08-24T08:45:44.681004Z","shell.execute_reply":"2021-08-24T08:45:44.692785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wave.shape","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:44.696087Z","iopub.execute_input":"2021-08-24T08:45:44.696440Z","iopub.status.idle":"2021-08-24T08:45:44.711097Z","shell.execute_reply.started":"2021-08-24T08:45:44.696407Z","shell.execute_reply":"2021-08-24T08:45:44.710294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import time","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:44.712256Z","iopub.execute_input":"2021-08-24T08:45:44.712585Z","iopub.status.idle":"2021-08-24T08:45:44.724488Z","shell.execute_reply.started":"2021-08-24T08:45:44.712544Z","shell.execute_reply":"2021-08-24T08:45:44.723389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Taken from https://www.kaggle.com/anjum48/continuous-wavelet-transform-cwt-in-pytorch\n\nclass CWT(nn.Module):\n    def __init__(\n        self,\n        widths,\n        wavelet=\"ricker\",\n        channels=1,\n        filter_len=2000,\n        bs=1,\n    ):\n        \"\"\"PyTorch implementation of a continuous wavelet transform.\n\n        Args:\n            widths (iterable): The wavelet scales to use, e.g. np.arange(1, 33)\n            wavelet (str, optional): Name of wavelet. Either \"ricker\" or \"morlet\".\n            Defaults to \"ricker\".\n            channels (int, optional): Number of audio channels in the input. Defaults to 3.\n            filter_len (int, optional): Size of the wavelet filter bank. Set to\n            the number of samples but can be smaller to save memory. Defaults to 2000.\n        \"\"\"\n        super().__init__()\n        self.widths = torch.from_numpy(widths)\n        self.wavelet = getattr(self, wavelet)\n        self.filter_len = filter_len\n        self.bs = bs\n        self.channels = channels\n        self.wavelet_bank = self._build_wavelet_bank()\n\n    def ricker(self, points, a):\n        # https://github.com/scipy/scipy/blob/v1.7.1/scipy/signal/wavelets.py#L262-L306\n        a = torch.Tensor([a])\n        A = 2 / (torch.sqrt(3 * a) * (np.pi ** 0.25))\n        wsq = a ** 2\n        vec = torch.arange(0, points) - (points - 1.0) / 2\n        xsq = vec ** 2\n        mod = 1 - xsq / wsq\n        gauss = torch.exp(-xsq / (2 * wsq))\n        total = A * mod * gauss\n        return total\n\n    def morlet(self, points, s):\n        s = torch.Tensor([s])\n        x = torch.arange(0, points) - (points - 1.0) / 2\n        x = x / s\n        # https://pywavelets.readthedocs.io/en/latest/ref/cwt.html#morlet-wavelet\n        wavelet = torch.exp(-(x ** 2.0) / 2.0) * torch.cos(5.0 * x)\n        output = torch.sqrt(1 / s) * wavelet\n        return output\n\n    def cmorlet(self, points, s, wavelet_width=1, center_freq=1):\n        # https://pywavelets.readthedocs.io/en/latest/ref/cwt.html#complex-morlet-wavelets\n        s = torch.Tensor([s])\n        x = torch.arange(0, points) - (points - 1.0) / 2\n        x = x / s\n        norm_constant = torch.sqrt(torch.Tensor([np.pi * wavelet_width]))\n        exp_term = torch.exp(-(x ** 2) / wavelet_width)\n        kernel_base = exp_term / norm_constant\n#         kernel = kernel_base * torch.exp(1j * 2 * np.pi * center_freq * x)\n        kernel_real = kernel_base * torch.cos(2 * np.pi * center_freq * x)\n        kernel_imag = kernel_base * torch.sin(2 * np.pi * center_freq * x)\n        return kernel_real, kernel_imag\n\n    def _build_wavelet_bank(self):\n        wavelet_bank_real = []\n        wavelet_bank_imag = []\n        for w in self.widths:\n            wavelet_bank = self.wavelet(self.filter_len, w)\n            wavelet_bank_real.append(wavelet_bank[0])\n            wavelet_bank_imag.append(wavelet_bank[1])\n#         wavelet_bank = [self.wavelet(self.filter_len, w) for w in self.widths]\n        wavelet_bank_real = torch.stack(wavelet_bank_real)\n        wavelet_bank_imag = torch.stack(wavelet_bank_imag)\n        wavelet_bank_real = wavelet_bank_real.view(\n            wavelet_bank_real.shape[0], 1, 1, wavelet_bank_real.shape[1]\n        )\n        wavelet_bank_imag = wavelet_bank_imag.view(\n            wavelet_bank_imag.shape[0], 1, 1, wavelet_bank_imag.shape[1]\n        )\n        wavelet_bank_real = torch.cat([wavelet_bank_real] * self.channels, 2)\n        wavelet_bank_imag = torch.cat([wavelet_bank_imag] * self.channels, 2)\n#         wavelet_bank_real = torch.cat([wavelet_bank_real] * self.bs, 1)\n#         wavelet_bank_imag = torch.cat([wavelet_bank_imag] * self.bs, 1)\n        return wavelet_bank_real, wavelet_bank_imag\n        \n\n#     def _build_wavelet_bank(self):\n#         \"\"\"This function builds a 2D wavelet filter using wavelets at different scales\n\n#         Returns:\n#             tensor: Tensor of shape (num_widths, 1, channels, filter_len)\n#         \"\"\"\n#         wavelet_bank = [\n#             torch.conj(torch.flip(self.wavelet(self.filter_len, w), [-1]))\n#             for w in self.widths\n#         ]\n#         wavelet_bank = torch.stack(wavelet_bank)\n#         wavelet_bank = wavelet_bank.view(\n#             wavelet_bank.shape[0], 1, 1, wavelet_bank.shape[1]\n#         )\n#         wavelet_bank = torch.cat([wavelet_bank] * self.channels, 2)\n#         return wavelet_bank\n\n    def forward(self, x):\n        \"\"\"Compute CWT arrays from a batch of multi-channel inputs\n\n        Args:\n            x (torch.tensor): Tensor of shape (batch_size, channels, time)\n\n        Returns:\n            torch.tensor: Tensor of shape (batch_size, channels, widths, time)\n        \"\"\"\n        x = x.unsqueeze(1)\n#         if self.wavelet_bank.is_complex():\n        if type(self.wavelet_bank) == tuple:\n#             wavelet_real = self.wavelet_bank.real.to(device=x.device, dtype=x.dtype)\n#             wavelet_imag = self.wavelet_bank.imag.to(device=x.device, dtype=x.dtype)\n            wavelet_real = self.wavelet_bank[0].to(device=x.device, dtype=x.dtype)\n            wavelet_imag = self.wavelet_bank[1].to(device=x.device, dtype=x.dtype)\n\n            output_real = nn.functional.conv2d(x, wavelet_real, padding=\"same\")\n            output_imag = nn.functional.conv2d(x, wavelet_imag, padding=\"same\")\n            output_real = torch.transpose(output_real, 1, 2)\n            output_imag = torch.transpose(output_imag, 1, 2)\n#             return torch.complex(output_real, output_imag)\n            return torch.sqrt(output_real**2 + output_imag**2)\n        else:\n            self.wavelet_bank = self.wavelet_bank.to(device=x.device, dtype=x.dtype)\n            output = nn.functional.conv2d(x, self.wavelet_bank, padding=\"same\")\n            return torch.transpose(output, 1, 2)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:44.726060Z","iopub.execute_input":"2021-08-24T08:45:44.726704Z","iopub.status.idle":"2021-08-24T08:45:44.753182Z","shell.execute_reply.started":"2021-08-24T08:45:44.726644Z","shell.execute_reply":"2021-08-24T08:45:44.752282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"widths = np.arange(25, 89)\npycwt = CWT(widths, \"cmorlet\", 3, 4096)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:44.754385Z","iopub.execute_input":"2021-08-24T08:45:44.754918Z","iopub.status.idle":"2021-08-24T08:45:44.805723Z","shell.execute_reply.started":"2021-08-24T08:45:44.754870Z","shell.execute_reply":"2021-08-24T08:45:44.804854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wavelet_bank_real = pycwt.wavelet_bank[0]\nwavelet_bank_real.shape","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:44.806986Z","iopub.execute_input":"2021-08-24T08:45:44.807521Z","iopub.status.idle":"2021-08-24T08:45:44.814079Z","shell.execute_reply.started":"2021-08-24T08:45:44.807469Z","shell.execute_reply":"2021-08-24T08:45:44.812859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bs = 16\ntorch.cat([wavelet_bank_real] * bs, 1).shape","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:44.815875Z","iopub.execute_input":"2021-08-24T08:45:44.816558Z","iopub.status.idle":"2021-08-24T08:45:44.856413Z","shell.execute_reply.started":"2021-08-24T08:45:44.816494Z","shell.execute_reply":"2021-08-24T08:45:44.855637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs = []\nfor i in range(1, 5): imgs.append(np.load(train_files[i]))\nimgs = torch.from_numpy(np.array(imgs))\nimgs.shape","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:44.857953Z","iopub.execute_input":"2021-08-24T08:45:44.858645Z","iopub.status.idle":"2021-08-24T08:45:44.932056Z","shell.execute_reply.started":"2021-08-24T08:45:44.858604Z","shell.execute_reply":"2021-08-24T08:45:44.931176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%timeit our_imgs = pycwt(imgs)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:45:44.933844Z","iopub.execute_input":"2021-08-24T08:45:44.934720Z","iopub.status.idle":"2021-08-24T08:46:50.654301Z","shell.execute_reply.started":"2021-08-24T08:45:44.934672Z","shell.execute_reply":"2021-08-24T08:46:50.653326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs = []\nfor i in range(1, 9): imgs.append(np.load(train_files[i]))\nimgs = torch.from_numpy(np.array(imgs))\nimgs.shape","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:46:50.656282Z","iopub.execute_input":"2021-08-24T08:46:50.657077Z","iopub.status.idle":"2021-08-24T08:46:50.762224Z","shell.execute_reply.started":"2021-08-24T08:46:50.657030Z","shell.execute_reply":"2021-08-24T08:46:50.760983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%timeit _ = pycwt(imgs)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:46:50.763727Z","iopub.execute_input":"2021-08-24T08:46:50.764063Z","iopub.status.idle":"2021-08-24T08:48:44.984185Z","shell.execute_reply.started":"2021-08-24T08:46:50.764027Z","shell.execute_reply":"2021-08-24T08:48:44.982827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"our_imgs = pycwt(imgs)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:51:41.984790Z","iopub.execute_input":"2021-08-24T08:51:41.985216Z","iopub.status.idle":"2021-08-24T08:51:59.933375Z","shell.execute_reply.started":"2021-08-24T08:51:41.985158Z","shell.execute_reply":"2021-08-24T08:51:59.932142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@njit(nogil=True)\ndef min_max_scaler_int8(wave):\n    return (wave - wave.min()) / (wave.max() - wave.min()) * 255","metadata":{"execution":{"iopub.status.busy":"2021-08-24T09:03:53.071683Z","iopub.execute_input":"2021-08-24T09:03:53.072348Z","iopub.status.idle":"2021-08-24T09:03:53.077457Z","shell.execute_reply.started":"2021-08-24T09:03:53.072293Z","shell.execute_reply":"2021-08-24T09:03:53.076625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def image_to_int8(image): \n#     g = (image - image.min()) / (image.max() - image.min())\n    return np.round_(min_max_scaler_int8(image)).astype(np.uint8)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T09:02:32.406124Z","iopub.execute_input":"2021-08-24T09:02:32.406568Z","iopub.status.idle":"2021-08-24T09:02:32.411643Z","shell.execute_reply.started":"2021-08-24T09:02:32.406506Z","shell.execute_reply":"2021-08-24T09:02:32.410773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m = our_imgs[0, 2].numpy().copy()","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:54:55.876683Z","iopub.execute_input":"2021-08-24T08:54:55.877113Z","iopub.status.idle":"2021-08-24T08:54:55.882013Z","shell.execute_reply.started":"2021-08-24T08:54:55.877078Z","shell.execute_reply":"2021-08-24T08:54:55.881142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"our_imgs[1, 1].numpy()","metadata":{"execution":{"iopub.status.busy":"2021-08-24T09:02:55.876953Z","iopub.execute_input":"2021-08-24T09:02:55.877861Z","iopub.status.idle":"2021-08-24T09:02:55.885494Z","shell.execute_reply.started":"2021-08-24T09:02:55.877817Z","shell.execute_reply":"2021-08-24T09:02:55.884631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\nImage.fromarray(image_to_int8(our_imgs[1, 1].numpy())).convert(\"RGB\").resize((400, 300)).save(\"data.png\")","metadata":{"execution":{"iopub.status.busy":"2021-08-24T09:07:00.169256Z","iopub.execute_input":"2021-08-24T09:07:00.169745Z","iopub.status.idle":"2021-08-24T09:07:00.406991Z","shell.execute_reply.started":"2021-08-24T09:07:00.169710Z","shell.execute_reply":"2021-08-24T09:07:00.406079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"With numpy: 2.19s.  \nWith pytorch (no GPU) also around 2.18s.  \nWithout using complex numbers: 740ms. (but with slightly different output).","metadata":{}},{"cell_type":"markdown","source":"with all: $2.16 s\\pm 16.4 ms$  \nwithout norm-const: $2.2 s \\pm 108 ms$  \nwithout exponential: $2.18 s \\pm 61.9 ms$  \nwithout kernel calc: $1.07 s \\pm 4.39 ms$","metadata":{}},{"cell_type":"code","source":"%timeit pycwt(torch.from_numpy(wave).view(1, 3, 4096))","metadata":{"execution":{"iopub.status.busy":"2021-08-24T06:52:24.687886Z","iopub.execute_input":"2021-08-24T06:52:24.688427Z","iopub.status.idle":"2021-08-24T06:52:39.388919Z","shell.execute_reply.started":"2021-08-24T06:52:24.688363Z","shell.execute_reply":"2021-08-24T06:52:39.387681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"min_max_scaler(butter_bandpass_filter(wave, 20, 500, fs, 4))","metadata":{"execution":{"iopub.status.busy":"2021-08-19T02:44:54.893902Z","iopub.execute_input":"2021-08-19T02:44:54.894289Z","iopub.status.idle":"2021-08-19T02:44:54.905636Z","shell.execute_reply.started":"2021-08-19T02:44:54.894246Z","shell.execute_reply":"2021-08-19T02:44:54.904702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Currently we are taking mean of all 3 waves. Perhaps there are other methods. ","metadata":{}},{"cell_type":"code","source":"def apply_qtransform(waves, transform=None, cuda=False):\n#     waves *= scipy.signal.tukey(4096, 0.2)\n    waves = min_max_scaler(butter_bandpass_filter(waves, 27.5, 466.16, fs, 4))\n    waves = np.ascontiguousarray(waves)\n#     waves = np.hstack(waves)\n    waves = torch.from_numpy(waves).float().view(1, 3, 4096)\n    if cuda: waves = waves.cuda()\n    image = torch.abs(pycwt(waves))\n    image = torch.mean(image, dim=1).squeeze()  # Get mean of all 3 different waves.\n#     image = transform(waves)\n    return image\n\n\nimgs = []\nfor i in tqdm(range(10)):\n    wave = np.load(train_files[i])\n#     img = apply_qtransform(wave, transform=CQT1992v2(sr=2048, fmin=21.83, fmax=1024, hop_length=64))\n    img = apply_qtransform(wave, cuda=True if torch.cuda.is_available() else False)\n    imgs.append(img)\nprint(img.shape)","metadata":{"execution":{"iopub.status.busy":"2021-08-19T03:24:20.075996Z","iopub.execute_input":"2021-08-19T03:24:20.076355Z","iopub.status.idle":"2021-08-19T03:24:42.059858Z","shell.execute_reply.started":"2021-08-19T03:24:20.076325Z","shell.execute_reply":"2021-08-19T03:24:42.05874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i in range(10):\n    plt.figure(dpi=150)\n    plt.imshow(imgs[i].cpu().numpy().squeeze(), aspect=\"auto\")","metadata":{"execution":{"iopub.status.busy":"2021-08-19T03:20:53.174277Z","iopub.execute_input":"2021-08-19T03:20:53.17632Z","iopub.status.idle":"2021-08-19T03:20:57.55527Z","shell.execute_reply.started":"2021-08-19T03:20:53.176266Z","shell.execute_reply":"2021-08-19T03:20:57.5542Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del imgs\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2021-08-19T03:23:18.574456Z","iopub.execute_input":"2021-08-19T03:23:18.574822Z","iopub.status.idle":"2021-08-19T03:23:18.799378Z","shell.execute_reply.started":"2021-08-19T03:23:18.574791Z","shell.execute_reply":"2021-08-19T03:23:18.798352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.mkdir(\"train/\")\nOUT_DIR = \"train/\"\n\nlabels = pd.read_csv(\"../input/g2net-gravitational-wave-detection/training_labels.csv\")\nlabels[\"file_path\"] = train_files\n\npd.set_option(\"display.max_colwidth\", None)\nlabels.head()","metadata":{"execution":{"iopub.status.busy":"2021-08-19T03:21:05.543922Z","iopub.execute_input":"2021-08-19T03:21:05.544282Z","iopub.status.idle":"2021-08-19T03:21:06.051465Z","shell.execute_reply.started":"2021-08-19T03:21:05.544252Z","shell.execute_reply":"2021-08-19T03:21:06.050419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ones_train = labels[labels[\"target\"] == 1][\"file_path\"].to_numpy()","metadata":{"execution":{"iopub.status.busy":"2021-08-19T03:21:06.052738Z","iopub.execute_input":"2021-08-19T03:21:06.053037Z","iopub.status.idle":"2021-08-19T03:21:06.139641Z","shell.execute_reply.started":"2021-08-19T03:21:06.053009Z","shell.execute_reply":"2021-08-19T03:21:06.138719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_images(file_path, out_dir):\n    file_name = file_path.split('/')[-1].split('.npy')[0]\n    waves = np.load(file_path).astype(np.float32) # (3, 4096)\n    image = apply_qtransform(wave, cuda=True).cpu()\n    plt.imsave(out_dir + file_name + \".png\", image.cpu().numpy().squeeze())","metadata":{"execution":{"iopub.status.busy":"2021-08-19T03:21:13.760021Z","iopub.execute_input":"2021-08-19T03:21:13.760461Z","iopub.status.idle":"2021-08-19T03:21:13.766937Z","shell.execute_reply.started":"2021-08-19T03:21:13.760424Z","shell.execute_reply":"2021-08-19T03:21:13.765862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Saving all the 1's in the 1's folder. \nimport joblib\nfrom tqdm.auto import tqdm\n\nfolder_name = \"train/ones/\"\n\nos.makedirs(folder_name, exist_ok=True)\n\n_ = joblib.Parallel(n_jobs=8, prefer=\"threads\")(\n    joblib.delayed(save_images)(file_path, out_dir=folder_name) for file_path in tqdm(ones_train)\n)","metadata":{"execution":{"iopub.status.busy":"2021-08-19T03:21:16.785772Z","iopub.execute_input":"2021-08-19T03:21:16.786133Z","iopub.status.idle":"2021-08-19T03:22:00.441037Z","shell.execute_reply.started":"2021-08-19T03:21:16.786087Z","shell.execute_reply":"2021-08-19T03:22:00.436331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"folder_name = \"train/zero/\"\nzeroes_train = labels[labels[\"target\"] == 0][\"file_path\"].to_numpy()\n\nos.makedirs(folder_name, exist_ok=True)\n\n_ = joblib.Parallel(n_jobs=8, prefer=\"threads\")(\n    joblib.delayed(save_images)(file_path, out_dir=folder_name) for file_path in tqdm(zeroes_train)\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport shutil\n\ndef move_to_destination(origin, destination, percentage_split):\n    num_images = int(len(os.listdir(origin))*percentage_split)\n    for image_name, image_number in zip(sorted(os.listdir(origin)), range(num_images)):\n        shutil.move(os.path.join(origin, image_name), destination)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.makedirs(\"./valid/ones\")\nos.makedirs(\"./valid/zero\")\nmove_to_destination(\"./train/ones\", \"./valid/ones\", 0.2)\nmove_to_destination(\"./train/zero\", \"./valid/zero\", 0.2)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nimport shutil\n\nshutil.make_archive(\"train/\", 'zip', \"train/\")\nshutil.rmtree(\"train/\")\n\nshutil.make_archive(\"valid/\", \"zip\", \"valid/\")\nshutil.rmtree(\"valid/\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"OUT_DIR = \"test/\"\nos.mkdir(\"test/\")\ntest_files = sorted(glob.glob(\"../input/g2net-gravitational-wave-detection/test/*/*/*/*.npy\"))\n\n_ = joblib.Parallel(n_jobs=8, prefer=\"threads\")(\n    joblib.delayed(save_images)(file_path, out_dir=OUT_DIR) for file_path in tqdm(test_files)\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nshutil.make_archive(\"test/\", 'zip', \"test/\")\nshutil.rmtree(\"test/\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}