{"metadata":{"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7392733,"sourceType":"datasetVersion","datasetId":4297749}],"dockerImageVersionId":30646,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.13"},"papermill":{"default_parameters":{},"duration":519.073705,"end_time":"2024-01-25T02:57:53.027402","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-01-25T02:49:13.953697","version":"2.4.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd, numpy as np, os\nimport random\nimport matplotlib.pyplot as plt, gc\nimport librosa\n\ntrain = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\nprint('Train shape', train.shape )\ndisplay( train.head() )\n\nCREATE_SPECTROGRAMS = True","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":0.816746,"end_time":"2024-01-25T02:49:18.638126","exception":false,"start_time":"2024-01-25T02:49:17.82138","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-20T14:31:49.953404Z","iopub.execute_input":"2024-04-20T14:31:49.9538Z","iopub.status.idle":"2024-04-20T14:31:50.734365Z","shell.execute_reply.started":"2024-04-20T14:31:49.953764Z","shell.execute_reply":"2024-04-20T14:31:50.733177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install --upgrade --no-deps scipy","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:31:50.736358Z","iopub.execute_input":"2024-04-20T14:31:50.736706Z","iopub.status.idle":"2024-04-20T14:31:53.440224Z","shell.execute_reply.started":"2024-04-20T14:31:50.736675Z","shell.execute_reply":"2024-04-20T14:31:53.438764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### v1\nLL = ( (Fp1 - F7) + (F7 - T3) + (T3 - T5) + (T5 - O1) )/4.\n### v2\nfrom https://www.kaggle.com/code/cdeotte/how-to-make-spectrogram-from-eeg v4\n\nLL Spec = ( spec(Fp1 - F7) + spec(F7 - T3) + spec(T3 - T5) + spec(T5 - O1) )/4.\n\n","metadata":{"papermill":{"duration":0.004862,"end_time":"2024-01-25T02:49:18.648533","exception":false,"start_time":"2024-01-25T02:49:18.643671","status":"completed"},"tags":[]}},{"cell_type":"code","source":"NAMES = ['LL','RL','LP','RP']\n\nFEATS = [['Fp1','F7','T3','T5','O1'],\n         ['Fp2','F8','T4','T6','O2'],\n         ['Fp1','F3','C3','P3','O1'],\n         ['Fp2','F4','C4','P4','O2']]\n\ndirectory_path = 'EEG_Spectrograms/'\nif not os.path.exists(directory_path):\n    os.makedirs(directory_path)","metadata":{"papermill":{"duration":0.017534,"end_time":"2024-01-25T02:49:18.671104","exception":false,"start_time":"2024-01-25T02:49:18.65357","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-20T14:31:53.442587Z","iopub.execute_input":"2024-04-20T14:31:53.442959Z","iopub.status.idle":"2024-04-20T14:31:53.451709Z","shell.execute_reply.started":"2024-04-20T14:31:53.442926Z","shell.execute_reply":"2024-04-20T14:31:53.450304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pywt\nprint(\"The wavelet functions we can use:\")\nprint(pywt.wavelist())\n\nUSE_WAVELET = None #or \"db8\" or anything below","metadata":{"papermill":{"duration":0.528032,"end_time":"2024-01-25T02:49:19.204498","exception":false,"start_time":"2024-01-25T02:49:18.676466","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-20T14:31:53.454866Z","iopub.execute_input":"2024-04-20T14:31:53.455197Z","iopub.status.idle":"2024-04-20T14:31:53.668986Z","shell.execute_reply.started":"2024-04-20T14:31:53.455169Z","shell.execute_reply":"2024-04-20T14:31:53.667877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 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","metadata":{"papermill":{"duration":0.020591,"end_time":"2024-01-25T02:49:19.230879","exception":false,"start_time":"2024-01-25T02:49:19.210288","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-20T14:31:53.670653Z","iopub.execute_input":"2024-04-20T14:31:53.671457Z","iopub.status.idle":"2024-04-20T14:31:53.682311Z","shell.execute_reply.started":"2024-04-20T14:31:53.671414Z","shell.execute_reply":"2024-04-20T14:31:53.681064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# From https://github.com/tomrunia/PyTorchWavelets/blob/master/wavelets_pytorch/wavelets.py\n\n## from G2net1 \nimport torch\nfrom scipy import signal\nfrom scipy import optimize\nimport torch.nn as nn\nfrom timm.layers.conv2d_same import conv2d_same\n\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\n\nclass CWT(nn.Module):\n    def __init__(\n        self,\n        dj=0.0625,\n        dt=1 / 200,\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        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:\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":{"papermill":{"duration":8.344473,"end_time":"2024-01-25T02:49:27.581249","exception":false,"start_time":"2024-01-25T02:49:19.236776","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-20T14:31:53.684666Z","iopub.execute_input":"2024-04-20T14:31:53.685574Z","iopub.status.idle":"2024-04-20T14:32:03.169229Z","shell.execute_reply.started":"2024-04-20T14:31:53.685528Z","shell.execute_reply":"2024-04-20T14:32:03.16793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\n# Function to create training data\ndef create_train_data():\n    path = '/kaggle/input/hms-harmful-brain-activity-classification/'\n    classes = [\n        \"seizure_vote\", \"lpd_vote\", \"gpd_vote\", \n        \"lrda_vote\", \"grda_vote\", \"other_vote\"\n    ]\n    df = pd.read_csv(f'{path}train.csv')\n    \n    # Create a new identifier combining multiple columns\n    id_cols = ['eeg_id', 'spectrogram_id', 'seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\n    df['new_id'] = df[id_cols].astype(str).agg('_'.join, axis=1)\n    \n    df['sum_votes'] = df[classes].sum(axis=1)\n    grouped_df = df.groupby('new_id').agg({\n        'eeg_id': 'first',\n        'eeg_label_offset_seconds': ['min', 'max'],  \n        'spectrogram_label_offset_seconds': ['min', 'max'],\n        'spectrogram_id': 'first',\n        'patient_id': 'first',\n        'expert_consensus': 'first',\n        'seizure_vote': 'first',\n        'lpd_vote': 'first',\n        'gpd_vote': 'first',\n        'lrda_vote': 'first',\n        'grda_vote': 'first',\n        'other_vote': 'first',\n        'sum_votes': 'first',\n    }).reset_index()\n\n    # Post-aggregation processing to adjust column names.\n    new_column_names = []\n    for col in grouped_df.columns:\n        if isinstance(col, tuple):\n            # For aggregated columns with multiple functions, add the function suffix to the column name.\n            new_column_names.append(f\"{col[0]}_{col[1]}\" if col[1] else col[0])\n        else:\n            new_column_names.append(col)\n    \n    grouped_df.columns = new_column_names\n    \n    # Removing \"_first\" from the column names.\n    grouped_df.columns = [col.replace('_first', '') for col in grouped_df.columns]\n    \n    return grouped_df\n\ntrain_df = create_train_data()\ntrain_df","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:32:03.170837Z","iopub.execute_input":"2024-04-20T14:32:03.171249Z","iopub.status.idle":"2024-04-20T14:32:05.506354Z","shell.execute_reply.started":"2024-04-20T14:32:03.171213Z","shell.execute_reply":"2024-04-20T14:32:05.505299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_spectogram_competition(spec_id, seconds_min):\n    spec = pd.read_parquet(f'/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/{spec_id}.parquet')\n    inicio = (seconds_min) // 2\n    img = spec.fillna(0).values[:, 1:].T.astype(\"float32\")\n    img = img[:, inicio:inicio+300]\n    \n    # Log transform and normalize\n    img = np.clip(img, np.exp(-4), np.exp(6))\n    img = np.log(img)\n    eps = 1e-6\n    img_mean = img.mean()\n    img_std = img.std()\n    img = (img - img_mean) / (img_std + eps)\n    \n    return img ","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:32:05.507986Z","iopub.execute_input":"2024-04-20T14:32:05.50874Z","iopub.status.idle":"2024-04-20T14:32:05.517994Z","shell.execute_reply.started":"2024-04-20T14:32:05.508695Z","shell.execute_reply":"2024-04-20T14:32:05.516359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"CWTs seem to be able to extract different information from STFTs, but my local CV has not improved...","metadata":{"papermill":{"duration":0.192232,"end_time":"2024-01-25T02:57:49.116455","exception":false,"start_time":"2024-01-25T02:57:48.924223","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def create_final_image(stft, cwt, kaggle_spec):   # 먼저 stft / cwt / 캐글 기본 스펙 으로 \n    \"\"\"Combine three images into a single final image.\"\"\"\n    # Initialize an empty image array for the first image composition\n    # stft_spec (128, 529)\n    # cwt (176, 527)\n    # spec (400, 300)\n    \n    single_channel_image1 = np.zeros((4*128, 529), dtype=\"float32\")\n    for i in range(4):\n        start = i * 128\n        end = start + 128\n        single_channel_image1[start:end, :] = stft[:, :, i]\n    # Initialize an empty image array for the second image composition\n\n    single_channel_image2 = np.zeros((4*176, 527), dtype=\"float32\")\n    for i in range(4):\n        start = i * 176\n        end = start + 176\n        single_channel_image2[start:end, :] = cwt[:, :, i]\n    # Resize images to fit the final composition\n    \n    resized_image1 = cv2.resize(single_channel_image1, (400, 800), interpolation=cv2.INTER_AREA)\n    resized_image2 = cv2.resize(single_channel_image2, (400, 800), interpolation=cv2.INTER_AREA)\n    resized_image3 = cv2.resize(kaggle_spec, (400, 800), interpolation=cv2.INTER_AREA)\n    '''\n    plt.figure(figsize=(30,20))\n    plt.subplot(1,3,1)\n    plt.imshow(resized_image1[::-1],aspect=\"auto\", cmap='jet')\n    plt.title(\"eeg → stft\")\n    plt.axis('off')  \n    plt.subplot(1,3,2)\n    plt.imshow(resized_image2,aspect=\"auto\", cmap='jet')\n    plt.title(\"eeg → CWT\")\n    plt.axis('off')\n    plt.subplot(1,3,3)\n    plt.imshow(resized_image3, aspect=\"auto\", cmap='jet')\n    plt.title(\"kaggle spectrogram\")\n    plt.axis('off') \n    '''\n    # Create the final image and place the resized images accordingly\n    '''\n    final_image = np.zeros((800, 1200), dtype='float32')\n    final_image[0:800, 0:400] = resized_image2\n    final_image[0:800, 400:800] = resized_image3\n    final_image[0:800, 800:1200] = resized_image1\n    final_image = final_image[::-1]  # Flip the final image vertically\n    '''\n    def check_and_concatenate_images(img1, img2, img3):\n        # Check if all images have the same number of rows and the same data type\n        if img1.shape[0] != img2.shape[0] or img2.shape[0] != img3.shape[0]:\n            raise ValueError(\"All images must have the same height.\")\n        if img1.dtype != img2.dtype or img2.dtype != img3.dtype:\n            raise ValueError(\"All images must have the same data type.\")\n        if len(img1.shape) != len(img2.shape) or len(img2.shape) != len(img3.shape):\n            raise ValueError(\"All images must have the same number of dimensions.\")\n        \n        # normalize\n        def normalize_image(image):\n            return (image - np.min(image)) / (np.max(image) - np.min(image))\n        img1 = normalize_image(img1) * 255\n        img2 = normalize_image(img2) * 255\n        img3 = normalize_image(img3) * 255\n        # Concatenate images horizontally\n        concatenated_image = cv2.hconcat([img1, img2, img3])\n        return concatenated_image\n\n    final_image=check_and_concatenate_images(resized_image1[::-1], resized_image2,resized_image3[::-1])\n    final_image = cv2.resize(final_image, (512, 512), interpolation=cv2.INTER_AREA)\n    return final_image","metadata":{"papermill":{"duration":0.17585,"end_time":"2024-01-25T02:57:49.471001","exception":false,"start_time":"2024-01-25T02:57:49.295151","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2024-04-20T14:32:05.519778Z","iopub.execute_input":"2024-04-20T14:32:05.520149Z","iopub.status.idle":"2024-04-20T14:32:05.542532Z","shell.execute_reply.started":"2024-04-20T14:32:05.520118Z","shell.execute_reply":"2024-04-20T14:32:05.540973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\nfrom scipy.signal.windows import hann\nfrom scipy.signal import ShortTimeFFT\n\ndef process_eegs(train_df, output_folder):\n    \"\"\"Process EEGs and save the final images.\"\"\"\n    # Ensure the output folder exists\n    if not os.path.exists(output_folder):\n        os.makedirs(output_folder)\n    '''\n    1)매 row마다 eeg_id 뽑아서 spec generate\n    2)\n    '''\n    # Iterate over the EEG data frame\n    for i in tqdm(range(len(train_df)), desc=\"Processing EEGs\"):\n        row = train_df.iloc[i]\n        eeg_id = row['eeg_id']\n        spec_id = row['spectrogram_id']\n        seconds_min = int(row.spectrogram_label_offset_seconds_min) # kaggle spec offset\n        start_second = int(row.eeg_label_offset_seconds_min)         # 보통 eeg offset\n\n        # Generate spectrogram images from Kaggle Spec\n        img_kaggle = create_spectogram_competition(spec_id, seconds_min)\n        '''\n        plt.figure(figsize=(40,30))        \n        plt.subplot(1, 3, 1)\n        plt.imshow(img_kaggle, aspect=\"auto\",origin='lower', cmap='jet')\n        plt.title(\"kaggle spectrogram\")\n        plt.axis('off')  \n        '''\n        # Load EEG data from file\n        eeg_data = pd.read_parquet(f'/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/{eeg_id}.parquet')\n        eeg_new_key = f'{eeg_id}_{seconds_min}_{start_second}'   # npz 파일 저장 이름\n\n\n\n        # 여기서부터는 eeg_to_spec 함수\n        start = start_second * 200\n        real_start = start \n        eeg = eeg_data.iloc[start:start + 10_000]\n\n        signals = []\n        img = np.zeros((128,529,4),dtype='float32')  # stft\n        #cwt_img = np.zeros((176,527,4),dtype='float32') # cwt 기본\n        cwt_img1 = np.zeros((176,527,4),dtype='float32') # cwt norm\n\n        pycwt = CWT(fmin=0, fmax=40, hop_length=10_000//512)\n        win_len=64\n        signals = []\n        #2\n        for k in range(4):\n            COLS = FEATS[k]\n\n            #1\n            for kk in range(4):\n\n                # COMPUTE PAIR DIFFERENCES\n                x = eeg[COLS[kk]].values - eeg[COLS[kk+1]].values\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                S = torch.tensor(x)[None,:]\n                # RAW SPECTROGRAM\n                hann_window = hann(win_len)\n                mfft = 10*(win_len)\n                hop = 10_000//512\n                sfft_hann = ShortTimeFFT(win=hann_window, hop=hop, fs=200, mfft=mfft)\n                Sx_hann = abs(sfft_hann.stft(x))\n                max_freq = 40\n                max_freq_idx = int(max_freq//sfft_hann.delta_f)\n                Sx_hann = Sx_hann[:max_freq_idx, :]\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_fft=1024, n_mels=128, fmin=0, fmax=20, win_length=128)\n\n                # LOG TRANSFORM\n                # Log transform and normalize\n                #width = (Sx_hann.shape[1]//32)*32\n                #Sx_hann = librosa.power_to_db(Sx_hann, ref=np.max).astype(np.float32)[:,:width]\n\n                # STANDARDIZE TO -1 TO 1\n                #Sx_hann = (Sx_hann+40)/40 \n                img[:,:,k] += Sx_hann\n                \n                out = pycwt(S).numpy()\n                #cwt_img[:,:,k] +=out[:,:,0]\n                #out1 = librosa.power_to_db(out, ref=np.max).astype(np.float32)\n                #out1 = (out1+40)/40 \n                # Log transform and normalize\n                cwt_img1[:,:,k] +=out[:,:,0]\n            #1\n            # AVERAGE THE 4 MONTAGE DIFFERENCES\n            img[:,:,k] /= 4.0\n            #cwt_img[:,:,k] /= 4.0\n            cwt_img1[:,:,k] /= 4.0\n\n        # Create the final combined image\n        #cwt_spec = np.concatenate([cwt_img[:,:,0],cwt_img[:,:,1],cwt_img[:,:,2],cwt_img[:,:,3]])\n        '''\n        plt.subplot(1,3,2)\n        plt.imshow(cwt_spec,aspect=\"auto\", cmap='jet')\n        plt.title(\"eeg → CWT\")\n        plt.axis('off')  \n        '''\n        #spec = np.concatenate([img[:,:,0],img[:,:,1],img[:,:,2],img[:,:,3]])\n        '''\n        plt.subplot(1, 3, 3)\n        plt.imshow(spec,aspect=\"auto\",origin='lower',cmap='jet')\n        plt.title(\"eeg → STFT\")\n        plt.axis('off')  \n        plt.tight_layout()\n        print('img',img.shape)\n        print('cwt_img', cwt_img.shape)\n        print('kaggle',img_kaggle.shape)\n        '''\n        \n        final_image = create_final_image(img, cwt_img1, img_kaggle)\n    \n        # Save the final image in compressed format\n        file_path = os.path.join(output_folder, f'{eeg_new_key}.npz')\n        np.savez_compressed(file_path, final_image=final_image)\n        if i < 3:  # Plot the final image for the first three iterations\n            # Plotting the final image\n            plt.figure(figsize=(10, 10))\n            plt.imshow(final_image, cmap='jet')\n            plt.title(\"Final Composite Image\")\n            plt.axis('off')\n            plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:32:05.548002Z","iopub.execute_input":"2024-04-20T14:32:05.548662Z","iopub.status.idle":"2024-04-20T14:32:05.574408Z","shell.execute_reply.started":"2024-04-20T14:32:05.548618Z","shell.execute_reply":"2024-04-20T14:32:05.572837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"process_eegs(train_df, 'images')","metadata":{"execution":{"iopub.status.busy":"2024-04-20T14:32:05.577Z","iopub.execute_input":"2024-04-20T14:32:05.577484Z","iopub.status.idle":"2024-04-20T14:32:30.901351Z","shell.execute_reply.started":"2024-04-20T14:32:05.577443Z","shell.execute_reply":"2024-04-20T14:32:30.899264Z"},"trusted":true},"execution_count":null,"outputs":[]}]}