{"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":7710760,"sourceType":"datasetVersion","datasetId":4502411}],"dockerImageVersionId":30648,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import IPython.display as ipd\nimport pandas as pd\nimport numpy as np\nimport seaborn as sns\nimport librosa\nimport soundfile as sf\nimport re\nimport scipy\nimport torch\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport h5py\nimport gc\nimport os\nimport seaborn as sns\nimport timeit\nimport albumentations as A\nimport timm\n\nfrom typing import Tuple, Dict, Any, Optional, Callable, Union\nfrom glob import glob\nfrom tqdm import tqdm\nfrom joblib import Parallel, delayed\nfrom tqdm import tqdm\nfrom collections import Counter\nfrom albumentations.pytorch import ToTensorV2\nfrom torchaudio.transforms import FrequencyMasking, TimeMasking\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom time import time\nfrom pprint import pprint\nfrom collections import OrderedDict\n\n%matplotlib inline","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-02-27T08:04:26.023599Z","iopub.execute_input":"2024-02-27T08:04:26.023957Z","iopub.status.idle":"2024-02-27T08:04:27.529595Z","shell.execute_reply.started":"2024-02-27T08:04:26.023927Z","shell.execute_reply":"2024-02-27T08:04:27.528806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Constants","metadata":{}},{"cell_type":"code","source":"NEURAL_SENSORS = (\n    \"Fp1\",\n    \"F3\",\n    \"C3\",\n    \"P3\",\n    \"F7\",\n    \"T3\",\n    \"T5\",\n    \"O1\",\n    \"Fz\",\n    \"Cz\",\n    \"Pz\",\n    \"Fp2\",\n    \"F4\",\n    \"C4\",\n    \"P4\",\n    \"F8\",\n    \"T4\",\n    \"T6\",\n    \"O2\",\n    \"EKG\",\n)\nSPEC_TYPES = (\"LL\", \"RL\", \"LP\", \"RP\")\nREVERSED_SPEC_TYPES = ('LL','LP','RP','RR')\nREVERSED_FEATS = [['Fp1','F7','T3','T5','O1'],\n         ['Fp1','F3','C3','P3','O1'],\n         ['Fp2','F8','T4','T6','O2'],\n         ['Fp2','F4','C4','P4','O2']]\nTARGETS = (\"Seizure\", \"LPD\", \"GPD\", \"LRDA\", \"GRDA\", \"Other\")\nTARGET2ID = {\"Seizure\": 0, \"LPD\": 1, \"GPD\": 2, \"LRDA\": 3, \"GRDA\": 4, \"Other\": 5}\nID2TARGET = {v: k for k, v in TARGET2ID.items()}\nDEFAULT_SAMPLE_RATE = 200","metadata":{"execution":{"iopub.status.busy":"2024-02-27T07:57:44.577052Z","iopub.execute_input":"2024-02-27T07:57:44.577529Z","iopub.status.idle":"2024-02-27T07:57:44.584881Z","shell.execute_reply.started":"2024-02-27T07:57:44.577504Z","shell.execute_reply":"2024-02-27T07:57:44.584014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare Data","metadata":{}},{"cell_type":"code","source":"class ProgressParallel(Parallel):\n    def __init__(self, use_tqdm=True, total=None, *args, **kwargs):\n        self._use_tqdm = use_tqdm\n        self._total = total\n        super().__init__(*args, **kwargs)\n\n    def __call__(self, *args, **kwargs):\n        with tqdm(disable=not self._use_tqdm, total=self._total) as self._pbar:\n            return Parallel.__call__(self, *args, **kwargs)\n\n    def print_progress(self):\n        if self._total is None:\n            self._pbar.total = self.n_dispatched_tasks\n        self._pbar.n = self.n_completed_tasks\n        self._pbar.refresh()\n\ndef compose_spec_from_df(input_df, spec_type, to_db=True, return_freq_and_time=False):\n    spec_cols = [col for col in input_df.columns if spec_type in col]\n    spec_cols = sorted(spec_cols, key=lambda x: float(x.split(\"_\")[1]))\n    spec = np.stack([input_df[el].values for el in spec_cols])\n    if to_db:\n        spec = librosa.amplitude_to_db(spec)\n    if return_freq_and_time:\n        freqs = [float(col.split(\"_\")[1]) for col in spec_cols]\n        times = list(input_df[\"time\"])\n        return spec, freqs, times\n\ndef process_and_save_spec(spec_src_path, spec_tgt_path):\n    sample_df = pd.read_parquet(spec_src_path, engine=\"pyarrow\")\n    prev_times, prev_freqs = None, None\n    with h5py.File(spec_tgt_path, \"w\") as data_file:\n        for spec_type in SPEC_TYPES:\n            spec, freqs, times = compose_spec_from_df(sample_df, spec_type, to_db=False, return_freq_and_time=True)\n            if prev_times is not None:\n                assert np.all(prev_times == times)\n                assert np.all(prev_freqs == freqs)\n            data_file.create_dataset(spec_type, data=spec)\n            prev_times, prev_freqs = times, freqs\n        data_file.create_dataset(\"freqs\", data=np.array(freqs))\n        data_file.create_dataset(\"times\", data=np.array(times))\n\ndef process_and_save_eeg(eeg_src_path, eeg_tgt_path):\n    sample_df = pd.read_parquet(eeg_src_path, engine=\"pyarrow\")\n    assert set(sample_df.columns) == set(NEURAL_SENSORS)\n    with h5py.File(eeg_tgt_path, \"w\") as data_file:\n        for sensor in NEURAL_SENSORS:\n            data_file.create_dataset(sensor, data=sample_df[sensor].values)\n            \ndef process_kaggle_spec(\n    h5py_path, middle, display=False\n):\n    middle = int(middle)\n    X = np.zeros((128, 256, 4),dtype='float32')\n    specs = read_h5py_file(h5py_path)\n    if display:\n        plt.figure(figsize=(10,10))\n    for k_id, k in enumerate(SPEC_TYPES):\n        # EXTRACT 300 ROWS OF SPECTROGRAM\n        img = specs[k][:, middle:middle+300]\n        \n        # LOG TRANSFORM SPECTROGRAM\n        img = np.clip(img, np.exp(-4), np.exp(8))\n        img = np.log(img)\n        \n        # STANDARDIZE PER IMAGE\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        \n        # CROP TO 256 TIME STEPS\n        X[14:-14, :, k_id] = img[:, 22:-22] / 2.0\n\n        if display:\n            eeg_id = os.path.splitext(os.path.basename(h5py_path))[0]\n            \n            plt.subplot(2,2,k_id+1)\n            plt.imshow(X[:,:,k_id],aspect='auto',origin='lower')\n            plt.title(f'EEG {eeg_id} - Spectrogram {k}')\n    return X\n\ndef spectrogram_from_eeg(h5py_path, display=False):\n    \n    # LOAD MIDDLE 50 SECONDS OF EEG SERIES\n    eeg = read_h5py_file(h5py_path)\n    middle = (len(eeg[NEURAL_SENSORS[0]])-10_000)//2\n    for k in NEURAL_SENSORS:\n        eeg[k] = eeg[k][middle:middle+10_000]\n        \n    # VARIABLE TO HOLD SPECTROGRAM\n    img = np.zeros((128,256,4),dtype='float32')\n    \n    if display: \n        plt.figure(figsize=(12,12))\n    signals = []\n    for k in range(4):\n        COLS = REVERSED_FEATS[k]\n        \n        for kk in range(4):\n        \n            # COMPUTE PAIR DIFFERENCES\n            x = eeg[COLS[kk]] - eeg[COLS[kk+1]]\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            eeg_id = os.path.splitext(os.path.basename(h5py_path))[0]\n            plt.subplot(2,2,k+1)\n            plt.imshow(img[:,:,k],aspect='auto',origin='lower')\n            plt.title(f'EEG {eeg_id} - Spectrogram {REVERSED_FEATS[k]}')\n    \n    if display: \n        plt.show()\n        plt.figure(figsize=(10,10))\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=REVERSED_FEATS[3-k])\n            offset += signals[3-k].max()\n        plt.legend()\n        plt.title(f'EEG {eeg_id} Signals')\n        plt.show()\n        \n    return img\n\ndef spectrogram_from_eeg_and_save(eeg_src_path, eeg_tgt_path):\n    \n    img = spectrogram_from_eeg(eeg_src_path, display=False)\n\n    with h5py.File(eeg_tgt_path, \"w\") as data_file:\n        for i, spec_type in enumerate(REVERSED_SPEC_TYPES):\n            data_file.create_dataset(spec_type, data=img[:,:,i]) \n\ndef read_h5py_file(file_path: str):\n    with h5py.File(file_path, \"r\") as data_file:\n        data = {key: data_file[key][:] for key in data_file.keys()}\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-02-27T07:57:44.586276Z","iopub.execute_input":"2024-02-27T07:57:44.586736Z","iopub.status.idle":"2024-02-27T07:57:44.620484Z","shell.execute_reply.started":"2024-02-27T07:57:44.586703Z","shell.execute_reply":"2024-02-27T07:57:44.619576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"eeg_spec_file_pathes = glob(\n    \"../input/hms-harmful-brain-activity-classification/test_spectrograms/*.parquet\"\n)\nprint(f\"Found {len(eeg_spec_file_pathes)} EEG Spectogram files\")\nos.makedirs(\n    \"./temp/test_spectrograms_npy\"\n)\nProgressParallel(n_jobs=4, total=len(eeg_spec_file_pathes))(\n    delayed(process_and_save_spec)(\n        spec_src_path=spec_src_path, \n        spec_tgt_path=os.path.join(\"./temp/test_spectrograms_npy\", os.path.basename(spec_src_path).replace(\".parquet\", \".h5\"))\n    ) for spec_src_path in eeg_spec_file_pathes\n);","metadata":{"execution":{"iopub.status.busy":"2024-02-27T07:57:45.499016Z","iopub.execute_input":"2024-02-27T07:57:45.499350Z","iopub.status.idle":"2024-02-27T07:57:46.625005Z","shell.execute_reply.started":"2024-02-27T07:57:45.499326Z","shell.execute_reply":"2024-02-27T07:57:46.623907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"eeg_file_pathes = glob(\n    \"../input/hms-harmful-brain-activity-classification/test_eegs/*.parquet\"\n)\nprint(f\"Found {len(eeg_file_pathes)} EEG files\")\n\nos.makedirs(\n    \"./temp/test_eegs_npy\"\n)\nProgressParallel(n_jobs=4, total=len(eeg_file_pathes))(\n    delayed(process_and_save_eeg)(\n        eeg_src_path=eeg_src_path,\n        eeg_tgt_path=os.path.join(\"./temp/test_eegs_npy\", os.path.basename(eeg_src_path).replace(\".parquet\", \".h5\"))\n    )\n    for eeg_src_path in eeg_file_pathes\n)","metadata":{"execution":{"iopub.status.busy":"2024-02-27T07:57:46.626979Z","iopub.execute_input":"2024-02-27T07:57:46.627309Z","iopub.status.idle":"2024-02-27T07:57:47.079113Z","shell.execute_reply.started":"2024-02-27T07:57:46.627279Z","shell.execute_reply":"2024-02-27T07:57:47.078038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"eeg_file_pathes = glob(\n    \"./temp/test_eegs_npy/*.h5\"\n)\nprint(f\"Found {len(eeg_file_pathes)} EEG files\")\n\nos.makedirs(\n    \"./temp/test_spectrograms_reversed_npy\"\n)\nProgressParallel(n_jobs=4, total=len(eeg_file_pathes))(\n    delayed(spectrogram_from_eeg_and_save)(\n        eeg_src_path=eeg_src_path,\n        eeg_tgt_path=eeg_src_path.replace(\"test_eegs_npy\", \"test_spectrograms_reversed_npy\")\n    )\n    for eeg_src_path in eeg_file_pathes\n);","metadata":{"execution":{"iopub.status.busy":"2024-02-27T07:58:20.750043Z","iopub.execute_input":"2024-02-27T07:58:20.750696Z","iopub.status.idle":"2024-02-27T07:58:30.251975Z","shell.execute_reply.started":"2024-02-27T07:58:20.750652Z","shell.execute_reply":"2024-02-27T07:58:30.251100Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!tree ./temp","metadata":{"execution":{"iopub.status.busy":"2024-02-27T07:58:34.770543Z","iopub.execute_input":"2024-02-27T07:58:34.770882Z","iopub.status.idle":"2024-02-27T07:58:35.720141Z","shell.execute_reply.started":"2024-02-27T07:58:34.770852Z","shell.execute_reply":"2024-02-27T07:58:35.719136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Meta","metadata":{}},{"cell_type":"code","source":"test_df = pd.read_csv(\"../input/hms-harmful-brain-activity-classification/test.csv\")\n# test_df[\"spectrogram_label_offset_seconds\"] = 0.0\ntest_df[\"middle\"] = 0.0\ntest_df = test_df.rename(columns={\"spectrogram_id\":\"spec_id\"})\ntest_df","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:00:09.461059Z","iopub.execute_input":"2024-02-27T08:00:09.461557Z","iopub.status.idle":"2024-02-27T08:00:09.484249Z","shell.execute_reply.started":"2024-02-27T08:00:09.461520Z","shell.execute_reply":"2024-02-27T08:00:09.483323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class CombinedDataset(torch.utils.data.Dataset):\n    def __init__(\n        self,\n        root_eeg,\n        root_spec,\n        df,\n        target_col=\"target\",\n        target_cols=None,\n        spec_id_col=\"spec_id\",\n        eeg_id_col=\"eeg_id\",\n        middle_second_col=\"middle\",\n        transform=None,\n        test_mode=False,\n        specs_to_use=\"all\"\n    ):\n        assert specs_to_use in [\"all\", \"eeg_spec\", \"original_spec\"]\n        \n        self.df = df.reset_index(drop=True)\n\n        self.target_col = target_col\n        if target_cols is None:\n            self.target_cols = [el.lower() + \"_vote\" for el in TARGETS]\n        else:\n            self.target_cols = target_cols\n        self.name_col = spec_id_col\n        self.eeg_id_col = eeg_id_col\n        self.test_mode = test_mode\n        self.middle_second_col = middle_second_col\n        self.specs_to_use = specs_to_use\n\n        self.root_eeg = root_eeg\n        self.root_spec = root_spec\n\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def _prepare_sample(\n        self,\n        spec_id,\n        eeg_id,\n        middle\n    ):\n        eeg_path = os.path.join(self.root_eeg, f\"{eeg_id}.h5\")\n        spec_path = os.path.join(self.root_spec, f\"{spec_id}.h5\")\n        middle = int(middle)\n\n        if self.specs_to_use == \"all\":\n            n_specs = 8\n        else:\n            n_specs = 4\n        \n        X = np.zeros((128, 256, n_specs),dtype='float32')\n\n        if self.specs_to_use in [\"all\", \"original_spec\"]:\n            specs = read_h5py_file(spec_path)\n\n            for k_id, k in enumerate(SPEC_TYPES):\n                # EXTRACT 300 ROWS OF SPECTROGRAM\n                # import ipdb; ipdb.set_trace()\n                img = specs[k][:, middle:middle+300]\n                \n                # LOG TRANSFORM SPECTROGRAM\n                img = np.clip(img, np.exp(-4), np.exp(8))\n                img = np.log(img)\n                \n                # STANDARDIZE PER IMAGE\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                \n                # CROP TO 256 TIME STEPS\n                X[14:-14, :, k_id] = img[:, 22:-22] / 2.0\n\n        if self.specs_to_use in [\"all\", \"eeg_spec\"]:\n            eeg_spec = read_h5py_file(eeg_path)\n\n            if self.specs_to_use == \"all\":\n                X[:, :, 4:] = np.stack([\n                    eeg_spec[k] for k in REVERSED_SPEC_TYPES\n                ], axis=-1)\n            else:\n                X[:, :, :] = np.stack([\n                    eeg_spec[k] for k in REVERSED_SPEC_TYPES\n                ], axis=-1)\n\n        \n        return X\n\n    def __getitem__(self, idx: int):\n        middle_second = self.df[self.middle_second_col].iloc[idx]\n        eeg_id = self.df[self.eeg_id_col].iloc[idx]\n        spec_id = self.df[self.name_col].iloc[idx]\n\n        if self.test_mode:\n            main_target = -1\n            all_targets = np.full(len(self.target_cols), -1.0)\n        else:\n            main_target = self.df[self.target_col].iloc[idx]\n            main_target = TARGET2ID[main_target]\n            all_targets = self.df[self.target_cols].iloc[idx].values\n\n        all_targets = torch.from_numpy(all_targets.astype(np.float32))\n        main_target = torch.tensor(main_target).long()\n        specs = self._prepare_sample(\n            spec_id=spec_id,\n            eeg_id=eeg_id,\n            middle=middle_second\n        )\n\n        if self.transform is not None:\n            specs = self.transform(image=specs)[\"image\"]\n\n        specs = specs.float()\n\n        return specs, main_target, all_targets, eeg_id","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:00:13.341338Z","iopub.execute_input":"2024-02-27T08:00:13.341668Z","iopub.status.idle":"2024-02-27T08:00:13.361748Z","shell.execute_reply.started":"2024-02-27T08:00:13.341642Z","shell.execute_reply":"2024-02-27T08:00:13.360784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = CombinedDataset(\n    root_eeg=\"./temp/test_spectrograms_reversed_npy\",\n    root_spec=\"./temp/test_spectrograms_npy\",\n    df=test_df,\n    transform=A.Compose([\n        ToTensorV2(transpose_mask=True),\n    ]),\n    test_mode=True\n)","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:01:23.990308Z","iopub.execute_input":"2024-02-27T08:01:23.990662Z","iopub.status.idle":"2024-02-27T08:01:23.996654Z","shell.execute_reply.started":"2024-02-27T08:01:23.990633Z","shell.execute_reply":"2024-02-27T08:01:23.995744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testing_specs, testing_main_target, testing_all_targets, testing_eeg_id = test_dataset[0]\n\nprint(\"Main target\", testing_main_target)\nprint(\"All targets\", testing_all_targets)\nprint(\"EEG Id\", testing_eeg_id)\n\nprint(\"Shape\", testing_specs.shape)\n\nfor idx, spec_type in enumerate(SPEC_TYPES+REVERSED_SPEC_TYPES):\n    plt.title(spec_type)\n    plt.imshow(testing_specs[idx].numpy())\n    plt.gca().invert_yaxis()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:01:45.800610Z","iopub.execute_input":"2024-02-27T08:01:45.801473Z","iopub.status.idle":"2024-02-27T08:01:47.834723Z","shell.execute_reply.started":"2024-02-27T08:01:45.801441Z","shell.execute_reply":"2024-02-27T08:01:47.833702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class NormalizeMelSpec(nn.Module):\n    def __init__(\n        self,\n        eps=1e-6,\n        normalize_standart=True,\n        normalize_minmax=True,\n    ):\n        super().__init__()\n        self.eps = eps\n        self.normalize_standart = normalize_standart\n        self.normalize_minmax = normalize_minmax\n\n    def forward(self, X):\n        if self.normalize_standart:\n            mean = X.mean((2, 3), keepdim=True)\n            std = X.std((2, 3), keepdim=True)\n            X = (X - mean) / (std + self.eps)\n            \n        if self.normalize_minmax:\n            norm_max = torch.amax(X, dim=(2, 3), keepdim=True)\n            norm_min = torch.amin(X, dim=(2, 3), keepdim=True)\n            X = (X - norm_min) / (norm_max - norm_min + self.eps)\n\n        return X\n\nclass CustomMasking(nn.Module):\n    def __init__(self, mask_max_length: int, mask_max_masks: int, p=1.0, inplace=True):\n        super().__init__()\n        assert isinstance(mask_max_masks, int) and mask_max_masks > 0\n        self.mask_max_masks = mask_max_masks\n        self.mask_max_length = mask_max_length\n        self.mask_module = None\n        self.p = p\n        self.inplace = inplace\n\n    def forward(self, x):\n        if not self.inplace:\n            output = x.clone()\n        for i in range(x.shape[0]):\n            if np.random.binomial(n=1, p=self.p):\n                n_applies = np.random.randint(low=1, high=self.mask_max_masks + 1)\n                for _ in range(n_applies):\n                    if self.inplace:\n                        x[i : i + 1] = self.mask_module(x[i : i + 1])\n                    else:\n                        output[i : i + 1] = self.mask_module(output[i : i + 1])\n        if self.inplace:\n            return x\n        else:\n            return output\n\n\nclass CustomTimeMasking(CustomMasking):\n    def __init__(self, mask_max_length: int, mask_max_masks: int, p=1.0, inplace=True):\n        super().__init__(mask_max_length=mask_max_length, mask_max_masks=mask_max_masks, p=p, inplace=inplace)\n        self.mask_module = TimeMasking(time_mask_param=mask_max_length)\n\n\nclass CustomFreqMasking(CustomMasking):\n    def __init__(self, mask_max_length: int, mask_max_masks: int, p=1.0, inplace=True):\n        super().__init__(mask_max_length=mask_max_length, mask_max_masks=mask_max_masks, p=p, inplace=inplace)\n        self.mask_module = FrequencyMasking(freq_mask_param=mask_max_length)\n\nclass SpecCNNClasifier(nn.Module):\n    def __init__(\n        self,\n        backbone: str,\n        device: str,\n        n_specs: int,\n        n_classes: int,\n        classifier_dropout: float = 0.5,\n        normalize_config: Dict[str, bool] = {\n            \"normalize_standart\": True,\n            \"normalize_minmax\": True,\n        },\n        pretrained: bool = True,\n        timm_kwargs: Optional[Dict] = None,\n        spec_augment_config: Optional[Dict[str, Any]] = None,\n    ):\n        super().__init__()\n        timm_kwargs = {} if timm_kwargs is None else timm_kwargs\n        self.device = device\n\n        self.instance_norm = NormalizeMelSpec(\n            **normalize_config\n        )\n        \n        if spec_augment_config is not None:\n            self.spec_augment = []\n            if \"freq_mask\" in spec_augment_config:\n                self.spec_augment.append(CustomFreqMasking(**spec_augment_config[\"freq_mask\"]))\n            if \"time_mask\" in spec_augment_config:\n                self.spec_augment.append(CustomTimeMasking(**spec_augment_config[\"time_mask\"]))\n            self.spec_augment = nn.Sequential(*self.spec_augment)\n        else:\n            self.spec_augment = None\n\n        \n        self.backbone = timm.create_model(\n            backbone,\n            features_only=True,\n            pretrained=pretrained,\n            in_chans=n_specs,\n            exportable=True,\n            **timm_kwargs,\n        )\n\n        self.pool = nn.AdaptiveAvgPool2d((1, 1))\n        self.classifier = nn.Sequential(\n            nn.Dropout(p=classifier_dropout),\n            nn.Linear(self.backbone.feature_info.channels()[-1], n_classes),\n        )\n        \n        self.to(self.device)\n\n    def forward(self, input, return_spec_feature=False, return_cnn_emb=False):\n        processed_spec = self.instance_norm(input)\n        if self.spec_augment is not None and self.training:\n            processed_spec = self.spec_augment(processed_spec)\n        if return_spec_feature:\n            return processed_spec\n            \n        emb = self.backbone(processed_spec)[-1]\n        if return_cnn_emb:\n            return emb\n\n        bs, ch, h, w = emb.shape\n        emb = self.pool(emb)\n        emb = emb.view(bs, ch)\n\n        logits = self.classifier(emb)\n\n        return {\"logits\": logits}","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:02:24.157002Z","iopub.execute_input":"2024-02-27T08:02:24.157884Z","iopub.status.idle":"2024-02-27T08:02:24.185018Z","shell.execute_reply.started":"2024-02-27T08:02:24.157851Z","shell.execute_reply":"2024-02-27T08:02:24.183990Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"test_loader = torch.utils.data.DataLoader(\n    test_dataset,\n    batch_size=64,\n    drop_last=False,\n    shuffle=False,\n    num_workers=4\n)","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:03:12.218216Z","iopub.execute_input":"2024-02-27T08:03:12.218580Z","iopub.status.idle":"2024-02-27T08:03:12.223988Z","shell.execute_reply.started":"2024-02-27T08:03:12.218550Z","shell.execute_reply":"2024-02-27T08:03:12.223043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def delete_prefix_from_chkp(chkp_dict: OrderedDict, prefix: str):\n    new_dict = OrderedDict()\n    for k in chkp_dict.keys():\n        if k.startswith(prefix):\n            new_dict[k[len(prefix) :]] = chkp_dict[k]\n        else:\n            new_dict[k] = chkp_dict[k]\n\n    return new_dict\n\ndef create_model_and_load_best_checkpoint(\n    model_class,\n    model_config,\n    model_device,\n    model_chkp_root,\n    model_chkp_regex,\n    sort_rule,\n    delete_prefix=None,\n):\n    basenames = os.listdir(model_chkp_root)\n    checkpoints = []\n    for el in basenames:\n        matches = re.findall(model_chkp_regex, el)\n        if not matches:\n            continue\n        parsed_dict = {key: value for key, value in matches}\n        parsed_dict[\"name\"] = el\n        checkpoints.append(parsed_dict)\n    print(\"All checkpoints\")\n    pprint(checkpoints)\n    checkpoints = sorted(checkpoints, key=sort_rule)\n    print(\"Sorted checkpoints\")\n    pprint(checkpoints)\n    best_checkpoint = os.path.join(model_chkp_root, checkpoints[0][\"name\"])\n    print(\"Best checkpoint\")\n    print(best_checkpoint)\n    t_chkp = torch.load(\n        best_checkpoint, \n        map_location=\"cpu\"\n    )[\"state_dict\"]\n    if delete_prefix is not None:\n        t_chkp = delete_prefix_from_chkp(t_chkp, delete_prefix)\n    t_model = model_class(**model_config, device=model_device)\n    t_model.load_state_dict(t_chkp)\n    t_model.eval()\n\n    return t_model","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:03:22.017965Z","iopub.execute_input":"2024-02-27T08:03:22.018352Z","iopub.status.idle":"2024-02-27T08:03:22.028441Z","shell.execute_reply.started":"2024-02-27T08:03:22.018322Z","shell.execute_reply":"2024-02-27T08:03:22.027445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = [create_model_and_load_best_checkpoint(\n        model_class=SpecCNNClasifier,\n        model_config=dict(\n            pretrained=False,\n            backbone=\"tf_efficientnet_b0.in1k\",\n            n_specs=8,\n            n_classes=6,\n            spec_augment_config={\n                \"freq_mask\": {\n                    \"mask_max_length\": 20,\n                    \"mask_max_masks\": 5,\n                    \"p\": 1.0,\n                    \"inplace\": True,\n                },\n                \"time_mask\": {\n                    \"mask_max_length\": 30,\n                    \"mask_max_masks\": 5,\n                    \"p\": 1.0,\n                    \"inplace\": True,\n                },\n            }\n        ),\n        model_device=\"cuda\",\n        model_chkp_root=f\"../input/ucu-hms-models/hms_baseline/fold_{m_i}/checkpoints\",\n        model_chkp_regex=r'(?P<key>\\w+)=(?P<value>[\\d.]+)(?=\\.ckpt|$)',\n        sort_rule=lambda x: float(x[\"valid_kl\"]),\n        delete_prefix=\"model.\"\n) for m_i in range(5)]","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:04:31.431949Z","iopub.execute_input":"2024-02-27T08:04:31.432340Z","iopub.status.idle":"2024-02-27T08:04:33.281735Z","shell.execute_reply.started":"2024-02-27T08:04:31.432308Z","shell.execute_reply":"2024-02-27T08:04:33.280893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.inference_mode()\ndef inference_function(\n    loader,\n    nn_models,\n    output_key,\n    device\n):\n    predicted_df = {\"eeg_id\": []}\n    for target in TARGETS:\n        predicted_df[target.lower() + \"_vote\"] = []\n\n    for batch in tqdm(loader):\n        specs, _, _, eeg_id = batch\n        pred_probs = []\n        for nn_model in nn_models:\n            local_probs = nn_model(specs.to(device))[output_key]\n            local_probs = torch.softmax(local_probs, dim=1)\n            local_probs = local_probs.detach().cpu().numpy()\n            pred_probs.append(local_probs)\n        pred_probs = np.stack(pred_probs, axis=0).mean(0)\n        predicted_df[\"eeg_id\"].append(eeg_id.cpu().numpy())\n        for i, target in enumerate(TARGETS):\n            predicted_df[target.lower() + \"_vote\"].append(pred_probs[:, i])\n\n    for key in predicted_df:\n        predicted_df[key] = np.concatenate(predicted_df[key])\n\n    predicted_df = pd.DataFrame(predicted_df)\n\n    return predicted_df","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:06:10.901527Z","iopub.execute_input":"2024-02-27T08:06:10.901883Z","iopub.status.idle":"2024-02-27T08:06:10.911166Z","shell.execute_reply.started":"2024-02-27T08:06:10.901854Z","shell.execute_reply":"2024-02-27T08:06:10.910135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predicted_df = inference_function(\n    loader=test_loader,\n    nn_models=model,\n    output_key=\"logits\",\n    device=\"cuda\"\n)","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:07:09.430973Z","iopub.execute_input":"2024-02-27T08:07:09.431392Z","iopub.status.idle":"2024-02-27T08:07:10.497054Z","shell.execute_reply.started":"2024-02-27T08:07:09.431359Z","shell.execute_reply":"2024-02-27T08:07:10.495963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predicted_df","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:07:13.144847Z","iopub.execute_input":"2024-02-27T08:07:13.145628Z","iopub.status.idle":"2024-02-27T08:07:13.159041Z","shell.execute_reply.started":"2024-02-27T08:07:13.145588Z","shell.execute_reply":"2024-02-27T08:07:13.158005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predicted_df[[el.lower() + \"_vote\" for el in TARGETS]].sum(axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:08:04.281234Z","iopub.execute_input":"2024-02-27T08:08:04.281607Z","iopub.status.idle":"2024-02-27T08:08:04.292656Z","shell.execute_reply.started":"2024-02-27T08:08:04.281568Z","shell.execute_reply":"2024-02-27T08:08:04.291773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# VERY IMPORTANT: Empty temp and save submission","metadata":{}},{"cell_type":"code","source":"ls ./","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:08:08.314949Z","iopub.execute_input":"2024-02-27T08:08:08.315338Z","iopub.status.idle":"2024-02-27T08:08:09.277413Z","shell.execute_reply.started":"2024-02-27T08:08:08.315305Z","shell.execute_reply":"2024-02-27T08:08:09.276229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rm ./* -rf","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:08:12.513322Z","iopub.execute_input":"2024-02-27T08:08:12.513757Z","iopub.status.idle":"2024-02-27T08:08:13.479406Z","shell.execute_reply.started":"2024-02-27T08:08:12.513721Z","shell.execute_reply":"2024-02-27T08:08:13.478266Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ls ./","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:08:17.190148Z","iopub.execute_input":"2024-02-27T08:08:17.190529Z","iopub.status.idle":"2024-02-27T08:08:18.147032Z","shell.execute_reply.started":"2024-02-27T08:08:17.190495Z","shell.execute_reply":"2024-02-27T08:08:18.145891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predicted_df.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-02-27T08:08:21.385571Z","iopub.execute_input":"2024-02-27T08:08:21.385962Z","iopub.status.idle":"2024-02-27T08:08:21.395343Z","shell.execute_reply.started":"2024-02-27T08:08:21.385925Z","shell.execute_reply":"2024-02-27T08:08:21.394332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}