{"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":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7392733,"sourceType":"datasetVersion","datasetId":4297749},{"sourceId":7816303,"sourceType":"datasetVersion","datasetId":4579068},{"sourceId":7903196,"sourceType":"datasetVersion","datasetId":4641843},{"sourceId":7929967,"sourceType":"datasetVersion","datasetId":4660964},{"sourceId":7987178,"sourceType":"datasetVersion","datasetId":4663320},{"sourceId":8041950,"sourceType":"datasetVersion","datasetId":4643243},{"sourceId":8057606,"sourceType":"datasetVersion","datasetId":4709890},{"sourceId":8170096,"sourceType":"datasetVersion","datasetId":4689282},{"sourceId":8303282,"sourceType":"datasetVersion","datasetId":4932730},{"sourceId":8305783,"sourceType":"datasetVersion","datasetId":4933986}],"dockerImageVersionId":30665,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## solution : https://www.kaggle.com/competitions/hms-harmful-brain-activity-classification/discussion/492560","metadata":{}},{"cell_type":"code","source":"%env XLA_PYTHON_CLIENT_PREALLOCATE=false\n%env XLA_PYTHON_CLIENT_ALLOCATOR=platform","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:54:58.796643Z","iopub.execute_input":"2024-05-03T21:54:58.797575Z","iopub.status.idle":"2024-05-03T21:54:58.811300Z","shell.execute_reply.started":"2024-05-03T21:54:58.797500Z","shell.execute_reply":"2024-05-03T21:54:58.810208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## yamash model","metadata":{"execution":{"iopub.status.busy":"2024-04-08T06:27:33.138035Z","iopub.execute_input":"2024-04-08T06:27:33.138644Z","iopub.status.idle":"2024-04-08T06:27:33.147977Z","shell.execute_reply.started":"2024-04-08T06:27:33.138608Z","shell.execute_reply":"2024-04-08T06:27:33.14714Z"}}},{"cell_type":"code","source":"import sys\nimport os\nimport gc\nimport copy\nimport yaml\nimport random\nimport shutil\nimport time\nimport typing as tp\nfrom glob import glob\nfrom pathlib import Path\nfrom collections import OrderedDict, defaultdict\nfrom logging import getLogger, INFO, StreamHandler, FileHandler, Formatter\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom tqdm import tqdm\nfrom scipy.signal import butter, lfilter, freqz\nfrom sklearn.model_selection import StratifiedGroupKFold\n\nimport torch\nfrom torch import nn\nfrom torch import optim\nfrom torch.optim import lr_scheduler, Adam, AdamW\nfrom torch.cuda import amp\nfrom torch.utils.data import DataLoader, Dataset, default_collate\nfrom torchvision.transforms import v2\n\nimport timm\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-03T21:54:58.813109Z","iopub.execute_input":"2024-05-03T21:54:58.813438Z","iopub.status.idle":"2024-05-03T21:55:05.798660Z","shell.execute_reply.started":"2024-05-03T21:54:58.813410Z","shell.execute_reply":"2024-05-03T21:55:05.797684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.autograd import Function\n\n\nclass EntmaxBisectFunction(Function):\n    @classmethod\n    def _gp(cls, x, alpha):\n        return x ** (alpha - 1)\n\n    @classmethod\n    def _gp_inv(cls, y, alpha):\n        return y ** (1 / (alpha - 1))\n\n    @classmethod\n    def _p(cls, X, alpha):\n        return cls._gp_inv(torch.clamp(X, min=0), alpha)\n\n    @classmethod\n    def forward(cls, ctx, X, alpha=1.5, dim=-1, n_iter=50, ensure_sum_one=True):\n\n        if not isinstance(alpha, torch.Tensor):\n            alpha = torch.tensor(alpha, dtype=X.dtype, device=X.device)\n\n        alpha_shape = list(X.shape)\n        alpha_shape[dim] = 1\n        alpha = alpha.expand(*alpha_shape)\n\n        ctx.alpha = alpha\n        ctx.dim = dim\n        d = X.shape[dim]\n\n        max_val, _ = X.max(dim=dim, keepdim=True)\n        X = X * (alpha - 1)\n        max_val = max_val * (alpha - 1)\n\n        # Note: when alpha < 1, tau_lo > tau_hi. This still works since dm < 0.\n        tau_lo = max_val - cls._gp(1, alpha)\n        tau_hi = max_val - cls._gp(1 / d, alpha)\n\n        # Note: f_lo should always be non-negative.\n        # f_lo = cls._p(X - tau_lo, alpha).sum(dim) - 1\n\n        dm = tau_hi - tau_lo\n\n        for it in range(n_iter):\n\n            dm /= 2\n            tau_m = tau_lo + dm\n            p_m = cls._p(X - tau_m, alpha)\n            f_m = p_m.sum(dim) - 1\n\n            mask = (f_m >= 0).unsqueeze(dim)\n            tau_lo = torch.where(mask, tau_m, tau_lo)\n\n        if ensure_sum_one:\n            p_m /= p_m.sum(dim=dim).unsqueeze(dim=dim)\n\n        ctx.save_for_backward(p_m)\n\n        return p_m\n\n    @classmethod\n    def backward(cls, ctx, dY):\n        Y, = ctx.saved_tensors\n\n        gppr = torch.where(Y > 0, Y ** (2 - ctx.alpha), Y.new_zeros(1))\n\n        dX = dY * gppr\n        q = dX.sum(ctx.dim) / gppr.sum(ctx.dim)\n        q = q.unsqueeze(ctx.dim)\n        dX -= q * gppr\n\n        d_alpha = None\n        if ctx.needs_input_grad[1]:\n\n            # alpha gradient computation\n            # d_alpha = (partial_y / partial_alpha) * dY\n            # NOTE: ensure alpha is not close to 1\n            # since there is an indetermination\n            # batch_size, _ = dY.shape\n\n            # shannon terms\n            S = torch.where(Y > 0, Y * torch.log(Y), Y.new_zeros(1))\n            # shannon entropy\n            ent = S.sum(ctx.dim).unsqueeze(ctx.dim)\n            Y_skewed = gppr / gppr.sum(ctx.dim).unsqueeze(ctx.dim)\n\n            d_alpha = dY * (Y - Y_skewed) / ((ctx.alpha - 1) ** 2)\n            d_alpha -= dY * (S - Y_skewed * ent) / (ctx.alpha - 1)\n            d_alpha = d_alpha.sum(ctx.dim).unsqueeze(ctx.dim)\n\n        return dX, d_alpha, None, None, None\n    \n    \ndef entmax_bisect(X, alpha=1.5, dim=-1, n_iter=50, ensure_sum_one=True):\n    return EntmaxBisectFunction.apply(X, alpha, dim, n_iter, ensure_sum_one)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:05.799859Z","iopub.execute_input":"2024-05-03T21:55:05.800429Z","iopub.status.idle":"2024-05-03T21:55:05.819062Z","shell.execute_reply.started":"2024-05-03T21:55:05.800396Z","shell.execute_reply":"2024-05-03T21:55:05.818062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    debug = False\n    \n    seed = 42\n    gpu_idx = 0\n\n    device = torch.device(f\"cuda:{gpu_idx}\")\n    \n    classes = [\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]\n    n_classes = len(classes)\n\n    n_folds = 5\n    \n    ## Dataset Preprocessing\n    bandpass_filter = {\"low\": 0.5, \"high\": 20, \"order\": 2}\n    rand_filter = {\"probab\": 0.1, \"low\": 10, \"high\": 20, \"band\": 1.0, \"order\": 2}\n    freq_channels = [(0.5, 4.5)]\n    filter_order = 2\n    random_close_zone = 0.0  # 0.2\n    \n    map_features = [\n        (\"Fp1\", \"F7\"),\n        (\"F7\", \"T3\"),\n        (\"T3\", \"T5\"),\n        (\"T5\", \"O1\"),\n        (\"Fp1\", \"F3\"),\n        (\"F3\", \"C3\"),\n        (\"C3\", \"P3\"),\n        (\"P3\", \"O1\"),\n        (\"Fp2\", \"F8\"),\n        (\"F8\", \"T4\"),\n        (\"T4\", \"T6\"),\n        (\"T6\", \"O2\"),\n        (\"Fp2\", \"F4\"),\n        (\"F4\", \"C4\"),\n        (\"C4\", \"P4\"),\n        (\"P4\", \"O2\"),\n        (\"Fz\", \"Cz\"),\n        (\"Cz\", \"Pz\")\n    ]\n    \n    eeg_features = [\"Fp1\",\"F3\",\"C3\",\"P3\",\"F7\",\"T3\",\"T5\",\"O1\",\"Fz\",\"Cz\",\"Pz\",\"Fp2\",\"F4\",\"C4\",\"P4\",\"F8\",\"T4\",\"T6\",\"O2\",\"EKG\"]\n    #eeg_features = [\"Fp1\", \"T3\", \"C3\", \"O1\", \"Fp2\", \"C4\", \"T4\", \"O2\"]\n    feature_to_index = {x: y for x, y in zip(eeg_features, range(len(eeg_features)))}\n    simple_features = []  # 'Fz', 'Cz', 'Pz', 'EKG'\n    \n    n_map_features = len(map_features)\n    in_channels = n_map_features + n_map_features * len(freq_channels) + len(simple_features)\n    \n    seq_length = 50  # Second's\n    sampling_rate = 200  # Hz\n    nsamples = seq_length * sampling_rate\n    out_samples = nsamples // 5\n    \n    input_size = 512\n    in_chans = 1\n\n    batch_size = 64\n    num_workers = 4","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:05.821651Z","iopub.execute_input":"2024-05-03T21:55:05.822005Z","iopub.status.idle":"2024-05-03T21:55:05.834657Z","shell.execute_reply.started":"2024-05-03T21:55:05.821973Z","shell.execute_reply":"2024-05-03T21:55:05.833888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_weights = [\n    {\n       'bandpass_filter':{'low':0.5, 'high':20, 'order':2}, \n       'model_name': \"convnext_atto_ols.a2_in1k\",\n       'coef': 0.075272, \n       'file_mask': \"/kaggle/input/hms-weights-yamash/022_023/*.bin\",\n       'tta_offset': [0],\n       'transform': 'test1',\n        'activation': 'entmax',\n        'alpha': 1.03,\n    },\n#     {\n#        'bandpass_filter':{'low':0.5, 'high':20, 'order':2}, \n#        'model_name': \"tf_efficientnet_b0\",\n#        'coef': 0.05, \n#        'file_mask': \"/kaggle/input/hms-weights-yamash/022_021/*.bin\",\n#        'tta_offset': [0],\n#        'transform': 'test1'\n#     },\n#     {\n#        'bandpass_filter':{'low':0.5, 'high':20, 'order':2}, \n#        'model_name': \"convnext_atto_ols.a2_in1k\",\n#        'coef': 0.0475, \n#        'file_mask': \"/kaggle/input/hms-weights-yamash/022_025_v2/*.bin\",\n#        'tta_offset': [0],\n#        'transform': 'test1'\n#     },\n#     {\n#         'bandpass_filter':{'low':0.5, 'high':20, 'order':2}, \n#         'model_name': \"convnext_atto_ols.a2_in1k\",\n#         'coef': 0.12221, \n#         'file_mask': \"/kaggle/input/hms-weights-yamash/026_001/*.bin\",\n#         'tta_offset': [0],\n#         'transform': 'test2'\n#     },\n    {\n        'bandpass_filter':{'low':0.5, 'high':20, 'order':2}, \n        'model_name': \"convnext_atto_ols.a2_in1k\",\n        'coef': 0.12221, \n        'file_mask': \"/kaggle/input/hms-weights-yamash/026_001_v2/*.bin\",\n        'tta_offset': [0],\n        'transform': 'test2',\n        'activation':'entmax',\n        'alpha': 1.02,\n    },\n#     {\n#        'bandpass_filter':{'low':0.5, 'high':20, 'order':2}, \n#        'model_name': \"tf_efficientnet_b0\",\n#        'coef': 0.08145, \n#        'file_mask': \"/kaggle/input/hms-weights-yamash/026_003/*.bin\",\n#        'tta_offset': [0],\n#        'transform': 'test2'\n#     },\n#     {\n#         'bandpass_filter':{'low':0.5, 'high':20, 'order':2}, \n#         'model_name': \"convnext_atto_ols.a2_in1k\",\n#         'coef': 0.12221, \n#         'file_mask': \"/kaggle/input/hms-weights-yamash/026_013/*.bin\",\n#         'tta_offset': [0],\n#         'transform': 'test2',\n#         'activation':'entmax',\n#         'alpha': 1.03,\n#     },\n    {\n        'bandpass_filter':{'low':0.5, 'high':40, 'order':2}, \n        'model_name': \"convnext_atto_ols.a2_in1k\",\n        'coef': 0.12221, \n        'file_mask': \"/kaggle/input/hms-weights-yamash/026_016/*.bin\",\n        'tta_offset': [0],\n        'transform': 'test2',\n        'activation': 'entmax',\n        'alpha': 1.04,\n    },\n    {\n        'bandpass_filter':{'low':0.5, 'high':40, 'order':2}, \n        'model_name': \"convnext_atto_ols.a2_in1k\",\n        'coef': 0.12221, \n        'file_mask': \"/kaggle/input/hms-weights-yamash/026_017/*.bin\",\n        'tta_offset': [0],\n        'transform': 'test3',\n        'activation':'entmax',\n        'alpha': 1.04,\n    },\n    {\n        'bandpass_filter':{'low':0.5, 'high':20, 'order':2}, \n        'model_name': \"inception_next_tiny\",\n        'coef': 0.12221, \n        'file_mask': \"/kaggle/input/hms-weights-yamash/030_002/*.bin\",\n        'tta_offset': [0],\n        'transform': 'test2',\n        'activation':'entmax',\n        'alpha': 1.02,\n    },\n#     {\n#         'bandpass_filter':{'low':0.5, 'high':20, 'order':2}, \n#         'model_name': \"inception_next_tiny\",\n#         'coef': 0.12221, \n#         'file_mask': \"/kaggle/input/hms-weights-yamash/030_003/*.bin\",\n#         'tta_offset': [0],\n#         'transform': 'test2'\n#     },\n#     {\n#         'bandpass_filter':{'low':0.5, 'high':20, 'order':2}, \n#         'model_name': \"convnext_atto_ols.a2_in1k\",\n#         'coef': 0.12221, \n#         'file_mask': \"/kaggle/input/hms-weights-yamash/033_002_alpha1.25/*.bin\",\n#         'tta_offset': [0],\n#         'transform': 'test2',\n#         'activation': 'entmax',\n#         'alpha': 1.25,\n#     },\n#     {\n#         'bandpass_filter':{'low':0.5, 'high':20, 'order':2}, \n#         'model_name': \"convnext_atto_ols.a2_in1k\",\n#         'coef': 0.12221, \n#         'file_mask': \"/kaggle/input/hms-weights-yamash/033_003_alpha1.5/*.bin\",\n#         'tta_offset': [0],\n#         'transform': 'test2',\n#         'activation': 'entmax',\n#         'alpha': 1.5,\n#     },\n]","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:05.836262Z","iopub.execute_input":"2024-05-03T21:55:05.836633Z","iopub.status.idle":"2024-05-03T21:55:05.850051Z","shell.execute_reply.started":"2024-05-03T21:55:05.836578Z","shell.execute_reply":"2024-05-03T21:55:05.849302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def quantize_data(data, classes):\n    mu_x = mu_law_encoding(data, classes)\n    return mu_x  # quantized\n\n\ndef mu_law_encoding(data, mu):\n    mu_x = np.sign(data) * np.log(1 + mu * np.abs(data)) / np.log(mu + 1)\n    return mu_x\n\n\ndef mu_law_expansion(data, mu):\n    s = np.sign(data) * (np.exp(np.abs(data) * np.log(mu + 1)) - 1) / mu\n    return s\n\n\ndef butter_bandpass(lowcut, highcut, fs, order=5):\n    return butter(order, [lowcut, highcut], fs=fs, btype=\"band\")\n\n\ndef butter_bandpass_filter(data, lowcut, highcut, fs, order=5):\n    b, a = butter_bandpass(lowcut, highcut, fs, order=order)\n    y = lfilter(b, a, data)\n    return y\n\n\ndef butter_lowpass_filter(data, cutoff_freq=20, sampling_rate=CFG.sampling_rate, order=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\n\n\ndef denoise_filter(x):\n    y = butter_bandpass_filter(x, CFG.lowcut, CFG.highcut, CFG.sampling_rate, order=6)\n    y = (y + np.roll(y, -1) + np.roll(y, -2) + np.roll(y, -3)) / 4\n    y = y[0:-1:4]\n    return y","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:05.851133Z","iopub.execute_input":"2024-05-03T21:55:05.851531Z","iopub.status.idle":"2024-05-03T21:55:05.863868Z","shell.execute_reply.started":"2024-05-03T21:55:05.851500Z","shell.execute_reply":"2024-05-03T21:55:05.863071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HMSDataset(torch.utils.data.Dataset):\n    def __init__(self, df, eegs,\n                 bandpass_filter = None,\n                 rand_filter = None,\n                 mode='train',\n                 transforms = None,\n                 tta_offset = 0,\n                 transform_type = \"test1\"\n                ):\n        self.df = df\n        self.eegs = eegs\n        self.bandpass_filter = bandpass_filter\n        self.rand_filter = rand_filter\n        self.mode = mode\n        self.transforms = transforms\n        self.tta_offset = tta_offset\n        self.transform_type = transform_type\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx: int):\n        row = self.df.iloc[idx]\n        eeg_id = row.eeg_id\n        \n        W = 10\n        random_offset = random.randint(-1000, 1000)\n        \n        if self.transform_type == \"test3\":\n            img8000 = np.zeros((W*CFG.in_channels, 8000, 1)).astype(np.float32)\n            img4000 = np.zeros((W*CFG.in_channels, 4000, 1)).astype(np.float32)\n            img2000 = np.zeros((W*CFG.in_channels, 2000, 1)).astype(np.float32)\n\n            X8000 = self.__data_generation(idx, 8000, random_offset)\n            X4000 = self.__data_generation(idx, 4000, random_offset)\n            X2000 = self.__data_generation(idx, 2000, random_offset)\n\n            for i in range(X2000.shape[1]):\n                img8000[i*W:(i+1)*W, :] = X8000[:,i][:,np.newaxis]\n                img4000[i*W:(i+1)*W, :] = X4000[:,i][:,np.newaxis]\n                img2000[i*W:(i+1)*W, :] = X2000[:,i][:,np.newaxis]\n\n            if self.transforms is not None:\n                img8000 = self.transforms[\"8000\"](image=img8000)['image']\n                img4000 = self.transforms[\"4000\"](image=img4000)['image']\n                img2000 = self.transforms[\"2000\"](image=img2000)['image']\n                img = torch.cat([img8000, img4000, img2000], axis=1)\n            else:\n                img = cv2.vconcat([cv2.resize(img8000,(512,112)),\n                                  cv2.resize(img4000,(512,200)),\n                                  cv2.resize(img2000,(512,200)),])\n        else:\n            img10000 = np.zeros((W*CFG.in_channels, 10000, 1)).astype(np.float32)\n            img5000 = np.zeros((W*CFG.in_channels, 5000, 1)).astype(np.float32)\n            img2000 = np.zeros((W*CFG.in_channels, 2000, 1)).astype(np.float32)\n\n            X10000 = self.__data_generation(idx, 10000, random_offset)\n            X5000 = self.__data_generation(idx, 5000, random_offset)\n            X2000 = self.__data_generation(idx, 2000, random_offset)\n\n            for i in range(X2000.shape[1]):\n                img10000[i*W:(i+1)*W, :] = X10000[:,i][:,np.newaxis]\n                img5000[i*W:(i+1)*W, :] = X5000[:,i][:,np.newaxis]\n                img2000[i*W:(i+1)*W, :] = X2000[:,i][:,np.newaxis]\n\n            if self.transforms is not None:\n                img10000 = self.transforms[\"10000\"](image=img10000)['image']\n                img5000 = self.transforms[\"5000\"](image=img5000)['image']\n                img2000 = self.transforms[\"2000\"](image=img2000)['image']\n                img = torch.cat([img10000, img5000, img2000], axis=1)\n            else:\n                img = cv2.vconcat([cv2.resize(img10000,(512,112)),\n                                  cv2.resize(img5000,(512,200)),\n                                  cv2.resize(img2000,(512,200)),])\n                \n        if self.mode != 'test':\n            y = row[CFG.classes].values.astype(np.float32)\n            return {\"data\": img, \"label\": y}\n        else:\n            return {\"data\": img}\n    \n    def __data_generation(self, index, out_samples=CFG.out_samples, random_offset=0):\n        X = np.zeros((out_samples, CFG.in_channels), dtype=\"float32\")\n        \n        random_offset = random_offset * out_samples // 2000\n        tta_offset = self.tta_offset\n        \n        if self.transform_type != \"test3\":\n            if out_samples > 6000:\n                random_offset = 0\n                tta_offset = 0\n\n        row = self.df.iloc[index]\n        data = self.eegs[row.eeg_id]\n\n        if CFG.nsamples != out_samples:\n            if self.mode != \"train\":\n                offset = (CFG.nsamples - out_samples) // 2 + tta_offset\n            else:             \n                #offset = ((CFG.nsamples - CFG.out_samples) * random.randint(0, 1000))\n                #offset = (CFG.nsamples - out_samples) // 2 + random.randint(-out_samples//2, out_samples//2)\n                offset = (CFG.nsamples - out_samples) // 2 + random_offset\n            data = data[offset:offset+out_samples,:]\n\n        # diff of eegs\n        for i, (feat_a, feat_b) in enumerate(CFG.map_features):\n            if self.mode == \"train\" and CFG.random_close_zone > 0 and random.uniform(0.0, 1.0) <= CFG.random_close_zone:\n                continue\n                \n            diff_feat = data[:, CFG.feature_to_index[feat_a]] - data[:, CFG.feature_to_index[feat_b]]\n\n            if not self.bandpass_filter is None:\n                diff_feat = butter_bandpass_filter(\n                    diff_feat,\n                    self.bandpass_filter[\"low\"],\n                    self.bandpass_filter[\"high\"],\n                    CFG.sampling_rate,\n                    order=self.bandpass_filter[\"order\"],\n                )\n                    \n            if (\n                self.mode == \"train\"\n                and not self.rand_filter is None\n                and random.uniform(0.0, 1.0) <= self.rand_filter[\"probab\"]\n            ):\n                lowcut = random.randint(\n                    self.rand_filter[\"low\"], self.rand_filter[\"high\"]\n                )\n                highcut = lowcut + self.rand_filter[\"band\"]\n                diff_feat = butter_bandpass_filter(\n                    diff_feat,\n                    lowcut,\n                    highcut,\n                    CFG.sampling_rate,\n                    order=self.rand_filter[\"order\"],\n                )\n\n            X[:, i] = diff_feat\n\n        # other frequency filtering\n        n = CFG.n_map_features\n        if len(CFG.freq_channels) > 0:\n            for i in range(CFG.n_map_features):\n                diff_feat = X[:, i]\n                for j, (lowcut, highcut) in enumerate(CFG.freq_channels):\n                    band_feat = butter_bandpass_filter(\n                        diff_feat, lowcut, highcut, CFG.sampling_rate, order=CFG.filter_order,\n                    )\n                    X[:, n] = band_feat\n                    n += 1\n                    \n        # single eeg \n        for spml_feat in CFG.simple_features:\n            feat_val = data[:, CFG.feature_to_index[spml_feat]]\n            \n            if not self.bandpass_filter is None:\n                feat_val = butter_bandpass_filter(\n                    feat_val,\n                    self.bandpass_filter[\"low\"],\n                    self.bandpass_filter[\"high\"],\n                    CFG.sampling_rate,\n                    order=self.bandpass_filter[\"order\"],\n                )\n\n            if (\n                self.mode == \"train\"\n                and not self.rand_filter is None\n                and random.uniform(0.0, 1.0) <= self.rand_filter[\"probab\"]\n            ):\n                lowcut = random.randint(\n                    self.rand_filter[\"low\"], self.rand_filter[\"high\"]\n                )\n                highcut = lowcut + self.rand_filter[\"band\"]\n                feat_val = butter_bandpass_filter(\n                    feat_val,\n                    lowcut,\n                    highcut,\n                    CFG.sampling_rate,\n                    order=self.rand_filter[\"order\"],\n                )\n\n            X[:, n] = feat_val\n            n += 1\n            \n        X = np.clip(X, -1024, 1024)\n        X = np.nan_to_num(X, nan=0) / 32.0\n        X = butter_lowpass_filter(X, order=CFG.filter_order)\n\n        return X","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:05.864947Z","iopub.execute_input":"2024-05-03T21:55:05.865258Z","iopub.status.idle":"2024-05-03T21:55:05.901418Z","shell.execute_reply.started":"2024-05-03T21:55:05.865229Z","shell.execute_reply":"2024-05-03T21:55:05.900478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def eeg_from_parquet(\n    parquet_path: str, display: bool = False, seq_length=CFG.seq_length\n) -> np.ndarray:\n    eeg = pd.read_parquet(parquet_path, columns=CFG.eeg_features)\n    rows = len(eeg)\n\n    offset = (rows - CFG.nsamples) // 2\n\n    eeg = eeg.iloc[offset : offset + CFG.nsamples]\n\n    if display:\n        plt.figure(figsize=(10, 5))\n        offset = 0\n\n    data = np.zeros((CFG.nsamples, len(CFG.eeg_features)))\n\n    for index, feature in enumerate(CFG.eeg_features):\n        x = eeg[feature].values.astype(\"float32\")\n\n\n        mean = np.nanmean(x)\n        nan_percentage = np.isnan(x).mean()\n\n        if nan_percentage < 1: \n            x = np.nan_to_num(x, nan=mean)\n        else:\n            x[:] = 0\n        data[:, index] = x\n\n        if display:\n            if index != 0:\n                offset += x.max()\n            plt.plot(range(CFG.nsamples), x - offset, label=feature)\n            offset -= x.min()\n\n    if display:\n        plt.legend()\n        name = parquet_path.split(\"/\")[-1].split(\".\")[0]\n        plt.yticks([])\n        plt.title(f\"EEG {name}\", size=16)\n        plt.show()\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:05.902401Z","iopub.execute_input":"2024-05-03T21:55:05.902717Z","iopub.status.idle":"2024-05-03T21:55:05.914777Z","shell.execute_reply.started":"2024-05-03T21:55:05.902694Z","shell.execute_reply":"2024-05-03T21:55:05.913953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HMSModel(nn.Module):\n    def __init__(self, model_name: str, pretrained: bool, in_channels: int, num_classes: int):\n        super().__init__()\n        \n        self.model = timm.create_model(\n            model_name=model_name, \n            pretrained=pretrained, \n            in_chans=in_channels,\n            num_classes = num_classes\n        )\n        \n    def forward(self, x):\n        x = self.model(x)  \n        return x","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:05.915782Z","iopub.execute_input":"2024-05-03T21:55:05.916043Z","iopub.status.idle":"2024-05-03T21:55:05.926854Z","shell.execute_reply.started":"2024-05-03T21:55:05.916021Z","shell.execute_reply":"2024-05-03T21:55:05.926078Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def inference_function(test_loader, model, device, activation=\"softmax\", alpha=1.03):\n    model.eval() \n    if activation == \"softmax\":\n        softmax = nn.Softmax(dim=1)\n    prediction_dict = {}\n    preds = []\n    with tqdm(test_loader, unit=\"test_batch\", desc=\"Inference\") as tqdm_test_loader:\n        for step, batch in enumerate(tqdm_test_loader):\n            X = batch.pop(\"data\").to(device, dtype=torch.float)\n            batch_size = X.size(0)\n            with torch.no_grad():\n                y_preds = model(X)\n            if activation == \"softmax\":\n                y_preds = softmax(y_preds)\n            elif activation == \"entmax\":\n                y_preds = entmax_bisect(y_preds, alpha=alpha, dim=1)\n            preds.append(y_preds.to(\"cpu\").numpy())\n\n    prediction_dict[\"predictions\"] = np.concatenate(\n        preds\n    )\n    return prediction_dict","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:05.930044Z","iopub.execute_input":"2024-05-03T21:55:05.930322Z","iopub.status.idle":"2024-05-03T21:55:05.937961Z","shell.execute_reply.started":"2024-05-03T21:55:05.930299Z","shell.execute_reply":"2024-05-03T21:55:05.937153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/test.csv\")\nprint(f\"Test dataframe shape is: {test_df.shape}\")\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:05.939121Z","iopub.execute_input":"2024-05-03T21:55:05.939431Z","iopub.status.idle":"2024-05-03T21:55:05.961234Z","shell.execute_reply.started":"2024-05-03T21:55:05.939409Z","shell.execute_reply":"2024-05-03T21:55:05.960246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_eeg_parquet_paths = glob(\"/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/\" + \"*.parquet\")\ntest_eeg_df = pd.read_parquet(test_eeg_parquet_paths[0])\ntest_eeg_features = test_eeg_df.columns\nprint(f\"There are {len(test_eeg_features)} raw eeg features\")\nprint(list(test_eeg_features))\ndel test_eeg_df\n_ = gc.collect()\n\n# %%time\nall_eegs = {}\neeg_ids = test_df.eeg_id.unique()\nfor i, eeg_id in tqdm(enumerate(eeg_ids)):\n    # Save EEG to Python dictionary of numpy arrays\n    eeg_path = \"/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/\" + str(eeg_id) + \".parquet\"\n    data = eeg_from_parquet(eeg_path)\n    all_eegs[eeg_id] = data","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:05.962395Z","iopub.execute_input":"2024-05-03T21:55:05.962948Z","iopub.status.idle":"2024-05-03T21:55:06.257905Z","shell.execute_reply.started":"2024-05-03T21:55:05.962917Z","shell.execute_reply":"2024-05-03T21:55:06.257006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_transforms = {\n    \n    \"test1\": {\n        \"2000\": A.Compose([\n            A.Resize(height=200, width=CFG.input_size),\n            ToTensorV2()], p=1.),\n        \"5000\": A.Compose([\n            A.Resize(height=200, width=CFG.input_size),\n            ToTensorV2()], p=1.),\n        \"10000\": A.Compose([\n            A.Resize(height=112, width=CFG.input_size),\n            ToTensorV2()], p=1.),\n    },\n    \"test2\": {\n        \"2000\": A.Compose([\n            A.Resize(height=256, width=CFG.input_size),\n            ToTensorV2()], p=1.),\n        \"5000\": A.Compose([\n            A.Resize(height=128, width=CFG.input_size),\n            ToTensorV2()], p=1.),\n        \"10000\": A.Compose([\n            A.Resize(height=128, width=CFG.input_size),\n            ToTensorV2()], p=1.),\n    },\n    \"test3\": {\n        \"2000\": A.Compose([\n            A.Resize(height=256, width=CFG.input_size),\n            ToTensorV2()], p=1.),\n        \"4000\": A.Compose([\n            A.Resize(height=128, width=CFG.input_size),\n            ToTensorV2()], p=1.),\n        \"8000\": A.Compose([\n            A.Resize(height=128, width=CFG.input_size),\n            ToTensorV2()], p=1.),\n    },\n    \n}","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:06.259106Z","iopub.execute_input":"2024-05-03T21:55:06.259491Z","iopub.status.idle":"2024-05-03T21:55:06.269648Z","shell.execute_reply.started":"2024-05-03T21:55:06.259457Z","shell.execute_reply":"2024-05-03T21:55:06.268760Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_preds = []\ncoef_sum = 0\ncoef_count = 0\n\nfor model_block in model_weights:\n    predictions = []\n    for offset in model_block['tta_offset']:\n        print(f\"TTA offset: {offset}\")\n        \n        test_dataset = HMSDataset(\n            df=test_df,\n            mode=\"test\",\n            eegs=all_eegs,\n            bandpass_filter=model_block['bandpass_filter'],\n            transforms=data_transforms[model_block['transform']],\n            tta_offset=offset,\n            transform_type=model_block['transform']\n        )\n\n        if len(predictions) == 0:\n            output = test_dataset[0]\n            X = output[\"data\"]\n            print(f\"X shape: {X.shape}\")\n\n        test_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        )\n\n        model = HMSModel(\n            model_name = model_block['model_name'], \n            pretrained = False,\n            in_channels = CFG.in_chans, \n            num_classes = CFG.n_classes\n        )\n\n        for weight_model_file in glob(model_block['file_mask']):\n            print(f\"{weight_model_file=}\")\n            checkpoint = torch.load(weight_model_file, map_location=CFG.device)\n            model.load_state_dict(checkpoint)\n            model.to(CFG.device)\n            prediction_dict = inference_function(test_loader, model, CFG.device, activation=model_block['activation'], alpha=model_block['alpha'])\n            predict = prediction_dict[\"predictions\"]\n            predictions.append(predict)\n            torch.cuda.empty_cache()\n            gc.collect()\n                \n    predictions = np.array(predictions)\n    predictions = np.mean(predictions, axis=0)\n    \n    tmp = pd.DataFrame({\"eeg_id\": test_df.eeg_id.values})\n    tmp[CFG.classes] = predictions\n    display(tmp.head())\n\n    #coef = model_block['coef']\n    #predictions = predictions * coef\n    #coef_sum += coef\n    #coef_count += 1\n\n    all_preds.append(predictions)\n    \n#display(coef_sum, coef_count)\n    \n#all_preds = np.array(all_preds)\n#coef_sum /= coef_count\n#all_preds /= coef_sum\n#all_preds = np.mean(all_preds, axis=0) \n\ny_1d_predictions = all_preds.copy()","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:06.271105Z","iopub.execute_input":"2024-05-03T21:55:06.271634Z","iopub.status.idle":"2024-05-03T21:55:27.117668Z","shell.execute_reply.started":"2024-05-03T21:55:06.271591Z","shell.execute_reply":"2024-05-03T21:55:27.116674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#y_sub_1d = pd.DataFrame({\"eeg_id\": test_df.eeg_id.values})\n#y_sub_1d[CFG.classes] = all_preds\n\n#y_sub_1d.to_csv(f\"submission.csv\", index=False)\n#y_sub_1d.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:27.118797Z","iopub.execute_input":"2024-05-03T21:55:27.119096Z","iopub.status.idle":"2024-05-03T21:55:27.123568Z","shell.execute_reply.started":"2024-05-03T21:55:27.119071Z","shell.execute_reply":"2024-05-03T21:55:27.122698Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## yamash: Spectrogram 2D","metadata":{}},{"cell_type":"code","source":"\nimport pywt, librosa\n\nclass CFG:\n    debug = False\n    \n    seed = 42\n    gpu_idx = 0\n\n    device = torch.device(f\"cuda:{gpu_idx}\")\n    \n    classes = [\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]\n    n_classes = len(classes)\n\n    n_folds = 5\n    train_folds = [0,1,2,3,4]\n    \n    in_chans = 1\n    input_size = 1024\n    model_name = \"tf_efficientnet_b2\"\n    model_dirs = [\n        \"/kaggle/input/hms-exp012-007-efficientnet-b2-cv0-3477\"\n    ]\n    \n    n_epoch = 5\n    train_batch_size = 16\n    valid_batch_size = 16\n\n\nUSE_WAVELET = None \n\nNAMES = ['LL','LP','RP','RR']\n\nFEATS = [['Fp1','F7','T3','T5','O1'],\n         ['Fp1','F3','C3','P3','O1'],\n         ['Fp2','F8','T4','T6','O2'],\n         ['Fp2','F4','C4','P4','O2']]\n\n# DENOISE FUNCTION\ndef maddest(d, axis=None):\n    return np.mean(np.absolute(d - np.mean(d, axis)), axis)\n\ndef denoise(x, wavelet='haar', level=1):    \n    coeff = pywt.wavedec(x, wavelet, mode=\"per\")\n    sigma = (1/0.6745) * maddest(coeff[-level])\n\n    uthresh = sigma * np.sqrt(2*np.log(len(x)))\n    coeff[1:] = (pywt.threshold(i, value=uthresh, mode='hard') for i in coeff[1:])\n\n    ret=pywt.waverec(coeff, wavelet, mode='per')\n    \n    return ret\n\ndef spectrogram_from_eeg(parquet_path, display=False):\n    \n    # LOAD MIDDLE 50 SECONDS OF EEG SERIES\n    eeg = pd.read_parquet(parquet_path)\n    middle = (len(eeg)-10_000)//2\n    eeg = eeg.iloc[middle:middle+10_000]\n    \n    # VARIABLE TO HOLD SPECTROGRAM\n    img = np.zeros((128,256,4),dtype='float32')\n    \n    if display: plt.figure(figsize=(10,7))\n    signals = []\n    for k in range(4):\n        COLS = FEATS[k]\n        \n        for kk in range(4):\n        \n            # COMPUTE PAIR DIFFERENCES\n            x = eeg[COLS[kk]].values - eeg[COLS[kk+1]].values\n\n            # FILL NANS\n            m = np.nanmean(x)\n            if np.isnan(x).mean()<1: x = np.nan_to_num(x,nan=m)\n            else: x[:] = 0\n\n            # DENOISE\n            if USE_WAVELET:\n                x = denoise(x, wavelet=USE_WAVELET)\n            signals.append(x)\n\n            # RAW SPECTROGRAM\n            mel_spec = librosa.feature.melspectrogram(y=x, sr=200, hop_length=len(x)//256, \n                  n_fft=1024, n_mels=128, fmin=0, fmax=20, win_length=128)\n\n            # LOG TRANSFORM\n            width = (mel_spec.shape[1]//32)*32\n            mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max).astype(np.float32)[:,:width]\n\n            # STANDARDIZE TO -1 TO 1\n            mel_spec_db = (mel_spec_db+40)/40 \n            img[:,:,k] += mel_spec_db\n                \n        # AVERAGE THE 4 MONTAGE DIFFERENCES\n        img[:,:,k] /= 4.0\n        \n        if display:\n            plt.subplot(2,2,k+1)\n            plt.imshow(img[:,:,k],aspect='auto',origin='lower')\n            plt.title(f'EEG {eeg_id} - Spectrogram {NAMES[k]}')\n            \n    if display: \n        plt.show()\n        plt.figure(figsize=(10,5))\n        offset = 0\n        for k in range(4):\n            if k>0: offset -= signals[3-k].min()\n            plt.plot(range(10_000),signals[k]+offset,label=NAMES[3-k])\n            offset += signals[3-k].max()\n        plt.legend()\n        plt.title(f'EEG {eeg_id} Signals')\n        plt.show()\n        print(); print('#'*25); print()\n        \n    return img\n\n\nALL_FEATS = ['Fp1','F7','T3','T5','O1','F3','C3','P3','Fp2','F8','T4','T6','O2','F4','C4','P4','Fz','Cz','Pz','EKG']\n\ndef spectrogram_from_eeg_2(parquet_path, display=False):\n    \n    # LOAD MIDDLE 50 SECONDS OF EEG SERIES\n    eeg = pd.read_parquet(parquet_path)\n    middle = (len(eeg)-10_000)//2  # 200Hz x 50sec = 10_000\n    eeg = eeg.iloc[middle:middle+10_000]\n    \n    average = np.mean(eeg.values[:,:-1], axis=1)  # except for EKG\n    \n    # VARIABLE TO HOLD SPECTROGRAM\n    img = np.zeros((128, 256, 20),dtype='float32')\n    \n    if display: plt.figure(figsize=(10,7))\n    signals = []\n\n    for k, f in enumerate(ALL_FEATS):\n        # COMPUTE PAIR DIFFERENCES\n        x = eeg[f].values - average\n\n        # FILL NANS\n        m = np.nanmean(x)\n        if np.isnan(x).mean()<1: x = np.nan_to_num(x,nan=m)\n        else: x[:] = 0\n\n        # DENOISE\n        if USE_WAVELET:\n            x = denoise(x, wavelet=USE_WAVELET)\n        signals.append(x)\n\n        # RAW SPECTROGRAM\n        mel_spec = librosa.feature.melspectrogram(y=x, sr=200, hop_length=len(x)//256, \n              n_fft=1024, n_mels=128, fmin=0, fmax=20, win_length=128)\n\n        # LOG TRANSFORM\n        width = (mel_spec.shape[1]//32)*32\n        mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max).astype(np.float32)[:,:width]\n\n        # STANDARDIZE TO -1 TO 1\n        mel_spec_db = (mel_spec_db+40)/40 \n        img[:,:,k] = mel_spec_db\n        \n    return img\n\n\nclass HMSDataset(torch.utils.data.Dataset):\n    def __init__(self, df, spec, eeg_spec, eeg_spec2, transforms=None):\n        self.df = df\n        self.spec = spec\n        self.eeg_spec = eeg_spec\n        self.eeg_spec2 = eeg_spec2\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx: int):\n        X = np.zeros((CFG.input_size, CFG.input_size, 1), dtype='float32')\n        row = self.df.iloc[idx]\n        \n        spec_id = row.spectrogram_id\n        eeg_id = row.eeg_id\n        r = 0 #int((row['min'] + row['max']) // 4)\n        \n        for k in range(4):\n            # kaggle spec\n            img = self.spec[spec_id][r:r+300,k*100:(k+1)*100].transpose()\n            img = np.clip(img, np.exp(-4), np.exp(8))\n            img = np.log(img)\n            ep = 1e-6\n            m = np.nanmean(img.flatten())\n            s = np.nanstd(img.flatten())\n            img = (img - m) / (s + ep)\n            img = np.nan_to_num(img, nan=0.0) \n            X[14:114, k*256:(k+1)*256, 0] = img[:,22:-22] / 2.0  # (1,1)->(1,4)\n\n            X[-114:-14, k*256:(k+1)*256, 0] = img[::-1,22:-22] / 2.0 # flipud & 隙間埋め\n\n        for k in range(4):\n            # eeg spec\n            X[128:256, k*256:(k+1)*256, 0] = self.eeg_spec[eeg_id][:,:,k] # (2,1)->(2,4)\n\n        for k in range(20):\n            r = k // 4\n            c = k % 4\n            X[(r+2)*128:(r+3)*128, c*256:(c+1)*256:, 0] = self.eeg_spec2[eeg_id][:,:,k]  #(3,1)->(7,4)\n                     \n        if self.transforms is not None:\n            X = self.transforms(image=X)['image']\n\n        return {\"data\": X}\n    \n    \ndata_transforms = {   \n    \"valid\": A.Compose([\n        A.Resize(height=CFG.input_size, width=CFG.input_size),\n        ToTensorV2()], p=1.),\n}\n\n\nclass HMSModel(nn.Module):\n    def __init__(self, model_name: str, pretrained: bool, in_channels: int, num_classes: int):\n        super().__init__()\n        \n        self.model = timm.create_model(\n            model_name=model_name, \n            pretrained=pretrained, \n            in_chans=in_channels,\n            num_classes = num_classes)\n\n    def forward(self, x):\n        x = self.model(x)  \n        return x\n    \n    \ndef run_inference_loop(model, loader, device):\n    model.to(device)\n    model.eval()\n    pred_list = []\n    with torch.no_grad():\n        for batch in tqdm(loader):\n            x = batch[\"data\"].to(device)\n            y = model(x)\n            pred_list.append(y.softmax(dim=1).detach().cpu().numpy())\n        \n    pred_arr = np.concatenate(pred_list)\n    del pred_list\n    return pred_arr\n\n\ntest_df = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/test.csv\")\ntest_df.head()\n\n\nPATH = '/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/'\nif CFG.debug:\n    PATH = '/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/'\nfiles = os.listdir(PATH)\nprint(f'There are {len(files)} spectrogram parquets')\n\nspectrograms = {}\nfor i,f in enumerate(files):\n    if i%100==0: print(i,', ',end='')\n    tmp = pd.read_parquet(f'{PATH}{f}')\n    name = int(f.split('.')[0])\n    spectrograms[name] = tmp.iloc[:,1:].values\n    \n    \n# READ ALL EEG SPECTROGRAMS\nPATH2 = '/kaggle/input/hms-harmful-brain-activity-classification/test_eegs/'\nDISPLAY = 0\nEEG_IDS2 = test_df.eeg_id.unique()\nif CFG.debug:\n    PATH2 = '/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/'\n    EEG_IDS2 = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\").eeg_id.unique()[:2700]  # ~2700 samples in hidden test data\neeg_spectrograms = {}\neeg_spectrograms2 = {}\n\nprint('Converting Test EEG to Spectrograms...'); print()\nfor i,eeg_id in enumerate(EEG_IDS2):\n    if (i%100==0)&(i!=0): print(i,', ',end='')\n        \n    # CREATE SPECTROGRAM FROM EEG PARQUET\n    img = spectrogram_from_eeg(f'{PATH2}{eeg_id}.parquet', i<DISPLAY)\n    eeg_spectrograms[eeg_id] = img\n    \n    img2 = spectrogram_from_eeg_2(f'{PATH2}{eeg_id}.parquet', i<DISPLAY)\n    eeg_spectrograms2[eeg_id] = img2\n    \n    \ntest_preds_arr = np.zeros((CFG.n_folds, len(test_df), CFG.n_classes))\n\ntest_dataset = HMSDataset(test_df, spectrograms, eeg_spectrograms, eeg_spectrograms2, data_transforms['valid'])\ntest_loader = DataLoader(test_dataset, batch_size=CFG.valid_batch_size, num_workers=4, shuffle=False, pin_memory=True)\n\npreds = []\n\nfor model_dir in CFG.model_dirs:\n    print(model_dir)\n    for fold in range(CFG.n_folds):\n        print(f\"[fold {fold}]\")\n\n        model_path = f\"{model_dir}/fold{fold}_best_loss.bin\"\n        model = HMSModel(model_name=CFG.model_name, pretrained=False, in_channels=CFG.in_chans, num_classes=CFG.n_classes).to(CFG.device)\n        model.load_state_dict(torch.load(model_path, map_location=CFG.device))\n\n        test_pred = run_inference_loop(model, test_loader, CFG.device)\n        test_preds_arr[fold] = test_pred\n\n        del model\n        torch.cuda.empty_cache()\n        gc.collect()\n        \n    mean_test_pred = test_preds_arr.mean(axis=0)\n    preds.append(mean_test_pred)\n    \n    \n#test_pred = np.mean(np.array(preds), axis=0)\n\ny_2d_predictions = preds.copy()\n\n#test_pred_df = pd.DataFrame(test_pred, columns=CFG.classes)\n#test_pred_df = pd.concat([test_df[[\"eeg_id\"]], test_pred_df], axis=1)\n\n#smpl_sub = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/sample_submission.csv\")\n#y_sub_2d = pd.merge(smpl_sub[[\"eeg_id\"]], test_pred_df, on=\"eeg_id\", how=\"left\")\n#y_sub_2d.head()\n\n","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:27.125219Z","iopub.execute_input":"2024-05-03T21:55:27.125579Z","iopub.status.idle":"2024-05-03T21:55:34.588725Z","shell.execute_reply.started":"2024-05-03T21:55:27.125548Z","shell.execute_reply":"2024-05-03T21:55:34.587647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Sugupoko Model","metadata":{}},{"cell_type":"code","source":"\"\"\"\nclass CFG:\n    debug = False\n    \n    seed = 42\n    gpu_idx = 0\n\n    device = torch.device(f\"cuda:{gpu_idx}\")\n    \n    classes = [\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]\n    n_classes = len(classes)\n\n    n_folds = 5\n    \n    ## Model\n    fixed_kernel_size = 5\n    #linear_layer_features = 448  # Full Signal = 10_000\n    #linear_layer_features = 352  # Half Signal = 5_000\n    linear_layer_features = 304   # 1/5  Signal = 2_000\n    kernels = [3, 5, 7, 9, 11]\n    \n    ## Dataset Preprocessing\n    bandpass_filter = None\n    # rand_filter = {\"probab\": 0.1, \"low\": 10, \"high\": 20, \"band\": 1.0, \"order\": 2}\n    rand_filter = {}\n    # freq_channels = [(8.0, 12.0), (0.5, 4.5)]\n    freq_channels = []\n    \n    filter_order = 2\n    random_close_zone = 0.0  # 0.2\n    \n    map_features = [\n        (\"Fp1\", \"F7\"),\n        (\"F7\", \"T3\"),\n        (\"T3\", \"T5\"),\n        (\"T5\", \"O1\"),\n        (\"Fp1\", \"F3\"),\n        (\"F3\", \"C3\"),\n        (\"C3\", \"P3\"),\n        (\"P3\", \"O1\"),\n        (\"Fp2\", \"F8\"),\n        (\"F8\", \"T4\"),\n        (\"T4\", \"T6\"),\n        (\"T6\", \"O2\"),\n        (\"Fp2\", \"F4\"),\n        (\"F4\", \"C4\"),\n        (\"C4\", \"P4\"),\n        (\"P4\", \"O2\"),\n        (\"Fz\", \"Cz\"),\n        (\"Cz\", \"Pz\")\n    ]\n    \n    eeg_features = [\"Fp1\",\"F3\",\"C3\",\"P3\",\"F7\",\"T3\",\"T5\",\"O1\",\"Fz\",\"Cz\",\"Pz\",\"Fp2\",\"F4\",\"C4\",\"P4\",\"F8\",\"T4\",\"T6\",\"O2\",\"EKG\"]\n    #eeg_features = [\"Fp1\", \"T3\", \"C3\", \"O1\", \"Fp2\", \"C4\", \"T4\", \"O2\"]\n    feature_to_index = {x: y for x, y in zip(eeg_features, range(len(eeg_features)))}\n    simple_features = []  # 'Fz', 'Cz', 'Pz', 'EKG'\n    \n    n_map_features = len(map_features)\n    in_channels = n_map_features #+ n_map_features * len(freq_channels) + len(simple_features)\n    \n    seq_length = 50  # Second's\n    sampling_rate = 200  # Hz\n    nsamples = seq_length * sampling_rate\n    out_samples = nsamples // 2\n    \n    input_size = 1024\n    in_chans = 1\n\n    batch_size = 4\n    num_workers = 4\n    \n    \ndata_transforms = {\n    \"valid\": A.Compose([\n        ToTensorV2()], p=1.),   \n}\n\nmodel_weights = [\n    {\n        'bandpass_filter':None, \n        'model_name': \"tf_efficientnet_b7\",\n        'file_data': \n        [\n            {'coef':1.0, 'file_mask':\"/kaggle/input/exp05-26-fix-stft-25sec-18band-b7-labelsmoothingof/*stage2.bin\"},\n        ]\n    },\n]\n\nimport torchaudio.transforms as T\n\nclass HMSDataset_2D(torch.utils.data.Dataset):\n    def __init__(self, df, eegs,\n                 downsample: int = None,\n                 bandpass_filter = None,\n                 rand_filter = None,\n                 mode='train',\n                 weighted=False,\n                 transforms = None,\n                ):\n        self.df = df\n        self.eegs = eegs\n        self.downsample = downsample\n        self.bandpass_filter = bandpass_filter\n        self.rand_filter = rand_filter\n        self.mode = mode\n        self.weighted = weighted\n        self.transforms = transforms\n        self.spec_imgsize_hw = [128, 256]\n        self.L = 50 * 200\n        self.L_input = CFG.out_samples\n        self.mel_spectrogram = T.MelSpectrogram(\n            sample_rate=200,\n            n_fft=1024,\n            win_length=128,\n            hop_length = self.L_input//self.spec_imgsize_hw[1],  # ここで適切な値に設定、256が widht想定。\n            n_mels=self.spec_imgsize_hw[0],\n            f_min=0.,\n            f_max=20.,\n            pad=0,\n            window_fn=torch.hann_window\n        ) \n            \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx: int):\n        row = self.df.iloc[idx]\n        eeg_id = row.eeg_id\n        \n        img = np.zeros((10*CFG.in_channels, CFG.out_samples, 1)).astype(np.float32)\n        \n        X, y = self.__data_generation(idx)\n\n        #####################\n        # 2D\n        img_spec = torch.zeros((X.shape[1], 257, 40), dtype=torch.float32)  # GPU対応\n        \n        for k in range((X.shape[1])):\n            x = torch.tensor(X[:,k], dtype=torch.float32)  # GPU対応\n\n            stft_result = torch.stft(\n                    x,\n                    n_fft=512,\n                    hop_length=128,\n                    win_length=512,\n                    window=torch.hann_window(512),\n                    center=True,\n                    pad_mode='reflect',\n                    normalized=False,\n                    onesided=True,  # 実際のスペクトルの半分のサイズの出力\n                    return_complex=True\n                )\n            \n            magnitude = torch.abs(stft_result)\n            # magnitude_std = (magnitude - magnitude.mean()) / magnitude.std()\n\n            # # 振幅スペクトルを正規化（0から1の範囲）\n            magnitude_normalized = (magnitude - magnitude.min()) / (magnitude.max() - magnitude.min())\n            # ログスケール変換（非常に小さい値を避けるために小さい定数を加算）\n            magnitude_log = torch.log(magnitude_normalized + 1e-6)\n            magnitude_log_normalized = (magnitude_log - magnitude_log.min()) / (magnitude_log.max() - magnitude_log.min())\n            magnitude_log_normalized = torch.nan_to_num(magnitude_log_normalized)\n\n            # 位相スペクトルを計算\n            # phase = torch.angle(stft_result)\n            # phase_normalized = phase / np.pi\n\n            img_spec[k, :,:] += magnitude_log_normalized\n            \n        reshaped_img = img_spec.view(2, 9, 257, 40)\n        img_spec = reshaped_img.permute(0, 2, 1, 3).reshape(2*257, 9*40).unsqueeze(0)\n        \n        #####################\n        # 1D\n#         if self.downsample is not None:\n#             X = X[::self.downsample, :]\n            \n#         for i in range(X.shape[1]):\n#             img[i*10:(i+1)*10, :] = X[:,i][:,np.newaxis]\n                       \n#         if self.transforms is not None:\n#             img = self.transforms(image=img)['image']\n                \n        #####################\n        #y = row[CFG.classes].values.astype(np.float32)\n        #if self.weighted:\n        #    w = row[\"num_votes\"].astype(np.float32) / 10.0\n        #else:\n        #    w = 1.0\n        \n#         img = torch.cat([img, img_spec], dim=0)\n        img = img_spec\n        return {\"data\": img}\n    \n    def __data_generation(self, index):\n        X = np.zeros((CFG.out_samples, CFG.in_channels), dtype=\"float32\")  # Size=(10000, 14)\n\n        row = self.df.iloc[index]\n        data = self.eegs[row.eeg_id]  # Size=(10000, 8)\n\n        if CFG.nsamples != CFG.out_samples:\n            if self.mode != \"train\":\n                offset = (CFG.nsamples - CFG.out_samples) // 2\n            else:             \n                offset = ((CFG.nsamples - CFG.out_samples) * random.randint(0, 1000)) // 1000\n            data = data[offset:offset+CFG.out_samples,:]\n\n        # diff of eegs\n        for i, (feat_a, feat_b) in enumerate(CFG.map_features):\n            if self.mode == \"train\" and CFG.random_close_zone > 0 and random.uniform(0.0, 1.0) <= CFG.random_close_zone:\n                continue\n                \n            diff_feat = data[:, CFG.feature_to_index[feat_a]] - data[:, CFG.feature_to_index[feat_b]]\n\n            if not self.bandpass_filter is None:\n                diff_feat = butter_bandpass_filter(\n                    diff_feat,\n                    self.bandpass_filter[\"low\"],\n                    self.bandpass_filter[\"high\"],\n                    CFG.sampling_rate,\n                    order=self.bandpass_filter[\"order\"],\n                )\n                    \n            if (\n                self.mode == \"train\"\n                and not self.rand_filter is None\n                and random.uniform(0.0, 1.0) <= self.rand_filter[\"probab\"]\n            ):\n                lowcut = random.randint(\n                    self.rand_filter[\"low\"], self.rand_filter[\"high\"]\n                )\n                highcut = lowcut + self.rand_filter[\"band\"]\n                diff_feat = butter_bandpass_filter(\n                    diff_feat,\n                    lowcut,\n                    highcut,\n                    CFG.sampling_rate,\n                    order=self.rand_filter[\"order\"],\n                )\n\n            X[:, i] = diff_feat\n\n        # other frequency filtering\n        n = CFG.n_map_features\n        if len(CFG.freq_channels) > 0:\n            for i in range(CFG.n_map_features):\n                diff_feat = X[:, i]\n                for j, (lowcut, highcut) in enumerate(CFG.freq_channels):\n                    band_feat = butter_bandpass_filter(\n                        diff_feat, lowcut, highcut, CFG.sampling_rate, order=CFG.filter_order,  # 6\n                    )\n                    X[:, n] = band_feat\n                    n += 1\n                    \n        # single eeg \n        for spml_feat in CFG.simple_features:\n            feat_val = data[:, CFG.feature_to_index[spml_feat]]\n            \n            if not self.bandpass_filter is None:\n                feat_val = butter_bandpass_filter(\n                    feat_val,\n                    self.bandpass_filter[\"low\"],\n                    self.bandpass_filter[\"high\"],\n                    CFG.sampling_rate,\n                    order=self.bandpass_filter[\"order\"],\n                )\n\n            if (\n                self.mode == \"train\"\n                and not self.rand_filter is None\n                and random.uniform(0.0, 1.0) <= self.rand_filter[\"probab\"]\n            ):\n                lowcut = random.randint(\n                    self.rand_filter[\"low\"], self.rand_filter[\"high\"]\n                )\n                highcut = lowcut + self.rand_filter[\"band\"]\n                feat_val = butter_bandpass_filter(\n                    feat_val,\n                    lowcut,\n                    highcut,\n                    CFG.sampling_rate,\n                    order=self.rand_filter[\"order\"],\n                )\n\n            X[:, n] = feat_val\n            n += 1\n            \n        X = np.clip(X, -1024, 1024)\n        X = np.nan_to_num(X, nan=0) / 32.0\n        X = butter_lowpass_filter(X, order=CFG.filter_order)\n\n        y = np.zeros(CFG.n_classes, dtype=\"float32\")  # Size=(6,)\n        if self.mode != \"test\":\n            y = row[CFG.classes].values.astype(np.float32)\n\n        return X, y\n        \ncoef_sum = 0\ncoef_count = 0\npredictions = []\nfiles = []\n    \nfor model_block in model_weights:\n    test_dataset = HMSDataset_2D(\n        df=test_df,\n        mode=\"test\",\n        eegs=all_eegs,\n        bandpass_filter=model_block['bandpass_filter'],\n        transforms=data_transforms['valid']\n    )\n\n    if len(predictions) == 0:\n        output = test_dataset[0]\n        X = output[\"data\"]\n        print(f\"X shape: {X.shape}\")\n                \n    test_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    )\n\n    model = HMSModel(\n        model_name = model_block['model_name'], \n        pretrained = False,\n        in_channels = CFG.in_chans, \n        num_classes = CFG.n_classes\n    )\n\n    for file_line in model_block['file_data']:\n        coef = file_line['coef']\n        for weight_model_file in glob(file_line['file_mask']):\n            files.append(weight_model_file)\n            checkpoint = torch.load(weight_model_file, map_location=CFG.device)\n            model.load_state_dict(checkpoint)\n            model.to(CFG.device)\n            prediction_dict = inference_function(test_loader, model, CFG.device)\n            predict = prediction_dict[\"predictions\"]\n            predict *= coef\n            coef_sum += coef\n            coef_count += 1\n            predictions.append(predict)\n            torch.cuda.empty_cache()\n            gc.collect()\n\npredictions = np.array(predictions)\ncoef_sum /= coef_count\npredictions /= coef_sum\npredictions_2d = np.mean(predictions, axis=0)\n\ndel model\ntorch.cuda.empty_cache()\ngc.collect()\n\ntest_pred = predictions_2d\n\nk_predictions = [test_pred.copy()]\n\n#test_pred_df = pd.DataFrame(test_pred, columns=CFG.classes)\n#test_pred_df = pd.concat([test_df[[\"eeg_id\"]], test_pred_df], axis=1)\n\n#smpl_sub = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/sample_submission.csv\")\n#k_sub = pd.merge(smpl_sub[[\"eeg_id\"]], test_pred_df, on=\"eeg_id\", how=\"left\")\n#k_sub.head()\n\n\"\"\"","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:34.590439Z","iopub.execute_input":"2024-05-03T21:55:34.590960Z","iopub.status.idle":"2024-05-03T21:55:34.613340Z","shell.execute_reply.started":"2024-05-03T21:55:34.590931Z","shell.execute_reply":"2024-05-03T21:55:34.612418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Sugupoko CWT Model","metadata":{}},{"cell_type":"code","source":"import sys\nimport os\nimport gc\nimport copy\nimport yaml\nimport random\nimport shutil\nimport time\nimport typing as tp\nfrom glob import glob\nfrom pathlib import Path\nfrom collections import OrderedDict, defaultdict\nfrom logging import getLogger, INFO, StreamHandler, FileHandler, Formatter\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom tqdm import tqdm\nfrom scipy.signal import butter, lfilter, freqz\nfrom sklearn.model_selection import StratifiedGroupKFold\n\nimport torch\nfrom torch import nn\nfrom torch import optim\nfrom torch.optim import lr_scheduler, Adam, AdamW\nfrom torch.cuda import amp\nfrom torch.utils.data import DataLoader, Dataset, default_collate\nfrom torchvision.transforms import v2\n\nimport timm\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport sys\nsys.path.append('/kaggle/input/kaggle-kl-div/')\nfrom kaggle_kl_div import score","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:34.614397Z","iopub.execute_input":"2024-05-03T21:55:34.614681Z","iopub.status.idle":"2024-05-03T21:55:34.635649Z","shell.execute_reply.started":"2024-05-03T21:55:34.614631Z","shell.execute_reply":"2024-05-03T21:55:34.634925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_weights = [\n    #{\n    #    'bandpass_filter':None, \n    #    'model_name': \"inception_next_tiny.sail_in1k\",\n    #    'file_data': \n    #    [\n    #        {'coef':0.16859119140863998, \n    #         'file_mask':\"/kaggle/input/hms-cwt-weights-v01/exp05-54_AMP_cwt_20to10sec_stride4_18band_inception_next_tiny_labelsmoothingOff/*stage2.bin\"},\n    #    ]\n    #},\n    #{\n    #    'bandpass_filter':None, \n    #    'model_name': \"maxvit_small_tf_512.in1k\",\n    #    'file_data': \n    #    [\n    #        {'coef':0.02, \n    #         'file_mask':\"/kaggle/input/hms-cwt-weights-v01/exp05-57_AMP_cwt_20to10sec_stride4_18band_maxvit_small_tf_512_labelsmoothingOff/*stage2.bin\"},\n    #    ]\n    #}, \n    #{\n    #    'bandpass_filter':None, \n    #    'model_name': \"tiny_vit_21m_512.dist_in22k_ft_in1k\",\n    #    'file_data': \n    #    [\n    #        {'coef':0.0196, \n    #         'file_mask':\"/kaggle/input/hms-cwt-weights-v01/exp05-58_AMP_cwt_20to10sec_stride4_18band_tiny_vit_21m_512.dist_in22k_ft_in1k_labelsmoothingOff/*stage2.bin\"},\n    #    ]\n    #}, \n    #{\n    #    'bandpass_filter':None, \n    #    'model_name': \"mixnet_l.ft_in1k\",\n    #    'file_data': \n    #    [\n    #        {'coef':0.0196, \n    #         'file_mask':\"/kaggle/input/hms-cwt-weights-v01/exp05-62_AMP_cwt_20to10sec_stride4_18band_mixnet_l_labelsmoothingOff/*stage2.bin\"},\n    #    ]\n    #},      \n    \n]\n\n# model_weights_25sec = [\n\n#     #{\n#     #    'bandpass_filter':None, \n#     #    'model_name': \"tiny_vit_21m_512.dist_in22k_ft_in1k\",\n#     #    'file_data': \n#     #    [\n#     #        {'coef':0.038416, \n#     #         'file_mask':\"/kaggle/input/hms-cwt-weights-v01/exp05-61_AMP_cwt_50to25sec_stride4_18band_tiny_vit_21m_512_labelsmoothingOff/*stage2.bin\"},\n#     #    ]\n#     #},\n#     #{\n#     #    'bandpass_filter':None, \n#     #    'model_name': \"maxvit_small_tf_512.in1k\",\n#     #    'file_data': \n#     #    [\n#     #        {'coef':0.0921984, \n#     #         'file_mask':\"/kaggle/input/hms-cwt-weights-v01/exp05-64_AMP_cwt_50to25sec_stride4_18band_maxvit_small_tf_512_labelsmoothingOff/*stage2.bin\"},\n#     #    ]\n#     #},\n#     #{\n#     #    'bandpass_filter':None, \n#     #    'model_name': \"inception_next_tiny\",\n#     #    'file_data': \n#     #    [\n#     #        {'coef':0.21614255308799996, \n#     #         'file_mask':\"/kaggle/input/hms-cwt-weights-v01/exp05-65_AMP_cwt_50to25sec_stride4_18band_inception_next_tiny_labelsmoothingOff/*stage2.bin\"},\n#     #    ]\n#     #}, \n#     {\n#         'bandpass_filter':None, \n#         'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n#         'file_data': \n#         [\n#             {'coef':0.21614255308799996, \n#              'file_mask':\"/kaggle/input/hms-cwt-weights-r02/exp05-72-2_AMP_cwt_50to25sec_stride4_18band_maxvit_base_tf_512_labelsmoothingOff/*stage2.bin\"},\n#         ]\n#     },  \n\n# ]\n\n\nmodel_weights_50sec = [\n    #{\n    #    'bandpass_filter':None, \n    #    'model_name': \"tiny_vit_21m_512.dist_in22k_ft_in1k\",\n    #    'file_data': \n    #    [\n    #        {'coef':0.09957427199999999, \n    #         'file_mask':\"/kaggle/input/hms-cwt-weights-v01/exp05-67_AMP_cwt_50sec_stride16_18band_tiny_vit_21m_512.dist_in22k_ft_in1k_labelsmoothingOff/*stage2.bin\"},\n    #    ]\n    #},\n    #{\n    #    'bandpass_filter':None, \n    #    'model_name': \"inception_next_tiny\",\n    #    'file_data': \n    #    [\n    #        {'coef':0.18985494528, \n    #         'file_mask':\"/kaggle/input/hms-cwt-weights-v01/exp05-68_AMP_cwt_50sec_stride16_18band_inception_next_tiny_labelsmoothingOff/*stage2.bin\"},\n    #    ]\n    #},   \n#     {\n#         'bandpass_filter':None, \n#         'model_name': \"maxvit_small_tf_512.in1k\",\n#         'file_data': \n#         [\n#             {'coef':0.18985494528, \n#              'file_mask':\"/kaggle/input/hms-cwt-weights-v01/exp05-69_AMP_cwt_50sec_stride16_18band_maxvit_small_tf_512_labelsmoothingOff/*stage2.bin\"},\n#         ]\n#     },    \n#     {\n#         'bandpass_filter':None, \n#         'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n#         'file_data': \n#         [\n#             {'coef':0.15562263822335998, \n#              'file_mask':\"/kaggle/input/hms-cwt-weights-v01/exp05-71_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/*stage2.bin\"},\n#         ]\n#     },\n#     {\n#         'bandpass_filter':None, \n#         'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n#         'file_data': \n#         [\n#             {'coef':0.15562263822335998, \n#              'file_mask':\"/kaggle/input/hms-cwt-weights-v01/exp05-71-2_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/*stage2.bin\"},\n#         ]\n#     },\n#     {\n#         'bandpass_filter':None, \n#         'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n#         'file_data': \n#         [\n#             {'coef':0.15562263822335998, \n#              'file_mask':\"/kaggle/input/hms-cwt-weights-r02/exp05-71-5_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/*stage2.bin\"},\n#         ]\n#     },\n#     {\n#         'bandpass_filter':None, \n#         'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n#         'file_data': \n#         [\n#             {'coef':0.15562263822335998, \n#              'file_mask':\"/kaggle/input/hms-cwt-weights-r02/exp05-71-6_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/*stage2.bin\"},\n#         ]\n#     },\n#     {\n#         'bandpass_filter':None, \n#         'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n#         'file_data': \n#         [\n#             {'coef':0.15562263822335998, \n#              'file_mask':\"/kaggle/input/hms-cwt-weights-r02/exp05-71-7_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/*stage2.bin\"},\n#         ]\n#     },\n]\n\nmodel_weights_50sec_to40Hz = [\n    {\n        'bandpass_filter':None, \n        'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n        'file_data': \n        [\n            {'coef':0.15562263822335998, \n             'file_mask':\"/kaggle/input/hms-cwt-weights-r02/exp05-78-2_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/*stage2.bin\"},\n        ]\n    },   \n#     {\n#         'bandpass_filter':None, \n#         'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n#         'file_data': \n#         [\n#             {'coef':0.15562263822335998, \n#              'file_mask':\"/kaggle/input/hms-cwt-weights-r02/exp05-78-3_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/*stage2.bin\"},\n#         ]\n#     },    \n     {\n        'bandpass_filter':None, \n        'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n        'file_data': \n        [\n            {'coef':0.15562263822335998, \n             'file_mask':\"/kaggle/input/hms-cwt-weights-r02/exp05-78-4_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/*stage2.bin\"},\n        ]\n    },   \n     {\n        'bandpass_filter':None, \n        'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n        'file_data': \n        [\n            {'coef':0.15562263822335998, \n             'file_mask':\"/kaggle/input/hms-cwt-weights-r02/exp05-78-5_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/*stage2.bin\"},\n        ]\n    },   \n#      {\n#         'bandpass_filter':None, \n#         'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n#         'file_data': \n#         [\n#             {'coef':0.15562263822335998, \n#              'file_mask':\"/kaggle/input/hms-cwt-weights-r02/exp05-78-6_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/*stage2.bin\"},\n#         ]\n#     },       \n     {\n        'bandpass_filter':None, \n        'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n        'file_data': \n        [\n            {'coef':0.15562263822335998, \n             'file_mask':\"/kaggle/input/hms-cwt-weights-r02/exp05-90_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_addData/*stage2.bin\"},\n        ]\n    },       \n]\n\n\n\nmodel_weights_50sec_to40Hz_and_1d = [\n    {\n        'bandpass_filter':None, \n        'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n        'file_data': \n        [\n            {'coef':0.15562263822335998, \n             'file_mask':\"/kaggle/input/hms-cwt-weights-r02/exp05-91_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_1dAnd2d/*stage2.bin\"},\n        ]\n    },   \n    {\n        'bandpass_filter':None, \n        'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n        'file_data': \n        [\n            {'coef':0.15562263822335998, \n             'file_mask':\"/kaggle/input/hms-cwt-weights-r02/exp05-91-2_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_1dAnd2d/*stage2.bin\"},\n        ]\n    },  \n    {\n        'bandpass_filter':None, \n        'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n        'file_data': \n        [\n            {'coef':0.15562263822335998, \n             'file_mask':\"/kaggle/input/hms-cwt-weights-r02/exp05-91-3_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_1dAnd2d/*stage2.bin\"},\n        ]\n    }, \n    {\n        'bandpass_filter':None, \n        'model_name': \"maxvit_base_tf_512.in21k_ft_in1k\",\n        'file_data': \n        [\n            {'coef':0.15562263822335998, \n             'file_mask':\"/kaggle/input/hms-cwt-weights-r02/exp05-91-5_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_1dAnd2d/*stage2.bin\"},\n        ]\n    }, \n]","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:34.636882Z","iopub.execute_input":"2024-05-03T21:55:34.637150Z","iopub.status.idle":"2024-05-03T21:55:34.654539Z","shell.execute_reply.started":"2024-05-03T21:55:34.637127Z","shell.execute_reply":"2024-05-03T21:55:34.653641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    debug = False\n    \n    CWT_STRIDE = 4\n            \n    seq_length = 20  # Second's\n    sampling_rate = 200  # Hz\n    nsamples = seq_length * sampling_rate\n    out_samples = 10 * sampling_rate\n    \n    \n    seed = 42\n    gpu_idx = 0\n\n    device = torch.device(f\"cuda:{gpu_idx}\")\n    \n    classes = [\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]\n    n_classes = len(classes)\n\n    n_folds = 5\n    \n    ## Model\n    fixed_kernel_size = 5\n    #linear_layer_features = 448  # Full Signal = 10_000\n    #linear_layer_features = 352  # Half Signal = 5_000\n    linear_layer_features = 304   # 1/5  Signal = 2_000\n    kernels = [3, 5, 7, 9, 11]\n    \n    ## Dataset Preprocessing\n    bandpass_filter = None\n    # rand_filter = {\"probab\": 0.1, \"low\": 10, \"high\": 20, \"band\": 1.0, \"order\": 2}\n    rand_filter = {}\n    # freq_channels = [(8.0, 12.0), (0.5, 4.5)]\n    freq_channels = []\n    \n    filter_order = 2\n    random_close_zone = 0.0  # 0.2\n    \n    map_features = [\n        (\"Fp1\", \"F7\"),\n        (\"F7\", \"T3\"),\n        (\"T3\", \"T5\"),\n        (\"T5\", \"O1\"),\n        (\"Fp1\", \"F3\"),\n        (\"F3\", \"C3\"),\n        (\"C3\", \"P3\"),\n        (\"P3\", \"O1\"),\n        (\"Fp2\", \"F8\"),\n        (\"F8\", \"T4\"),\n        (\"T4\", \"T6\"),\n        (\"T6\", \"O2\"),\n        (\"Fp2\", \"F4\"),\n        (\"F4\", \"C4\"),\n        (\"C4\", \"P4\"),\n        (\"P4\", \"O2\"),\n        (\"Fz\", \"Cz\"),\n        (\"Cz\", \"Pz\")\n    ]\n    \n    eeg_features = [\"Fp1\",\"F3\",\"C3\",\"P3\",\"F7\",\"T3\",\"T5\",\"O1\",\"Fz\",\"Cz\",\"Pz\",\"Fp2\",\"F4\",\"C4\",\"P4\",\"F8\",\"T4\",\"T6\",\"O2\",\"EKG\"]\n    #eeg_features = [\"Fp1\", \"T3\", \"C3\", \"O1\", \"Fp2\", \"C4\", \"T4\", \"O2\"]\n    feature_to_index = {x: y for x, y in zip(eeg_features, range(len(eeg_features)))}\n    simple_features = []  # 'Fz', 'Cz', 'Pz', 'EKG'\n    \n    n_map_features = len(map_features)\n    in_channels = n_map_features #+ n_map_features * len(freq_channels) + len(simple_features)\n\n    input_size = 512\n    in_chans = 1\n\n    batch_size = 4\n    num_workers = 4","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:34.656010Z","iopub.execute_input":"2024-05-03T21:55:34.656351Z","iopub.status.idle":"2024-05-03T21:55:34.668725Z","shell.execute_reply.started":"2024-05-03T21:55:34.656320Z","shell.execute_reply":"2024-05-03T21:55:34.667948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_transforms = {\n    \"valid\": A.Compose([\n        A.Resize(height=CFG.input_size, width=CFG.input_size),\n        ToTensorV2()], p=1.),\n    \n}","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:34.669694Z","iopub.execute_input":"2024-05-03T21:55:34.670007Z","iopub.status.idle":"2024-05-03T21:55:34.682013Z","shell.execute_reply.started":"2024-05-03T21:55:34.669984Z","shell.execute_reply":"2024-05-03T21:55:34.681086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def quantize_data(data, classes):\n    mu_x = mu_law_encoding(data, classes)\n    return mu_x  # quantized\n\n\ndef mu_law_encoding(data, mu):\n    mu_x = np.sign(data) * np.log(1 + mu * np.abs(data)) / np.log(mu + 1)\n    return mu_x\n\n\ndef mu_law_expansion(data, mu):\n    s = np.sign(data) * (np.exp(np.abs(data) * np.log(mu + 1)) - 1) / mu\n    return s\n\n\ndef butter_bandpass(lowcut, highcut, fs, order=5):\n    return butter(order, [lowcut, highcut], fs=fs, btype=\"band\")\n\n\ndef butter_bandpass_filter(data, lowcut, highcut, fs, order=5):\n    b, a = butter_bandpass(lowcut, highcut, fs, order=order)\n    y = lfilter(b, a, data)\n    return y\n\n\ndef butter_lowpass_filter(data, cutoff_freq=20, sampling_rate=CFG.sampling_rate, order=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\n\n\ndef denoise_filter(x):\n    y = butter_bandpass_filter(x, CFG.lowcut, CFG.highcut, CFG.sampling_rate, order=6)\n    y = (y + np.roll(y, -1) + np.roll(y, -2) + np.roll(y, -3)) / 4\n    y = y[0:-1:4]\n    return y","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:34.683161Z","iopub.execute_input":"2024-05-03T21:55:34.683741Z","iopub.status.idle":"2024-05-03T21:55:34.695619Z","shell.execute_reply.started":"2024-05-03T21:55:34.683710Z","shell.execute_reply":"2024-05-03T21:55:34.694813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def eeg_from_parquet(\n    parquet_path: str, display: bool = False, seq_length=CFG.seq_length, CFG=CFG,\n) -> np.ndarray:\n    eeg = pd.read_parquet(parquet_path, columns=CFG.eeg_features)\n    rows = len(eeg)\n\n    offset = (rows - CFG.nsamples) // 2\n\n    eeg = eeg.iloc[offset : offset + CFG.nsamples]\n\n    if display:\n        plt.figure(figsize=(10, 5))\n        offset = 0\n\n    data = np.zeros((CFG.nsamples, len(CFG.eeg_features)))\n\n    for index, feature in enumerate(CFG.eeg_features):\n        x = eeg[feature].values.astype(\"float32\")\n\n\n        mean = np.nanmean(x)\n        nan_percentage = np.isnan(x).mean()\n\n        if nan_percentage < 1: \n            x = np.nan_to_num(x, nan=mean)\n        else:\n            x[:] = 0\n        data[:, index] = x\n\n        if display:\n            if index != 0:\n                offset += x.max()\n            plt.plot(range(CFG.nsamples), x - offset, label=feature)\n            offset -= x.min()\n\n    if display:\n        plt.legend()\n        name = parquet_path.split(\"/\")[-1].split(\".\")[0]\n        plt.yticks([])\n        plt.title(f\"EEG {name}\", size=16)\n        plt.show()\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:34.696567Z","iopub.execute_input":"2024-05-03T21:55:34.696814Z","iopub.status.idle":"2024-05-03T21:55:34.707040Z","shell.execute_reply.started":"2024-05-03T21:55:34.696792Z","shell.execute_reply":"2024-05-03T21:55:34.706253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def inference_function(test_loader, model, device):\n    model.eval() \n#     softmax = nn.Softmax(dim=1)\n    prediction_dict = {}\n    preds = []\n    with tqdm(test_loader, unit=\"test_batch\", desc=\"Inference\") as tqdm_test_loader:\n        for step, batch in enumerate(tqdm_test_loader):\n            X = batch.pop(\"data\").to(device, dtype=torch.float)\n            batch_size = X.size(0)\n            with torch.no_grad():\n                y_preds = model(X)\n#             y_preds = softmax(y_preds)\n            y_preds = entmax_bisect(y_preds, alpha=1.03, dim=1)\n            preds.append(y_preds.to(\"cpu\").numpy())\n\n    prediction_dict[\"predictions\"] = np.concatenate(\n        preds\n    )\n    return prediction_dict","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:34.708258Z","iopub.execute_input":"2024-05-03T21:55:34.708524Z","iopub.status.idle":"2024-05-03T21:55:34.719897Z","shell.execute_reply.started":"2024-05-03T21:55:34.708502Z","shell.execute_reply":"2024-05-03T21:55:34.719159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mode = \"test\"","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:34.720877Z","iopub.execute_input":"2024-05-03T21:55:34.721127Z","iopub.status.idle":"2024-05-03T21:55:34.728836Z","shell.execute_reply.started":"2024-05-03T21:55:34.721105Z","shell.execute_reply":"2024-05-03T21:55:34.728048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(f\"/kaggle/input/hms-harmful-brain-activity-classification/{mode}.csv\")\nprint(f\"Test dataframe shape is: {test_df.shape}\")\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:34.729868Z","iopub.execute_input":"2024-05-03T21:55:34.730189Z","iopub.status.idle":"2024-05-03T21:55:34.745607Z","shell.execute_reply.started":"2024-05-03T21:55:34.730158Z","shell.execute_reply":"2024-05-03T21:55:34.744765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_eeg_parquet_paths = glob(f\"/kaggle/input/hms-harmful-brain-activity-classification/{mode}_eegs/\" + \"*.parquet\")\ntest_eeg_df = pd.read_parquet(test_eeg_parquet_paths[0])\ntest_eeg_features = test_eeg_df.columns\nprint(f\"There are {len(test_eeg_features)} raw eeg features\")\nprint(list(test_eeg_features))\ndel test_eeg_df\n_ = gc.collect()\n\n# %%time\nall_eegs = {}\neeg_ids = test_df.eeg_id.unique()\nfor i, eeg_id in tqdm(enumerate(eeg_ids)):\n    # Save EEG to Python dictionary of numpy arrays\n    eeg_path = f\"/kaggle/input/hms-harmful-brain-activity-classification/{mode}_eegs/\" + str(eeg_id) + \".parquet\"\n    data = eeg_from_parquet(eeg_path, CFG=CFG)\n    all_eegs[eeg_id] = data","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:34.746580Z","iopub.execute_input":"2024-05-03T21:55:34.746846Z","iopub.status.idle":"2024-05-03T21:55:34.999220Z","shell.execute_reply.started":"2024-05-03T21:55:34.746789Z","shell.execute_reply":"2024-05-03T21:55:34.998274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CWT(nn.Module):\n    def __init__(\n        self,\n        wavelet_width,\n        fs,\n        lower_freq,\n        upper_freq,\n        n_scales,\n        size_factor=1.0,\n        border_crop=0,\n        stride=1\n    ):\n        super().__init__()\n\n        self.initial_wavelet_width = wavelet_width\n        self.fs = fs\n        self.lower_freq = lower_freq\n        self.upper_freq = upper_freq\n        self.size_factor = size_factor\n        self.n_scales = n_scales\n        self.wavelet_width = wavelet_width\n        self.border_crop = border_crop\n        self.stride = stride\n        wavelet_bank_real, wavelet_bank_imag = self._build_wavelet_kernel()\n        self.wavelet_bank_real = nn.Parameter(wavelet_bank_real, requires_grad=False)\n        self.wavelet_bank_imag = nn.Parameter(wavelet_bank_imag, requires_grad=False)\n\n        self.kernel_size = self.wavelet_bank_real.size(3)\n\n    def _build_wavelet_kernel(self):\n        s_0 = 1 / self.upper_freq\n        s_n = 1 / self.lower_freq\n\n        base = np.power(s_n / s_0, 1 / (self.n_scales - 1))\n        scales = s_0 * np.power(base, np.arange(self.n_scales))\n\n        frequencies = 1 / scales\n        truncation_size = scales.max() * np.sqrt(4.5 * self.initial_wavelet_width) * self.fs\n        one_side = int(self.size_factor * truncation_size)\n        kernel_size = 2 * one_side + 1\n\n        k_array = np.arange(kernel_size, dtype=np.float32) - one_side\n        t_array = k_array / self.fs\n\n        wavelet_bank_real = []\n        wavelet_bank_imag = []\n\n        for scale in scales:\n            norm_constant = np.sqrt(np.pi * self.wavelet_width) * scale * self.fs / 2.0\n            scaled_t = t_array / scale\n            exp_term = np.exp(-(scaled_t ** 2) / self.wavelet_width)\n            kernel_base = exp_term / norm_constant\n            kernel_real = kernel_base * np.cos(2 * np.pi * scaled_t)\n            kernel_imag = kernel_base * np.sin(2 * np.pi * scaled_t)\n            wavelet_bank_real.append(kernel_real)\n            wavelet_bank_imag.append(kernel_imag)\n\n        wavelet_bank_real = np.stack(wavelet_bank_real, axis=0)\n        wavelet_bank_imag = np.stack(wavelet_bank_imag, axis=0)\n\n        wavelet_bank_real = torch.from_numpy(wavelet_bank_real).unsqueeze(1).unsqueeze(2)\n        wavelet_bank_imag = torch.from_numpy(wavelet_bank_imag).unsqueeze(1).unsqueeze(2)\n        return wavelet_bank_real, wavelet_bank_imag\n\n    def forward(self, x):\n        border_crop = self.border_crop // self.stride\n        start = border_crop\n        end = (-border_crop) if border_crop > 0 else None\n\n        # x [n_batch, n_channels, time_len]\n        out_reals = []\n        out_imags = []\n\n        in_width = x.size(2)\n        out_width = int(np.ceil(in_width / self.stride))\n        pad_along_width = np.max((out_width - 1) * self.stride + self.kernel_size - in_width, 0)\n        padding = pad_along_width // 2 + 1\n\n        for i in range(x.size(1)):\n            # [n_batch, 1, 1, time_len]\n            x_ = x[:, i, :].unsqueeze(1).unsqueeze(2)\n            out_real = nn.functional.conv2d(x_, self.wavelet_bank_real, stride=(1, self.stride), padding=(0, padding))\n            out_imag = nn.functional.conv2d(x_, self.wavelet_bank_imag, stride=(1, self.stride), padding=(0, padding))\n            out_real = out_real.transpose(2, 1)\n            out_imag = out_imag.transpose(2, 1)\n            out_reals.append(out_real)\n            out_imags.append(out_imag)\n\n        out_real = torch.cat(out_reals, axis=1)\n        out_imag = torch.cat(out_imags, axis=1)\n\n        out_real = out_real[:, :, :, start:end]\n        out_imag = out_imag[:, :, :, start:end]\n\n        scalograms = torch.sqrt(out_real ** 2 + out_imag ** 2)\n        return scalograms","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.006963Z","iopub.execute_input":"2024-05-03T21:55:35.007268Z","iopub.status.idle":"2024-05-03T21:55:35.027068Z","shell.execute_reply.started":"2024-05-03T21:55:35.007244Z","shell.execute_reply":"2024-05-03T21:55:35.026185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchaudio.transforms as T\n\nclass HMSDataset_preprocesser(torch.utils.data.Dataset):\n    def __init__(self, df, eegs, CFG,\n                 downsample: int = None,\n                 bandpass_filter = None,\n                 rand_filter = None,\n                 mode='train',\n                 weighted=False,\n                 transforms = None,\n                ):\n        self.df = df\n        self.eegs = eegs\n        self.CFG = CFG\n        self.downsample = downsample\n        self.bandpass_filter = bandpass_filter\n        self.rand_filter = rand_filter\n        self.mode = mode\n        self.weighted = weighted\n        self.transforms = transforms\n        self.spec_imgsize_hw = [128, 256]\n        self.L       = CFG.nsamples    # 4000\n        self.L_input = CFG.out_samples # 2000\n        self.cwt = CWT(wavelet_width=7, fs=200, \\\n                       lower_freq=0.5, upper_freq=20, \\\n                       n_scales=40, border_crop = 1, stride = self.CFG.CWT_STRIDE)\n                   \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx: int):\n        row = self.df.iloc[idx]\n        eeg_id = row.eeg_id\n        \n        img = np.zeros((10*self.CFG.in_channels, self.CFG.out_samples, 1)).astype(np.float32)\n        \n        X, y = self.__data_generation(idx)\n        #####################\n        # 2D\n        # img_spec = torch.zeros((X.shape[1], self.spec_imgsize_hw[0], self.spec_imgsize_hw[1]), dtype=torch.float32)  # GPU対応\n        X_in = torch.tensor(X.transpose(1, 0), dtype=torch.float32).unsqueeze(0)\n        img_spec = self.cwt(X_in)\n        img_spec = img_spec[0]\n        img_spec = torch.cat([img_spec[i] for i in range(img_spec.shape[0])], dim=0)\n        \n        if self.transforms is not None:\n            img_spec = self.transforms(image=img_spec.numpy() )['image']\n            \n        return {\"data\": img_spec, \"idx\": idx}\n    \n    def __data_generation(self, index):\n        X = np.zeros(( self.L, self.CFG.in_channels), dtype=\"float32\")  # Size=(10000, 14)\n\n        row = self.df.iloc[index]\n        data = self.eegs[row.eeg_id]  # Size=(10000, 8)\n\n        # 10000から4000切り出し\n        if self.L != self.L:\n            offset = (self.CFG.nsamples - self.L ) // 2\n            data = data[offset:offset+self.L,:]\n\n        # diff of eegs\n        for i, (feat_a, feat_b) in enumerate(self.CFG.map_features):\n            if self.mode == \"train\" and self.CFG.random_close_zone > 0 and random.uniform(0.0, 1.0) <= self.CFG.random_close_zone:\n                continue\n                \n            diff_feat = data[:, self.CFG.feature_to_index[feat_a]] - data[:, self.CFG.feature_to_index[feat_b]]\n\n            if not self.bandpass_filter is None:\n                diff_feat = butter_bandpass_filter(\n                    diff_feat,\n                    self.bandpass_filter[\"low\"],\n                    self.bandpass_filter[\"high\"],\n                    self.CFG.sampling_rate,\n                    order=self.bandpass_filter[\"order\"],\n                )\n                    \n            if (\n                self.mode == \"train\"\n                and not self.rand_filter is None\n                and random.uniform(0.0, 1.0) <= self.rand_filter[\"probab\"]\n            ):\n                lowcut = random.randint(\n                    self.rand_filter[\"low\"], self.rand_filter[\"high\"]\n                )\n                highcut = lowcut + self.rand_filter[\"band\"]\n                diff_feat = butter_bandpass_filter(\n                    diff_feat,\n                    lowcut,\n                    highcut,\n                    self.CFG.sampling_rate,\n                    order=self.rand_filter[\"order\"],\n                )\n\n            X[:, i] = diff_feat\n\n        # other frequency filtering\n        n = self.CFG.n_map_features\n        if len(self.CFG.freq_channels) > 0:\n            for i in range(self.CFG.n_map_features):\n                diff_feat = X[:, i]\n                for j, (lowcut, highcut) in enumerate(self.CFG.freq_channels):\n                    band_feat = butter_bandpass_filter(\n                        diff_feat, lowcut, highcut, self.CFG.sampling_rate, order=self.CFG.filter_order,  # 6\n                    )\n                    X[:, n] = band_feat\n                    n += 1\n                    \n        # single eeg \n        for spml_feat in self.CFG.simple_features:\n            feat_val = data[:, self.CFG.feature_to_index[spml_feat]]\n            \n            if not self.bandpass_filter is None:\n                feat_val = butter_bandpass_filter(\n                    feat_val,\n                    self.bandpass_filter[\"low\"],\n                    self.bandpass_filter[\"high\"],\n                    self.CFG.sampling_rate,\n                    order=self.bandpass_filter[\"order\"],\n                )\n\n            if (\n                self.mode == \"train\"\n                and not self.rand_filter is None\n                and random.uniform(0.0, 1.0) <= self.rand_filter[\"probab\"]\n            ):\n                lowcut = random.randint(\n                    self.rand_filter[\"low\"], self.rand_filter[\"high\"]\n                )\n                highcut = lowcut + self.rand_filter[\"band\"]\n                feat_val = butter_bandpass_filter(\n                    feat_val,\n                    lowcut,\n                    highcut,\n                    self.CFG.sampling_rate,\n                    order=self.rand_filter[\"order\"],\n                )\n\n            X[:, n] = feat_val\n            n += 1\n            \n        X = np.clip(X, -1024, 1024)\n        X = np.nan_to_num(X, nan=0) / 32.0\n#         X = butter_lowpass_filter(X, order=self.CFG.filter_order)\n\n        y = np.zeros(self.CFG.n_classes, dtype=\"float32\")  # Size=(6,)\n\n        return X, y","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.028292Z","iopub.execute_input":"2024-05-03T21:55:35.028585Z","iopub.status.idle":"2024-05-03T21:55:35.375636Z","shell.execute_reply.started":"2024-05-03T21:55:35.028561Z","shell.execute_reply":"2024-05-03T21:55:35.374520Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# cwt_path = \"/kaggle/working/cwt_all_20sec_stride4/\"\n# os.makedirs(cwt_path, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.379004Z","iopub.execute_input":"2024-05-03T21:55:35.379333Z","iopub.status.idle":"2024-05-03T21:55:35.384059Z","shell.execute_reply.started":"2024-05-03T21:55:35.379304Z","shell.execute_reply":"2024-05-03T21:55:35.383108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dataset = HMSDataset_preprocesser(df=test_df, eegs=all_eegs, CFG=CFG, mode=\"test\", transforms=None)\n\n# pre_loader = DataLoader(\n#     dataset,\n#     batch_size=4,\n#     shuffle=False,\n#     num_workers=4, pin_memory=True, drop_last=False\n# )\n\n# if len(model_weights) > 0:\n#     for data in tqdm(pre_loader):\n#         for idx, cwtidx in enumerate(data[\"idx\"].numpy()):\n#             np.save(f\"{cwt_path}cwt_all_{cwtidx:06}.npy\", data[\"data\"][idx].numpy())","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.385178Z","iopub.execute_input":"2024-05-03T21:55:35.385501Z","iopub.status.idle":"2024-05-03T21:55:35.392734Z","shell.execute_reply.started":"2024-05-03T21:55:35.385475Z","shell.execute_reply":"2024-05-03T21:55:35.391999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_df['spec_path'] = test_df.index.map(lambda x: f\"{cwt_path}cwt_all_{x:06}.npy\")","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.394938Z","iopub.execute_input":"2024-05-03T21:55:35.395618Z","iopub.status.idle":"2024-05-03T21:55:35.404098Z","shell.execute_reply.started":"2024-05-03T21:55:35.395593Z","shell.execute_reply":"2024-05-03T21:55:35.403246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HMSModel(nn.Module):\n    def __init__(self, model_name: str, pretrained: bool, in_channels: int, num_classes: int):\n        super().__init__()\n        \n        self.model = timm.create_model(\n            model_name=model_name, \n            pretrained=pretrained, \n            in_chans=in_channels,\n            num_classes = num_classes\n        )\n        \n    def forward(self, x):\n        x = self.model(x)  \n        return x","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.404991Z","iopub.execute_input":"2024-05-03T21:55:35.405260Z","iopub.status.idle":"2024-05-03T21:55:35.414292Z","shell.execute_reply.started":"2024-05-03T21:55:35.405237Z","shell.execute_reply":"2024-05-03T21:55:35.413423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchaudio.transforms as T\n\nclass HMSDataset_2D(torch.utils.data.Dataset):\n    def __init__(self, df, eegs, CFG,\n                 downsample: int = None,\n                 bandpass_filter = None,\n                 rand_filter = None,\n                 mode='train',\n                 weighted=False,\n                 transforms = None,\n                ):\n        self.df = df\n        self.eegs = eegs\n        self.CFG = CFG\n        self.downsample = downsample\n        self.bandpass_filter = bandpass_filter\n        self.rand_filter = rand_filter\n        self.mode = mode\n        self.weighted = weighted\n        self.transforms = transforms\n        self.spec_imgsize_hw = [128, 256]\n\n        self.L           = CFG.nsamples    # 4000\n        self.L_input     = CFG.out_samples # 2000\n        \n        self.L_stride          = CFG.nsamples//CFG.CWT_STRIDE\n        self.L_input_stride    = CFG.out_samples//CFG.CWT_STRIDE\n        \n            \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx: int):\n        row = self.df.iloc[idx]\n        eeg_id = row.eeg_id\n        \n        img = np.zeros((10*self.CFG.in_channels, self.CFG.out_samples, 1)).astype(np.float32)\n        \n        X, y, offset = self.__data_generation(idx)\n\n        #####################\n        # 2D\n        img_spec = torch.from_numpy(np.load(row.spec_path))\n\n        if self.L != self.L_input:\n            offset_stride = offset//self.CFG.CWT_STRIDE\n            img_spec = img_spec[:, offset_stride: offset_stride+self.L_input_stride]\n\n        if self.transforms is not None:\n            img_spec = self.transforms(image=img_spec.numpy() )['image']\n        #####################\n        # 1D\n#         if self.downsample is not None:\n#             X = X[::self.downsample, :]\n            \n#         for i in range(X.shape[1]):\n#             img[i*10:(i+1)*10, :] = X[:,i][:,np.newaxis]\n                       \n#         if self.transforms is not None:\n#             img = self.transforms(image=img)['image']\n                \n        #####################\n        #y = row[self.CFG.classes].values.astype(np.float32)\n        #if self.weighted:\n        #    w = row[\"num_votes\"].astype(np.float32) / 10.0\n        #else:\n        #    w = 1.0\n        \n#         img = torch.cat([img, img_spec], dim=0)\n        img = img_spec\n        return {\"data\": img}\n    \n    def __data_generation(self, index):\n        X = np.zeros((self.CFG.out_samples, self.CFG.in_channels), dtype=\"float32\")  # Size=(10000, 14)\n\n        row = self.df.iloc[index]\n        data = self.eegs[row.eeg_id]  # Size=(10000, 8)\n        \n        offset=0\n        if self.L != self.L_input:\n            offset = (self.L - self.L_input) // 2\n            data = data[offset:offset+self.L_input,:]\n\n        # diff of eegs\n        for i, (feat_a, feat_b) in enumerate(self.CFG.map_features):\n            if self.mode == \"train\" and self.CFG.random_close_zone > 0 and random.uniform(0.0, 1.0) <= self.CFG.random_close_zone:\n                continue\n                \n            diff_feat = data[:, self.CFG.feature_to_index[feat_a]] - data[:, self.CFG.feature_to_index[feat_b]]\n\n            if not self.bandpass_filter is None:\n                diff_feat = butter_bandpass_filter(\n                    diff_feat,\n                    self.bandpass_filter[\"low\"],\n                    self.bandpass_filter[\"high\"],\n                    self.CFG.sampling_rate,\n                    order=self.bandpass_filter[\"order\"],\n                )\n                    \n            if (\n                self.mode == \"train\"\n                and not self.rand_filter is None\n                and random.uniform(0.0, 1.0) <= self.rand_filter[\"probab\"]\n            ):\n                lowcut = random.randint(\n                    self.rand_filter[\"low\"], self.rand_filter[\"high\"]\n                )\n                highcut = lowcut + self.rand_filter[\"band\"]\n                diff_feat = butter_bandpass_filter(\n                    diff_feat,\n                    lowcut,\n                    highcut,\n                    self.CFG.sampling_rate,\n                    order=self.rand_filter[\"order\"],\n                )\n\n            X[:, i] = diff_feat\n\n        # other frequency filtering\n        n = self.CFG.n_map_features\n        if len(self.CFG.freq_channels) > 0:\n            for i in range(self.CFG.n_map_features):\n                diff_feat = X[:, i]\n                for j, (lowcut, highcut) in enumerate(self.CFG.freq_channels):\n                    band_feat = butter_bandpass_filter(\n                        diff_feat, lowcut, highcut, self.CFG.sampling_rate, order=self.CFG.filter_order,  # 6\n                    )\n                    X[:, n] = band_feat\n                    n += 1\n                    \n        # single eeg \n        for spml_feat in self.CFG.simple_features:\n            feat_val = data[:, self.CFG.feature_to_index[spml_feat]]\n            \n            if not self.bandpass_filter is None:\n                feat_val = butter_bandpass_filter(\n                    feat_val,\n                    self.bandpass_filter[\"low\"],\n                    self.bandpass_filter[\"high\"],\n                    self.CFG.sampling_rate,\n                    order=self.bandpass_filter[\"order\"],\n                )\n\n            if (\n                self.mode == \"train\"\n                and not self.rand_filter is None\n                and random.uniform(0.0, 1.0) <= self.rand_filter[\"probab\"]\n            ):\n                lowcut = random.randint(\n                    self.rand_filter[\"low\"], self.rand_filter[\"high\"]\n                )\n                highcut = lowcut + self.rand_filter[\"band\"]\n                feat_val = butter_bandpass_filter(\n                    feat_val,\n                    lowcut,\n                    highcut,\n                    CFG.sampling_rate,\n                    order=self.rand_filter[\"order\"],\n                )\n\n            X[:, n] = feat_val\n            n += 1\n            \n        X = np.clip(X, -1024, 1024)\n        X = np.nan_to_num(X, nan=0) / 32.0\n#         X = butter_lowpass_filter(X, order=self.CFG.filter_order)\n\n        y = np.zeros(self.CFG.n_classes, dtype=\"float32\")  # Size=(6,)\n        if self.mode != \"test\":\n            y = row[self.CFG.classes].values.astype(np.float32)\n\n        return X, y, offset","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.415708Z","iopub.execute_input":"2024-05-03T21:55:35.415970Z","iopub.status.idle":"2024-05-03T21:55:35.443805Z","shell.execute_reply.started":"2024-05-03T21:55:35.415948Z","shell.execute_reply":"2024-05-03T21:55:35.442866Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# coef_sum = 0\n# coef_count = 0\n# predictions = []\n# files = []\n    \n# for model_block in model_weights:\n#     test_dataset = HMSDataset_2D(\n#         df=test_df,\n#         mode=\"test\",\n#         eegs=all_eegs,\n#         CFG=CFG, \n#         bandpass_filter=model_block['bandpass_filter'],\n#         transforms=data_transforms['valid']\n#     )\n\n#     if len(predictions) == 0:\n#         output = test_dataset[0]\n#         X = output[\"data\"]\n#         print(f\"X shape: {X.shape}\")\n                \n#     test_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#     )\n\n#     model = HMSModel(\n#         model_name = model_block['model_name'], \n#         pretrained = False,\n#         in_channels = CFG.in_chans, \n#         num_classes = CFG.n_classes\n#     )\n\n#     single_model_preds = []\n#     for file_line in model_block['file_data']:\n#         coef = file_line['coef']\n#         for weight_model_file in glob(file_line['file_mask']):\n#             files.append(weight_model_file)\n#             checkpoint = torch.load(weight_model_file, map_location=CFG.device)\n#             model.load_state_dict(checkpoint)\n#             model.to(CFG.device)\n#             prediction_dict = inference_function(test_loader, model, CFG.device)\n#             predict = prediction_dict[\"predictions\"]\n#             #predict *= coef\n#             #coef_sum += coef\n#             #coef_count += 1\n#             single_model_preds.append(predict)\n#             torch.cuda.empty_cache()\n#             gc.collect()\n            \n#     single_model_preds = np.mean(single_model_preds, axis=0)\n#     predictions.append(single_model_preds.copy())\n\n# k_cwt10sec_predictions = predictions.copy()\n# #predictions = np.array(predictions)\n# #coef_sum /= coef_count\n# #predictions /= coef_sum\n# #predictions_2d_10sec = np.mean(predictions, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.445161Z","iopub.execute_input":"2024-05-03T21:55:35.445519Z","iopub.status.idle":"2024-05-03T21:55:35.457611Z","shell.execute_reply.started":"2024-05-03T21:55:35.445490Z","shell.execute_reply":"2024-05-03T21:55:35.456832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import shutil\n# shutil.rmtree(cwt_path)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.458637Z","iopub.execute_input":"2024-05-03T21:55:35.458972Z","iopub.status.idle":"2024-05-03T21:55:35.470321Z","shell.execute_reply.started":"2024-05-03T21:55:35.458942Z","shell.execute_reply":"2024-05-03T21:55:35.469448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    debug = False\n    \n    CWT_STRIDE = 8\n            \n    seq_length = 50  # Second's\n    sampling_rate = 200  # Hz\n    nsamples = seq_length * sampling_rate\n    out_samples = 25 * sampling_rate\n    \n    \n    seed = 42\n    gpu_idx = 0\n\n    device = torch.device(f\"cuda:{gpu_idx}\")\n    \n    classes = [\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]\n    n_classes = len(classes)\n\n    n_folds = 5\n    \n    ## Model\n    fixed_kernel_size = 5\n    #linear_layer_features = 448  # Full Signal = 10_000\n    #linear_layer_features = 352  # Half Signal = 5_000\n    linear_layer_features = 304   # 1/5  Signal = 2_000\n    kernels = [3, 5, 7, 9, 11]\n    \n    ## Dataset Preprocessing\n    bandpass_filter = None\n    # rand_filter = {\"probab\": 0.1, \"low\": 10, \"high\": 20, \"band\": 1.0, \"order\": 2}\n    rand_filter = {}\n    # freq_channels = [(8.0, 12.0), (0.5, 4.5)]\n    freq_channels = []\n    \n    filter_order = 2\n    random_close_zone = 0.0  # 0.2\n    \n    map_features = [\n        (\"Fp1\", \"F7\"),\n        (\"F7\", \"T3\"),\n        (\"T3\", \"T5\"),\n        (\"T5\", \"O1\"),\n        (\"Fp1\", \"F3\"),\n        (\"F3\", \"C3\"),\n        (\"C3\", \"P3\"),\n        (\"P3\", \"O1\"),\n        (\"Fp2\", \"F8\"),\n        (\"F8\", \"T4\"),\n        (\"T4\", \"T6\"),\n        (\"T6\", \"O2\"),\n        (\"Fp2\", \"F4\"),\n        (\"F4\", \"C4\"),\n        (\"C4\", \"P4\"),\n        (\"P4\", \"O2\"),\n        (\"Fz\", \"Cz\"),\n        (\"Cz\", \"Pz\")\n    ]\n    \n    eeg_features = [\"Fp1\",\"F3\",\"C3\",\"P3\",\"F7\",\"T3\",\"T5\",\"O1\",\"Fz\",\"Cz\",\"Pz\",\"Fp2\",\"F4\",\"C4\",\"P4\",\"F8\",\"T4\",\"T6\",\"O2\",\"EKG\"]\n    #eeg_features = [\"Fp1\", \"T3\", \"C3\", \"O1\", \"Fp2\", \"C4\", \"T4\", \"O2\"]\n    feature_to_index = {x: y for x, y in zip(eeg_features, range(len(eeg_features)))}\n    simple_features = []  # 'Fz', 'Cz', 'Pz', 'EKG'\n    \n    n_map_features = len(map_features)\n    in_channels = n_map_features #+ n_map_features * len(freq_channels) + len(simple_features)\n\n    input_size = 512\n    in_chans = 1\n\n    batch_size = 4\n    num_workers = 4","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.471505Z","iopub.execute_input":"2024-05-03T21:55:35.471835Z","iopub.status.idle":"2024-05-03T21:55:35.487594Z","shell.execute_reply.started":"2024-05-03T21:55:35.471806Z","shell.execute_reply":"2024-05-03T21:55:35.486279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(f\"/kaggle/input/hms-harmful-brain-activity-classification/{mode}.csv\")\nprint(f\"Test dataframe shape is: {test_df.shape}\")\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.488655Z","iopub.execute_input":"2024-05-03T21:55:35.489095Z","iopub.status.idle":"2024-05-03T21:55:35.508236Z","shell.execute_reply.started":"2024-05-03T21:55:35.489062Z","shell.execute_reply":"2024-05-03T21:55:35.507259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_eeg_parquet_paths = glob(f\"/kaggle/input/hms-harmful-brain-activity-classification/{mode}_eegs/\" + \"*.parquet\")\ntest_eeg_df = pd.read_parquet(test_eeg_parquet_paths[0])\ntest_eeg_features = test_eeg_df.columns\nprint(f\"There are {len(test_eeg_features)} raw eeg features\")\nprint(list(test_eeg_features))\ndel test_eeg_df\n_ = gc.collect()\n\n# %%time\nall_eegs = {}\neeg_ids = test_df.eeg_id.unique()\nfor i, eeg_id in tqdm(enumerate(eeg_ids)):\n    # Save EEG to Python dictionary of numpy arrays\n    eeg_path = f\"/kaggle/input/hms-harmful-brain-activity-classification/{mode}_eegs/\" + str(eeg_id) + \".parquet\"\n    data = eeg_from_parquet(eeg_path, CFG=CFG)\n    all_eegs[eeg_id] = data","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.509720Z","iopub.execute_input":"2024-05-03T21:55:35.510004Z","iopub.status.idle":"2024-05-03T21:55:35.758463Z","shell.execute_reply.started":"2024-05-03T21:55:35.509981Z","shell.execute_reply":"2024-05-03T21:55:35.757543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# cwt_path = \"/kaggle/working/cwt_all_50sec_stride8/\"\n# os.makedirs(cwt_path, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.759553Z","iopub.execute_input":"2024-05-03T21:55:35.759836Z","iopub.status.idle":"2024-05-03T21:55:35.763939Z","shell.execute_reply.started":"2024-05-03T21:55:35.759805Z","shell.execute_reply":"2024-05-03T21:55:35.763044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dataset = HMSDataset_preprocesser(df=test_df, eegs=all_eegs, CFG=CFG, mode=\"test\", transforms=None)\n\n# pre_loader = DataLoader(\n#     dataset,\n#     batch_size=4,\n#     shuffle=False,\n#     num_workers=4, pin_memory=True, drop_last=False\n# )\n\n# for data in tqdm(pre_loader):\n#     for idx, cwtidx in enumerate(data[\"idx\"].numpy()):\n#         np.save(f\"{cwt_path}cwt_all_{cwtidx:06}.npy\", data[\"data\"][idx].numpy())","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.765032Z","iopub.execute_input":"2024-05-03T21:55:35.765488Z","iopub.status.idle":"2024-05-03T21:55:35.773484Z","shell.execute_reply.started":"2024-05-03T21:55:35.765439Z","shell.execute_reply":"2024-05-03T21:55:35.772724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_df['spec_path'] = test_df.index.map(lambda x: f\"{cwt_path}cwt_all_{x:06}.npy\")","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.774586Z","iopub.execute_input":"2024-05-03T21:55:35.775493Z","iopub.status.idle":"2024-05-03T21:55:35.782609Z","shell.execute_reply.started":"2024-05-03T21:55:35.775468Z","shell.execute_reply":"2024-05-03T21:55:35.781679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# coef_sum = 0\n# coef_count = 0\n# predictions = []\n# files = []\n    \n# for model_block in model_weights_25sec:\n#     test_dataset = HMSDataset_2D(\n#         df=test_df,\n#         mode=\"test\",\n#         CFG=CFG, \n#         eegs=all_eegs,\n#         bandpass_filter=model_block['bandpass_filter'],\n#         transforms=data_transforms['valid']\n#     )\n\n#     if len(predictions) == 0:\n#         output = test_dataset[0]\n#         X = output[\"data\"]\n#         print(f\"X shape: {X.shape}\")\n                \n#     test_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#     )\n\n#     model = HMSModel(\n#         model_name = model_block['model_name'], \n#         pretrained = False,\n#         in_channels = CFG.in_chans, \n#         num_classes = CFG.n_classes\n#     )\n\n#     single_model_preds = []\n#     for file_line in model_block['file_data']:\n#         coef = file_line['coef']\n#         for weight_model_file in glob(file_line['file_mask']):\n#             files.append(weight_model_file)\n#             checkpoint = torch.load(weight_model_file, map_location=CFG.device)\n#             model.load_state_dict(checkpoint)\n#             model.to(CFG.device)\n#             prediction_dict = inference_function(test_loader, model, CFG.device)\n#             predict = prediction_dict[\"predictions\"]\n#             #predict *= coef\n#             #coef_sum += coef\n#             #coef_count += 1\n#             single_model_preds.append(predict)\n#             torch.cuda.empty_cache()\n#             gc.collect()\n            \n#     single_model_preds = np.mean(single_model_preds, axis=0)\n#     predictions.append(single_model_preds.copy())\n\n# k_cwt25sec_predictions = predictions.copy()\n# #predictions = np.array(predictions)\n# #coef_sum /= coef_count\n# #predictions /= coef_sum\n# #predictions_2d_25sec = np.mean(predictions, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.783750Z","iopub.execute_input":"2024-05-03T21:55:35.784012Z","iopub.status.idle":"2024-05-03T21:55:35.792615Z","shell.execute_reply.started":"2024-05-03T21:55:35.783990Z","shell.execute_reply":"2024-05-03T21:55:35.791902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import shutil\n# shutil.rmtree(cwt_path)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.793870Z","iopub.execute_input":"2024-05-03T21:55:35.794194Z","iopub.status.idle":"2024-05-03T21:55:35.804380Z","shell.execute_reply.started":"2024-05-03T21:55:35.794164Z","shell.execute_reply":"2024-05-03T21:55:35.803514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 50sec 0.5~20hz","metadata":{}},{"cell_type":"code","source":"class CFG:\n    debug = False\n    \n    CWT_STRIDE = 16\n            \n    seq_length = 50  # Second's\n    sampling_rate = 200  # Hz\n    nsamples = seq_length * sampling_rate\n    out_samples = 50 * sampling_rate\n    \n    \n    seed = 42\n    gpu_idx = 0\n\n    device = torch.device(f\"cuda:{gpu_idx}\")\n    \n    classes = [\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]\n    n_classes = len(classes)\n\n    n_folds = 5\n    \n    ## Model\n    fixed_kernel_size = 5\n    #linear_layer_features = 448  # Full Signal = 10_000\n    #linear_layer_features = 352  # Half Signal = 5_000\n    linear_layer_features = 304   # 1/5  Signal = 2_000\n    kernels = [3, 5, 7, 9, 11]\n    \n    ## Dataset Preprocessing\n    bandpass_filter = None\n    # rand_filter = {\"probab\": 0.1, \"low\": 10, \"high\": 20, \"band\": 1.0, \"order\": 2}\n    rand_filter = {}\n    # freq_channels = [(8.0, 12.0), (0.5, 4.5)]\n    freq_channels = []\n    \n    filter_order = 2\n    random_close_zone = 0.0  # 0.2\n    \n    map_features = [\n        (\"Fp1\", \"F7\"),\n        (\"F7\", \"T3\"),\n        (\"T3\", \"T5\"),\n        (\"T5\", \"O1\"),\n        (\"Fp1\", \"F3\"),\n        (\"F3\", \"C3\"),\n        (\"C3\", \"P3\"),\n        (\"P3\", \"O1\"),\n        (\"Fp2\", \"F8\"),\n        (\"F8\", \"T4\"),\n        (\"T4\", \"T6\"),\n        (\"T6\", \"O2\"),\n        (\"Fp2\", \"F4\"),\n        (\"F4\", \"C4\"),\n        (\"C4\", \"P4\"),\n        (\"P4\", \"O2\"),\n        (\"Fz\", \"Cz\"),\n        (\"Cz\", \"Pz\")\n    ]\n    \n    eeg_features = [\"Fp1\",\"F3\",\"C3\",\"P3\",\"F7\",\"T3\",\"T5\",\"O1\",\"Fz\",\"Cz\",\"Pz\",\"Fp2\",\"F4\",\"C4\",\"P4\",\"F8\",\"T4\",\"T6\",\"O2\",\"EKG\"]\n    #eeg_features = [\"Fp1\", \"T3\", \"C3\", \"O1\", \"Fp2\", \"C4\", \"T4\", \"O2\"]\n    feature_to_index = {x: y for x, y in zip(eeg_features, range(len(eeg_features)))}\n    simple_features = []  # 'Fz', 'Cz', 'Pz', 'EKG'\n    \n    n_map_features = len(map_features)\n    in_channels = n_map_features #+ n_map_features * len(freq_channels) + len(simple_features)\n\n    input_size = 512\n    in_chans = 1\n\n    batch_size = 4\n    num_workers = 4","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.805535Z","iopub.execute_input":"2024-05-03T21:55:35.805796Z","iopub.status.idle":"2024-05-03T21:55:35.817261Z","shell.execute_reply.started":"2024-05-03T21:55:35.805768Z","shell.execute_reply":"2024-05-03T21:55:35.816466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(f\"/kaggle/input/hms-harmful-brain-activity-classification/{mode}.csv\")\nprint(f\"Test dataframe shape is: {test_df.shape}\")\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.818451Z","iopub.execute_input":"2024-05-03T21:55:35.818772Z","iopub.status.idle":"2024-05-03T21:55:35.834594Z","shell.execute_reply.started":"2024-05-03T21:55:35.818742Z","shell.execute_reply":"2024-05-03T21:55:35.833667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_eeg_parquet_paths = glob(f\"/kaggle/input/hms-harmful-brain-activity-classification/{mode}_eegs/\" + \"*.parquet\")\ntest_eeg_df = pd.read_parquet(test_eeg_parquet_paths[0])\ntest_eeg_features = test_eeg_df.columns\nprint(f\"There are {len(test_eeg_features)} raw eeg features\")\nprint(list(test_eeg_features))\ndel test_eeg_df\n_ = gc.collect()\n\n# %%time\nall_eegs = {}\neeg_ids = test_df.eeg_id.unique()\nfor i, eeg_id in tqdm(enumerate(eeg_ids)):\n    # Save EEG to Python dictionary of numpy arrays\n    eeg_path = f\"/kaggle/input/hms-harmful-brain-activity-classification/{mode}_eegs/\" + str(eeg_id) + \".parquet\"\n    data = eeg_from_parquet(eeg_path, CFG=CFG)\n    all_eegs[eeg_id] = data","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:35.835702Z","iopub.execute_input":"2024-05-03T21:55:35.835961Z","iopub.status.idle":"2024-05-03T21:55:36.072594Z","shell.execute_reply.started":"2024-05-03T21:55:35.835940Z","shell.execute_reply":"2024-05-03T21:55:36.071643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# cwt_path = \"/kaggle/working/cwt_all_50sec_stride16/\"\n# os.makedirs(cwt_path, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:36.073677Z","iopub.execute_input":"2024-05-03T21:55:36.074020Z","iopub.status.idle":"2024-05-03T21:55:36.077918Z","shell.execute_reply.started":"2024-05-03T21:55:36.073986Z","shell.execute_reply":"2024-05-03T21:55:36.077034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dataset = HMSDataset_preprocesser(df=test_df, eegs=all_eegs, CFG=CFG, mode=\"test\", transforms=None)\n\n# pre_loader = DataLoader(\n#     dataset,\n#     batch_size=4,\n#     shuffle=False,\n#     num_workers=4, pin_memory=True, drop_last=False\n# )\n\n# for data in tqdm(pre_loader):\n#     for idx, cwtidx in enumerate(data[\"idx\"].numpy()):\n#         np.save(f\"{cwt_path}cwt_all_{cwtidx:06}.npy\", data[\"data\"][idx].numpy())","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:36.079166Z","iopub.execute_input":"2024-05-03T21:55:36.079735Z","iopub.status.idle":"2024-05-03T21:55:36.088019Z","shell.execute_reply.started":"2024-05-03T21:55:36.079702Z","shell.execute_reply":"2024-05-03T21:55:36.087099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_df['spec_path'] = test_df.index.map(lambda x: f\"{cwt_path}cwt_all_{x:06}.npy\")","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:36.089069Z","iopub.execute_input":"2024-05-03T21:55:36.089417Z","iopub.status.idle":"2024-05-03T21:55:36.097649Z","shell.execute_reply.started":"2024-05-03T21:55:36.089385Z","shell.execute_reply":"2024-05-03T21:55:36.096780Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# coef_sum = 0\n# coef_count = 0\n# predictions = []\n# files = []\n    \n# for model_block in model_weights_50sec:\n#     test_dataset = HMSDataset_2D(\n#         df=test_df,\n#         mode=\"test\",\n#         CFG=CFG, \n#         eegs=all_eegs,\n#         bandpass_filter=model_block['bandpass_filter'],\n#         transforms=data_transforms['valid']\n#     )\n\n#     if len(predictions) == 0:\n#         output = test_dataset[0]\n#         X = output[\"data\"]\n#         print(f\"X shape: {X.shape}\")\n                \n#     test_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#     )\n\n#     model = HMSModel(\n#         model_name = model_block['model_name'], \n#         pretrained = False,\n#         in_channels = CFG.in_chans, \n#         num_classes = CFG.n_classes\n#     )\n\n#     single_model_preds = []\n#     for file_line in model_block['file_data']:\n#         coef = file_line['coef']\n#         for weight_model_file in glob(file_line['file_mask']):\n#             files.append(weight_model_file)\n#             checkpoint = torch.load(weight_model_file, map_location=CFG.device)\n#             model.load_state_dict(checkpoint)\n#             model.to(CFG.device)\n#             prediction_dict = inference_function(test_loader, model, CFG.device)\n#             predict = prediction_dict[\"predictions\"]\n#             #predict *= coef\n#             #coef_sum += coef\n#             #coef_count += 1\n#             single_model_preds.append(predict)\n#             torch.cuda.empty_cache()\n#             gc.collect()\n            \n#     single_model_preds = np.mean(single_model_preds, axis=0)\n#     predictions.append(single_model_preds.copy())\n\n# k_cwt50sec_predictions = predictions.copy()\n# #predictions = np.array(predictions)\n# #coef_sum /= coef_count\n# #predictions /= coef_sum\n# #predictions_2d_50sec = np.mean(predictions, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:36.098669Z","iopub.execute_input":"2024-05-03T21:55:36.098964Z","iopub.status.idle":"2024-05-03T21:55:36.108043Z","shell.execute_reply.started":"2024-05-03T21:55:36.098941Z","shell.execute_reply":"2024-05-03T21:55:36.107037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import shutil\n# shutil.rmtree(cwt_path)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:36.109163Z","iopub.execute_input":"2024-05-03T21:55:36.109504Z","iopub.status.idle":"2024-05-03T21:55:36.119928Z","shell.execute_reply.started":"2024-05-03T21:55:36.109469Z","shell.execute_reply":"2024-05-03T21:55:36.119218Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 50sec 40hz","metadata":{}},{"cell_type":"code","source":"cwt_path = \"/kaggle/working/cwt_all_50sec_40hz_stride16/\"\nos.makedirs(cwt_path, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:36.120997Z","iopub.execute_input":"2024-05-03T21:55:36.121293Z","iopub.status.idle":"2024-05-03T21:55:36.128990Z","shell.execute_reply.started":"2024-05-03T21:55:36.121266Z","shell.execute_reply":"2024-05-03T21:55:36.128175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchaudio.transforms as T\n\nclass HMSDataset_preprocesser(torch.utils.data.Dataset):\n    def __init__(self, df, eegs, CFG,\n                 downsample: int = None,\n                 bandpass_filter = None,\n                 rand_filter = None,\n                 mode='train',\n                 weighted=False,\n                 transforms = None,\n                ):\n        self.df = df\n        self.eegs = eegs\n        self.CFG = CFG\n        self.downsample = downsample\n        self.bandpass_filter = bandpass_filter\n        self.rand_filter = rand_filter\n        self.mode = mode\n        self.weighted = weighted\n        self.transforms = transforms\n        self.spec_imgsize_hw = [128, 256]\n        self.L       = CFG.nsamples    # 4000\n        self.L_input = CFG.out_samples # 2000\n        self.cwt = CWT(wavelet_width=7, fs=200, \\\n                       lower_freq=0.5, upper_freq=40, \\\n                       n_scales=40, border_crop = 1, stride = self.CFG.CWT_STRIDE)\n                   \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx: int):\n        row = self.df.iloc[idx]\n        eeg_id = row.eeg_id\n        \n        img = np.zeros((10*self.CFG.in_channels, self.CFG.out_samples, 1)).astype(np.float32)\n        \n        X, y = self.__data_generation(idx)\n        #####################\n        # 2D\n        # img_spec = torch.zeros((X.shape[1], self.spec_imgsize_hw[0], self.spec_imgsize_hw[1]), dtype=torch.float32)  # GPU対応\n        X_in = torch.tensor(X.transpose(1, 0), dtype=torch.float32).unsqueeze(0)\n        img_spec = self.cwt(X_in)\n        img_spec = img_spec[0]\n        img_spec = torch.cat([img_spec[i] for i in range(img_spec.shape[0])], dim=0)\n        \n        if self.transforms is not None:\n            img_spec = self.transforms(image=img_spec.numpy() )['image']\n            \n        return {\"data\": img_spec, \"idx\": idx}\n    \n    def __data_generation(self, index):\n        X = np.zeros(( self.L, self.CFG.in_channels), dtype=\"float32\")  # Size=(10000, 14)\n\n        row = self.df.iloc[index]\n        data = self.eegs[row.eeg_id]  # Size=(10000, 8)\n\n        # 10000から4000切り出し\n        if self.L != self.L:\n            offset = (self.CFG.nsamples - self.L ) // 2\n            data = data[offset:offset+self.L,:]\n\n        # diff of eegs\n        for i, (feat_a, feat_b) in enumerate(self.CFG.map_features):\n            if self.mode == \"train\" and self.CFG.random_close_zone > 0 and random.uniform(0.0, 1.0) <= self.CFG.random_close_zone:\n                continue\n                \n            diff_feat = data[:, self.CFG.feature_to_index[feat_a]] - data[:, self.CFG.feature_to_index[feat_b]]\n\n            if not self.bandpass_filter is None:\n                diff_feat = butter_bandpass_filter(\n                    diff_feat,\n                    self.bandpass_filter[\"low\"],\n                    self.bandpass_filter[\"high\"],\n                    self.CFG.sampling_rate,\n                    order=self.bandpass_filter[\"order\"],\n                )\n                    \n            if (\n                self.mode == \"train\"\n                and not self.rand_filter is None\n                and random.uniform(0.0, 1.0) <= self.rand_filter[\"probab\"]\n            ):\n                lowcut = random.randint(\n                    self.rand_filter[\"low\"], self.rand_filter[\"high\"]\n                )\n                highcut = lowcut + self.rand_filter[\"band\"]\n                diff_feat = butter_bandpass_filter(\n                    diff_feat,\n                    lowcut,\n                    highcut,\n                    self.CFG.sampling_rate,\n                    order=self.rand_filter[\"order\"],\n                )\n\n            X[:, i] = diff_feat\n\n        # other frequency filtering\n        n = self.CFG.n_map_features\n        if len(self.CFG.freq_channels) > 0:\n            for i in range(self.CFG.n_map_features):\n                diff_feat = X[:, i]\n                for j, (lowcut, highcut) in enumerate(self.CFG.freq_channels):\n                    band_feat = butter_bandpass_filter(\n                        diff_feat, lowcut, highcut, self.CFG.sampling_rate, order=self.CFG.filter_order,  # 6\n                    )\n                    X[:, n] = band_feat\n                    n += 1\n                    \n        # single eeg \n        for spml_feat in self.CFG.simple_features:\n            feat_val = data[:, self.CFG.feature_to_index[spml_feat]]\n            \n            if not self.bandpass_filter is None:\n                feat_val = butter_bandpass_filter(\n                    feat_val,\n                    self.bandpass_filter[\"low\"],\n                    self.bandpass_filter[\"high\"],\n                    self.CFG.sampling_rate,\n                    order=self.bandpass_filter[\"order\"],\n                )\n\n            if (\n                self.mode == \"train\"\n                and not self.rand_filter is None\n                and random.uniform(0.0, 1.0) <= self.rand_filter[\"probab\"]\n            ):\n                lowcut = random.randint(\n                    self.rand_filter[\"low\"], self.rand_filter[\"high\"]\n                )\n                highcut = lowcut + self.rand_filter[\"band\"]\n                feat_val = butter_bandpass_filter(\n                    feat_val,\n                    lowcut,\n                    highcut,\n                    self.CFG.sampling_rate,\n                    order=self.rand_filter[\"order\"],\n                )\n\n            X[:, n] = feat_val\n            n += 1\n            \n        X = np.clip(X, -1024, 1024)\n        X = np.nan_to_num(X, nan=0) / 32.0\n#         X = butter_lowpass_filter(X, order=self.CFG.filter_order)\n\n        y = np.zeros(self.CFG.n_classes, dtype=\"float32\")  # Size=(6,)\n\n        return X, y","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:36.130332Z","iopub.execute_input":"2024-05-03T21:55:36.130702Z","iopub.status.idle":"2024-05-03T21:55:36.158144Z","shell.execute_reply.started":"2024-05-03T21:55:36.130672Z","shell.execute_reply":"2024-05-03T21:55:36.157193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = HMSDataset_preprocesser(df=test_df, eegs=all_eegs, CFG=CFG, mode=\"test\", transforms=None)\n\npre_loader = DataLoader(\n    dataset,\n    batch_size=4,\n    shuffle=False,\n    num_workers=4, pin_memory=True, drop_last=False\n)\n\nfor data in tqdm(pre_loader):\n    for idx, cwtidx in enumerate(data[\"idx\"].numpy()):\n        np.save(f\"{cwt_path}cwt_all_{cwtidx:06}.npy\", data[\"data\"][idx].numpy())","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:36.159287Z","iopub.execute_input":"2024-05-03T21:55:36.159553Z","iopub.status.idle":"2024-05-03T21:55:37.328875Z","shell.execute_reply.started":"2024-05-03T21:55:36.159530Z","shell.execute_reply":"2024-05-03T21:55:37.327723Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df['spec_path'] = test_df.index.map(lambda x: f\"{cwt_path}cwt_all_{x:06}.npy\")","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:37.330175Z","iopub.execute_input":"2024-05-03T21:55:37.330494Z","iopub.status.idle":"2024-05-03T21:55:37.336982Z","shell.execute_reply.started":"2024-05-03T21:55:37.330465Z","shell.execute_reply":"2024-05-03T21:55:37.336088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coef_sum = 0\ncoef_count = 0\npredictions = []\nfiles = []\n    \nfor model_block in model_weights_50sec_to40Hz:\n    test_dataset = HMSDataset_2D(\n        df=test_df,\n        mode=\"test\",\n        CFG=CFG, \n        eegs=all_eegs,\n        bandpass_filter=model_block['bandpass_filter'],\n        transforms=data_transforms['valid']\n    )\n\n    if len(predictions) == 0:\n        output = test_dataset[0]\n        X = output[\"data\"]\n        print(f\"X shape: {X.shape}\")\n                \n    test_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    )\n\n    model = HMSModel(\n        model_name = model_block['model_name'], \n        pretrained = False,\n        in_channels = CFG.in_chans, \n        num_classes = CFG.n_classes\n    )\n\n    single_model_preds = []\n    for file_line in model_block['file_data']:\n        coef = file_line['coef']\n        for weight_model_file in glob(file_line['file_mask']):\n            files.append(weight_model_file)\n            checkpoint = torch.load(weight_model_file, map_location=CFG.device)\n            model.load_state_dict(checkpoint)\n            model.to(CFG.device)\n            prediction_dict = inference_function(test_loader, model, CFG.device)\n            predict = prediction_dict[\"predictions\"]\n            #predict *= coef\n            #coef_sum += coef\n            #coef_count += 1\n            single_model_preds.append(predict)\n            torch.cuda.empty_cache()\n            gc.collect()\n            \n    single_model_preds = np.mean(single_model_preds, axis=0)\n    predictions.append(single_model_preds.copy())\n\nk_cwt50sec_to40Hz_predictions = predictions.copy()\n#predictions = np.array(predictions)\n#coef_sum /= coef_count\n#predictions /= coef_sum\n#predictions_2d_50sec = np.mean(predictions, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:55:37.338468Z","iopub.execute_input":"2024-05-03T21:55:37.339107Z","iopub.status.idle":"2024-05-03T21:57:22.704124Z","shell.execute_reply.started":"2024-05-03T21:55:37.339075Z","shell.execute_reply":"2024-05-03T21:57:22.703025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 1d and 2D","metadata":{}},{"cell_type":"code","source":"class CFG:\n    debug = False\n    \n    CWT_STRIDE = 16\n            \n    seq_length = 50  # Second's\n    sampling_rate = 200  # Hz\n    nsamples = seq_length * sampling_rate\n    out_samples = 50 * sampling_rate\n    \n    \n    seed = 42\n    gpu_idx = 0\n\n    device = torch.device(f\"cuda:{gpu_idx}\")\n    \n    classes = [\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]\n    n_classes = len(classes)\n\n    n_folds = 5\n    \n    ## Model\n    fixed_kernel_size = 5\n    #linear_layer_features = 448  # Full Signal = 10_000\n    #linear_layer_features = 352  # Half Signal = 5_000\n    linear_layer_features = 304   # 1/5  Signal = 2_000\n    kernels = [3, 5, 7, 9, 11]\n    \n    ## Dataset Preprocessing\n    bandpass_filter = {\"low\": 0.5, \"high\": 40, \"order\": 2}\n    # bandpass_filter = None\n    rand_filter = {\"probab\": 0, \"low\": 10, \"high\": 20, \"band\": 1.0, \"order\": 2}\n    freq_channels = [(8.0, 12.0), (0.5, 4.5)]\n    \n    filter_order = 2\n    random_close_zone = 0.0  # 0.2\n    \n    map_features = [\n        (\"Fp1\", \"F7\"),\n        (\"F7\", \"T3\"),\n        (\"T3\", \"T5\"),\n        (\"T5\", \"O1\"),\n        (\"Fp1\", \"F3\"),\n        (\"F3\", \"C3\"),\n        (\"C3\", \"P3\"),\n        (\"P3\", \"O1\"),\n        (\"Fp2\", \"F8\"),\n        (\"F8\", \"T4\"),\n        (\"T4\", \"T6\"),\n        (\"T6\", \"O2\"),\n        (\"Fp2\", \"F4\"),\n        (\"F4\", \"C4\"),\n        (\"C4\", \"P4\"),\n        (\"P4\", \"O2\"),\n        (\"Fz\", \"Cz\"),\n        (\"Cz\", \"Pz\")\n    ]\n    \n    eeg_features = [\"Fp1\",\"F3\",\"C3\",\"P3\",\"F7\",\"T3\",\"T5\",\"O1\",\"Fz\",\"Cz\",\"Pz\",\"Fp2\",\"F4\",\"C4\",\"P4\",\"F8\",\"T4\",\"T6\",\"O2\",\"EKG\"]\n    #eeg_features = [\"Fp1\", \"T3\", \"C3\", \"O1\", \"Fp2\", \"C4\", \"T4\", \"O2\"]\n    feature_to_index = {x: y for x, y in zip(eeg_features, range(len(eeg_features)))}\n    simple_features = []  # 'Fz', 'Cz', 'Pz', 'EKG'\n    \n    n_map_features = len(map_features)\n    in_channels = n_map_features + n_map_features * len(freq_channels) + len(simple_features)\n\n    input_size = 512\n    in_chans = 1\n\n    batch_size = 4\n    num_workers = 4","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:57:22.705635Z","iopub.execute_input":"2024-05-03T21:57:22.705938Z","iopub.status.idle":"2024-05-03T21:57:22.718708Z","shell.execute_reply.started":"2024-05-03T21:57:22.705911Z","shell.execute_reply":"2024-05-03T21:57:22.717720Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HMSModel(nn.Module):\n    def __init__(self, model_name: str, pretrained: bool, in_channels: int, num_classes: int):\n        super().__init__()\n\n        self.model = timm.create_model(\n            model_name=model_name, \n            pretrained=pretrained, \n            in_chans=in_channels//2, # Assuming you want to split channels\n            drop_rate = 0.1,\n            drop_path_rate = 0.2,\n            num_classes=0 # Assuming this disables the final classifier\n        )\n        self.model_1d = timm.create_model(\n            model_name=\"convnext_atto_ols.a2_in1k\", \n            pretrained=pretrained, \n            in_chans=in_channels//2, # Adjust based on your needs\n            drop_rate = 0.1,\n            drop_path_rate = 0.2,\n            num_classes=0 # Assuming this disables the final classifier\n        )\n        self.classifier = nn.Linear(1088, num_classes)\n        \n    def forward(self, x):\n        # Assuming x is split correctly into two parts: x0 and x1\n        # This might need adjustment based on how your data is structured\n        x0 = self.model(x[:, 0, :, :].unsqueeze(1))  # Add a channel dimension back\n        x1 = self.model_1d(x[:, 1, :, :].unsqueeze(1))  # Add a channel dimension back\n\n        # Combine the features from both models\n        combined_features = torch.cat([x0, x1], dim=1)\n        \n        # Pass the combined features through the classifier\n        output = self.classifier(combined_features)\n        \n        return output","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:57:22.719885Z","iopub.execute_input":"2024-05-03T21:57:22.720150Z","iopub.status.idle":"2024-05-03T21:57:22.733603Z","shell.execute_reply.started":"2024-05-03T21:57:22.720127Z","shell.execute_reply":"2024-05-03T21:57:22.732741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchaudio.transforms as T\n\nclass HMSDataset_1D2D(torch.utils.data.Dataset):\n    def __init__(self, df, eegs, CFG,\n                 downsample: int = None,\n                 bandpass_filter = None,\n                 rand_filter = None,\n                 mode='train',\n                 weighted=False,\n                 transforms = None,\n                ):\n        self.df = df\n        self.eegs = eegs\n        self.CFG = CFG\n        self.downsample = downsample\n        self.bandpass_filter = bandpass_filter\n        self.rand_filter = rand_filter\n        self.mode = mode\n        self.weighted = weighted\n        self.transforms = transforms\n        self.spec_imgsize_hw = [128, 256]\n\n        self.L           = CFG.nsamples    # 4000\n        self.L_input     = CFG.out_samples # 2000\n        \n        self.L_stride          = CFG.nsamples//CFG.CWT_STRIDE\n        self.L_input_stride    = CFG.out_samples//CFG.CWT_STRIDE\n        \n            \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx: int):\n        row = self.df.iloc[idx]\n        eeg_id = row.eeg_id\n        \n        img = np.zeros((10*self.CFG.in_channels, self.CFG.out_samples, 1)).astype(np.float32)\n        \n        X, y, offset = self.__data_generation(idx)\n\n        #####################\n        # 2D\n        img_spec = torch.from_numpy(np.load(row.spec_path))\n\n        if self.L != self.L_input:\n            offset_stride = offset//self.CFG.CWT_STRIDE\n            img_spec = img_spec[:, offset_stride: offset_stride+self.L_input_stride]\n\n        if self.transforms is not None:\n            img_spec = self.transforms(image=img_spec.numpy() )['image']\n        #####################\n        # 1D\n        for i in range(X.shape[1]):\n            img[i*10:(i+1)*10, :] = X[:,i][:,np.newaxis]\n                       \n        if self.transforms is not None:\n            img = self.transforms(image=img)['image']\n                \n        img_spec = torch.cat([img_spec, img], axis=0)\n        #####################\n        #y = row[self.CFG.classes].values.astype(np.float32)\n        #if self.weighted:\n        #    w = row[\"num_votes\"].astype(np.float32) / 10.0\n        #else:\n        #    w = 1.0\n        \n#         img = torch.cat([img, img_spec], dim=0)\n        img = img_spec\n        return {\"data\": img}\n    \n    def __data_generation(self, index):\n        X = np.zeros((self.CFG.out_samples, self.CFG.in_channels), dtype=\"float32\")  # Size=(10000, 14)\n\n        row = self.df.iloc[index]\n        data = self.eegs[row.eeg_id]  # Size=(10000, 8)\n        \n        offset=0\n        if self.L != self.L_input:\n            offset = (self.L - self.L_input) // 2\n            data = data[offset:offset+self.L_input,:]\n\n        # diff of eegs\n        for i, (feat_a, feat_b) in enumerate(self.CFG.map_features):\n            if self.mode == \"train\" and self.CFG.random_close_zone > 0 and random.uniform(0.0, 1.0) <= self.CFG.random_close_zone:\n                continue\n                \n            diff_feat = data[:, self.CFG.feature_to_index[feat_a]] - data[:, self.CFG.feature_to_index[feat_b]]\n\n            if not self.bandpass_filter is None:\n                diff_feat = butter_bandpass_filter(\n                    diff_feat,\n                    self.bandpass_filter[\"low\"],\n                    self.bandpass_filter[\"high\"],\n                    self.CFG.sampling_rate,\n                    order=self.bandpass_filter[\"order\"],\n                )\n                    \n            if (\n                self.mode == \"train\"\n                and not self.rand_filter is None\n                and random.uniform(0.0, 1.0) <= self.rand_filter[\"probab\"]\n            ):\n                lowcut = random.randint(\n                    self.rand_filter[\"low\"], self.rand_filter[\"high\"]\n                )\n                highcut = lowcut + self.rand_filter[\"band\"]\n                diff_feat = butter_bandpass_filter(\n                    diff_feat,\n                    lowcut,\n                    highcut,\n                    self.CFG.sampling_rate,\n                    order=self.rand_filter[\"order\"],\n                )\n\n            X[:, i] = diff_feat\n\n        # other frequency filtering\n        n = self.CFG.n_map_features\n        if len(self.CFG.freq_channels) > 0:\n            for i in range(self.CFG.n_map_features):\n                diff_feat = X[:, i]\n                for j, (lowcut, highcut) in enumerate(self.CFG.freq_channels):\n                    band_feat = butter_bandpass_filter(\n                        diff_feat, lowcut, highcut, self.CFG.sampling_rate, order=self.CFG.filter_order,  # 6\n                    )\n                    X[:, n] = band_feat\n                    n += 1\n                    \n        # single eeg \n        for spml_feat in self.CFG.simple_features:\n            feat_val = data[:, self.CFG.feature_to_index[spml_feat]]\n            \n            if not self.bandpass_filter is None:\n                feat_val = butter_bandpass_filter(\n                    feat_val,\n                    self.bandpass_filter[\"low\"],\n                    self.bandpass_filter[\"high\"],\n                    self.CFG.sampling_rate,\n                    order=self.bandpass_filter[\"order\"],\n                )\n\n            if (\n                self.mode == \"train\"\n                and not self.rand_filter is None\n                and random.uniform(0.0, 1.0) <= self.rand_filter[\"probab\"]\n            ):\n                lowcut = random.randint(\n                    self.rand_filter[\"low\"], self.rand_filter[\"high\"]\n                )\n                highcut = lowcut + self.rand_filter[\"band\"]\n                feat_val = butter_bandpass_filter(\n                    feat_val,\n                    lowcut,\n                    highcut,\n                    CFG.sampling_rate,\n                    order=self.rand_filter[\"order\"],\n                )\n\n            X[:, n] = feat_val\n            n += 1\n            \n        X = np.clip(X, -1024, 1024)\n        X = np.nan_to_num(X, nan=0) / 32.0\n#         X = butter_lowpass_filter(X, order=self.CFG.filter_order)\n\n        y = np.zeros(self.CFG.n_classes, dtype=\"float32\")  # Size=(6,)\n        if self.mode != \"test\":\n            y = row[self.CFG.classes].values.astype(np.float32)\n\n        return X, y, offset","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:57:22.734743Z","iopub.execute_input":"2024-05-03T21:57:22.735055Z","iopub.status.idle":"2024-05-03T21:57:22.765001Z","shell.execute_reply.started":"2024-05-03T21:57:22.735032Z","shell.execute_reply":"2024-05-03T21:57:22.764050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coef_sum = 0\ncoef_count = 0\npredictions = []\nfiles = []\n    \nfor model_block in model_weights_50sec_to40Hz_and_1d:\n    test_dataset = HMSDataset_1D2D(\n        df=test_df,\n        mode=\"test\",\n        CFG=CFG, \n        eegs=all_eegs,\n        bandpass_filter=model_block['bandpass_filter'],\n        transforms=data_transforms['valid']\n    )\n\n    if len(predictions) == 0:\n        output = test_dataset[0]\n        X = output[\"data\"]\n        print(f\"X shape: {X.shape}\")\n                \n    test_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    )\n\n    model = HMSModel(\n        model_name = model_block['model_name'], \n        pretrained = False,\n        in_channels = 2, \n        num_classes = CFG.n_classes\n    )\n\n    single_model_preds = []\n    for file_line in model_block['file_data']:\n        coef = file_line['coef']\n        for weight_model_file in glob(file_line['file_mask']):\n            files.append(weight_model_file)\n            checkpoint = torch.load(weight_model_file, map_location=CFG.device)\n            model.load_state_dict(checkpoint)\n            model.to(CFG.device)\n            prediction_dict = inference_function(test_loader, model, CFG.device)\n            predict = prediction_dict[\"predictions\"]\n            #predict *= coef\n            #coef_sum += coef\n            #coef_count += 1\n            single_model_preds.append(predict)\n            torch.cuda.empty_cache()\n            gc.collect()\n            \n    single_model_preds = np.mean(single_model_preds, axis=0)\n    predictions.append(single_model_preds.copy())\n\nk_cwt50sec_to40Hz_and_1d_predictions = predictions.copy()\n#predictions = np.array(predictions)\n#coef_sum /= coef_count\n#predictions /= coef_sum\n#predictions_2d_50sec = np.mean(predictions, axis=0)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:57:22.766102Z","iopub.execute_input":"2024-05-03T21:57:22.766397Z","iopub.status.idle":"2024-05-03T21:59:12.363075Z","shell.execute_reply.started":"2024-05-03T21:57:22.766374Z","shell.execute_reply":"2024-05-03T21:59:12.362242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil\nshutil.rmtree(cwt_path)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.364305Z","iopub.execute_input":"2024-05-03T21:59:12.364599Z","iopub.status.idle":"2024-05-03T21:59:12.369848Z","shell.execute_reply.started":"2024-05-03T21:59:12.364574Z","shell.execute_reply":"2024-05-03T21:59:12.369026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#sub = pd.DataFrame({\"eeg_id\": test_df.eeg_id.values})\n#sub[CFG.classes] = predictions\n\n#sub.to_csv(f\"submission.csv\", index=False)\n#sub.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.370973Z","iopub.execute_input":"2024-05-03T21:59:12.371381Z","iopub.status.idle":"2024-05-03T21:59:12.379026Z","shell.execute_reply.started":"2024-05-03T21:59:12.371349Z","shell.execute_reply":"2024-05-03T21:59:12.378240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Saito Model","metadata":{}},{"cell_type":"code","source":"del img2\ntorch.cuda.empty_cache()\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.380133Z","iopub.execute_input":"2024-05-03T21:59:12.380477Z","iopub.status.idle":"2024-05-03T21:59:12.619308Z","shell.execute_reply.started":"2024-05-03T21:59:12.380447Z","shell.execute_reply":"2024-05-03T21:59:12.618441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport sys\nsys.path.append(\"/kaggle/input/kaggle-kl-div\")\n\nimport cv2\nimport datetime as dt\nfrom functools import partial\nimport gc\nfrom glob import glob\nimport math\nfrom matplotlib import pyplot as plt\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nimport random\nimport time\nfrom tqdm.auto import tqdm\nfrom typing import Dict, List, Union\nimport warnings\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD, AdamW\nfrom torch.utils.data import DataLoader, Dataset\nimport timm\n\nimport librosa\nfrom scipy.signal import butter, lfilter, freqz\n\nfrom kaggle_kl_div import score\n\n# === jax ===\nimport jax\nimport jax.numpy as jnp\nfrom jax import jit","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.620539Z","iopub.execute_input":"2024-05-03T21:59:12.620825Z","iopub.status.idle":"2024-05-03T21:59:12.629330Z","shell.execute_reply.started":"2024-05-03T21:59:12.620801Z","shell.execute_reply":"2024-05-03T21:59:12.628448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0,1\"\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint('Using', torch.cuda.device_count(), 'GPU(s)')","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.630581Z","iopub.execute_input":"2024-05-03T21:59:12.631347Z","iopub.status.idle":"2024-05-03T21:59:12.640226Z","shell.execute_reply.started":"2024-05-03T21:59:12.631314Z","shell.execute_reply":"2024-05-03T21:59:12.639399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Paths:\n    INPUT_DIR = Path(\"/kaggle/input/hms-harmful-brain-activity-classification\")\n    WORK_DIR = Path(\"/kaggle/working\")\n    TEST_CSV = INPUT_DIR / \"test.csv\"\n    TEST_RAW_EEGS = INPUT_DIR / \"test_eegs\"\n    TEST_SPECTROGRAMS = INPUT_DIR / \"test_spectrograms\"\n    FILE_RAW_EEG = TEST_RAW_EEGS / \"eegs.npy\"","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.641362Z","iopub.execute_input":"2024-05-03T21:59:12.641831Z","iopub.status.idle":"2024-05-03T21:59:12.650539Z","shell.execute_reply.started":"2024-05-03T21:59:12.641798Z","shell.execute_reply":"2024-05-03T21:59:12.649629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Cfg:\n    target_cols = [\n        \"seizure_vote\",\n        \"lpd_vote\",\n        \"gpd_vote\",\n        \"lrda_vote\",\n        \"grda_vote\",\n        \"other_vote\",\n    ]\n\n    eeg_features = [\n        \"Fp1\", \"F7\", \"T3\", \"T5\", \"O1\", \"F3\", \"C3\", \"P3\",\n        \"Fp2\", \"F8\", \"T4\", \"T6\", \"O2\", \"F4\", \"C4\", \"P4\",\n        \"Fz\", \"Cz\", \"Pz\",\n    ]\n    feature_to_index = {x: y for x, y in zip(eeg_features, range(len(eeg_features)))}\n    map_features = [\n        (\"Fp1\", [\"F7\",]),\n        (\"F7\", [\"T3\",]),\n        (\"T3\", [\"T5\",]),\n        (\"T5\", [\"O1\",]),\n        (\"Fp1\", [\"F3\",]),\n        (\"F3\", [\"C3\",]),\n        (\"C3\", [\"P3\",]),\n        (\"P3\", [\"O1\",]),\n        (\"P4\", [\"O2\",]),\n        (\"C4\", [\"P4\",]),\n        (\"F4\", [\"C4\",]),\n        (\"Fp2\", [\"F4\",]),\n        (\"T6\", [\"O2\",]),\n        (\"T4\", [\"T6\",]),\n        (\"F8\", [\"T4\",]),\n        (\"Fp2\", [\"F8\",]),\n    ]\n\n    montages_ch = {\n#         \"kaggle\": 4,\n        \"anterior_posterior\": 4,\n    }\n    mode_montage = \"anterior_posterior\"\n    if mode_montage == \"anterior_posterior\":\n        names_montage = [\"Fp1-F7\", \"F7-T3\", \"T3-T5\", \"T5-O1\", \"Fp1-F3\", \"F3-C3\", \"C3-P3\", \"P3-O1\",\n                 \"Fp2-F4\", \"F4-C4\", \"C4-P4\", \"P4-O2\", \"Fp2-F8\", \"F8-T4\", \"T4-T6\", \"T6-O2\",]\n        names_spec = ['LL','LP','RP','RL']\n        feats_montage = [\n            ['Fp1','F7','T3','T5','O1'],\n            ['Fp1','F3','C3','P3','O1'],\n            ['Fp2','F4','C4','P4','O2'],\n            ['Fp2','F8','T4','T6','O2'],\n        ]\n        names_spec_16ch = names_montage\n        feats_montage_16ch = [\n            ['Fp1','F7'],\n            ['F7','T3'],\n            ['T3','T5'],\n            ['T5','O1'],\n            ['Fp1','F3'],\n            ['F3','C3'],\n            ['C3','P3'],\n            ['P3','O1'],\n            ['Fp2','F4'],\n            ['F4','C4'],\n            ['C4','P4'],\n            ['P4','O2'],\n            ['Fp2','F8'],\n            ['F8','T4'],\n            ['T4','T6'],\n            ['T6','O2'],\n        ]\n        \n    batch_size = 32\n    num_workers = 0\n\n    seq_length = 50  # sec, FIXED\n    sampling_rate = 200  # Hz, FIXED\n    n_samples = seq_length * sampling_rate\n    n_samples_spec = n_samples // 10\n    \n    n_map_features = len(map_features)\n    \n    # Model configs\n    VERSION = [\n        {\n           \"name\": \"1dcnn_v062_stage2\",\n           \"model_2dcnn\": \"gcvit_xtiny\",\n           \"model_ver\": \"v2\",\n           \"input_type\": \"eeg_spec_rnn\",\n           \"spec_type\": \"superlet\",\n           \"imsize\": 224,\n           \"n_freqs\": [8, 3, 3],\n           \"n_freqs_spec\": 7,\n           \"fs_spec\": 40,\n           \"in_ch_eeg\": 48,\n           \"in_ch_spec\": 14,\n           \"montage_ch\": 16,\n           \"out_samples\": 2240,\n           \"out_samples_spec\": 512,\n           \"coef\": 0.0,\n           \"norm_spec\": False,\n           \"interpolation\": cv2.INTER_LINEAR,\n           \"symmetric_lr_ch\": False,\n           \"bandpass_filter\": {\"low\": 0.5, \"high\": 36, \"order\": 4, \"lowpass_order\": None},\n           \"freq_channels\": [(0.5, 4.0), (2.5, 16.0)],\n       },\n       {\n           \"name\": \"1dcnn_v061_stage2\",\n           \"model_2dcnn\": \"convnextv2_atto\",\n           \"model_ver\": \"v2\",\n           \"input_type\": \"eegs_spec\",\n           \"spec_type\": \"superlet\",\n           \"imsize\": 256,\n           \"n_freqs\": [8, 4, 4],\n           \"n_freqs_spec\": 7,\n           \"fs_spec\": 40,\n           \"in_ch_eeg\": 48,\n           \"in_ch_spec\": 16,\n           \"montage_ch\": 16,\n           \"out_samples\": 2048,\n           \"out_samples_spec\": 768,\n           \"coef\": 0.0,\n           \"norm_spec\": False,\n           \"interpolation\": cv2.INTER_LINEAR,\n           \"symmetric_lr_ch\": False,\n           \"bandpass_filter\": {\"low\": 0.5, \"high\": 33, \"order\": 4, \"lowpass_order\": None},\n           \"freq_channels\": [(0.5, 4.0), (2.5, 16.0)],\n       },\n       {\n           \"name\": \"1dcnn_v060_stage2\",\n           \"model_2dcnn\": \"maxvit_rmlp_pico_rw_256\",\n           \"model_ver\": \"v2\",\n           \"input_type\": \"eegs_spec\",\n           \"spec_type\": \"superlet\",\n           \"imsize\": 256,\n           \"n_freqs\": [9, 4, 3],\n           \"n_freqs_spec\": 7,\n           \"fs_spec\": 40,\n           \"in_ch_eeg\": 48,\n           \"in_ch_spec\": 16,\n           \"montage_ch\": 16,\n           \"out_samples\": 2048,\n           \"out_samples_spec\": 768,\n           \"coef\": 0.0,\n           \"norm_spec\": False,\n           \"interpolation\": cv2.INTER_LINEAR,\n           \"symmetric_lr_ch\": False,\n           \"bandpass_filter\": {\"low\": 0.5, \"high\": 40, \"order\": 4, \"lowpass_order\": None},\n           \"freq_channels\": [(0.5, 4.0), (2.5, 16.0)],\n       },\n       {\n           \"name\": \"1dcnn_v059_stage2\",\n           \"model_2dcnn\": \"inception_next_tiny\",\n           \"model_ver\": \"v2\",\n           \"input_type\": \"eegs_spec\",\n           \"spec_type\": \"superlet\",\n           \"imsize\": 256,\n           \"n_freqs\": [16],\n           \"n_freqs_spec\": 7,\n           \"fs_spec\": 40,\n           \"in_ch_eeg\": 48,\n           \"in_ch_spec\": 16,\n           \"montage_ch\": 16,\n           \"out_samples\": 2048,\n           \"out_samples_spec\": 768,\n           \"coef\": 0.0,\n           \"norm_spec\": False,\n           \"interpolation\": cv2.INTER_LINEAR,\n           \"symmetric_lr_ch\": False,\n           \"bandpass_filter\": {\"low\": 0.5, \"high\": 20, \"order\": 4, \"lowpass_order\": None},\n           \"freq_channels\": [],\n       },\n       {\n           \"name\": \"1dcnn_v057_stage2\",\n           \"model_2dcnn\": \"swinv2_tiny_window16_256\",\n           \"model_ver\": \"v2\",\n           \"input_type\": \"eegs_spec\",\n           \"spec_type\": \"superlet\",\n           \"imsize\": 256,\n           \"n_freqs\": [10, 3, 3],\n           \"n_freqs_spec\": 7,\n           \"fs_spec\": 40,\n           \"in_ch_eeg\": 48,\n           \"in_ch_spec\": 16,\n           \"montage_ch\": 16,\n           \"out_samples\": 2048,\n           \"out_samples_spec\": 768,\n           \"coef\": 0.0,\n           \"norm_spec\": False,\n           \"interpolation\": cv2.INTER_LINEAR,\n           \"symmetric_lr_ch\": False,\n           \"bandpass_filter\": {\"low\": 0.5, \"high\": 36, \"order\": 4, \"lowpass_order\": None},\n           \"freq_channels\": [(0.5, 4.0), (2.5, 16.0)],\n       },\n       {\n           \"name\": \"1dcnn_v056_stage2\",\n           \"model_2dcnn\": \"poolformerv2_s12\",\n           \"model_ver\": \"v2\",\n           \"input_type\": \"eegs_spec\",\n           \"spec_type\": \"superlet\",\n           \"imsize\": 224,\n           \"n_freqs\": [8, 3, 3],\n           \"n_freqs_spec\": 7,\n           \"fs_spec\": 40,\n           \"in_ch_eeg\": 48,\n           \"in_ch_spec\": 56,\n           \"montage_ch\": 4,\n           \"out_samples\": 2240,\n           \"out_samples_spec\": 768,\n           \"coef\": 0.0,\n           \"norm_spec\": False,\n           \"interpolation\": cv2.INTER_LINEAR,\n           \"symmetric_lr_ch\": False,\n           \"bandpass_filter\": {\"low\": 0.5, \"high\": 40, \"order\": 4, \"lowpass_order\": None},\n           \"freq_channels\": [(0.5, 4.0), (2.5, 16.0)],\n       },\n       {\n           \"name\": \"1dcnn_v055_stage2\",\n           \"model_2dcnn\": \"gcvit_xxtiny\",\n           \"model_ver\": \"v2\",\n           \"input_type\": \"eegs_spec\",\n           \"spec_type\": \"superlet\",\n           \"imsize\": 224,\n           \"n_freqs\": [7, 4, 3],\n           \"n_freqs_spec\": 7,\n           \"fs_spec\": 40,\n           \"in_ch_eeg\": 48,\n           \"in_ch_spec\": 56,\n           \"montage_ch\": 4,\n           \"out_samples\": 2240,\n           \"out_samples_spec\": 768,\n           \"coef\": 0.0,\n           \"norm_spec\": False,\n           \"interpolation\": cv2.INTER_LINEAR,\n           \"symmetric_lr_ch\": False,\n           \"bandpass_filter\": {\"low\": 0.5, \"high\": 40, \"order\": 4, \"lowpass_order\": None},\n           \"freq_channels\": [(0.5, 4.0), (2.5, 16.0)],\n       },\n       {\n           \"name\": \"1dcnn_v054_stage2\",\n           \"model_2dcnn\": \"caformer_s18\",\n           \"model_ver\": \"v2\",\n           \"input_type\": \"eegs_spec\",\n           \"spec_type\": \"superlet\",\n           \"imsize\": 224,\n           \"n_freqs\": [8, 3, 3],\n           \"n_freqs_spec\": 7,\n           \"fs_spec\": 40,\n           \"in_ch_eeg\": 48,\n           \"in_ch_spec\": 56,\n           \"montage_ch\": 4,\n           \"out_samples\": n_samples // 5,\n           \"out_samples_spec\": 768,\n           \"coef\": 0.0,\n           \"norm_spec\": False,\n           \"interpolation\": cv2.INTER_LINEAR,\n           \"symmetric_lr_ch\": False,\n           \"bandpass_filter\": {\"low\": 0.5, \"high\": 35, \"order\": 4, \"lowpass_order\": None},\n           \"freq_channels\": [(0.5, 4.0), (2.5, 16.0)],\n       },\n       {\n           \"name\": \"1dcnn_v053_stage2\",\n           \"model_2dcnn\": \"swinv2_tiny_window8_256\",\n           \"model_ver\": \"v2\",\n           \"input_type\": \"eeg_rnn_spec_1ddw\",\n           \"spec_type\": \"superlet\",\n           \"imsize\": 256,\n           \"n_freqs\": [9, 4, 3],\n           \"n_freqs_spec\": 3,\n           \"fs_spec\": 40,\n           \"in_ch_eeg\": 48,\n           \"in_ch_spec\": 16,\n           \"montage_ch\": 4,\n           \"out_samples\": n_samples // 5,\n           \"out_samples_spec\": 512,\n           \"coef\": 0.0,\n           \"norm_spec\": False,\n           \"interpolation\": cv2.INTER_LINEAR,\n           \"symmetric_lr_ch\": False,\n           \"bandpass_filter\": {\"low\": 0.5, \"high\": 40, \"order\": 4, \"lowpass_order\": None},\n           \"freq_channels\": [(0.5, 4.0), (2.5, 16.0)],\n       },\n       {\n           \"name\": \"1dcnn_v051_stage2\",\n           \"model_2dcnn\": \"swinv2_tiny_window16_256\",\n           \"model_ver\": \"v2\",\n           \"input_type\": \"eegs_spec\",\n           \"spec_type\": \"superlet\",\n           \"imsize\": 256,\n           \"n_freqs\": [10, 3, 3],\n           \"n_freqs_spec\": 7,\n           \"fs_spec\": 40,\n           \"in_ch_eeg\": 48,\n           \"in_ch_spec\": 64,\n           \"montage_ch\": 4,\n           \"out_samples\": n_samples // 5,\n           \"out_samples_spec\": 768,\n           \"coef\": 0.0,\n           \"norm_spec\": False,\n           \"interpolation\": cv2.INTER_LINEAR,\n           \"symmetric_lr_ch\": False,\n           \"bandpass_filter\": {\"low\": 0.5, \"high\": 35, \"order\": 4, \"lowpass_order\": None},\n           \"freq_channels\": [(0.5, 4.0), (2.5, 16.0)],\n       },\n#        {\n#            \"name\": \"1dcnn_v050_stage2\",\n#            \"model_2dcnn\": \"maxvit_rmlp_tiny_rw_256\",\n#            \"model_ver\": \"v2\",\n#            \"input_type\": \"eegs_spec\",\n#            \"spec_type\": \"superlet\",\n#            \"imsize\": 256,\n#            \"n_freqs\": [6, 5, 5],\n#            \"n_freqs_spec\": 7,\n#            \"fs_spec\": 40,\n#            \"in_ch_eeg\": 48,\n#            \"in_ch_spec\": 64,\n#            \"montage_ch\": 4,\n#            \"out_samples\": n_samples // 5,\n#            \"out_samples_spec\": 768,\n#            \"coef\": 0.0,\n#            \"norm_spec\": False,\n#            \"interpolation\": cv2.INTER_LINEAR,\n#            \"symmetric_lr_ch\": False,\n#            \"bandpass_filter\": {\"low\": 0.5, \"high\": 40, \"order\": 4, \"lowpass_order\": None},\n#            \"freq_channels\": [(0.5, 4.0), (2.5, 16.0)],\n#        },\n       {\n           \"name\": \"1dcnn_v048_stage2\",\n           \"model_2dcnn\": \"nextvit_small\",\n           \"model_ver\": \"v2\",\n           \"input_type\": \"eegs_spec\",\n           \"spec_type\": \"superlet\",\n           \"imsize\": 256,\n           \"n_freqs\": [8, 4, 4],\n           \"n_freqs_spec\": 7,\n           \"fs_spec\": 40,\n           \"in_ch_eeg\": 48,\n           \"in_ch_spec\": 64,\n           \"montage_ch\": 4,\n           \"out_samples\": n_samples // 5,\n           \"out_samples_spec\": 512,\n           \"coef\": 0.0,\n           \"norm_spec\": False,\n           \"interpolation\": cv2.INTER_LINEAR,\n           \"symmetric_lr_ch\": False,\n           \"bandpass_filter\": {\"low\": 0.5, \"high\": 30, \"order\": 4, \"lowpass_order\": None},\n           \"freq_channels\": [(0.5, 4.0), (2.5, 16.0)],\n       },\n       {\n           \"name\": \"1dcnn_v047_stage2\",\n           \"model_2dcnn\": \"inception_next_tiny\",\n           \"model_ver\": \"v2\",\n           \"input_type\": \"eegs_spec\",\n           \"spec_type\": \"superlet\",\n           \"imsize\": 256,\n           \"n_freqs\": [8, 4, 4],\n           \"n_freqs_spec\": 7,\n           \"fs_spec\": 40,\n           \"in_ch_eeg\": 48,\n           \"in_ch_spec\": 64,\n           \"montage_ch\": 4,\n           \"out_samples\": n_samples // 5,\n           \"out_samples_spec\": 512,\n           \"coef\": 0.0,\n           \"norm_spec\": False,\n           \"interpolation\": cv2.INTER_LINEAR,\n           \"symmetric_lr_ch\": False,\n           \"bandpass_filter\": {\"low\": 0.5, \"high\": 30, \"order\": 4, \"lowpass_order\": None},\n           \"freq_channels\": [(0.5, 4.0), (2.5, 16.0)],\n       },\n       {\n           \"name\": \"1dcnn_v046_2_stage2\",\n           \"model_2dcnn\": \"convnextv2_atto\",\n           \"model_ver\": \"v2\",\n           \"input_type\": \"eegs_rnn_spec_1ddw\",\n           \"spec_type\": \"superlet\",\n           \"imsize\": 256,\n           \"n_freqs\": [16, 8, 6, 2],\n           \"n_freqs_spec\": 7,\n           \"fs_spec\": 40,\n           \"in_ch_eeg\": 64,\n           \"in_ch_spec\": 16,\n           \"montage_ch\": 4,\n           \"out_samples\": n_samples // 5,\n           \"out_samples_spec\": 256,\n           \"coef\": 0.0,\n           \"norm_spec\": False,\n           \"interpolation\": cv2.INTER_LINEAR,\n           \"symmetric_lr_ch\": True,\n           \"bandpass_filter\": {\"low\": 0.5, \"high\": 20, \"order\": 4, \"lowpass_order\": None},\n           \"freq_channels\": [(0.5, 4.0), (2.5, 16.0), (16.0, 20.0)],\n       },\n       {\n           \"name\": \"1dcnn_v042_stage2\",\n           \"model_2dcnn\": \"inception_next_tiny\",\n           \"model_ver\": \"v1\",\n           \"input_type\": \"eegs_rnn_spec_1ddw\",\n           \"spec_type\": \"stft\",\n           \"imsize\": 256,\n           \"n_freqs\": 16,\n           \"n_freqs_spec\": 7,\n           \"fs_spec\": 40,\n           \"in_ch_eeg\": 16,\n           \"in_ch_spec\": 16,\n           \"montage_ch\": 4,\n           \"out_samples\": n_samples // 5,\n           \"out_samples_spec\": 256,\n           \"out_ch_cnn\": 2304,\n           \"out_ch_gru\": 256,\n           \"coef\": 0.13426578,\n           \"norm_spec\": False,\n           \"interpolation\": cv2.INTER_NEAREST,\n           \"symmetric_lr_ch\": True,\n           \"bandpass_filter\": {\"low\": 0.5, \"high\": 20, \"order\": 4, \"lowpass_order\": None},\n           \"freq_channels\": [],\n       },\n    ]","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.651950Z","iopub.execute_input":"2024-05-03T21:59:12.652289Z","iopub.status.idle":"2024-05-03T21:59:12.703933Z","shell.execute_reply.started":"2024-05-03T21:59:12.652261Z","shell.execute_reply":"2024-05-03T21:59:12.703189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.listdir(\"/kaggle/input/hms-weights-saito\")","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.704881Z","iopub.execute_input":"2024-05-03T21:59:12.705173Z","iopub.status.idle":"2024-05-03T21:59:12.719046Z","shell.execute_reply.started":"2024-05-03T21:59:12.705148Z","shell.execute_reply":"2024-05-03T21:59:12.718132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg = Cfg()\n\nsum_coef = 0\nfor v in cfg.VERSION:\n    sum_coef += v[\"coef\"]\n    \nprint(f\"sum_coef: {sum_coef}\")\n\nmodel_weights = []\nfor ver in cfg.VERSION:\n    ver[\"file_mask\"] = f\"/kaggle/input/hms-weights-saito/hms-weights-saito/{ver['name']}/{ver['name']}/*_best.pth\"\n    model_weights.append(ver)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.720017Z","iopub.execute_input":"2024-05-03T21:59:12.720327Z","iopub.status.idle":"2024-05-03T21:59:12.727058Z","shell.execute_reply.started":"2024-05-03T21:59:12.720297Z","shell.execute_reply":"2024-05-03T21:59:12.726108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def butter_bandpass(lowcut, highcut, fs, order=5):\n    return butter(order, [lowcut, highcut], fs=fs, btype=\"band\")\n\n\ndef butter_bandpass_filter(data, lowcut, highcut, fs, order=5):\n    b, a = butter_bandpass(lowcut, highcut, fs, order=order)\n    y = lfilter(b, a, data)\n    return y\n\n\ndef butter_lowpass_filter(\n    data, cutoff_freq=20, sampling_rate=cfg.sampling_rate, order=4\n):\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-05-03T21:59:12.728274Z","iopub.execute_input":"2024-05-03T21:59:12.728600Z","iopub.status.idle":"2024-05-03T21:59:12.735662Z","shell.execute_reply.started":"2024-05-03T21:59:12.728577Z","shell.execute_reply":"2024-05-03T21:59:12.734949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def eeg_from_parquet(\n    parquet_path: str,\n    display: bool = False,\n    seq_length=cfg.seq_length,\n) -> np.ndarray:\n    eeg = pd.read_parquet(parquet_path, columns=cfg.eeg_features)\n    rows = len(eeg)\n\n    offset = (rows - cfg.n_samples) // 2\n\n    eeg = eeg.iloc[offset : offset + cfg.n_samples]\n\n    if display:\n        plt.figure(figsize=(10, 5))\n        offset = 0\n\n    data = np.zeros((cfg.n_samples, len(cfg.eeg_features)))\n\n    for index, feature in enumerate(cfg.eeg_features):\n        x = eeg[feature].values.astype(\"float32\")\n\n        mean = np.nanmean(x)\n        nan_percentage = np.isnan(x).mean()\n        \n        if nan_percentage < 1:\n            x = np.nan_to_num(x, nan=mean)\n        else:\n            x[:] = 0\n        data[:, index] = x\n\n        if display:\n            if index != 0:\n                offset += x.max()\n            plt.plot(range(cfg.n_samples), x - offset, label=feature)\n            offset -= x.min()\n\n    if display:\n        plt.legend()\n        name = parquet_path.split(\"/\")[-1].split(\".\")[0]\n        plt.yticks([])\n        plt.title(f\"EEG {name}\", size=16)\n        plt.show()\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.736798Z","iopub.execute_input":"2024-05-03T21:59:12.737150Z","iopub.status.idle":"2024-05-03T21:59:12.748738Z","shell.execute_reply.started":"2024-05-03T21:59:12.737119Z","shell.execute_reply":"2024-05-03T21:59:12.748047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def spectrogram_from_eeg(\n    parquet_path: Path,\n    names: list,\n    feats: list,\n    t_start: int,\n    fs: int = cfg.sampling_rate,\n    mode_montage: str = \"anterior_posterior\",\n    display: bool = False,\n    cfg_save: dict = {\n        \"enable\": False,\n        \"save_dir\": None,\n        \"target\": None,\n    },\n):\n    eeg = pd.read_parquet(parquet_path)\n    \n    width_img = t_start * fs // 10 + 1000\n    img = np.zeros((100, width_img, len(names)), dtype=\"float32\")\n\n    if display:\n        plt.figure(figsize=(10, 5))\n    signals = []\n    for k in range(len(names)):\n        COLS = feats[k]\n\n        # Set loop count\n        if mode_montage == \"anterior_posterior\":\n            loop_count = len(COLS) - 1\n        elif mode_montage == \"laplacian\":\n            loop_count = len(COLS)\n\n        for kk in range(loop_count):\n            # COMPUTE PAIR DIFFERENCES\n            if mode_montage == \"anterior_posterior\":\n                x = eeg[COLS[kk]].values - eeg[COLS[kk + 1]].values\n            elif mode_montage == \"laplacian\":\n                x = calc_laplacian_montage(eeg, COLS[kk])\n\n            # FILL NANS\n            m = np.nanmean(x)\n            if np.isnan(x).mean() < 1:\n                x = np.nan_to_num(x, nan=m)\n            else:\n                x[:] = 0\n\n            x = np.nan_to_num(x, nan=0) / 32.0\n            x = butter_lowpass_filter(x, order=4)\n            signals.append(x)\n\n            # RAW SPECTROGRAM\n            mel_spec = librosa.feature.melspectrogram(\n                y=x,\n                sr=fs,\n                hop_length=len(x)//width_img,\n                n_fft=1000,\n                win_length=200,\n                n_mels=100,\n                fmin=0.5,\n                fmax=20.5,\n            )\n\n            # LOG TRANSFORM\n            mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max).astype(np.float32)[:,:width_img]\n\n            # STANDARDIZE TO -1 TO 1\n            mel_spec_db = (mel_spec_db + 40) / 40\n            img[:, :, k] += mel_spec_db\n\n        # AVERAGE THE MONTAGE DIFFERENCES\n        img[:, :, k] /= loop_count\n\n        if display:\n            W = int(np.sqrt(len(names)))\n            H = int(np.ceil(len(names) / W))\n            plt.subplot(W, H, k + 1)\n            plt.subplots_adjust(wspace=0.4, hspace=0.5)\n            plt.imshow(img[:, :, k], aspect=\"auto\", origin=\"lower\", cmap=\"jet\")\n            plt.title(f\"EEG {parquet_path.stem} - Spec {names[k]}\", fontsize=10)\n            if cfg_save[\"enable\"]:\n                cfg_save[\"save_dir\"].joinpath(\"spec\").mkdir(parents=True, exist_ok=True)\n                plt.savefig(\n                    cfg_save[\"save_dir\"].joinpath(parquet_path.stem + \"_spec.png\"),\n                    format=\"png\",\n                    dpi=200,\n                )\n\n    if display:\n        fig = plt.figure(figsize=(5 * W, 2.5 * H))\n        ax = fig.add_subplot(1, 1, 1)\n        offset = 0\n        for k in range(len(names)):\n            # idx = len(names) - 1 - k\n            if k > 0:\n                offset += signals[k - 1].max() - signals[k].min()\n\n            linestyle = \"-\"\n            \n            ax.plot(\n                range(len(eeg)),\n                signals[k] + offset,\n                linestyle=linestyle,\n                label=names[k],\n                linewidth=0.5,\n            )\n            # offset += signals[k].max()\n        handles, labels = ax.get_legend_handles_labels()\n\n        # Show legend in reverse order\n        ax.legend(handles=handles[::-1], labels=labels[::-1])\n        ax.set_title(f\"EEG {parquet_path.stem} Signals\")\n        if cfg_save[\"enable\"]:\n            cfg_save[\"save_dir\"].joinpath(\"signal\").mkdir(parents=True, exist_ok=True)\n            plt.savefig(\n                cfg_save[\"save_dir\"].joinpath(parquet_path.stem + \"_signal.png\"),\n                format=\"png\",\n                dpi=200,\n            )\n        plt.show()\n\n    return img","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.749866Z","iopub.execute_input":"2024-05-03T21:59:12.750146Z","iopub.status.idle":"2024-05-03T21:59:12.771133Z","shell.execute_reply.started":"2024-05-03T21:59:12.750124Z","shell.execute_reply":"2024-05-03T21:59:12.770261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_spectrogram_from_eeg(montages_ch=4):\n    paths_eegs = list(Paths.TEST_RAW_EEGS.glob(\"*.parquet\"))\n    all_eeg_specs = {}\n    counter = 0\n\n    for file_path in tqdm(paths_eegs):\n        eeg_id = file_path.stem\n        eeg_spectrogram = spectrogram_from_eeg(\n            file_path,\n            cfg.names_spec,\n            cfg.feats_montage,\n            t_start=0,\n            fs=200,\n            mode_montage=cfg.mode_montage,\n            display=(counter < 1),\n        )\n        all_eeg_specs[int(eeg_id)] = eeg_spectrogram\n        counter += 1\n    \n    return all_eeg_specs","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.772318Z","iopub.execute_input":"2024-05-03T21:59:12.772653Z","iopub.status.idle":"2024-05-03T21:59:12.782709Z","shell.execute_reply.started":"2024-05-03T21:59:12.772624Z","shell.execute_reply":"2024-05-03T21:59:12.781825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Copyright (c) 2021 Irhum Shafkat\n# Released under the MIT license\n# https://github.com/irhum/superlets/blob/main/LICENSE\n\ndef get_bc(cycles, freq, k_sd=5):\n    return cycles/(k_sd * freq)\n\ndef cxmorelet(freq, cycles, sampling_freq):\n    t = jnp.linspace(-1, 1, sampling_freq*2)\n\n    bc = get_bc(cycles, freq)\n    norm = 1/(bc * jnp.sqrt(2*jnp.pi))\n    gauss = jnp.exp(-t**2/(2*bc**2))\n    sine = jnp.exp(1j*2*jnp.pi*freq*t)\n\n    wavelet = norm * gauss * sine\n    return wavelet / jnp.sum(jnp.abs(wavelet))\n\n@partial(jax.jit, static_argnums=3)\n@partial(jax.vmap, in_axes=(None, 0, None, None))\ndef wavelet_transform(signal, freq, cycles, sampling_freq):    \n    wavelet = cxmorelet(freq, cycles, sampling_freq)\n    return jax.scipy.signal.convolve(signal, wavelet, mode=\"same\")\n\n\n@partial(jax.jit, static_argnums=3)\n@partial(jax.vmap, in_axes=(None, None, 0, None))\ndef superlet_transform_helper(signal, freqs, order, sampling_freq):\n    return wavelet_transform(signal, freqs, order, sampling_freq) * jnp.sqrt(2)\n\ndef order_to_cycles(base_cycle, max_order, mode):\n    if mode == \"add\":\n        return jnp.arange(0, max_order) + base_cycle\n    elif mode == \"mul\":\n        return jnp.arange(1, max_order+1) * base_cycle\n    else: raise ValueError(\"mode should be one of \\\"mul\\\" or \\\"add\\\"\")\n\ndef get_order(f, f_min: int, f_max: int, o_min: int, o_max: int):\n    return o_min + round((o_max - o_min) * (f - f_min) / (f_max - f_min))\n\n@partial(jax.vmap, in_axes=(0, None))\ndef get_mask(order, max_order):\n    return jnp.arange(1, max_order+1) > order\n\n@jax.jit\ndef norm_geomean(X, root_pows, eps):\n    X = jnp.log(X + eps).sum(axis=0)\n\n    return jnp.exp(X / jnp.array(root_pows).reshape(-1, 1))\n\n# @jax.jit \ndef adaptive_superlet_transform(signal, freqs, sampling_freq: int, base_cycle: int, min_order: int, max_order: int, eps=1e-12, mode=\"mul\"):\n    \"\"\"Computes the adaptive superlet transform of the provided signal\n\n    Args:\n        signal (jnp.ndarray): 1D array containing the signal data\n        freqs (jnp.ndarray): 1D sorted array containing the frequencies to compute the wavelets at\n        sampling_freq (int): Sampling frequency of the signal \n\n        base_cycle (int): The number of cycles corresponding to order=1\n        min_order (int): The minimum upper limit of orders to be used for a frequency in the adaptive superlet.\n        max_order (int): The maximum upper limit of orders to be used for a frequency in the adaptive superlet.\n        \n        eps (float, optional): Epsilon value to be used for numerical stability in the geometric mean. Defaults to 1e-12.\n        mode (str, optional): \"add\" or \"mul\", corresponding to the use of additive or multiplicative adaptive superlets. Defaults to \"mul\".\n\n    Returns:\n        jnp.ndarray: 2D array (Frequency x Time) representing the computed scalogram\n    \"\"\"\n    cycles = order_to_cycles(base_cycle, max_order, mode)\n    orders = get_order(freqs, min(freqs), max(freqs), min_order, max_order)\n    mask = get_mask(orders, max_order)\n\n    out = superlet_transform_helper(signal, freqs, cycles, sampling_freq)\n    out = out.at[mask.T].set(1)\n    return norm_geomean(out, orders, eps)\n\ndef superlet_from_eeg(\n    parquet_path: Path,\n    names: list,\n    names_spec: list,\n    feats: list,\n    t_start: int,\n    fs: int = 200,\n    mode_montage: str = \"anterior_posterior\",\n    display: bool = False,\n):\n\n    # Load eegs\n    eeg = pd.read_parquet(parquet_path)\n    t_start *= fs  # Seconds -> Number of samples\n    eeg = eeg.iloc[t_start:t_start+10000]  # 10000 samples\n    width_img = 1000  # Downsampling to 1/10\n    height_img = 32\n\n    img = np.zeros((height_img, width_img, len(feats)), dtype=\"float32\")\n\n    # Settings for superlet\n    min_freq, max_freq = 0.5, 20.0\n    freqs = jnp.linspace(min_freq, 20.0, height_img)\n    base_cycle, min_order, max_order = [1, 1, height_img // 2]\n\n    if display:\n        plt.figure(figsize=(3.5 * len(feats) ** 0.5, 2 * len(feats) ** 0.5))\n    signals = []  # for display\n    for k in range(len(feats)):\n        COLS = feats[k]\n\n        # Set loop count\n        if mode_montage == \"anterior_posterior\":\n            loop_count = len(COLS) - 1\n\n        for kk in range(loop_count):\n            # Calc diff\n            if mode_montage == \"anterior_posterior\":\n                x = eeg[COLS[kk]].values - eeg[COLS[kk + 1]].values\n\n            # Fill nans & clip\n            m = np.nanmean(x)\n            if np.isnan(x).mean() < 1:\n                x = np.nan_to_num(x, nan=m)\n            else:\n                x[:] = 0\n            x = np.clip(x, -1024, 1024)\n            signals.append(x[4000:6000])  # for display\n\n            # Superlet\n            scalogram = adaptive_superlet_transform(\n                x, freqs, sampling_freq=fs,\n                base_cycle=base_cycle, min_order=min_order, max_order=max_order, mode=\"add\",\n            )\n            scalogram = np.array(jnp.abs(scalogram)).astype(\"float32\")\n            scalogram = cv2.resize(scalogram, (width_img, height_img))\n            img[:, :, k] += scalogram\n\n        # Average montage diff\n        img[:, :, k] /= loop_count\n        \n        if display:\n            W = int(np.sqrt(len(feats)))\n            H = int(np.ceil(len(feats)) / W)\n            plt.rcParams[\"font.size\"] = 10\n            plt.subplot(W, H, k + 1)\n            plt.subplots_adjust(wspace=0.3, hspace=0.5)\n            plt.imshow(img[:, 400:600, k], aspect=\"auto\", origin=\"lower\", cmap=\"jet\")\n            plt.title(f\"EEG {parquet_path.stem} - Spec {names_spec[k]}\", fontsize=10)\n            plt.yticks([31 * i / 4 for i in range(5)], [(20 - 0.5) * i / 4 + 0.5 for i in range(5)])\n\n\n    if display:\n        fig = plt.figure(figsize=(4 * W, 2 * H))\n        ax = fig.add_subplot(1, 1, 1)\n        offset = 0\n        for k in range(16):\n            if k > 0:\n                offset += signals[k - 1].max() - signals[k].min()\n            ax.plot(\n                range(2000),\n                signals[k] + offset,\n                linestyle=\"-\",\n                label=names[k],\n                linewidth=0.5,\n            )\n        handles, labels = ax.get_legend_handles_labels()\n\n        # Show legend in reverse order\n        ax.legend(handles=handles[::-1], labels=labels[::-1])\n        ax.set_title(f\"EEG {parquet_path.stem} Signals\")\n\n        plt.show()\n\n    return img","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.784145Z","iopub.execute_input":"2024-05-03T21:59:12.784492Z","iopub.status.idle":"2024-05-03T21:59:12.818995Z","shell.execute_reply.started":"2024-05-03T21:59:12.784462Z","shell.execute_reply":"2024-05-03T21:59:12.818107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_superlet_from_eeg(montages_ch=4):\n    paths_eegs = list(Paths.TEST_RAW_EEGS.glob(\"*.parquet\"))\n    all_eeg_superlets = {}\n    counter = 0\n    \n    names_spec = cfg.names_spec_16ch if montages_ch == 16 else cfg.names_spec\n    feats_montage = cfg.feats_montage_16ch if montages_ch == 16 else cfg.feats_montage\n\n    for file_path in tqdm(paths_eegs):\n        eeg_id = file_path.stem\n        eeg_spectrogram = superlet_from_eeg(\n            file_path,\n            cfg.names_montage,\n            names_spec,\n            feats_montage,\n            t_start=0,\n            fs=200,\n            mode_montage=cfg.mode_montage,\n            display=(counter < 1),\n        )\n        all_eeg_superlets[int(eeg_id)] = eeg_spectrogram\n        counter += 1\n    \n    return all_eeg_superlets","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.820151Z","iopub.execute_input":"2024-05-03T21:59:12.820455Z","iopub.status.idle":"2024-05-03T21:59:12.831584Z","shell.execute_reply.started":"2024-05-03T21:59:12.820424Z","shell.execute_reply":"2024-05-03T21:59:12.830672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EEGDataset(Dataset):\n    def __init__(\n        self,\n        df: pd.DataFrame,\n        batch_size: int,\n        imsize: int,\n        n_samples: int,\n        n_samples_spec: int,\n        in_ch_eeg: int,\n        in_ch_spec: int,\n        out_samples: int, \n        out_samples_spec: int, \n        eegs: Dict[int, np.ndarray],\n        eeg_specs: Dict[int, np.ndarray],\n        montage_ch: int,\n        downsample: int = None,\n        norm_spec: bool = False,\n        interpolation: int = cv2.INTER_NEAREST,\n        symmetric_lr_ch: bool = True,\n        bandpass_filter: Dict[str, Union[int, float]] = None,\n        freq_channels = [],\n    ):\n        self.df = df\n        self.batch_size = batch_size\n        self.imsize = imsize\n        self.n_samples = n_samples\n        self.n_samples_spec = n_samples_spec\n        self.in_ch_eeg = in_ch_eeg\n        self.in_ch_spec = in_ch_spec\n        self.out_samples = out_samples\n        self.out_samples_spec = out_samples_spec\n        self.eegs = eegs\n        self.eeg_specs = eeg_specs\n        self.montage_ch = montage_ch\n        self.downsample = downsample\n        self.norm_spec = norm_spec  \n        self.interpolation = interpolation\n        self.symmetric_lr_ch = symmetric_lr_ch\n        self.bandpass_filter = bandpass_filter\n        self.freq_channels = freq_channels\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __norm(self, img, smooth=1e-6):\n        vmax = img.max()\n        vmin = img.min()\n        return (img - vmin) / (vmax - vmin + smooth)\n\n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        X = self.__data_generation(row)\n        if self.out_samples_spec:\n            X_spec = self.__data_generation_spec(row)\n        else:\n            X_spec = None\n        if self.downsample is not None:\n            X = X[::self.downsample, :]\n            \n        X = torch.tensor(X, dtype=torch.float32)\n        if self.out_samples_spec:\n            X_spec = torch.tensor(X_spec, dtype=torch.float32)\n            \n        if self.out_samples_spec:\n            return {\n                \"eegs\": X,\n                \"specs\": X_spec,\n                \"eeg_ids\": row[\"eeg_id\"],\n            }\n        else:\n            return {\n                \"eegs\": X,\n                \"eeg_ids\": row[\"eeg_id\"],\n            }\n\n    def __diff_features(self, tgt_feat, ref_feats):\n        diff = tgt_feat - np.mean(ref_feats, axis=1)\n        return diff\n    \n    def __data_generation(self, row):\n        data = self.eegs[row.eeg_id]\n        X = np.zeros(\n            (self.n_samples, self.in_ch_eeg), dtype=\"float32\"\n        )\n\n        for i, (feat_a, feat_b) in enumerate(cfg.map_features):  # e.g. [Fp1, T3]\n            diff_feat = self.__diff_features(data[:, cfg.feature_to_index[feat_a]], data[:, [cfg.feature_to_index[f] for f in feat_b]])\n\n            if self.bandpass_filter:\n                diff_feat = butter_bandpass_filter(\n                    diff_feat,\n                    self.bandpass_filter[\"low\"],\n                    self.bandpass_filter[\"high\"],\n                    cfg.sampling_rate,\n                    order=self.bandpass_filter[\"order\"],\n                )\n\n            X[:, i] = diff_feat\n            \n        n = cfg.n_map_features\n        if len(self.freq_channels) > 0:\n            for j, (lowcut, highcut) in enumerate(self.freq_channels):\n                for i in range(cfg.n_map_features):\n                    diff_feat = X[:, i]\n                    band_feat = butter_bandpass_filter(\n                        diff_feat, lowcut, highcut, cfg.sampling_rate, order=self.bandpass_filter[\"order\"],\n                    )\n                    X[:, n] = band_feat\n                    n += 1\n\n        X = np.clip(X, -1024, 1024)\n        X = np.nan_to_num(X, nan=0) / 32.0\n        if  self.bandpass_filter[\"lowpass_order\"]:\n            X = butter_lowpass_filter(X, order=self.bandpass_filter[\"lowpass_order\"])\n\n        trim_start = (cfg.n_samples - self.out_samples) // 2\n        X = X[trim_start:trim_start+self.out_samples,:]\n\n        return X\n\n    def __data_generation_spec(self, row):\n        X = np.zeros((self.in_ch_spec, self.imsize, self.montage_ch), dtype='float32')\n\n        n_ch = 0\n        img = self.eeg_specs[row.eeg_id].astype(np.float32)\n#         start = (cfg.n_samples_spec - self.out_samples_spec) // 2\n        start = (cfg.n_samples_spec - self.out_samples_spec) // 2\n        img = img[:, start:start+self.out_samples_spec]\n\n        if self.norm_spec:\n            img = self.__norm(img)\n        img = cv2.resize(img, (X.shape[1], X.shape[0]), interpolation=self.interpolation)\n        X[:, :, n_ch:n_ch+self.montage_ch] = img\n        if self.symmetric_lr_ch:\n            X[:, :, n_ch+self.montage_ch//2:n_ch+self.montage_ch] = X[::-1, :, n_ch+self.montage_ch//2:n_ch+self.montage_ch]\n        n_ch += self.montage_ch\n\n        X = np.concatenate([X[:, :, i] for i in range(n_ch)], axis=0)  # HW\n        \n        return X","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.832946Z","iopub.execute_input":"2024-05-03T21:59:12.833249Z","iopub.status.idle":"2024-05-03T21:59:12.858603Z","shell.execute_reply.started":"2024-05-03T21:59:12.833215Z","shell.execute_reply":"2024-05-03T21:59:12.857791Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EEGNet(nn.Module):\n    def __init__(\n        self,\n        model_2dcnn: str,  # e.g. \"maxvit_rmlp_nano_rw_256\"\n        input_type: str,  # \"eegs_rnn\" or \"eegs_rnn_specs\" or \"eegs_rnn_specs_1ddw\"\n        in_channels: int = 1,  # Fixed, gray-scale image\n        eeg_channels: int = 16,\n        spec_channels: int = 64,\n        fs: int = 200,\n        fs_spec: int = 40,\n        n_samples: int = 2000,\n        n_freqs: int = 16,  # Number of freq filters\n        n_freqs_spec: int = 7,\n        n_dw_features: int = 3,\n        p_dropouts: dict = {\"dw\": 0.2, \"sep\": 0.2},\n        out_ch_cnn: int = 512,\n        out_ch_gru: int = 640,\n    ):\n        super().__init__()\n        \n        self.model_2dcnn = model_2dcnn\n        self.input_type = input_type\n        self.num_classes = 6  # Fixed\n\n        # ===========================================\n        # 1. Conv2D Block (frequency filter)\n        # ===========================================\n        ksize_freq = (1, fs)\n        padsize_freq = (ksize_freq[1] - 1) // 2\n        self.pad_freq = nn.ZeroPad2d(\n            (padsize_freq, padsize_freq + (ksize_freq[1] - 1) % 2, 0, 0)\n        )  # (left, right, top, bottom)\n        out_ch_conv2d = n_freqs\n        self.conv2d_freq = nn.Conv2d(\n            in_channels=in_channels,  # EEG Image (channels x time)\n            out_channels=out_ch_conv2d,\n            kernel_size=ksize_freq,\n            padding=0,  # padding is done in previous layer\n        )\n        self.bn_freq = nn.BatchNorm2d(out_ch_conv2d)\n        \n        # ===========================================\n        # 2. CNN model\n        # ===========================================\n        if \"maxvit\" in self.model_2dcnn.lower() or \"maxxvit\" in self.model_2dcnn.lower():\n            self.avgpool_cnn = nn.AvgPool2d((1, 8))\n            self.pad_cnn = nn.ZeroPad2d(\n                (3, 3, 0, 0)\n            )  # (left, right, top, bottom)\n        else:\n            self.avgpool_cnn = nn.AvgPool2d((1, 4))\n\n        self.order_in_ch_cnn = []\n        for ch in range(eeg_channels):\n            self.order_in_ch_cnn += [eeg_channels * i + ch for i in range(n_freqs)]\n            \n        # ===== Channel chunk =====\n        self.cnn_ch_chunk = self.__timm_create_model()\n        self.features_cnn_ch_chunk, self.fc_ch_chunk = self.__create_features_fc_layer(self.cnn_ch_chunk)\n\n        # ===== Freq chunk =====\n        self.cnn_freq_chunk = self.__timm_create_model()\n        self.features_cnn_freq_chunk, self.fc_freq_chunk = self.__create_features_fc_layer(self.cnn_freq_chunk)\n\n        # ===========================================\n        # 3. GRU layer\n        # ===========================================\n        if \"rnn\" in self.input_type.lower():\n            self.gru = nn.GRU(\n                input_size=eeg_channels*out_ch_conv2d,\n                hidden_size=out_ch_gru,\n                num_layers=1,\n                bidirectional=False,\n                batch_first=True,\n            )\n            self.fc_gru = nn.Sequential(\n                nn.Linear(out_ch_gru, self.num_classes)\n            )\n\n        # ===========================================\n        # 4. CNN for spec\n        # ===========================================\n        if \"spec\" in self.input_type.lower():\n            if \"spec_1ddw\" in self.input_type.lower():\n                ksize_freq_spec = (1, fs_spec)\n                padsize_freq_spec = (ksize_freq_spec[1] - 1) // 2\n                self.pad_freq_spec = nn.ZeroPad2d(\n                    (padsize_freq_spec, padsize_freq_spec + (ksize_freq_spec[1] - 1) % 2, 0, 0)\n                )  # (left, right, top, bottom)\n                out_ch_conv2d_spec = n_freqs_spec\n                self.conv2d_freq_spec = nn.Conv2d(\n                    in_channels=in_channels,  # EEG Image (channels x time)\n                    out_channels=out_ch_conv2d_spec,\n                    kernel_size=ksize_freq_spec,\n                    padding=0,  # padding is done in previous layer\n                )\n                self.bn_freq_spec = nn.BatchNorm2d(out_ch_conv2d_spec)\n                self.order_in_spec = []\n                for ch in range(spec_channels):\n                    self.order_in_spec += [spec_channels * i + ch for i in range(n_freqs_spec+1)]\n            \n            self.cnn_spec = self.__timm_create_model()\n            self.features_cnn_spec, self.fc_spec = self.__create_features_fc_layer(self.cnn_spec)\n\n        # ===========================================\n        # 5. Fully-connected(fc) layer\n        # ===========================================\n        if \"rnn\" in self.input_type.lower() and \"spec\" in self.input_type.lower():\n            in_ch_linear = out_ch_cnn * 3 + out_ch_gru\n        elif \"rnn\" in self.input_type.lower() and not \"spec\" in self.input_type.lower():\n            in_ch_linear = out_ch_cnn * 2 + out_ch_gru\n        self.fc = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(in_ch_linear, self.num_classes)\n        )\n        \n    def __timm_create_model(self):\n        return timm.create_model(\n            self.model_2dcnn,\n            pretrained=False,\n            drop_rate=0.1,\n            drop_path_rate=0.2,\n            in_chans=1,\n            num_classes=self.num_classes,\n        )\n    \n    def __create_features_fc_layer(self, cnn):\n        if \"maxvit\" in self.model_2dcnn.lower() or \"maxxvit\" in self.model_2dcnn.lower() or \"nextvit\" in self.model_2dcnn.lower():\n            features = nn.Sequential(\n                *list(cnn.children())[:-1],\n                list(cnn.children())[-1].global_pool,\n                list(cnn.children())[-1].drop,\n            )\n            fc = nn.Sequential(\n                list(cnn.children())[-1].fc,\n            )\n        elif \"convnextv2\" in self.model_2dcnn.lower():\n            features = nn.Sequential(\n                *list(cnn.children())[:-1],\n                list(cnn.children())[-1].global_pool,\n                list(cnn.children())[-1].norm,\n                list(cnn.children())[-1].flatten,\n                list(cnn.children())[-1].drop,\n            )\n            fc = nn.Sequential(\n                list(cnn.children())[-1].fc,\n            )\n        elif \"inception_next\" in self.model_2dcnn.lower():\n            features = nn.Sequential(\n                *list(cnn.children())[:-1],\n                list(cnn.children())[-1].global_pool,\n                list(cnn.children())[-1].fc1,\n                list(cnn.children())[-1].act,\n                list(cnn.children())[-1].norm,\n            )\n            fc = nn.Sequential(\n                list(cnn.children())[-1].fc2,\n                list(cnn.children())[-1].drop,\n            )\n        elif \"mixnet\" in self.model_2dcnn.lower():\n            features = nn.Sequential(\n                *list(cnn.children())[:-1],\n            )\n            fc = nn.Sequential(\n                list(cnn.children())[-1],\n            )\n        else:  # EfficientNet\n            features = nn.Sequential(*list(cnn.children())[:-1])\n            fc = nn.Sequential(\n                nn.Flatten(),\n                nn.Linear(out_ch_cnn, self.num_classes)\n            )\n\n        return features, fc\n\n    def __create_eeg_image(self, eegs):\n        return eegs.permute(0, 2, 1).unsqueeze(1)\n\n    def forward(self, x):\n        eeg = self.__create_eeg_image(x[\"eegs\"])\n        \n        # 1. Conv2D Block (frequency filter)\n        out = self.conv2d_freq(self.pad_freq(eeg))\n        out = self.bn_freq(out)\n\n        # 2. CNN model\n        out = self.avgpool_cnn(out)\n        if \"maxvit\" in self.model_2dcnn.lower() or \"maxxvit\" in self.model_2dcnn.lower():\n            out = self.pad_cnn(out)\n        in_ch_chunk = torch.cat([out[:, i] for i in range(out.shape[1])], dim=1).unsqueeze(1)\n        in_freq_chunk = torch.cat([in_ch_chunk[:, :, ch:ch+1] for ch in self.order_in_ch_cnn], dim=2)\n\n        out_ch_chunk = self.features_cnn_ch_chunk(in_ch_chunk)\n        out_freq_chunk = self.features_cnn_freq_chunk(in_freq_chunk)\n        out_cnn = torch.cat([out_ch_chunk, out_freq_chunk], dim=1)\n\n        # 3. GRU layer\n        if \"rnn\" in self.input_type.lower():\n            out_gru, _ = self.gru(in_freq_chunk.squeeze(1).permute(0, 2, 1))  # B, Time, Channel\n            out_gru = out_gru[:, -1, :]\n\n        # 4. CNN for spec\n        if \"spec\" in self.input_type.lower():\n            x_spec = x[\"specs\"].unsqueeze(1)\n            if \"spec_1ddw\" in self.input_type.lower():\n                out_spec = self.conv2d_freq_spec(self.pad_freq_spec(x_spec))\n                out_spec = self.bn_freq_spec(out_spec)\n                \n                out_spec = torch.cat([out_spec[:, i] for i in range(out_spec.shape[1])], dim=1).unsqueeze(1)\n                out_spec = torch.cat([x_spec, out_spec], dim=2)\n                out_spec = torch.cat([out_spec[:, :, ch:ch+1] for ch in self.order_in_spec], dim=2)\n                out_spec = self.features_cnn_spec(out_spec)\n            else:\n                out_spec = self.features_cnn_spec(x_spec)\n\n        # 5. Fully-connected(fc) layer\n        if \"rnn\" in self.input_type.lower() and \"spec\" in self.input_type.lower():\n            out = self.fc(torch.cat([out_cnn, out_gru, out_spec], dim=1))\n        elif \"rnn\" in self.input_type.lower() and not \"spec\" in self.input_type.lower():\n            out = self.fc(torch.cat([out_cnn, out_gru], dim=1))\n        \n        return {\n            \"main\": out,\n        }","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.859932Z","iopub.execute_input":"2024-05-03T21:59:12.860477Z","iopub.status.idle":"2024-05-03T21:59:12.900832Z","shell.execute_reply.started":"2024-05-03T21:59:12.860445Z","shell.execute_reply":"2024-05-03T21:59:12.899883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EEGNetV2(nn.Module):\n    def __init__(\n        self,\n        model_2dcnn,  # e.g. \"maxvit_rmlp_nano_rw_256\"\n        input_type,  # e.g. \"eegs_rnn_specs\"\n        imsize: int, \n        in_channels=1,\n        eeg_channels=16,\n        spec_channels=64,\n        n_samples=2000,\n        n_freqs=[16],  # Number of freq filters\n        n_freqs_spec=7,\n        n_dw_features=3,\n        p_dropouts={\"dw\": 0.2, \"sep\": 0.2},\n        fs=200,\n        fs_spec=40,\n    ):\n        super().__init__()\n        \n        self.model_2dcnn = model_2dcnn\n        self.input_type = input_type\n        self.imsize = imsize\n        self.num_classes = 6\n        self.fs = fs   # Sampling rate\n        self.fs_spec = fs_spec   # Sampling rate\n\n        if any(s in self.model_2dcnn.lower() for s in [\"maxvit\", \"maxxvit\"]):\n            if \"pico\" in self.model_2dcnn.lower():\n                out_ch_cnn=256\n            elif \"tiny\" in self.model_2dcnn.lower():  \n                out_ch_cnn=512\n            elif \"small\" in self.model_2dcnn.lower():  \n                out_ch_cnn=768\n        elif \"nextvit\" in self.model_2dcnn.lower():\n            out_ch_cnn=1024\n        elif \"convnextv2\" in self.model_2dcnn.lower():\n            if \"atto\" in self.model_2dcnn.lower():\n                out_ch_cnn=320\n            elif \"pico\" in self.model_2dcnn.lower():\n                out_ch_cnn=512\n        elif \"inception_next\" in self.model_2dcnn.lower():\n            out_ch_cnn=2304\n        elif \"mixnet\" in self.model_2dcnn.lower():\n            out_ch_cnn=1536\n        elif \"swin\" in self.model_2dcnn.lower():\n            if \"tiny\" in self.model_2dcnn.lower():\n                out_ch_cnn=768\n            elif \"base\" in self.model_2dcnn.lower():\n                out_ch_cnn=1024\n        elif \"caformer\" in self.model_2dcnn.lower():\n            out_ch_cnn=512\n        elif \"gcvit\" in self.model_2dcnn.lower():\n            out_ch_cnn=512\n        elif \"sequencer\" in self.model_2dcnn.lower():\n            out_ch_cnn=384\n        elif \"tinyvit\" in self.model_2dcnn.lower():\n            out_ch_cnn=320\n        elif \"poolformer\" in self.model_2dcnn.lower():\n            out_ch_cnn=512\n        else:  # EfficientNet\n            out_ch_cnn=512\n        out_ch_gru=out_ch_cnn\n        \n        # ===========================================\n        # 1. Conv2D Block (frequency filter)\n        # ===========================================\n        out_ch_conv2d = 0\n        self.dw_blocks = nn.ModuleList()\n        for i in range(len(n_freqs)):\n            ksize_freq = (1, self.fs)\n            padsize_freq = (ksize_freq[1] - 1) // 2\n            pad_freq = nn.ZeroPad2d(\n                (padsize_freq, padsize_freq + (ksize_freq[1] - 1) % 2, 0, 0)\n            )  # (left, right, top, bottom)\n            out_ch = n_freqs[i]\n            out_ch_conv2d += out_ch\n            conv2d_freq = nn.Conv2d(\n                in_channels=in_channels,  # EEG Image (channels x time)\n                out_channels=out_ch,\n                kernel_size=ksize_freq,\n                padding=0,  # padding is done in previous layer\n            )\n            bn_freq = nn.BatchNorm2d(out_ch)\n            silu = nn.SiLU()\n            self.dw_blocks.append(\n                nn.Sequential(\n                    pad_freq,\n                    conv2d_freq,\n                    bn_freq,\n                    silu,\n                )\n            )\n    \n        # ===========================================\n        # 2. CNN model\n        # ===========================================\n        if any(s in self.model_2dcnn.lower() for s in [\"maxvit\", \"maxxvit\", \"swin\"]):\n            self.avgpool_cnn = nn.AvgPool2d((1, 8))\n            self.pad_cnn = nn.ZeroPad2d(\n                (3, 3, 0, 0)\n            )  # (left, right, top, bottom)\n        elif \"coatnet\" in self.model_2dcnn.lower() or \"gcvit\" in self.model_2dcnn.lower():\n            self.avgpool_cnn = nn.AvgPool2d((1, 10))\n        else:\n            self.avgpool_cnn = nn.AvgPool2d((1, 4))\n\n        self.order_in_ch_cnn = []\n        for ch in range(eeg_channels):\n            self.order_in_ch_cnn += [eeg_channels * i + ch for i in range(np.sum(n_freqs))]\n            \n        # ===== Channel chunk =====\n        self.cnn_ch_chunk = self.__timm_create_model()\n        if any(s in self.model_2dcnn.lower() for s in [\"maxvit\", \"maxxvit\", \"nextvit\", \"gcvit\", \"swin\", \"coatnet\", \"sequencer\"]):\n            self.features_cnn_ch_chunk = nn.Sequential(\n                *list(self.cnn_ch_chunk.children())[:-1],\n                list(self.cnn_ch_chunk.children())[-1].global_pool,\n                list(self.cnn_ch_chunk.children())[-1].drop,\n            )\n            self.fc_ch_chunk = nn.Sequential(\n                list(self.cnn_ch_chunk.children())[-1].fc,\n            )\n        elif any(s in self.model_2dcnn.lower() for s in [\"convnextv2\", \"convformer\", \"caformer\", \"tiny_vit\", \"poolformer\"]):\n            self.features_cnn_ch_chunk = nn.Sequential(\n                *list(self.cnn_ch_chunk.children())[:-1],\n                list(self.cnn_ch_chunk.children())[-1].global_pool,\n                list(self.cnn_ch_chunk.children())[-1].norm,\n                list(self.cnn_ch_chunk.children())[-1].flatten,\n                list(self.cnn_ch_chunk.children())[-1].drop,\n            )\n            self.fc_ch_chunk = nn.Sequential(\n                list(self.cnn_ch_chunk.children())[-1].fc,\n            )\n        elif \"inception_next\" in self.model_2dcnn.lower():\n            self.features_cnn_ch_chunk = nn.Sequential(\n                *list(self.cnn_ch_chunk.children())[:-1],\n                list(self.cnn_ch_chunk.children())[-1].global_pool,\n                list(self.cnn_ch_chunk.children())[-1].fc1,\n                list(self.cnn_ch_chunk.children())[-1].act,\n                list(self.cnn_ch_chunk.children())[-1].norm,\n            )\n            self.fc_ch_chunk = nn.Sequential(\n                list(self.cnn_ch_chunk.children())[-1].fc2,\n                list(self.cnn_ch_chunk.children())[-1].drop,\n            )\n        elif \"mixnet\" in self.model_2dcnn.lower():\n            self.features_cnn_ch_chunk = nn.Sequential(\n                *list(self.cnn_ch_chunk.children())[:-1],\n            )\n            self.fc_ch_chunk = nn.Sequential(\n                list(self.cnn_ch_chunk.children())[-1],\n            )\n        elif \"swin\" in self.model_2dcnn.lower():\n            self.features_cnn_ch_chunk = nn.Sequential(\n                *list(self.cnn_ch_chunk.children())[:-1],\n                list(self.cnn_ch_chunk.children())[-1].global_pool,\n                list(self.cnn_ch_chunk.children())[-1].drop,\n            )\n            self.fc_ch_chunk = nn.Sequential(\n                list(self.cnn_ch_chunk.children())[-1].fc,\n            )\n        else:  # EfficientNet\n            self.features_cnn_ch_chunk = nn.Sequential(*list(self.cnn_ch_chunk.children())[:-1])\n            self.fc_ch_chunk = nn.Sequential(\n                nn.Flatten(),\n                nn.Linear(out_ch_cnn, self.num_classes)\n            )\n \n\n        # ===== Freq chunk =====\n        self.cnn_freq_chunk = self.__timm_create_model()\n        if any(s in self.model_2dcnn.lower() for s in [\"maxvit\", \"maxxvit\", \"nextvit\", \"gcvit\", \"swin\", \"coatnet\", \"sequencer\"]):\n            self.features_cnn_freq_chunk = nn.Sequential(\n                *list(self.cnn_freq_chunk.children())[:-1],\n                list(self.cnn_freq_chunk.children())[-1].global_pool,\n                list(self.cnn_freq_chunk.children())[-1].drop,\n            )\n            self.fc_freq_chunk = nn.Sequential(\n                list(self.cnn_freq_chunk.children())[-1].fc,\n            )\n        elif any(s in self.model_2dcnn.lower() for s in [\"convnextv2\", \"convformer\", \"caformer\", \"tiny_vit\", \"poolformer\"]):\n            self.features_cnn_freq_chunk = nn.Sequential(\n                *list(self.cnn_freq_chunk.children())[:-1],\n                list(self.cnn_freq_chunk.children())[-1].global_pool,\n                list(self.cnn_freq_chunk.children())[-1].norm,\n                list(self.cnn_freq_chunk.children())[-1].flatten,\n                list(self.cnn_freq_chunk.children())[-1].drop,\n            )\n            self.fc_freq_chunk = nn.Sequential(\n                list(self.cnn_freq_chunk.children())[-1].fc,\n            )\n        elif \"inception_next\" in self.model_2dcnn.lower():\n            self.features_cnn_freq_chunk = nn.Sequential(\n                *list(self.cnn_freq_chunk.children())[:-1],\n                list(self.cnn_freq_chunk.children())[-1].global_pool,\n                list(self.cnn_freq_chunk.children())[-1].fc1,\n                list(self.cnn_freq_chunk.children())[-1].act,\n                list(self.cnn_freq_chunk.children())[-1].norm,\n            )\n            self.fc_freq_chunk = nn.Sequential(\n                list(self.cnn_freq_chunk.children())[-1].fc2,\n                list(self.cnn_freq_chunk.children())[-1].drop,\n            )\n        elif \"mixnet\" in self.model_2dcnn.lower():\n            self.features_cnn_freq_chunk = nn.Sequential(\n                *list(self.cnn_freq_chunk.children())[:-1],\n            )\n            self.fc_freq_chunk = nn.Sequential(\n                list(self.cnn_freq_chunk.children())[-1],\n            )\n        elif \"swin\" in self.model_2dcnn.lower():\n            self.features_cnn_freq_chunk = nn.Sequential(\n                *list(self.cnn_freq_chunk.children())[:-1],\n                list(self.cnn_freq_chunk.children())[-1].global_pool,\n                list(self.cnn_freq_chunk.children())[-1].drop,\n            )\n            self.fc_freq_chunk = nn.Sequential(\n                list(self.cnn_freq_chunk.children())[-1].fc,\n            )\n        else:  # EfficientNet\n            self.features_cnn_freq_chunk = nn.Sequential(*list(self.cnn_freq_chunk.children())[:-1])\n            self.fc_freq_chunk = nn.Sequential(\n                nn.Flatten(),\n                nn.Linear(out_ch_cnn, self.num_classes)\n            )\n\n        # ===========================================\n        # 3. GRU layer\n        # ===========================================\n        if \"rnn\" in self.input_type.lower():\n            self.gru = nn.GRU(\n                input_size=eeg_channels*out_ch_conv2d,\n                hidden_size=out_ch_gru,\n                num_layers=1,\n                bidirectional=False,\n                batch_first=True,\n            )\n            self.fc_gru = nn.Sequential(\n                nn.Linear(out_ch_gru, self.num_classes)\n            )\n\n        # ===========================================\n        # 4. CNN for spec\n        # ===========================================\n        if \"spec\" in self.input_type.lower():\n            if \"spec_1ddw\" in self.input_type.lower():\n                ksize_freq_spec = (1, fs_spec)\n                padsize_freq_spec = (ksize_freq_spec[1] - 1) // 2\n                self.pad_freq_spec = nn.ZeroPad2d(\n                    (padsize_freq_spec, padsize_freq_spec + (ksize_freq_spec[1] - 1) % 2, 0, 0)\n                )  # (left, right, top, bottom)\n                out_ch_conv2d_spec = n_freqs_spec\n                self.conv2d_freq_spec = nn.Conv2d(\n                    in_channels=in_channels,  # EEG Image (channels x time)\n                    out_channels=out_ch_conv2d_spec,\n                    kernel_size=ksize_freq_spec,\n                    padding=0,  # padding is done in previous layer\n                )\n                self.bn_freq_spec = nn.BatchNorm2d(out_ch_conv2d_spec)\n                self.silu = nn.SiLU()\n                self.order_in_spec = []\n                for ch in range(spec_channels):\n                    self.order_in_spec += [spec_channels * i + ch for i in range(n_freqs_spec+1)]\n            \n            self.cnn_spec = self.__timm_create_model()\n            if any(s in self.model_2dcnn.lower() for s in [\"maxvit\", \"maxxvit\", \"nextvit\", \"gcvit\", \"swin\", \"coatnet\", \"sequencer\"]):\n                self.features_cnn_spec = nn.Sequential(\n                    *list(self.cnn_spec.children())[:-1],\n                    list(self.cnn_spec.children())[-1].global_pool,\n                    list(self.cnn_spec.children())[-1].drop,\n                )\n                self.fc_spec = nn.Sequential(\n                    list(self.cnn_spec.children())[-1].fc,\n                )\n            elif any(s in self.model_2dcnn.lower() for s in [\"convnextv2\", \"convformer\", \"caformer\", \"tiny_vit\", \"poolformer\"]):\n                self.features_cnn_spec = nn.Sequential(\n                    *list(self.cnn_spec.children())[:-1],\n                    list(self.cnn_spec.children())[-1].global_pool,\n                    list(self.cnn_spec.children())[-1].norm,\n                    list(self.cnn_spec.children())[-1].flatten,\n                    list(self.cnn_spec.children())[-1].drop,\n                )\n                self.fc_spec = nn.Sequential(\n                    list(self.cnn_spec.children())[-1].fc,\n                )\n            elif \"inception_next\" in self.model_2dcnn.lower():\n                self.features_cnn_spec = nn.Sequential(\n                    *list(self.cnn_spec.children())[:-1],\n                    list(self.cnn_spec.children())[-1].global_pool,\n                    list(self.cnn_spec.children())[-1].fc1,\n                    list(self.cnn_spec.children())[-1].act,\n                    list(self.cnn_spec.children())[-1].norm,\n                )\n                self.fc_spec = nn.Sequential(\n                    list(self.cnn_spec.children())[-1].fc2,\n                    list(self.cnn_spec.children())[-1].drop,\n                )\n            elif \"mixnet\" in self.model_2dcnn.lower():\n                self.features_cnn_spec = nn.Sequential(\n                    *list(self.cnn_spec.children())[:-1],\n                )\n                self.fc_spec = nn.Sequential(\n                    list(self.cnn_spec.children())[-1],\n                )\n            elif \"swin\" in self.model_2dcnn.lower():\n                self.features_cnn_spec = nn.Sequential(\n                    *list(self.cnn_spec.children())[:-1],\n                    list(self.cnn_spec.children())[-1].global_pool,\n                    list(self.cnn_spec.children())[-1].drop,\n                )\n                self.fc_spec = nn.Sequential(\n                    list(self.cnn_spec.children())[-1].fc,\n                )\n            else:  # EfficientNet\n                self.features_cnn_spec = nn.Sequential(*list(self.cnn_spec.children())[:-1])\n                self.fc_spec = nn.Sequential(\n                    nn.Flatten(),\n                    nn.Linear(out_ch_cnn, self.num_classes)\n                )\n\n        # ===========================================\n        # 5. Fully-connected(fc) layer\n        # ===========================================\n        if \"rnn\" in self.input_type.lower() and \"spec\" in self.input_type.lower():\n            in_ch_linear = out_ch_cnn * 3 + out_ch_gru\n        elif \"rnn\" in self.input_type.lower() and not \"spec\" in self.input_type.lower():\n            in_ch_linear = out_ch_cnn * 2 + out_ch_gru\n        elif \"spec\" in self.input_type.lower() and not \"rnn\" in self.input_type.lower():\n            in_ch_linear = out_ch_cnn * 3\n        else:\n            in_ch_linear = out_ch_cnn * 2\n        self.fc = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(in_ch_linear, self.num_classes)\n        )\n        \n    def __timm_create_model(self):\n        return timm.create_model(\n            self.model_2dcnn,\n            pretrained=False,\n            drop_rate=0.1,\n            drop_path_rate=0.2,\n            in_chans=1,\n            num_classes=self.num_classes,\n        )\n\n    def __create_eeg_image(self, eegs):\n        return eegs.permute(0, 2, 1).unsqueeze(1)\n\n    def forward(self, x):\n        eeg = self.__create_eeg_image(x[\"eegs\"])\n        \n        # 1. Conv2D Block (frequency filter)\n        outs = []\n        for i, dw in enumerate(self.dw_blocks):\n            out = dw(eeg[:, :, i*16:(i+1)*16])\n            outs.append(out)\n        out = torch.cat(outs, dim=1)\n\n        # 2. CNN model\n        out = self.avgpool_cnn(out)\n        if any(s in self.model_2dcnn.lower() for s in [\"maxvit\", \"maxxvit\", \"swin\"]) and out.shape[3] % self.imsize != 0:\n            out = self.pad_cnn(out)\n        in_ch_chunk = torch.cat([out[:, i] for i in range(out.shape[1])], dim=1).unsqueeze(1)\n        in_freq_chunk = torch.cat([in_ch_chunk[:, :, ch:ch+1] for ch in self.order_in_ch_cnn], dim=2)\n\n        out_ch_chunk = self.features_cnn_ch_chunk(in_ch_chunk)\n        out_freq_chunk = self.features_cnn_freq_chunk(in_freq_chunk)\n        out_cnn = torch.cat([out_ch_chunk, out_freq_chunk], dim=1)\n\n        # 3. GRU layer\n        if \"rnn\" in self.input_type.lower():\n            out_gru, _ = self.gru(in_freq_chunk.squeeze(1).permute(0, 2, 1))  # B, Time, Channel\n            out_gru = out_gru[:, -1, :]\n\n        # 4. CNN for spec\n        if \"spec\" in self.input_type.lower():\n            x_spec = x[\"specs\"].unsqueeze(1)\n            if \"spec_1ddw\" in self.input_type.lower():\n                out_spec = self.conv2d_freq_spec(self.pad_freq_spec(x_spec))\n                out_spec = self.bn_freq_spec(out_spec)\n                out_spec = self.silu(out_spec)\n                \n                out_spec = torch.cat([out_spec[:, i] for i in range(out_spec.shape[1])], dim=1).unsqueeze(1)\n                out_spec = torch.cat([x_spec, out_spec], dim=2)\n                out_spec = torch.cat([out_spec[:, :, ch:ch+1] for ch in self.order_in_spec], dim=2)\n                out_spec = self.features_cnn_spec(out_spec)\n            else:\n                out_spec = self.features_cnn_spec(x_spec)\n\n        # 5. Fully-connected(fc) layer\n        if \"rnn\" in self.input_type.lower() and \"spec\" in self.input_type.lower():\n            out = self.fc(torch.cat([out_cnn, out_gru, out_spec], dim=1))\n        elif \"rnn\" in self.input_type.lower() and not \"spec\" in self.input_type.lower():\n            out = self.fc(torch.cat([out_cnn, out_gru], dim=1))\n        elif \"spec\" in self.input_type.lower() and not \"rnn\" in self.input_type.lower():\n            out = self.fc(torch.cat([out_cnn, out_spec], dim=1))\n        else:\n            out = self.fc(out_cnn)\n        return {\n            \"main\": out,\n        }","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.902298Z","iopub.execute_input":"2024-05-03T21:59:12.903120Z","iopub.status.idle":"2024-05-03T21:59:12.980679Z","shell.execute_reply.started":"2024-05-03T21:59:12.903096Z","shell.execute_reply":"2024-05-03T21:59:12.979722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def inference_function(test_loader, model, device, model_type):\n    model.eval()  # set model in evaluation mode\n#     softmax = nn.Softmax(dim=1)\n    prediction_dict = {}\n    preds = []\n    with tqdm(test_loader, unit=\"test_batch\", desc=\"Inference\") as tqdm_test_loader:\n        for step, X in enumerate(tqdm_test_loader):\n            for k, v in X.items():\n                X[k] = v.to(device)\n#             X = batch.pop(\"eegs\").to(device)\n#             X_specs = batch.pop(\"specs\").to(device)\n            with torch.no_grad():\n                y_preds = model(X)\n#                 y_preds = softmax(y_preds[\"main\"]).to(\"cpu\").numpy()\n                y_preds = entmax_bisect(y_preds[\"main\"], alpha=1.03, dim=1).to(\"cpu\").numpy()\n            preds.append(y_preds)\n\n    prediction_dict[\"predictions\"] = np.concatenate(\n        preds\n    )  # np.array() of shape (fold_size, target_cols)\n    return prediction_dict","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.981906Z","iopub.execute_input":"2024-05-03T21:59:12.982260Z","iopub.status.idle":"2024-05-03T21:59:12.993956Z","shell.execute_reply.started":"2024-05-03T21:59:12.982232Z","shell.execute_reply":"2024-05-03T21:59:12.993051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(Paths.TEST_CSV)\nprint(f\"Test dataframe shape is: {test_df.shape}\")\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:12.995124Z","iopub.execute_input":"2024-05-03T21:59:12.995480Z","iopub.status.idle":"2024-05-03T21:59:13.014953Z","shell.execute_reply.started":"2024-05-03T21:59:12.995449Z","shell.execute_reply":"2024-05-03T21:59:13.014104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_eeg_parquet_paths = list(Paths.TEST_RAW_EEGS.glob(\"*.parquet\"))\ntest_eeg_df = pd.read_parquet(test_eeg_parquet_paths[0])\ntest_eeg_features = test_eeg_df.columns\nprint(f\"There are {len(test_eeg_features)} raw eeg features\")\nprint(list(test_eeg_features))\ndel test_eeg_df\n_ = gc.collect()\n\n# %%time\nall_eegs = {}\neeg_ids = test_df.eeg_id.unique()\nfor i, eeg_id in tqdm(enumerate(eeg_ids)):\n    # Save EEG to Python dictionary of numpy arrays\n    eeg_path = Paths.TEST_RAW_EEGS / (str(eeg_id) + \".parquet\")\n    data = eeg_from_parquet(eeg_path)\n    all_eegs[eeg_id] = data","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:13.027308Z","iopub.execute_input":"2024-05-03T21:59:13.027585Z","iopub.status.idle":"2024-05-03T21:59:13.307428Z","shell.execute_reply.started":"2024-05-03T21:59:13.027561Z","shell.execute_reply":"2024-05-03T21:59:13.306443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coef_sum = 0\ncoef_count = 0\npredictions = []\nfiles = []\n\nfor model_block in model_weights:\n    if \"v062\" in model_block['name']:  # v057~v059\n        all_eeg_specs = create_superlet_from_eeg(montages_ch=16)\n    elif \"v056\" in model_block['name']:  # v046~v056\n        del all_eeg_specs\n        gc.collect()\n        all_eeg_specs = create_superlet_from_eeg(montages_ch=4)\n    elif \"v042\" in model_block['name']:  # v042\n        del all_eeg_specs\n        gc.collect()\n        all_eeg_specs = create_spectrogram_from_eeg()\n    test_dataset = EEGDataset(\n        df=test_df,\n        batch_size=cfg.batch_size,\n        imsize=model_block['imsize'],\n        n_samples=cfg.n_samples,\n        n_samples_spec=cfg.n_samples_spec,\n        in_ch_eeg=model_block['in_ch_eeg'],\n        in_ch_spec=model_block['in_ch_spec'],\n        out_samples=model_block['out_samples'],\n        out_samples_spec=model_block['out_samples_spec'],\n        eegs=all_eegs,\n        eeg_specs=all_eeg_specs,\n        montage_ch=model_block['montage_ch'],\n        norm_spec=model_block['norm_spec'],\n        interpolation=model_block['interpolation'],\n        symmetric_lr_ch=model_block['symmetric_lr_ch'],\n        bandpass_filter=model_block['bandpass_filter'],\n        freq_channels=model_block['freq_channels'],\n    )       \n    if len(predictions) == 0:\n        output = test_dataset[0]\n        X = output[\"eegs\"]\n        print(f\"X shape: {X.shape}\")\n                \n    test_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    )\n    \n    print(model_block[\"model_2dcnn\"], model_block[\"input_type\"])\n    if model_block[\"model_ver\"] == \"v2\":\n        model = EEGNetV2(\n            model_2dcnn=model_block[\"model_2dcnn\"],\n            input_type=model_block[\"input_type\"],\n            imsize=model_block['imsize'],\n            eeg_channels=cfg.n_map_features,\n            spec_channels=model_block[\"in_ch_spec\"] * 4,\n            n_freqs=model_block[\"n_freqs\"],\n            n_freqs_spec=model_block[\"n_freqs_spec\"],\n            fs=cfg.sampling_rate,\n            fs_spec=model_block[\"fs_spec\"],\n        )\n    else:\n        model = EEGNet(\n            model_2dcnn=model_block[\"model_2dcnn\"],\n            input_type=model_block[\"input_type\"],\n            eeg_channels=cfg.n_map_features,\n            spec_channels=model_block[\"in_ch_spec\"] * 4,\n            n_freqs=model_block[\"n_freqs\"],\n            n_freqs_spec=model_block[\"n_freqs_spec\"],\n            fs=cfg.sampling_rate,\n            fs_spec=model_block[\"fs_spec\"],\n            out_ch_cnn=model_block[\"out_ch_cnn\"],\n            out_ch_gru=model_block[\"out_ch_gru\"],\n        )\n\n\n    coef = model_block['coef']\n    single_model_preds = []\n    for weight_model_file in glob(model_block['file_mask']):\n        files.append(weight_model_file)\n        checkpoint = torch.load(weight_model_file, map_location=device)\n        model.load_state_dict(checkpoint[\"model\"], strict=True)\n        model.to(device)\n        prediction_dict = inference_function(test_loader, model, device, model_block[\"input_type\"])\n        predict = prediction_dict[\"predictions\"]\n#         predict *= coef\n#         coef_sum += coef\n#         coef_count += 1\n        single_model_preds.append(predict)\n        torch.cuda.empty_cache()\n        gc.collect()\n        \n    predictions.append(np.mean(np.array(single_model_preds), axis=0))\n        \ns_predictions = predictions.copy()\n\n# predictions = np.array(predictions)\n# coef_sum /= coef_count\n# predictions /= coef_sum\n# predictions = np.mean(predictions, axis=0)\n\n# print(coef_count)\n# display(files)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T21:59:13.308872Z","iopub.execute_input":"2024-05-03T21:59:13.309152Z","iopub.status.idle":"2024-05-03T22:02:13.958120Z","shell.execute_reply.started":"2024-05-03T21:59:13.309128Z","shell.execute_reply":"2024-05-03T22:02:13.957109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#sub = pd.DataFrame({\"eeg_id\": test_df.eeg_id.values})\n#sub[cfg.target_cols] = predictions\n\n#sub.to_csv(f\"submission.csv\", index=False)\n#print(f\"Submission shape: {sub.shape}\")\n#sub.head()\n\n#s_sub = sub.copy()","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:02:13.959439Z","iopub.execute_input":"2024-05-03T22:02:13.959793Z","iopub.status.idle":"2024-05-03T22:02:13.964527Z","shell.execute_reply.started":"2024-05-03T22:02:13.959760Z","shell.execute_reply":"2024-05-03T22:02:13.963480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del all_eegs, all_eeg_specs, model, test_dataset, test_loader, single_model_preds\ntorch.cuda.empty_cache()\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:02:13.965657Z","iopub.execute_input":"2024-05-03T22:02:13.965932Z","iopub.status.idle":"2024-05-03T22:02:14.224608Z","shell.execute_reply.started":"2024-05-03T22:02:13.965909Z","shell.execute_reply":"2024-05-03T22:02:14.223589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Fujii Model","metadata":{}},{"cell_type":"code","source":"import sys\nimport os\nimport gc\nimport copy\nimport yaml\nimport random\nimport shutil\nimport time\nimport typing as tp\nfrom glob import glob\nfrom pathlib import Path\nfrom collections import OrderedDict, defaultdict\nfrom logging import getLogger, INFO, StreamHandler, FileHandler, Formatter\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom tqdm import tqdm\nfrom scipy.signal import butter, lfilter, freqz\nfrom sklearn.model_selection import StratifiedGroupKFold\n\nimport torch\nfrom torch import nn\nfrom torch import optim\nfrom torch.optim import lr_scheduler, Adam, AdamW\nfrom torch.cuda import amp\nfrom torch.utils.data import DataLoader, Dataset, default_collate\nfrom torchvision.transforms import v2\n\nimport timm\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nsys.path.append(\"/kaggle/input/kaggle-kl-div\")\nfrom kaggle_kl_div import score","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:02:14.225728Z","iopub.execute_input":"2024-05-03T22:02:14.226021Z","iopub.status.idle":"2024-05-03T22:02:14.234088Z","shell.execute_reply.started":"2024-05-03T22:02:14.225996Z","shell.execute_reply":"2024-05-03T22:02:14.233178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0,1\"","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:02:14.235153Z","iopub.execute_input":"2024-05-03T22:02:14.235446Z","iopub.status.idle":"2024-05-03T22:02:14.248372Z","shell.execute_reply.started":"2024-05-03T22:02:14.235421Z","shell.execute_reply":"2024-05-03T22:02:14.247586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    print('Using', torch.cuda.device_count(), 'GPU(s)')\n    \n    classes = [\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]\n    n_classes = len(classes)\n\n    eeg_features = [\"Fp1\",\"F3\",\"C3\",\"P3\",\"F7\",\"T3\",\"T5\",\"O1\",\"Fz\",\"Cz\",\"Pz\",\"Fp2\",\"F4\",\"C4\",\"P4\",\"F8\",\"T4\",\"T6\",\"O2\",\"EKG\"]\n    #eeg_features = [\"Fp1\",\"F3\",\"C3\",\"P3\",\"F7\",\"T3\",\"T5\",\"O1\",\"Fp2\",\"F4\",\"C4\",\"P4\",\"F8\",\"T4\",\"T6\",\"O2\", \"Fz\", \"Cz\", \"Pz\"]\n    feature_to_index = {x: y for x, y in zip(eeg_features, range(len(eeg_features)))}\n    map_features = [\n        (\"Fp1\", \"F7\"),\n        (\"F7\", \"T3\"),\n        (\"T3\", \"T5\"),\n        (\"T5\", \"O1\"),\n        (\"Fp1\", \"F3\"),\n        (\"F3\", \"C3\"),\n        (\"C3\", \"P3\"),\n        (\"P3\", \"O1\"),\n        (\"Fp2\", \"F8\"),\n        (\"F8\", \"T4\"),\n        (\"T4\", \"T6\"),\n        (\"T6\", \"O2\"),\n        (\"Fp2\", \"F4\"),\n        (\"F4\", \"C4\"),\n        (\"C4\", \"P4\"),\n        (\"P4\", \"O2\"),\n        (\"Fz\", \"Cz\"),\n        (\"Cz\", \"Pz\")\n    ]\n    \n    n_map_features = len(map_features)\n    #freq_channels = [(0.5, 4.0), (2.0, 6.0), (4.0, 8.0), (6.0, 10.0), (8.0, 13.0), (10.0, 15.0), (13.0, 20.0)]\n    freq_channels = []\n    in_channels = n_map_features + n_map_features * len(freq_channels)\n    seq_length = 50  # Second's\n    sampling_rate = 200  # Hz\n    nsamples = seq_length * sampling_rate\n        \n    batch_size = 32\n    num_workers = 0\n    \n#     stride = 4\n#     input_type = 'paul'\n#     bandpass_filter = None\n    \n    \n    # Model condigs\n    VERSION = [\n        {\n           \"file_mask\": \"/kaggle/input/hms-weights-fujii/done_51_20240403_1841/*stage2.bin\",\n           \"model\": \"maxvit_base_tf_512.in21k_ft_in1k\",\n           \"tta_offset\":[0],\n           \"input_size\":(512,512),\n           \"in_chan\":1,\n           \"out_samples\":5000, \n           \"stride\":8,\n           \"input_type\":'paul',\n           \"bandpass_filter\":None,\n           \"m\":4, \n        },\n        {\n           \"file_mask\": \"/kaggle/input/hms-weights-fujii/done_52_20240403_1829/*stage2.bin\",\n           \"model\": \"maxvit_base_tf_512.in21k_ft_in1k\",\n           \"tta_offset\":[0],\n           \"input_size\":(512,512),\n           \"in_chan\":1,\n           \"out_samples\":2000, \n           \"stride\":4,\n           \"input_type\":'paul',\n           \"bandpass_filter\":None,\n           \"m\":4, \n        },\n        {\n           \"file_mask\": \"/kaggle/input/hms-weights-fujii/done_71_20240406_2039/*stage2.bin\",\n           \"model\": \"maxvit_base_tf_512.in21k_ft_in1k\",\n           \"tta_offset\":[0],\n           \"input_size\":(512,512),\n           \"in_chan\":1,\n           \"out_samples\":5000, \n           \"stride\":8,\n           \"input_type\":'paul',\n           \"bandpass_filter\":None,\n           \"m\":16, \n        },\n    ]\n    \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    \n    ","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:02:14.249811Z","iopub.execute_input":"2024-05-03T22:02:14.250412Z","iopub.status.idle":"2024-05-03T22:02:14.265366Z","shell.execute_reply.started":"2024-05-03T22:02:14.250380Z","shell.execute_reply":"2024-05-03T22:02:14.264527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def eeg_from_parquet(\n    parquet_path: str, display: bool = False, seq_length=CFG.seq_length\n) -> np.ndarray:\n    eeg = pd.read_parquet(parquet_path, columns=CFG.eeg_features)\n    rows = len(eeg)\n\n    offset = (rows - CFG.nsamples) // 2\n\n    eeg = eeg.iloc[offset : offset + CFG.nsamples]\n\n    if display:\n        plt.figure(figsize=(10, 5))\n        offset = 0\n\n    data = np.zeros((CFG.nsamples, len(CFG.eeg_features)))\n\n    for index, feature in enumerate(CFG.eeg_features):\n        x = eeg[feature].values.astype(\"float32\")\n\n\n        mean = np.nanmean(x)\n        nan_percentage = np.isnan(x).mean()\n\n        if nan_percentage < 1: \n            x = np.nan_to_num(x, nan=mean)\n        else:\n            x[:] = 0\n        data[:, index] = x\n\n        if display:\n            if index != 0:\n                offset += x.max()\n            plt.plot(range(CFG.nsamples), x - offset, label=feature)\n            offset -= x.min()\n\n    if display:\n        plt.legend()\n        name = parquet_path.split(\"/\")[-1].split(\".\")[0]\n        plt.yticks([])\n        plt.title(f\"EEG {name}\", size=16)\n        plt.show()\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:02:14.266300Z","iopub.execute_input":"2024-05-03T22:02:14.266602Z","iopub.status.idle":"2024-05-03T22:02:14.279696Z","shell.execute_reply.started":"2024-05-03T22:02:14.266579Z","shell.execute_reply":"2024-05-03T22:02:14.278719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from scipy import optimize\nfrom scipy.special import factorial, gamma, hermitenorm\nfrom timm.models.layers import conv2d_same\n\n## From https://github.com/tomrunia/PyTorchWavelets/blob/master/wavelets_pytorch/wavelets.py\nclass Morlet(object):\n    def __init__(self, w0=6):\n        \"\"\"w0 is the nondimensional frequency constant. If this is\n        set too low then the wavelet does not sample very well: a\n        value over 5 should be ok; Terrence and Compo set it to 6.\n        \"\"\"\n        self.w0 = w0\n        if w0 == 6:\n            # value of C_d from TC98\n            self.C_d = 0.776\n\n    def __call__(self, *args, **kwargs):\n        return self.time(*args, **kwargs)\n\n    def time(self, t, s=1.0, complete=True):\n        \"\"\"\n        Complex Morlet wavelet, centred at zero.\n        Parameters\n        ----------\n        t : float\n            Time. If s is not specified, this can be used as the\n            non-dimensional time t/s.\n        s : float\n            Scaling factor. Default is 1.\n        complete : bool\n            Whether to use the complete or the standard version.\n        Returns\n        -------\n        out : complex\n            Value of the Morlet wavelet at the given time\n        See Also\n        --------\n        scipy.signal.gausspulse\n        Notes\n        -----\n        The standard version::\n            pi**-0.25 * exp(1j*w*x) * exp(-0.5*(x**2))\n        This commonly used wavelet is often referred to simply as the\n        Morlet wavelet.  Note that this simplified version can cause\n        admissibility problems at low values of `w`.\n        The complete version::\n            pi**-0.25 * (exp(1j*w*x) - exp(-0.5*(w**2))) * exp(-0.5*(x**2))\n        The complete version of the Morlet wavelet, with a correction\n        term to improve admissibility. For `w` greater than 5, the\n        correction term is negligible.\n        Note that the energy of the return wavelet is not normalised\n        according to `s`.\n        The fundamental frequency of this wavelet in Hz is given\n        by ``f = 2*s*w*r / M`` where r is the sampling rate.\n        \"\"\"\n        w = self.w0\n\n        x = t / s\n\n        output = np.exp(1j * w * x)\n\n        if complete:\n            output -= np.exp(-0.5 * (w ** 2))\n\n        output *= np.exp(-0.5 * (x ** 2)) * np.pi ** (-0.25)\n\n        return output\n\n    # Fourier wavelengths\n    def fourier_period(self, s):\n        \"\"\"Equivalent Fourier period of Morlet\"\"\"\n        return 4 * np.pi * s / (self.w0 + (2 + self.w0 ** 2) ** 0.5)\n\n    def scale_from_period(self, period):\n        \"\"\"\n        Compute the scale from the fourier period.\n        Returns the scale\n        \"\"\"\n        # Solve 4 * np.pi * scale / (w0 + (2 + w0 ** 2) ** .5)\n        #  for s to obtain this formula\n        coeff = np.sqrt(self.w0 * self.w0 + 2)\n        return (period * (coeff + self.w0)) / (4.0 * np.pi)\n\n    # Frequency representation\n    def frequency(self, w, s=1.0):\n        \"\"\"Frequency representation of Morlet.\n        Parameters\n        ----------\n        w : float\n            Angular frequency. If `s` is not specified, i.e. set to 1,\n            this can be used as the non-dimensional angular\n            frequency w * s.\n        s : float\n            Scaling factor. Default is 1.\n        Returns\n        -------\n        out : complex\n            Value of the Morlet wavelet at the given frequency\n        \"\"\"\n        x = w * s\n        # Heaviside mock\n        Hw = np.array(w)\n        Hw[w <= 0] = 0\n        Hw[w > 0] = 1\n        return np.pi ** -0.25 * Hw * np.exp((-((x - self.w0) ** 2)) / 2)\n\n    def coi(self, s):\n        \"\"\"The e folding time for the autocorrelation of wavelet\n        power at each scale, i.e. the timescale over which an edge\n        effect decays by a factor of 1/e^2.\n        This can be worked out analytically by solving\n            |Y_0(T)|^2 / |Y_0(0)|^2 = 1 / e^2\n        \"\"\"\n        return 2 ** 0.5 * s\n    \nclass Paul(object):\n    def __init__(self, m=4):\n        \"\"\"Initialise a Paul wavelet function of order `m`.\"\"\"\n        self.m = m\n\n    def __call__(self, *args, **kwargs):\n        return self.time(*args, **kwargs)\n\n    def time(self, t, s=1.0):\n        \"\"\"\n        Complex Paul wavelet, centred at zero.\n        Parameters\n        ----------\n        t : float\n            Time. If `s` is not specified, i.e. set to 1, this can be\n            used as the non-dimensional time t/s.\n        s : float\n            Scaling factor. Default is 1.\n        Returns\n        -------\n        out : complex\n            Value of the Paul wavelet at the given time\n        The Paul wavelet is defined (in time) as::\n            (2 ** m * i ** m * m!) / (pi * (2 * m)!) \\\n                    * (1 - i * t / s) ** -(m + 1)\n        \"\"\"\n        m = self.m\n        x = t / s\n\n        const = (2 ** m * 1j ** m * factorial(m)) / (np.pi * factorial(2 * m)) ** 0.5\n        functional_form = (1 - 1j * x) ** -(m + 1)\n\n        output = const * functional_form\n\n        return output\n\n    # Fourier wavelengths\n    def fourier_period(self, s):\n        \"\"\"Equivalent Fourier period of Paul\"\"\"\n        return 4 * np.pi * s / (2 * self.m + 1)\n\n    def scale_from_period(self, period):\n        raise NotImplementedError()\n\n    # Frequency representation\n    def frequency(self, w, s=1.0):\n        \"\"\"Frequency representation of Paul.\n        Parameters\n        ----------\n        w : float\n            Angular frequency. If `s` is not specified, i.e. set to 1,\n            this can be used as the non-dimensional angular\n            frequency w * s.\n        s : float\n            Scaling factor. Default is 1.\n        Returns\n        -------\n        out : complex\n            Value of the Paul wavelet at the given frequency\n        \"\"\"\n        m = self.m\n        x = w * s\n        # Heaviside mock\n        Hw = 0.5 * (np.sign(x) + 1)\n\n        # prefactor\n        const = 2 ** m / (m * factorial(2 * m - 1)) ** 0.5\n\n        functional_form = Hw * (x) ** m * np.exp(-x)\n\n        output = const * functional_form\n\n        return output\n\n    def coi(self, s):\n        \"\"\"The e folding time for the autocorrelation of wavelet\n        power at each scale, i.e. the timescale over which an edge\n        effect decays by a factor of 1/e^2.\n        This can be worked out analytically by solving\n            |Y_0(T)|^2 / |Y_0(0)|^2 = 1 / e^2\n        \"\"\"\n        return s / 2 ** 0.5\n    \nclass DOG(object):\n    def __init__(self, m=2):\n        \"\"\"Initialise a Derivative of Gaussian wavelet of order `m`.\"\"\"\n        if m == 2:\n            # value of C_d from TC98\n            self.C_d = 3.541\n        elif m == 6:\n            self.C_d = 1.966\n        else:\n            pass\n        self.m = m\n\n    def __call__(self, *args, **kwargs):\n        return self.time(*args, **kwargs)\n\n    def time(self, t, s=1.0):\n        \"\"\"\n        Return a Derivative of Gaussian wavelet,\n        When m = 2, this is also known as the \"Mexican hat\", \"Marr\"\n        or \"Ricker\" wavelet.\n        It models the function::\n            ``A d^m/dx^m exp(-x^2 / 2)``,\n        where ``A = (-1)^(m+1) / (gamma(m + 1/2))^.5``\n        and   ``x = t / s``.\n        Note that the energy of the return wavelet is not normalised\n        according to `s`.\n        Parameters\n        ----------\n        t : float\n            Time. If `s` is not specified, this can be used as the\n            non-dimensional time t/s.\n        s : scalar\n            Width parameter of the wavelet.\n        Returns\n        -------\n        out : float\n            Value of the DOG wavelet at the given time\n        Notes\n        -----\n        The derivative of the Gaussian has a polynomial representation:\n        from http://en.wikipedia.org/wiki/Gaussian_function:\n        \"Mathematically, the derivatives of the Gaussian function can be\n        represented using Hermite functions. The n-th derivative of the\n        Gaussian is the Gaussian function itself multiplied by the n-th\n        Hermite polynomial, up to scale.\"\n        http://en.wikipedia.org/wiki/Hermite_polynomial\n        Here, we want the 'probabilists' Hermite polynomial (He_n),\n        which is computed by scipy.special.hermitenorm\n        \"\"\"\n        x = t / s\n        m = self.m\n\n        # compute the Hermite polynomial (used to evaluate the\n        # derivative of a Gaussian)\n        He_n = hermitenorm(m)\n        # gamma = scipy.special.gamma\n\n        const = (-1) ** (m + 1) / gamma(m + 0.5) ** 0.5\n        function = He_n(x) * np.exp(-(x ** 2) / 2) * np.exp(-1j * x)\n\n        return const * function\n\n    def fourier_period(self, s):\n        \"\"\"Equivalent Fourier period of derivative of Gaussian\"\"\"\n        return 2 * np.pi * s / (self.m + 0.5) ** 0.5\n\n    def scale_from_period(self, period):\n        raise NotImplementedError()\n\n    def frequency(self, w, s=1.0):\n        \"\"\"Frequency representation of derivative of Gaussian.\n        Parameters\n        ----------\n        w : float\n            Angular frequency. If `s` is not specified, i.e. set to 1,\n            this can be used as the non-dimensional angular\n            frequency w * s.\n        s : float\n            Scaling factor. Default is 1.\n        Returns\n        -------\n        out : complex\n            Value of the derivative of Gaussian wavelet at the\n            given time\n        \"\"\"\n        m = self.m\n        x = s * w\n        # gamma = scipy.special.gamma\n        const = -(1j ** m) / gamma(m + 0.5) ** 0.5\n        function = x ** m * np.exp(-(x ** 2) / 2)\n        return const * function\n\n    def coi(self, s):\n        \"\"\"The e folding time for the autocorrelation of wavelet\n        power at each scale, i.e. the timescale over which an edge\n        effect decays by a factor of 1/e^2.\n        This can be worked out analytically by solving\n            |Y_0(T)|^2 / |Y_0(0)|^2 = 1 / e^2\n        \"\"\"\n        return 2 ** 0.5 * s\n\n\nclass Ricker(DOG):\n    def __init__(self):\n        \"\"\"The Ricker, aka Marr / Mexican Hat, wavelet is a\n        derivative of Gaussian order 2.\n        \"\"\"\n        DOG.__init__(self, m=2)\n        # value of C_d from TC98\n        self.C_d = 3.541\n    \nclass CWT2(nn.Module):\n    def __init__(\n        self,\n        dj=0.0625,\n        dt=1 / 2048,\n        wavelet=Morlet(),\n        fmin: int = 20,\n        fmax: int = 500,\n        output_format=\"Magnitude\",\n        trainable=False,\n        hop_length: int = 1,\n    ):\n        super().__init__()\n        self.wavelet = wavelet\n\n        self.dt = dt\n        self.dj = dj\n        self.fmin = fmin\n        self.fmax = fmax\n        self.output_format = output_format\n        self.trainable = trainable  # TODO make kernel a trainable parameter\n        self.stride = (1, hop_length)\n        # self.padding = 0  # \"same\"\n\n        self._scale_minimum = self.compute_minimum_scale()\n\n        self.signal_length = None\n        self._channels = None\n\n        self._scales = None\n        self._kernel = None\n        self._kernel_real = None\n        self._kernel_imag = None\n\n    def compute_optimal_scales(self):\n        \"\"\"\n        Determines the optimal scale distribution (see. Torrence & Combo, Eq. 9-10).\n        :return: np.ndarray, collection of scales\n        \"\"\"\n        if self.signal_length is None:\n            raise ValueError(\n                \"Please specify signal_length before computing optimal scales.\"\n            )\n        J = int(\n            (1 / self.dj) * np.log2(self.signal_length * self.dt / self._scale_minimum)\n        )\n        scales = self._scale_minimum * 2 ** (self.dj * np.arange(0, J + 1))\n\n        # Remove high and low frequencies\n        frequencies = np.array([1 / self.wavelet.fourier_period(s) for s in scales])\n        if self.fmin:\n            frequencies = frequencies[frequencies >= self.fmin]\n            scales = scales[0 : len(frequencies)]\n        if self.fmax:\n            frequencies = frequencies[frequencies <= self.fmax]\n            scales = scales[len(scales) - len(frequencies) : len(scales)]\n\n        return scales\n\n    def compute_minimum_scale(self):\n        \"\"\"\n        Choose s0 so that the equivalent Fourier period is 2 * dt.\n        See Torrence & Combo Sections 3f and 3h.\n        :return: float, minimum scale level\n        \"\"\"\n        dt = self.dt\n\n        def func_to_solve(s):\n            return self.wavelet.fourier_period(s) - 2 * dt\n\n        return optimize.fsolve(func_to_solve, 1)[0]\n\n    def _build_filters(self):\n        self._filters = []\n        for scale_idx, scale in enumerate(self._scales):\n            # Number of points needed to capture wavelet\n            M = 10 * scale / self.dt\n            # Times to use, centred at zero\n            t = torch.arange((-M + 1) / 2.0, (M + 1) / 2.0) * self.dt\n            if len(t) % 2 == 0:\n                t = t[0:-1]  # requires odd filter size\n            # Sample wavelet and normalise\n            norm = (self.dt / scale) ** 0.5\n            filter_ = norm * self.wavelet(t, scale)\n            self._filters.append(torch.conj(torch.flip(filter_, [-1])))\n\n        self._pad_filters()\n\n    def _pad_filters(self):\n        filter_len = self._filters[-1].shape[0]\n        padded_filters = []\n\n        for f in self._filters:\n            pad = (filter_len - f.shape[0]) // 2\n            padded_filters.append(nn.functional.pad(f, (pad, pad)))\n\n        self._filters = padded_filters\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        self._build_filters()\n        wavelet_bank = torch.stack(self._filters)\n        wavelet_bank = wavelet_bank.view(\n            wavelet_bank.shape[0], 1, 1, wavelet_bank.shape[1]\n        )\n        # See comment by tez6c32\n        # https://www.kaggle.com/anjum48/continuous-wavelet-transform-cwt-in-pytorch/comments#1499878\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        if self.signal_length is None:\n            self.signal_length = x.shape[-1]\n            self.channels = x.shape[-2]\n            self._scales = self.compute_optimal_scales()\n            self._kernel = self._build_wavelet_bank()\n\n            if self._kernel.is_complex():\n                self._kernel_real = self._kernel.real\n                self._kernel_imag = self._kernel.imag\n\n        x = x.unsqueeze(1)\n\n        if self._kernel.is_complex():\n            if (\n                x.dtype != self._kernel_real.dtype\n                or x.device != self._kernel_real.device\n            ):\n                self._kernel_real = self._kernel_real.to(device=x.device, dtype=x.dtype)\n                self._kernel_imag = self._kernel_imag.to(device=x.device, dtype=x.dtype)\n\n            # Strides > 1 not yet supported for \"same\" padding\n            # output_real = nn.functional.conv2d(\n            #     x, self._kernel_real, padding=self.padding, stride=self.stride\n            # )\n            # output_imag = nn.functional.conv2d(\n            #     x, self._kernel_imag, padding=self.padding, stride=self.stride\n            # )\n            output_real = conv2d_same(x, self._kernel_real, stride=self.stride)\n            output_imag = conv2d_same(x, self._kernel_imag, stride=self.stride)\n            output_real = torch.transpose(output_real, 1, 2)\n            output_imag = torch.transpose(output_imag, 1, 2)\n\n            if self.output_format == \"Magnitude\":\n                return torch.sqrt(output_real ** 2 + output_imag ** 2)\n            else:\n                return torch.stack([output_real, output_imag], -1)\n\n        else:\n            if x.device != self._kernel.device or x.dtype != self._kernel.dtype:\n                self._kernel = self._kernel.to(device=x.device, dtype=x.dtype)\n\n            # output = nn.functional.conv2d(\n            #     x, self._kernel, padding=self.padding, stride=self.stride\n            # )\n            output = conv2d_same(x, self._kernel, stride=self.stride)\n            return torch.transpose(output, 1, 2)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:02:14.281134Z","iopub.execute_input":"2024-05-03T22:02:14.281482Z","iopub.status.idle":"2024-05-03T22:02:14.337073Z","shell.execute_reply.started":"2024-05-03T22:02:14.281458Z","shell.execute_reply":"2024-05-03T22:02:14.336313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchaudio.transforms as T\nimport torch.nn.functional as F\n                    \n\nclass HMSDataset_preprocesser(torch.utils.data.Dataset):\n    def __init__(self, df, eegs,\n                 bandpass_filter = None,\n                 mode='test',\n                 transforms = None,\n                 out_samples = 5000,\n                 input_type='paul',\n                 tta_offset=0,\n                 stride = 4,\n                 m=4,\n                ):\n        self.df = df\n        self.eegs = eegs\n        self.bandpass_filter = bandpass_filter\n        self.mode = mode\n        self.transforms = transforms\n        self.out_samples = out_samples\n        self.stride = stride\n        self.tta_offset = tta_offset\n        self.L = 50 * CFG.sampling_rate\n        self.m = m\n        if input_type == 'morl':\n            self.cwt2 = CWT2(\n                #dj=0.0625,\n                #dj=0.098,\n                dj=0.125,\n                dt=1 / 200,\n                wavelet=Morlet(),\n                fmin = 0.5,\n                fmax = 40,\n                output_format=\"Magnitude\",\n                trainable=False,\n                hop_length=self.stride,\n            ) \n        elif input_type == 'dog':\n            self.cwt2 = CWT2(\n                #dj=0.0625,\n                #dj=0.098,\n                dj=0.125,\n                dt=1 / 200,\n                wavelet=DOG(m=self.m),\n                fmin = 0.5,\n                fmax = 40,\n                output_format=\"Magnitude\",\n                trainable=False,\n                hop_length=self.stride,\n            )\n        elif input_type == 'paul':\n            self.cwt2 = CWT2(\n                #dj=0.0625,\n                #dj=0.098,\n                dj=0.125,\n                dt=1 / 200,\n                wavelet=Paul(m=self.m),\n                fmin = 0.5,\n                fmax = 40,\n                output_format=\"Magnitude\",\n                trainable=False,\n                hop_length=self.stride,\n            )\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx: int):\n        \n        X, y = self.__data_generation(idx)\n                       \n        if self.transforms is not None:\n            X = self.transforms(image=X)['image']\n        \n        return {\"data\": X, \"idx\": idx}\n    \n    def __data_generation(self, index):\n        X = np.zeros((int(self.out_samples), CFG.in_channels), dtype=\"float32\")  # Size=(10000, 14)\n\n         # eeg_idを指定\n        row = self.df.iloc[index]\n        data = self.eegs[row.eeg_id]\n\n        if CFG.nsamples != self.out_samples:\n            offset = (CFG.nsamples - self.out_samples) // 2 + self.tta_offset\n            data = data[offset:offset+self.out_samples,:]\n                    \n        # ######## EEG ########\n\n        # diff of eegs\n        for i, (feat_a, feat_b) in enumerate(CFG.map_features):\n                \n            diff_feat = data[:, CFG.feature_to_index[feat_a]] - data[:, CFG.feature_to_index[feat_b]]\n\n            if not self.bandpass_filter is None:\n                diff_feat = butter_bandpass_filter(\n                    diff_feat,\n                    self.bandpass_filter[\"low\"],\n                    self.bandpass_filter[\"high\"],\n                    CFG.sampling_rate,\n                    order=self.bandpass_filter[\"order\"],\n                )\n            \n            X[:, i] = diff_feat\n\n        X = np.clip(X, -1024, 1024)\n        X = np.nan_to_num(X, nan=0) / 32.0\n        \n        X_in = torch.tensor(X.transpose(1, 0), dtype=torch.float32).unsqueeze(0)\n        \n        X_cwt = self.cwt2(X_in)[0]\n        X_cwt = np.concatenate([X_cwt[i] for i in range(X_cwt.shape[0])], axis=0)\n        \n        y = np.zeros(CFG.n_classes, dtype=\"float32\")  # Size=(6,)\n        if self.mode != \"test\":\n            y = row[CFG.classes].values.astype(np.float32)\n\n        return X_cwt.astype(\"float32\"), y","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:02:14.338809Z","iopub.execute_input":"2024-05-03T22:02:14.339077Z","iopub.status.idle":"2024-05-03T22:02:14.358039Z","shell.execute_reply.started":"2024-05-03T22:02:14.339054Z","shell.execute_reply":"2024-05-03T22:02:14.357126Z"},"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\ntest_eeg_parquet_paths = glob(CFG.test_eeg + \"*.parquet\")\ntest_eeg_df = pd.read_parquet(test_eeg_parquet_paths[0])\ntest_eeg_features = test_eeg_df.columns\nprint(f\"There are {len(test_eeg_features)} raw eeg features\")\nprint(list(test_eeg_features))\ndel test_eeg_df\n_ = gc.collect()\n\n# %%time\nall_eegs = {}\neeg_ids = test_df.eeg_id.unique()\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","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:02:14.359116Z","iopub.execute_input":"2024-05-03T22:02:14.359475Z","iopub.status.idle":"2024-05-03T22:02:14.627020Z","shell.execute_reply.started":"2024-05-03T22:02:14.359444Z","shell.execute_reply":"2024-05-03T22:02:14.626049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchaudio.transforms as T\nimport torch.nn.functional as F\n\nclass HMSDataset(torch.utils.data.Dataset):\n    def __init__(self, df,\n                 mode='test',\n                 transforms = None,\n                ):\n        self.df = df\n        self.mode = mode\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx: int):\n        \n        X, y = self.__data_generation(idx)\n                       \n        if self.transforms is not None:\n            X = self.transforms(image=X)['image']\n        \n        return {\"data\": X, \"label\": y}\n    \n    def __data_generation(self, index):\n\n         # eeg_idを指定\n        row = self.df.iloc[index]\n        \n        X_cwt = np.load(row.spec_path)\n        y = np.zeros(CFG.n_classes, dtype=\"float32\")  # Size=(6,)\n        if self.mode != \"test\":\n            y = row[CFG.classes].values.astype(np.float32)\n\n        return X_cwt.astype(\"float32\"), y","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:02:14.628042Z","iopub.execute_input":"2024-05-03T22:02:14.628338Z","iopub.status.idle":"2024-05-03T22:02:14.637483Z","shell.execute_reply.started":"2024-05-03T22:02:14.628308Z","shell.execute_reply":"2024-05-03T22:02:14.636660Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HMSModel(nn.Module):\n    def __init__(self, model_name: str, pretrained: bool, in_channels: int, num_classes: int):\n        super().__init__()\n        \n        self.model = timm.create_model(\n            model_name=model_name, \n            pretrained=pretrained, \n            in_chans=in_channels,\n            drop_rate = 0.1,\n            drop_path_rate = 0.2,\n            num_classes = num_classes\n        )\n        \n    def forward(self, x):\n        x = self.model(x)  \n        return x\n\ndef inference_function(test_loader, model, device):\n    model.eval() \n#     softmax = nn.Softmax(dim=1)\n    prediction_dict = {}\n    preds = []\n    with tqdm(test_loader, unit=\"test_batch\", desc=\"Inference\") as tqdm_test_loader:\n        for step, batch in enumerate(tqdm_test_loader):\n            X = batch.pop(\"data\").to(device, dtype=torch.float)\n            batch_size = X.size(0)\n            with torch.no_grad():\n                y_preds = model(X)\n#             y_preds = softmax(y_preds)\n            y_preds = entmax_bisect(y_preds, alpha=1.03, dim=1)\n            preds.append(y_preds.to(\"cpu\").numpy())\n\n    prediction_dict[\"predictions\"] = np.concatenate(\n        preds\n    )\n    return prediction_dict","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:02:14.638528Z","iopub.execute_input":"2024-05-03T22:02:14.638787Z","iopub.status.idle":"2024-05-03T22:02:14.648311Z","shell.execute_reply.started":"2024-05-03T22:02:14.638755Z","shell.execute_reply":"2024-05-03T22:02:14.647569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_preds = []\n\nfor ver in CFG.VERSION:\n    \n    data_transforms = {\n        \"test\": A.Compose([\n            A.Resize(height=ver[\"input_size\"][0], width=ver[\"input_size\"][1]),\n            ToTensorV2()], p=1.),\n    }\n    \n    predictions = []\n    for offset in ver['tta_offset']:\n        print(f\"TTA offset: {offset}\")\n        #### pre CWT #####\n        cwt_path = \"/kaggle/working/cwt_paul/\"\n        os.makedirs(cwt_path, exist_ok=True)\n\n        test_dataset = HMSDataset_preprocesser(\n                    df=test_df,\n                    mode=\"test\",\n                    eegs=all_eegs,\n                    bandpass_filter=ver[\"bandpass_filter\"],\n                    transforms=data_transforms[\"test\"],\n                    out_samples=ver[\"out_samples\"],\n                    input_type=ver[\"input_type\"],\n                    stride=ver['stride'],\n                    tta_offset=offset,\n                    m=ver['m'],\n                )\n\n        pre_loader = DataLoader(\n            test_dataset,\n            batch_size=4,\n            shuffle=False,\n            num_workers=4, pin_memory=True, drop_last=False\n        )\n\n        for data in tqdm(pre_loader):\n            for idx, cwtidx in enumerate(data[\"idx\"].numpy()):\n                np.save(f\"{cwt_path}cwt_all_{cwtidx:06}.npy\", data[\"data\"][idx].numpy())\n                \n        del test_dataset, pre_loader\n\n        test_df['spec_path'] = test_df.index.map(lambda x: f\"{cwt_path}cwt_all_{x:06}.npy\")\n        #### pre CWT ####\n        \n        test_dataset = HMSDataset(\n            df=test_df,\n            mode=\"test\",\n            transforms=None,\n        )\n\n        if len(predictions) == 0:\n            output = test_dataset[0]\n            X = output[\"data\"]\n            print(f\"X shape: {X.shape}\")\n\n        test_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        )\n\n        model = HMSModel(\n            model_name = ver['model'], \n            pretrained = False,\n            in_channels = ver['in_chan'], \n            num_classes = CFG.n_classes\n        )\n\n        for weight_model_file in glob(ver['file_mask']):\n            print(f\"{weight_model_file=}\")\n            checkpoint = torch.load(weight_model_file, map_location=CFG.device)\n            model.load_state_dict(checkpoint)\n            model.to(CFG.device)\n            prediction_dict = inference_function(test_loader, model, CFG.device)\n            predict = prediction_dict[\"predictions\"]\n            predictions.append(predict)\n            torch.cuda.empty_cache()\n            gc.collect()\n            \n        del model\n        torch.cuda.empty_cache()\n        gc.collect()\n            \n        shutil.rmtree(cwt_path)\n                \n    predictions = np.array(predictions)\n    predictions = np.mean(predictions, axis=0)\n    \n    tmp = pd.DataFrame({\"eeg_id\": test_df.eeg_id.values})\n    tmp[CFG.classes] = predictions\n    display(tmp.head())\n\n    all_preds.append(predictions)\n\nf_predictions = all_preds.copy()\n    ","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:02:14.649395Z","iopub.execute_input":"2024-05-03T22:02:14.649653Z","iopub.status.idle":"2024-05-03T22:03:30.242556Z","shell.execute_reply.started":"2024-05-03T22:02:14.649631Z","shell.execute_reply":"2024-05-03T22:03:30.241576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Hill Climb & Non-negative Regression","metadata":{}},{"cell_type":"code","source":"oof_list = [\n #yamash model\n '/kaggle/input/hms-weights-yamash/022_023/oof2_vote10.csv',\n '/kaggle/input/hms-weights-yamash/026_001_v2/oof2_vote10.csv',\n '/kaggle/input/hms-weights-yamash/026_016/oof2_vote10.csv',\n '/kaggle/input/hms-weights-yamash/026_017/oof2_vote10.csv',\n '/kaggle/input/hms-weights-yamash/030_002/oof2_vote10.csv',\n    \n #sugupoko model 50sec\n#  '/kaggle/input/hms-cwt-weights-r02/exp05-71-6_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/oof2_vote10.csv',\n#  '/kaggle/input/hms-cwt-weights-r02/exp05-71-7_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/oof2_vote10.csv',\n #sugupoko model 50sec　0.5~40Hz\n '/kaggle/input/hms-cwt-weights-r02/exp05-78-2_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/oof2_vote10.csv',\n '/kaggle/input/hms-cwt-weights-r02/exp05-78-4_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/oof2_vote10.csv',\n '/kaggle/input/hms-cwt-weights-r02/exp05-78-5_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_labelsmoothingOff/oof2_vote10.csv',\n '/kaggle/input/hms-cwt-weights-r02/exp05-90_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_addData/oof2_vote10.csv',\n # sugupoko 50sec 0.5~40hz 1d and 2d model\n '/kaggle/input/hms-cwt-weights-r02/exp05-91_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_1dAnd2d/oof2_vote10.csv',\n '/kaggle/input/hms-cwt-weights-r02/exp05-91-2_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_1dAnd2d/oof2_vote10.csv',\n '/kaggle/input/hms-cwt-weights-r02/exp05-91-3_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_1dAnd2d/oof2_vote10.csv',\n '/kaggle/input/hms-cwt-weights-r02/exp05-91-5_AMP_cwt_50sec_stride16_18band_maxvit_base_tf_512_1dAnd2d/oof2_vote10.csv',\n    \n # saito model\n '/kaggle/input/hms-weights-saito/hms-weights-saito/1dcnn_v062_stage2/1dcnn_v062_stage2/submission_yamash.csv',\n '/kaggle/input/hms-weights-saito/hms-weights-saito/1dcnn_v061_stage2/1dcnn_v061_stage2/submission_yamash.csv',\n '/kaggle/input/hms-weights-saito/hms-weights-saito/1dcnn_v060_stage2/1dcnn_v060_stage2/submission_yamash.csv',\n '/kaggle/input/hms-weights-saito/hms-weights-saito/1dcnn_v059_stage2/1dcnn_v059_stage2/submission_yamash.csv',\n '/kaggle/input/hms-weights-saito/hms-weights-saito/1dcnn_v057_stage2/1dcnn_v057_stage2/submission_yamash.csv',\n '/kaggle/input/hms-weights-saito/hms-weights-saito/1dcnn_v056_stage2/1dcnn_v056_stage2/submission_yamash.csv',\n '/kaggle/input/hms-weights-saito/hms-weights-saito/1dcnn_v055_stage2/1dcnn_v055_stage2/submission_yamash.csv',\n '/kaggle/input/hms-weights-saito/hms-weights-saito/1dcnn_v054_stage2/1dcnn_v054_stage2/submission_yamash.csv',\n '/kaggle/input/hms-weights-saito/hms-weights-saito/1dcnn_v053_stage2/1dcnn_v053_stage2/submission_yamash.csv',\n '/kaggle/input/hms-weights-saito/hms-weights-saito/1dcnn_v051_stage2/1dcnn_v051_stage2/submission_yamash.csv',\n '/kaggle/input/hms-weights-saito/hms-weights-saito/1dcnn_v048_stage2/1dcnn_v048_stage2/submission_yamash.csv',\n '/kaggle/input/hms-weights-saito/hms-weights-saito/1dcnn_v047_stage2/1dcnn_v047_stage2/submission_yamash.csv',\n '/kaggle/input/hms-weights-saito/hms-weights-saito/1dcnn_v046_2_stage2/1dcnn_v046_2_stage2/submission_yamash.csv',\n '/kaggle/input/hms-weights-saito/hms-weights-saito/1dcnn_v042_stage2/1dcnn_v042_stage2/submission_yamash.csv',\n\n # fujii model\n '/kaggle/input/hms-weights-fujii/done_51_20240403_1841/oof2_vote10.csv',\n '/kaggle/input/hms-weights-fujii/done_52_20240403_1829/oof2_vote10.csv',\n '/kaggle/input/hms-weights-fujii/done_71_20240406_2039/oof2_vote10.csv',\n]","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:30.243954Z","iopub.execute_input":"2024-05-03T22:03:30.244272Z","iopub.status.idle":"2024-05-03T22:03:30.252168Z","shell.execute_reply.started":"2024-05-03T22:03:30.244247Z","shell.execute_reply":"2024-05-03T22:03:30.251255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(oof_list), \":\")\nprint(len(y_1d_predictions))\n#print(len(k_cwt50sec_predictions))\nprint(len(k_cwt50sec_to40Hz_predictions))\nprint(len(k_cwt50sec_to40Hz_and_1d_predictions))\nprint(len(s_predictions))\nprint(len(f_predictions))\nprint(len(y_1d_predictions) + len(k_cwt50sec_to40Hz_predictions) + len(k_cwt50sec_to40Hz_and_1d_predictions)+ len(s_predictions) + len(f_predictions))","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:30.253226Z","iopub.execute_input":"2024-05-03T22:03:30.253496Z","iopub.status.idle":"2024-05-03T22:03:30.267329Z","shell.execute_reply.started":"2024-05-03T22:03:30.253473Z","shell.execute_reply":"2024-05-03T22:03:30.266475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess_oof(df):\n    if \"seizure_vote_pred\" not in df.columns:\n        df = df.drop(\"Unnamed: 0\", axis=1)\n        return df\n    df= df.sort_values(\"eeg_id\")\n    df = df.drop([\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"], axis=1)\n    df = df.rename(columns={\"seizure_vote_pred\": \"seizure_vote\", \n                            \"lpd_vote_pred\": \"lpd_vote\",\n                            \"gpd_vote_pred\": \"gpd_vote\",\n                            \"lrda_vote_pred\": \"lrda_vote\",\n                            \"grda_vote_pred\": \"grda_vote\",\n                            \"other_vote_pred\": \"other_vote\"})\n    df = df.merge(label_id_df[[\"eeg_id\", \"label_id\"]], on=\"eeg_id\")\n    df = df.drop(\"eeg_id\", axis=1)\n    df[classes] /= df[classes].values.sum(axis=1, keepdims=True)\n    return df","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:30.268299Z","iopub.execute_input":"2024-05-03T22:03:30.268569Z","iopub.status.idle":"2024-05-03T22:03:30.276647Z","shell.execute_reply.started":"2024-05-03T22:03:30.268547Z","shell.execute_reply":"2024-05-03T22:03:30.275879Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"classes = [\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]\n\ntrain_df = pd.read_csv(\"/kaggle/input/t1-train-groupkfold/t1_train_groupkfold.csv\")\n\ntrain_df[\"num_votes\"] = train_df[classes].sum(axis=1)\n\ndf = train_df.groupby('eeg_id')[['spectrogram_id','spectrogram_label_offset_seconds']].agg({\n    'spectrogram_id':'first',\n    'spectrogram_label_offset_seconds':'min'\n})\ndf.columns = ['spectrogram_id','min']\n\naux = train_df.groupby('eeg_id')[['spectrogram_id','spectrogram_label_offset_seconds']].agg({\n    'spectrogram_label_offset_seconds':'max'\n})\ndf['max'] = aux\n\naux = train_df.groupby('eeg_id')[['patient_id']].agg('first')\ndf['patient_id'] = aux\n\naux = train_df.groupby('eeg_id')[classes].agg('sum')\nfor label in classes:\n    df[label] = aux[label].values\n    \ny_data = df[classes].values\ny_data = y_data / y_data.sum(axis=1,keepdims=True)\ndf[classes] = y_data\n\naux = train_df.groupby('eeg_id')[['expert_consensus']].agg('first')\ndf['target'] = aux\n\nfold = train_df.groupby('eeg_id')[['fold']].agg('first')\ndf['fold'] = fold\n\nlabel_id = train_df.groupby('eeg_id')[['label_id']].agg('first')\ndf['label_id'] = label_id\n\nnum_votes = train_df.groupby('eeg_id')[['num_votes']].agg('first')\ndf['num_votes'] = num_votes\n\ndf = df.reset_index()\n\ntrain = df\n\ntrue = train[[\"label_id\"] + classes].copy()\ntrue_vote10 = true.loc[train['num_votes'].values >=10].copy().reset_index(drop=True)\n\nlabel_id_df = train.loc[train['num_votes'].values >=10].copy().reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:30.277624Z","iopub.execute_input":"2024-05-03T22:03:30.277889Z","iopub.status.idle":"2024-05-03T22:03:30.649283Z","shell.execute_reply.started":"2024-05-03T22:03:30.277867Z","shell.execute_reply":"2024-05-03T22:03:30.648262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# csv読み込み + singleモデルでの性能表示\noof_name_list = []\ndf_list = []\nfor oof in tqdm(oof_list):\n    oof_name_list.append(oof.split(\"/\")[-2] + \"/\" + oof.split(\"/\")[-1])\n    df_list.append(preprocess_oof(pd.read_csv(oof).reset_index(drop=True)))\n    \nbest_single_cv = np.inf\nbest_model_idx = 0\nfor idx, (name, df) in enumerate(zip(oof_name_list, df_list)):\n    cv_score = score(solution=true_vote10.copy(), submission=df.copy(), row_id_column_name='label_id')\n    print(f'{name:60s}: {cv_score}')    \n    if cv_score < best_single_cv:\n        best_single_cv = cv_score\n        best_model_idx = idx\nprint(\"\\nBest single model index: \", best_model_idx, \", CV:\", best_single_cv)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:30.650458Z","iopub.execute_input":"2024-05-03T22:03:30.650751Z","iopub.status.idle":"2024-05-03T22:03:32.248427Z","shell.execute_reply.started":"2024-05-03T22:03:30.650726Z","shell.execute_reply":"2024-05-03T22:03:32.247479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Hill Climbing\n\n# SINGLE_BEST = score(solution=true_vote10.copy(), submission=df_list[best_model_idx].copy(), row_id_column_name='label_id')\n# print(\"1 \", SINGLE_BEST)\n\n# N = len(df_list)  # num of models\n# hill_climb_coef_array = np.zeros((N)).astype(np.float32)\n    \n# selected_models = [\n#     {\n#         \"model_idx\": best_model_idx,\n#         \"current_score\": SINGLE_BEST,\n#         \"weight\": 1.0,\n#         \"name\": oof_name_list[best_model_idx]\n#     }\n# ]\n# add_model_info = {\n#     \"model_idx\": -1,\n#     \"current_score\": np.inf,\n#     \"weight\": 0.0,\n#     \"name\": \"\"\n# }\n# potential_new_best_cv_score = SINGLE_BEST\n# current_best_ensemble = df_list[best_model_idx][classes].values\n\n# selected_idx = [best_model_idx]\n# STOP = False\n\n# tmp_df = df_list[best_model_idx].copy()\n\n# search_range = [-0.5, 0.5]  # default:[0.0, 0.5]\n\n# while not STOP:\n#     add_model_info[\"model_idx\"] = -1\n#     add_model_info[\"current_score\"] = np.inf\n#     add_model_info[\"weight\"] = 0.0\n#     add_model_info[\"name\"] = \"\"\n#     saved_best_ensemble = None\n#     for k in tqdm(range(len(df_list))):\n#         for wgt in np.arange(search_range[0], search_range[1], 0.02):\n#             potential_ensemble = (1-wgt) * current_best_ensemble + wgt * df_list[k][classes].values\n        \n#             potential_ensemble = np.clip(potential_ensemble, 0.0, np.inf)\n#             potential_ensemble = potential_ensemble / potential_ensemble.sum(axis=1,keepdims=True)\n            \n#             tmp_df[classes] = potential_ensemble\n#             cv_score = score(solution=true_vote10.copy(), submission=tmp_df.copy(), row_id_column_name='label_id')\n#             if cv_score < potential_new_best_cv_score and k not in selected_idx:\n#                 potential_new_best_cv_score = cv_score\n#                 saved_best_ensemble = potential_ensemble.copy()\n#                 add_model_info[\"model_idx\"] = k\n#                 add_model_info[\"current_score\"] = potential_new_best_cv_score\n#                 add_model_info[\"weight\"] = wgt\n#                 add_model_info[\"name\"] = oof_name_list[k]\n#     if add_model_info[\"current_score\"] > potential_new_best_cv_score:\n#         STOP = True\n#     else:\n#         selected_idx.append(add_model_info[\"model_idx\"])\n#         current_best_ensemble = saved_best_ensemble\n#         for idx in range(len(selected_models)):\n#             selected_models[idx][\"weight\"] *= (1.0 - add_model_info[\"weight\"])\n#         selected_models.append(add_model_info.copy())\n#         print(len(selected_models), potential_new_best_cv_score)\n\n# import pprint\n# pprint.pprint(selected_models)\n\n# for m in selected_models:\n#     hill_climb_coef_array[m[\"model_idx\"]] = m[\"weight\"]\n\n# display(hill_climb_coef_array)\n \n# hill_climb_coef = np.zeros((6, N*6)).astype(np.float32)\n# for i, coef in enumerate(hill_climb_coef_array):\n#     print(coef)\n#     for k in range(6):\n#         hill_climb_coef[k, i*6+k] = coef\n# print(np.array2string(hill_climb_coef))\n\n\n# list_1 = oof_name_list\n# list_2 = selected_models\n\n# # Extract names from list_2 for easier comparison\n# names_in_list_2 = [d['name'] for d in list_2]\n\n# # Find elements in list_1 not in list_2\n# not_in_list_2 = [item for item in list_1 if item not in names_in_list_2]\n\n# print(\"not contributed model idx: \", not_in_list_2)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:32.249568Z","iopub.execute_input":"2024-05-03T22:03:32.249895Z","iopub.status.idle":"2024-05-03T22:03:32.257369Z","shell.execute_reply.started":"2024-05-03T22:03:32.249868Z","shell.execute_reply":"2024-05-03T22:03:32.256287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Non-negative regression\nfrom sklearn.linear_model import LinearRegression\n\nif False:\n    N = len(df_list)  # num of models\n    M = len(df_list[0])  # num of samples\n\n    X = np.zeros((M, N*6))\n    y = np.zeros((M, 6))\n\n    for i in range(6):\n        y[:, i] = true_vote10[classes[i]].values\n\n    for idx, d in enumerate(df_list):\n        for i in range(6):\n            X[:, idx*6+i] = d[classes[i]].values\n\n    df_pred = df_list[0].copy()\n    regr_list = []\n    for i in range(6):\n        regr = LinearRegression(positive=True, fit_intercept=False)\n        regr.fit(X, y[:,i])\n        df_pred[classes[i]] = regr.predict(X)\n        regr_list.append(regr.coef_)\n\n    regr_coef = np.vstack(regr_list)\n    print(regr_coef.shape)\n    print(np.array2string(regr_coef, separator=','))\n\n    y_data = df_pred[classes].values\n    y_data = y_data / y_data.sum(axis=1,keepdims=True)\n    df_pred[classes] = y_data\n\n    cv_score = score(solution=true_vote10.copy(), submission=df_pred.copy(), row_id_column_name='label_id')\n    print(cv_score)\n\n    regression_coef = regr_coef \n\n    # not contributed models\n    not_contributed = []\n    for i in range(N):\n        idx = [x for x in range(i*6, (i+1)*6)]\n        if np.sum(regr_coef[:, idx]) < 1e-6:\n            not_contributed.append(i)\n    print(\"not contributed model idx: \", not_contributed)\nelse:\n    N = len(df_list)  # num of models\n    M = len(df_list[0])  # num of samples\n    print(N)\n    X = np.zeros((M, N*6))\n    y = np.zeros((M, 6))\n    \n    regr_coef = np.zeros((6, N*6))\n\n    for i in range(6):\n        y[:, i] = true_vote10[classes[i]].values\n\n    for idx, d in enumerate(df_list):\n        for i in range(6):\n            X[:, idx*6+i] = d[classes[i]].values\n\n    df_pred = df_list[0].copy()\n    regr_list = []\n    for i in range(6):\n        regr = LinearRegression(positive=True, fit_intercept=False)\n        regr.fit(X[:,i::6], y[:,i])\n        df_pred[classes[i]] = regr.predict(X[:,i::6])\n        regr_coef[i, i::6] = regr.coef_\n\n    print(regr_coef.shape)\n    print(np.array2string(regr_coef, separator=','))\n\n    y_data = df_pred[classes].values\n    y_data = y_data / y_data.sum(axis=1,keepdims=True)\n    df_pred[classes] = y_data\n\n    cv_score = score(solution=true_vote10.copy(), submission=df_pred.copy(), row_id_column_name='label_id')\n    print(cv_score)\n\n    regression_coef = regr_coef \n\n    # not contributed models\n    not_contributed = []\n    for i in range(N):\n        idx = [x for x in range(i*6, (i+1)*6)]\n        if np.sum(regr_coef[:, idx]) < 1e-6:\n            not_contributed.append(i)\n    print(\"not contributed model idx: \", not_contributed)","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:32.258626Z","iopub.execute_input":"2024-05-03T22:03:32.259044Z","iopub.status.idle":"2024-05-03T22:03:32.433571Z","shell.execute_reply.started":"2024-05-03T22:03:32.259013Z","shell.execute_reply":"2024-05-03T22:03:32.432346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"for p in y_1d_predictions:\n    display(p[:min(len(p),3), :])","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:32.435127Z","iopub.execute_input":"2024-05-03T22:03:32.436333Z","iopub.status.idle":"2024-05-03T22:03:32.455546Z","shell.execute_reply.started":"2024-05-03T22:03:32.436238Z","shell.execute_reply":"2024-05-03T22:03:32.454156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for p in y_2d_predictions:\n    display(p[:min(len(p),3), :])","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:32.457442Z","iopub.execute_input":"2024-05-03T22:03:32.458261Z","iopub.status.idle":"2024-05-03T22:03:32.466614Z","shell.execute_reply.started":"2024-05-03T22:03:32.458216Z","shell.execute_reply":"2024-05-03T22:03:32.465733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#for p in k_predictions:\n#    display(p[:min(len(p),3), :])","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:32.468052Z","iopub.execute_input":"2024-05-03T22:03:32.468755Z","iopub.status.idle":"2024-05-03T22:03:32.474395Z","shell.execute_reply.started":"2024-05-03T22:03:32.468723Z","shell.execute_reply":"2024-05-03T22:03:32.473496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#for p in k_cwt10sec_predictions:\n#    display(p[:min(len(p),3), :])","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:32.475811Z","iopub.execute_input":"2024-05-03T22:03:32.476440Z","iopub.status.idle":"2024-05-03T22:03:32.483534Z","shell.execute_reply.started":"2024-05-03T22:03:32.476408Z","shell.execute_reply":"2024-05-03T22:03:32.482617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for p in k_cwt25sec_predictions:\n#     display(p[:min(len(p),3), :])","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:32.485071Z","iopub.execute_input":"2024-05-03T22:03:32.485726Z","iopub.status.idle":"2024-05-03T22:03:32.492792Z","shell.execute_reply.started":"2024-05-03T22:03:32.485693Z","shell.execute_reply":"2024-05-03T22:03:32.491741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for p in k_cwt50sec_predictions:\n#     display(p[:min(len(p),3), :])","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:32.496894Z","iopub.execute_input":"2024-05-03T22:03:32.497829Z","iopub.status.idle":"2024-05-03T22:03:32.505974Z","shell.execute_reply.started":"2024-05-03T22:03:32.497772Z","shell.execute_reply":"2024-05-03T22:03:32.505118Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for p in k_cwt50sec_to40Hz_predictions:\n    display(p[:min(len(p),3), :])","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:32.507234Z","iopub.execute_input":"2024-05-03T22:03:32.507917Z","iopub.status.idle":"2024-05-03T22:03:32.521018Z","shell.execute_reply.started":"2024-05-03T22:03:32.507891Z","shell.execute_reply":"2024-05-03T22:03:32.519984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for p in k_cwt50sec_to40Hz_and_1d_predictions:\n    display(p[:min(len(p),3), :])","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:32.522106Z","iopub.execute_input":"2024-05-03T22:03:32.522451Z","iopub.status.idle":"2024-05-03T22:03:32.532865Z","shell.execute_reply.started":"2024-05-03T22:03:32.522429Z","shell.execute_reply":"2024-05-03T22:03:32.532035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for p in s_predictions:\n    display(p[:min(len(p),3), :])","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:32.533915Z","iopub.execute_input":"2024-05-03T22:03:32.534163Z","iopub.status.idle":"2024-05-03T22:03:32.563125Z","shell.execute_reply.started":"2024-05-03T22:03:32.534141Z","shell.execute_reply":"2024-05-03T22:03:32.562250Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for p in f_predictions:\n    display(p[:min(len(p),3), :])","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:32.564177Z","iopub.execute_input":"2024-05-03T22:03:32.564525Z","iopub.status.idle":"2024-05-03T22:03:32.574388Z","shell.execute_reply.started":"2024-05-03T22:03:32.564494Z","shell.execute_reply":"2024-05-03T22:03:32.573495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/test.csv\")\nclasses = [\"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \"lrda_vote\", \"grda_vote\", \"other_vote\"]\n\nfinal_sub = pd.DataFrame({\"eeg_id\": test_df.eeg_id.values})\n\nif False:  # hill climbing\n    print(\"Hill Climb\")\n    all_predictions = np.hstack(y_1d_predictions + k_cwt50sec_to40Hz_predictions +k_cwt50sec_to40Hz_and_1d_predictions+ s_predictions + f_predictions)\n    final_sub[classes] = all_predictions @ hill_climb_coef.transpose()\nelse:  # non-negative regression\n    print(\"Non-negative Regression\")\n    all_predictions = np.hstack(y_1d_predictions + k_cwt50sec_to40Hz_predictions +k_cwt50sec_to40Hz_and_1d_predictions+ s_predictions + f_predictions)\n    final_sub[classes] = all_predictions @ regression_coef.transpose()\n    \ny_data = final_sub[classes].values\ny_data = y_data / y_data.sum(axis=1,keepdims=True)\nfinal_sub[classes] = y_data\n\nfinal_sub.to_csv(\"submission.csv\", index=False)\nfinal_sub.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:32.575421Z","iopub.execute_input":"2024-05-03T22:03:32.575759Z","iopub.status.idle":"2024-05-03T22:03:32.601315Z","shell.execute_reply.started":"2024-05-03T22:03:32.575734Z","shell.execute_reply":"2024-05-03T22:03:32.600416Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df","metadata":{"execution":{"iopub.status.busy":"2024-05-03T22:03:32.602314Z","iopub.execute_input":"2024-05-03T22:03:32.602572Z","iopub.status.idle":"2024-05-03T22:03:32.610038Z","shell.execute_reply.started":"2024-05-03T22:03:32.602550Z","shell.execute_reply":"2024-05-03T22:03:32.609095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}