{"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":"This notebook shows how Constant Wavelet Transform from waves can be calculated with usage of PyTorchWavelets package (https://github.com/tomrunia/PyTorchWavelets).\n\nNote: due to unfixed bug of the official version, I use fixed version from https://github.com/ar4/PyTorchWavelets/blob/master/wavelets_pytorch/transform.py\n\nPyTorchWavelets is a SciPy/PyTorch implementation for the wavelet analysis outlined in Torrence and Compo (BAMS, 1998). \n\nHave any questions or suggestions? Please comment below.\n\n**<font color='red'>And if you liked this notebook, please upvote it!</font>**\n\n**Changelog**\n* v2 - number of processed samples can be now easily changed via num_samples variable\n* v1 - initial version","metadata":{}},{"cell_type":"markdown","source":"## Import packages","metadata":{}},{"cell_type":"code","source":"!git clone https://github.com/ar4/PyTorchWavelets.git > /dev/null\n%cd PyTorchWavelets\n!pip install -r requirements.txt > /dev/null\n!python setup.py install > /dev/null","metadata":{"execution":{"iopub.status.busy":"2021-08-22T19:43:52.453140Z","iopub.execute_input":"2021-08-22T19:43:52.453799Z","iopub.status.idle":"2021-08-22T19:44:03.420541Z","shell.execute_reply.started":"2021-08-22T19:43:52.453741Z","shell.execute_reply":"2021-08-22T19:44:03.419038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport matplotlib.pyplot as plt\nfrom scipy import signal\nfrom scipy.cluster.vq import whiten\nimport torch\nfrom torch.utils.data import Dataset\nfrom wavelets_pytorch.transform import WaveletTransform # Use WaveletTransformTorch to use with PyTorch\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n%matplotlib inline","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-08-22T19:44:03.422830Z","iopub.execute_input":"2021-08-22T19:44:03.423249Z","iopub.status.idle":"2021-08-22T19:44:03.434324Z","shell.execute_reply.started":"2021-08-22T19:44:03.423207Z","shell.execute_reply":"2021-08-22T19:44:03.432182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_samples = 4 # first N samples to process","metadata":{"execution":{"iopub.status.busy":"2021-08-22T19:44:03.437680Z","iopub.execute_input":"2021-08-22T19:44:03.438069Z","iopub.status.idle":"2021-08-22T19:44:03.448048Z","shell.execute_reply.started":"2021-08-22T19:44:03.438032Z","shell.execute_reply":"2021-08-22T19:44:03.447122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define dataset","metadata":{}},{"cell_type":"markdown","source":"Let's define a dataset to work with.","metadata":{}},{"cell_type":"code","source":"class G2NetDataset(Dataset):\n    def __init__(self, paths, targets, use_filter=True): \n        self.paths = paths\n        self.targets = targets\n        self.use_filter = use_filter\n        if self.use_filter:\n            self.bHP, self.aHP = signal.butter(8, (20, 500), btype='bandpass', fs=2048)\n\n    def __len__(self):\n        return len(self.paths)\n    \n    def __getitem__(self, index):      \n        waves = np.load(self.paths[index])\n        waves = np.concatenate(waves, axis=0)\n        if self.use_filter:\n            waves *= signal.tukey(4096*3, 0.2)\n            waves = signal.filtfilt(self.bHP, self.aHP, waves)\n        waves = waves / np.max(waves)\n        targets = self.targets[index]\n                \n        return {\n            \"waves\": torch.tensor(waves, dtype=torch.float),\n            \"target\": torch.tensor(targets, dtype=torch.long),\n        }","metadata":{"execution":{"iopub.status.busy":"2021-08-22T19:44:03.449806Z","iopub.execute_input":"2021-08-22T19:44:03.450272Z","iopub.status.idle":"2021-08-22T19:44:03.464111Z","shell.execute_reply.started":"2021-08-22T19:44:03.450226Z","shell.execute_reply":"2021-08-22T19:44:03.463013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Read training labels","metadata":{}},{"cell_type":"markdown","source":"Now we read training labels data, and get npy paths.","metadata":{}},{"cell_type":"code","source":"ROOT_DIR = '/kaggle/input/g2net-gravitational-wave-detection'\ndf = pd.read_csv(os.path.join(ROOT_DIR, 'training_labels.csv'))\ndf['path'] = df['id'].apply(lambda x: f'{ROOT_DIR}/train/{x[0]}/{x[1]}/{x[2]}/{x}.npy')","metadata":{"execution":{"iopub.status.busy":"2021-08-22T19:44:03.465768Z","iopub.execute_input":"2021-08-22T19:44:03.466173Z","iopub.status.idle":"2021-08-22T19:44:04.273429Z","shell.execute_reply.started":"2021-08-22T19:44:03.466136Z","shell.execute_reply":"2021-08-22T19:44:04.272576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Demonstrate CWT usage","metadata":{}},{"cell_type":"markdown","source":"Let's calculate CWT for 4 first signals with and without usage of a bandpass filter (20-500Hz), and plot results!","metadata":{}},{"cell_type":"code","source":"transform = WaveletTransform(dt=0.1)  \n\nds = G2NetDataset(df['path'], df['target'], use_filter=False)\nds_f = G2NetDataset(df['path'], df['target'], use_filter=True)\n\nwaves = []\nwaves_f = []\ncwts = []\ncwts_f = []\nfor i in range(num_samples):\n    waves.append(ds.__getitem__(i)['waves'])\n    waves_f.append(ds_f.__getitem__(i)['waves'])\n    cwts.append(transform.power(waves[i]).squeeze())\n    cwts_f.append(transform.power(waves_f[i]).squeeze())","metadata":{"execution":{"iopub.status.busy":"2021-08-22T19:44:04.274891Z","iopub.execute_input":"2021-08-22T19:44:04.275444Z","iopub.status.idle":"2021-08-22T19:44:09.428087Z","shell.execute_reply.started":"2021-08-22T19:44:04.275410Z","shell.execute_reply":"2021-08-22T19:44:09.426951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Without a filter","metadata":{}},{"cell_type":"code","source":"fig, axs = plt.subplots(num_samples)\nfig.set_figheight(15)\nfig.set_figwidth(15)\nfor i in range(num_samples):\n    nid = df['id'][i]\n    ntarget = df['target'][i]\n    axs[i].title.set_text(f'{nid}.npy, target: {ntarget}')\n    axs[i].plot(waves[i])","metadata":{"execution":{"iopub.status.busy":"2021-08-22T19:44:09.429506Z","iopub.execute_input":"2021-08-22T19:44:09.429913Z","iopub.status.idle":"2021-08-22T19:44:11.378342Z","shell.execute_reply.started":"2021-08-22T19:44:09.429878Z","shell.execute_reply":"2021-08-22T19:44:11.377295Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axs = plt.subplots(num_samples)\nfig.set_figheight(15)\nfig.set_figwidth(15)\nfor i in range(num_samples):\n    nid = df['id'][i]\n    ntarget = df['target'][i]\n    axs[i].title.set_text(f'{nid}.npy, target: {ntarget}')\n    axs[i].pcolormesh(cwts[i])","metadata":{"execution":{"iopub.status.busy":"2021-08-22T19:44:42.724844Z","iopub.execute_input":"2021-08-22T19:44:42.725221Z","iopub.status.idle":"2021-08-22T19:44:48.466406Z","shell.execute_reply.started":"2021-08-22T19:44:42.725187Z","shell.execute_reply":"2021-08-22T19:44:48.465613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### With a filter with Tukey window","metadata":{}},{"cell_type":"code","source":"fig, axs = plt.subplots(num_samples)\nfig.set_figheight(15)\nfig.set_figwidth(15)\nfor i in range(num_samples):\n    nid = df['id'][i]\n    ntarget = df['target'][i]\n    axs[i].title.set_text(f'{nid}.npy, target: {ntarget}')\n    axs[i].plot(waves_f[i])","metadata":{"execution":{"iopub.status.busy":"2021-08-22T19:44:14.588673Z","iopub.execute_input":"2021-08-22T19:44:14.589058Z","iopub.status.idle":"2021-08-22T19:44:16.496574Z","shell.execute_reply.started":"2021-08-22T19:44:14.589022Z","shell.execute_reply":"2021-08-22T19:44:16.495336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axs = plt.subplots(num_samples)\nfig.set_figheight(15)\nfig.set_figwidth(15)\nfor i in range(num_samples):\n    nid = df['id'][i]\n    ntarget = df['target'][i]\n    axs[i].title.set_text(f'{nid}.npy, target: {ntarget}')\n    axs[i].pcolormesh(cwts_f[i])","metadata":{"execution":{"iopub.status.busy":"2021-08-22T19:44:16.498351Z","iopub.execute_input":"2021-08-22T19:44:16.499108Z","iopub.status.idle":"2021-08-22T19:44:22.313273Z","shell.execute_reply.started":"2021-08-22T19:44:16.499059Z","shell.execute_reply":"2021-08-22T19:44:22.311394Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"You can use WaveletTransformTorch() as your model block to convert waves to CWT on-the-fly in PyTorch models.","metadata":{"execution":{"iopub.status.busy":"2021-07-03T14:29:22.62631Z","iopub.execute_input":"2021-07-03T14:29:22.626623Z","iopub.status.idle":"2021-07-03T14:29:22.632544Z","shell.execute_reply.started":"2021-07-03T14:29:22.626594Z","shell.execute_reply":"2021-07-03T14:29:22.63147Z"}}}]}