{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"dockerImageVersionId":30673,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Main\nThis notebook show how to use [torchlibrosa](https://github.com/qiuqiangkong/torchlibrosa) create spectrogram from raw eeg by nn.Conv Module.","metadata":{}},{"cell_type":"markdown","source":"### Config","metadata":{}},{"cell_type":"code","source":"class CFG:\n    PATH = '/kaggle/input/hms-harmful-brain-activity-classification/'\n    test_eeg = \"/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/\"\n    test_csv = \"/kaggle/input/hms-harmful-brain-activity-classification/test.csv\"\n    seed = 2024\n    batch_size = 32\n    num_workers = 1","metadata":{"execution":{"iopub.status.busy":"2024-03-30T09:10:03.980385Z","iopub.execute_input":"2024-03-30T09:10:03.980762Z","iopub.status.idle":"2024-03-30T09:10:03.985998Z","shell.execute_reply.started":"2024-03-30T09:10:03.980734Z","shell.execute_reply":"2024-03-30T09:10:03.984559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nfrom glob import glob\nimport random\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\nimport torch\nimport warnings \nwarnings.filterwarnings('ignore')\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nfrom matplotlib import pyplot as plt\n\nfrom typing import Dict\nfrom torch.utils.data import DataLoader, Dataset","metadata":{"execution":{"iopub.status.busy":"2024-03-30T09:10:04.215098Z","iopub.execute_input":"2024-03-30T09:10:04.215508Z","iopub.status.idle":"2024-03-30T09:10:04.222957Z","shell.execute_reply.started":"2024-03-30T09:10:04.215474Z","shell.execute_reply":"2024-03-30T09:10:04.221606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(CFG.test_csv)\nprint(f\"Test dataframe shape is: {test_df.shape}\")\ntest_df.head()\n\neeg_parquet_paths = glob(CFG.test_eeg+ \"*.parquet\")\neeg_df = pd.read_parquet(eeg_parquet_paths[0])\neeg_features = eeg_df.columns\nprint(f'There are {len(eeg_features)} raw eeg features')\nprint(list(eeg_features))\neeg_features = ['Fp1', 'F3', 'C3', 'P3', 'F7', 'T3', 'T5', 'O1', 'Fz', 'Cz', 'Pz', 'Fp2', 'F4', 'C4', 'P4', 'F8', 'T4', 'T6', 'O2', 'EKG']\nfeature_to_index = {x:y for x,y in zip(eeg_features, range(len(eeg_features)))}\n","metadata":{"execution":{"iopub.status.busy":"2024-03-30T09:10:04.472208Z","iopub.execute_input":"2024-03-30T09:10:04.472650Z","iopub.status.idle":"2024-03-30T09:10:04.684119Z","shell.execute_reply.started":"2024-03-30T09:10:04.472619Z","shell.execute_reply":"2024-03-30T09:10:04.683176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def eeg_from_parquet(parquet_path: str) -> np.ndarray:\n    \"\"\"\n    This function reads a parquet file and extracts the middle 50 seconds of readings. Then it fills NaN values\n    with the mean value (ignoring NaNs).\n    :param parquet_path: path to parquet file.\n    :param display: whether to display EEG plots or not.\n    :return data: np.array of shape  (time_steps, eeg_features) -> (10_000, 8)\n    \"\"\"\n    # === Extract middle 50 seconds ===\n    eeg = pd.read_parquet(parquet_path, columns=eeg_features)\n    rows = len(eeg)\n    offset = (rows - 10_000) // 2 # 50 * 200 = 10_000\n    eeg = eeg.iloc[offset:offset+10_000] # middle 50 seconds, has the same amount of readings to left and right\n    # === Convert to numpy ===\n    data = np.zeros((10_000, len(eeg_features))) # create placeholder of same shape with zeros\n    for index, feature in enumerate(eeg_features):\n        x = eeg[feature].values.astype('float32') # convert to float32\n        mean = np.nanmean(x) # arithmetic mean along the specified axis, ignoring NaNs\n        nan_percentage = np.isnan(x).mean() # percentage of NaN values in feature\n        # === Fill nan values ===\n        if nan_percentage < 1: # if some values are nan, but not all\n            x = np.nan_to_num(x, nan=mean)\n        else: # if all values are nan\n            x[:] = 0\n        data[:, index] = x\n   \n    return data\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-03-30T09:10:05.718058Z","iopub.execute_input":"2024-03-30T09:10:05.719176Z","iopub.status.idle":"2024-03-30T09:10:05.729074Z","shell.execute_reply.started":"2024-03-30T09:10:05.719135Z","shell.execute_reply":"2024-03-30T09:10:05.727674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load data","metadata":{}},{"cell_type":"code","source":"CREATE_EEGS = False\nall_eegs = {}\nvisualize = 1\neeg_paths = glob(CFG.test_eeg + \"*.parquet\")\neeg_ids = test_df.eeg_id.unique()\n\nfor i, eeg_id in tqdm(enumerate(eeg_ids)):  \n    # Save EEG to Python dictionary of numpy arrays\n    eeg_path = CFG.test_eeg + str(eeg_id) + \".parquet\"\n    data = eeg_from_parquet(eeg_path)              \n    all_eegs[eeg_id] = data\n\nfrom scipy.signal import butter, lfilter\n\ndef butter_lowpass_filter(data, cutoff_freq: int = 20, sampling_rate: int = 200, order: int = 4):\n    nyquist = 0.5 * sampling_rate\n    normal_cutoff = cutoff_freq / nyquist\n    b, a = butter(order, normal_cutoff, btype='low', analog=False)\n    filtered_data = lfilter(b, a, data, axis=0)\n    return filtered_data","metadata":{"execution":{"iopub.status.busy":"2024-03-30T09:10:06.143644Z","iopub.execute_input":"2024-03-30T09:10:06.144069Z","iopub.status.idle":"2024-03-30T09:10:07.687358Z","shell.execute_reply.started":"2024-03-30T09:10:06.144037Z","shell.execute_reply":"2024-03-30T09:10:07.686357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"use_channels_names = [ 'FP1', 'F3', 'C3', 'P3', 'F7', 'T7', 'P7', 'O1', 'FZ', 'CZ', 'PZ', 'FP2', 'F4', 'C4', 'P4', 'F8', 'T8', 'P8', 'O2' ]\n\ndef ch_name_to_index(ls):\n    return [use_channels_names.index(x) for x in ls]\n\nuse_channels_names_two_ref_16ch_tuple = [\n    ('FP1', 'F7'),('F7', 'T7'),('T7', 'P7'),('P7', 'O1'),\n    ('FP1', 'F3'),('F3', 'C3'),('C3', 'P3'),('P3', 'O1'),\n    ('FP2', 'F8'),('F8', 'T8'),('T8', 'P8'),('P8', 'O2'),\n    ('FP2', 'F4'),('F4', 'C4'),('C4', 'P4'),('P4', 'O2'),\n]\nuse_channels_names_two_ref_16ch_pos = [x[0] for x in use_channels_names_two_ref_16ch_tuple]\nuse_channels_names_two_ref_16ch_neg = [x[1] for x in use_channels_names_two_ref_16ch_tuple]\n\nuse_channels_names_two_ref_16ch_pos_idx = ch_name_to_index(use_channels_names_two_ref_16ch_pos)\nuse_channels_names_two_ref_16ch_neg_idx = ch_name_to_index(use_channels_names_two_ref_16ch_neg)\n","metadata":{"execution":{"iopub.status.busy":"2024-03-30T09:10:07.689062Z","iopub.execute_input":"2024-03-30T09:10:07.689589Z","iopub.status.idle":"2024-03-30T09:10:07.698624Z","shell.execute_reply.started":"2024-03-30T09:10:07.689560Z","shell.execute_reply":"2024-03-30T09:10:07.697461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EEGDataset(Dataset):\n    def __init__(\n        self, df: pd.DataFrame, config, mode: str = 'train',\n        eegs: Dict[int, np.ndarray] = all_eegs, downsample: int = None\n    ): \n        self.df = df\n        self.config = config\n        self.mode = mode\n        self.eegs = eegs\n        self.downsample = downsample\n        \n    def __len__(self):\n        \"\"\"\n        Length of dataset.\n        \"\"\"\n        return len(self.df)\n        \n    def __getitem__(self, index):\n        \"\"\"\n        Get one item.\n        \"\"\"\n        X, y = self.__data_generation(index)\n        if self.downsample is not None:\n            X = X[::self.downsample,:]\n        output = {\n            \"eeg\": torch.tensor(X, dtype=torch.float32),\n            \"labels\": torch.tensor(y, dtype=torch.float32)\n        }\n        return output\n                        \n    def __data_generation(self, index):\n        row = self.df.iloc[index]\n        y = np.zeros(6, dtype='float32')\n        data = self.eegs[row.eeg_id]\n        x = np.array(data, dtype='float32')\n        x = np.nan_to_num(x, nan=0)\n        x = x[:,:-1]\n        x = x[:,use_channels_names_two_ref_16ch_pos_idx] - x[:,use_channels_names_two_ref_16ch_neg_idx]\n        \n        x = np.clip(x,-1024,1024)\n        x = np.nan_to_num(x, nan=0)\n        x = x / 1024.0\n\n        # === Butter Low-pass Filter ===\n        x = butter_lowpass_filter(x)\n        \n        x = torch.from_numpy(x).transpose_(0,1).float()\n        x = x.squeeze_(0).transpose_(0,1)\n        x = x.detach_().to(torch.float)\n        if self.mode != 'test':\n            y = row[self.config.target_cols].values.astype(np.float32)\n            \n        return x, y\n\ntest_dataset = EEGDataset(test_df, CFG, mode='test')\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=CFG.batch_size,\n    shuffle=False,\n    num_workers=CFG.num_workers,\n    pin_memory=True,\n    drop_last=False\n)\noutput = test_dataset[0]\nX = output[\"eeg\"]\nprint(f\"X shape: {X.shape}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-03-30T09:10:07.699989Z","iopub.execute_input":"2024-03-30T09:10:07.701066Z","iopub.status.idle":"2024-03-30T09:10:07.784268Z","shell.execute_reply.started":"2024-03-30T09:10:07.701027Z","shell.execute_reply":"2024-03-30T09:10:07.783377Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### [torchlibrosa](https://github.com/qiuqiangkong/torchlibrosa)\n\"TorchLibrosa\" feature extractor the same as librosa.feature.melspectrogram()","metadata":{}},{"cell_type":"code","source":"!pip install torchlibrosa","metadata":{"execution":{"iopub.status.busy":"2024-03-30T09:10:08.903912Z","iopub.execute_input":"2024-03-30T09:10:08.904664Z","iopub.status.idle":"2024-03-30T09:10:25.368541Z","shell.execute_reply.started":"2024-03-30T09:10:08.904629Z","shell.execute_reply":"2024-03-30T09:10:25.367399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torchlibrosa as tl\nimport matplotlib.pyplot as plt\n\nclass SpecModule(torch.nn.Module):\n    def __init__(self,  \n                        sample_rate = 200,\n                        win_length = 256,\n                        hop_length = 10000//256,\n                        n_mels = 256,\n                        n_fft = 3092,\n                        \n                        *args, **kwargs) -> None:\n        super().__init__(*args, **kwargs)\n        \n        self.feature_extractor = torch.nn.Sequential(\n            tl.Spectrogram(\n                n_fft = n_fft,\n                hop_length=hop_length,\n                win_length=win_length,\n                pad_mode=\"constant\",\n            ),\n            tl.LogmelFilterBank(\n                sr=sample_rate,\n                n_fft=n_fft,\n                n_mels=n_mels,\n                fmin=0,\n                fmax=20,\n                ref=1,\n                is_log=True, # Default is true\n            )\n            )\n    def forward(self,x):\n        # x: (B,C,T)\n        B,_,_ = x.shape\n        x = x.flatten(0,1)\n        x = self.feature_extractor(x)\n        \n        _,_,b,c = x.shape\n        x = x.reshape((B,4,4,b,c))\n        x = x.transpose_(-1,-2)\n        x = x.mean(dim=2)\n        return x\n    \nx = torch.zeros((3,16,10000))\nm = SpecModule()\nprint(m(x).shape)","metadata":{"execution":{"iopub.status.busy":"2024-03-30T09:10:25.371694Z","iopub.execute_input":"2024-03-30T09:10:25.372117Z","iopub.status.idle":"2024-03-30T09:10:40.825166Z","shell.execute_reply.started":"2024-03-30T09:10:25.372079Z","shell.execute_reply":"2024-03-30T09:10:40.823991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m = SpecModule()\nm.eval()\nm.to(device)\nplt.figure(figsize=(10,10))\nfor step, batch in enumerate(test_loader):\n    x = batch.pop(\"eeg\").to(device) # send inputs to `device`   \n    x = x.transpose_(1,2)\n    x = m(x).clone().detach().cpu()\n    print(x.shape)\n    for k in range(4):\n        img = x[0,k,:]\n        print(img.mean(),img.std(),img.max(),img.min())\n        mn = x.min()\n        mx = x.max()\n        img = (img-mn)/(mx-mn)\n        plt.subplot(2,2,k+1)\n        plt.imshow(img,aspect='auto',origin='lower')\n        \n        plt.ylabel('Frequencies (Hz)',size=14)\n        plt.xlabel('Time (sec)',size=16)\n    break","metadata":{"execution":{"iopub.status.busy":"2024-03-30T09:10:40.826600Z","iopub.execute_input":"2024-03-30T09:10:40.827025Z","iopub.status.idle":"2024-03-30T09:10:46.304553Z","shell.execute_reply.started":"2024-03-30T09:10:40.826971Z","shell.execute_reply":"2024-03-30T09:10:46.303329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}