{"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":7933212,"sourceType":"datasetVersion","datasetId":4461503}],"dockerImageVersionId":30674,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport cv2\nfrom matplotlib import pyplot as plt\nimport albumentations as A\n\nimport torch as tc\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport timm\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom scipy.special import kl_div\nfrom scipy.signal import butter, lfilter\n\nfrom tqdm import tqdm","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# GLOBAL","metadata":{}},{"cell_type":"code","source":"path = \"/kaggle/input/hms-harmful-brain-activity-classification\"\nbatch_size = 16\ndevice = \"cuda\" if tc.cuda.is_available() else \"cpu\"\ndevice\ndomain = \"test\"\nread_all_eegs = False\nread_all_specs = False\nread_all_l2i = False\nis_augument = False\nis_test_random = True","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(path+f\"/{domain}.csv\")\n\nif domain==\"train\":\n    train_df =  df.copy()\n    \n    train_df = train_df[(train_df.iloc[:,-6:].sum(1)>6)]\n    \n    label_cols = train_df.columns[-6:]\n    eeg_ids = train_df.eeg_id.unique()\n    train_df = df.groupby('eeg_id')[['patient_id']].agg('first')\n    aux = df.groupby('eeg_id')[label_cols].agg('sum')\n    si = df.groupby('eeg_id')[[\"spectrogram_id\", \"spectrogram_label_offset_seconds\"]].agg('first')\n\n    for k in si:\n        train_df[k] = si[k].values\n\n    for label in label_cols:\n        train_df[label] = aux[label].values\n\n    y_data = train_df[label_cols].values\n    y_data = y_data / y_data.sum(axis=1,keepdims=True)\n    train_df[label_cols] = y_data\n\n    train_df = train_df.reset_index()\n    train_df = train_df.loc[train_df.eeg_id.isin(eeg_ids)]\n    print(f\"Train dataframe with unique eeg_id has shape: {train_df.shape}\")\n    display(train_df.head())\n\n    \n    train_df.iloc[:,-6:].sum(0).plot.bar()\n    plt.show()\n    df = train_df\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def softmax(d,dim):\n    ed = np.exp(d)\n    return ed/(ed.sum(axis=dim, keepdims=True))\n\ndef butter_lowpass_filter(data, cutoff_freq=40, sampling_rate=200, 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\ndef butter_highpass_filter(data, cutoff_freq=1, sampling_rate=200, order=4):\n    nyquist = 0.5 * sampling_rate\n    normal_cutoff = cutoff_freq / nyquist\n    b, a = butter(order, normal_cutoff, btype='high', analog=False)\n    filtered_data = lfilter(b, a, data, axis=0)\n    return filtered_data\n\ndef get_eeg_from_parquet(idx, metadata, eeg_ids):\n    eeg_id = eeg_ids[idx]\n    eeg_path = path + \"/train_eegs/\" + str(eeg_id) + \".parquet\"\n    eeg = pd.read_parquet(eeg_path).iloc[0:10_000,:]\n    eeg = eeg.values.T\n    return eeg\n\ndef get_spec_from_parquet(idx, metadata, eeg_ids):\n    eeg_id = eeg_ids[idx]\n    eeg_path = path + \"/train_eegs/\" + str(eeg_id) + \".parquet\"\n    occurances_df = metadata.loc[metadata[\"eeg_id\"]==eeg_id,:]\n    \n    spec_id = occurances_df.loc[:,\"spectrogram_id\"].values[0]\n    spec_offset = int(occurances_df.loc[:,\"spectrogram_label_offset_seconds\"].values[0]//2)\n    spec_path = path + \"/train_spectrograms/\" + str(spec_id) + \".parquet\"\n    x = pd.read_parquet(spec_path)\n    x = x.values[spec_offset:300+spec_offset,1:].reshape(300,4,100).transpose((1,2,0))\n    return x\n\ndef process_spec(x):\n    x = np.clip(x,np.exp(-4),np.exp(8))\n    x = np.log(x)\n    \n    mu = np.nanmean(x)\n    sigma = np.nanstd(x)\n    x = (x-mu)/(sigma+1e-5)\n    x = np.nan_to_num(x, nan=0.0)\n    return x\n    \ndef get_eeg_frame():\n    class TrainingDatasetEEG(Dataset):\n        def __init__(self, metadata):\n            self.metadata = metadata\n            self.eeg_ids = np.array(sorted(metadata[\"eeg_id\"].unique()), dtype=np.int64).squeeze()\n            print()\n            \n        def __len__(self):\n            return self.eeg_ids.shape[0]\n\n        def __getitem__(self, idx):\n            eeg = get_eeg_from_parquet(idx, self.metadata, self.eeg_ids)\n\n            is_corrupted = 0\n            for i in range(eeg.shape[0]):\n                nans = np.isnan(eeg[i]).sum()\n                if nans/1e4 > 0.7:\n                    is_corrupted = 1\n            \n            eeg = np.nan_to_num(eeg, nan=0)\n            eeg = eeg.astype(np.float32)\n            eeg = tc.from_numpy(eeg).type(tc.float32)\n            \n            return eeg,is_corrupted\n    \n    bs=4\n    train_dataset = TrainingDatasetEEG(df)\n    train_dataloader = DataLoader(train_dataset, batch_size=bs, shuffle=False, num_workers=4)\n    print(len(train_dataset))\n    print(len(train_dataloader))\n    all_eegs = np.zeros((df[\"eeg_id\"].unique().shape[0],20,10000))\n    all_is_corrupted_list = np.zeros((df[\"eeg_id\"].unique().shape[0],))\n    j = 0\n    for i, (eegs,is_corrupted_list) in enumerate(tqdm(train_dataloader), 0):\n        for (eeg,is_corrupted) in zip(eegs,is_corrupted_list):\n            all_eegs[j] = eeg\n            all_is_corrupted_list[j] = is_corrupted\n            j+=1\n    return all_eegs, all_is_corrupted_list\n\ndef get_spec_frame():\n    class TrainingDatasetEEG(Dataset):\n        def __init__(self, metadata):\n            self.metadata = metadata\n            self.eeg_ids = np.array(sorted(metadata[\"eeg_id\"].unique()), dtype=np.int64).squeeze()\n            print()\n            \n        def __len__(self):\n            return self.eeg_ids.shape[0]\n\n        def __getitem__(self, idx):\n            spec = get_spec_from_parquet(idx, self.metadata, self.eeg_ids)\n            spec = process_spec(spec)\n            sepc = tc.from_numpy(np.ascontiguousarray(spec)).type(tc.float32)\n            return spec\n    \n    bs=4\n    train_dataset = TrainingDatasetEEG(df)\n    train_dataloader = DataLoader(train_dataset, batch_size=bs, shuffle=False, num_workers=4)\n    print(len(train_dataset))\n    print(len(train_dataloader))\n    all_specs = np.zeros((df[\"eeg_id\"].unique().shape[0],4,100,300))\n    j = 0\n    for i, specs in enumerate(tqdm(train_dataloader), 0):\n        for spec in specs:\n            all_specs[j] = spec\n            j+=1\n    return all_specs\n\ndef line2img(x):\n    C,L = x.shape\n\n    mx=x.max(axis=1, keepdims=True)\n    mn=x.min(axis=1, keepdims=True)\n    nx = (x-mn)/(mx-mn+1e-5)\n    \n    sw = 3\n    H,W = 128-sw,1024-sw\n    \n    h_inds = (nx*(H-1)).astype(np.int32)\n    w_inds = np.linspace(0,W-1,L)[None,:].repeat(C, axis=0).astype(np.int32)\n    chans = np.array(range(C))[:,None]\n    \n    img = np.zeros((C,H+sw,W+sw), dtype=np.uint8)\n    for i in range(sw):\n        for j in range(sw):\n            img[chans,h_inds+i,w_inds+j] = 1\n            \n    img = cv2.resize(img.transpose(1,2,0), (512,64), interpolation=cv2.INTER_AREA).transpose(2,0,1)\n    return img\n\ndef eeg2img():\n    names = ['Fp1','F3','C3','P3','F7','T3','T5','O1','Fz','Cz','Pz','Fp2','F4','C4','P4','F8','T4','T6','O2','EKG']\n    pairs = list(zip(range(20),names))\n    pairs = {k.lower():v for v,k in pairs}\n    cen_locs = ['fp1,fp1,fp1,fp2,fp2,fp2'.split(\",\"),'f7,f3,fz,fz,f4,f8'.split(\",\"), 't3,c3,cz,cz,c4,t4'.split(\",\"), 't5,p3,pz,pz,p4,t6'.split(\",\"),'o1,o1,o1,o2,o2,o2'.split(\",\")]\n    cen_ilocs = [pairs[i] for i in sum(cen_locs,[])]\n    cen_ilocs = np.array(cen_ilocs).reshape(5,6)\n    cen_ilocs.astype(np.int32)\n    \n    def _eeg2img(idx, metadata, eeg_ids):\n        eeg = get_eeg_from_parquet(idx, metadata, eeg_ids)\n        \n        x2_ = eeg[cen_ilocs,4000:6000] #5,6,L\n        x2_ = np.split(x2_, 5, axis=0) #list of 5 x (1,6,L)\n        x2_ = np.concatenate([x2_[i]-x2_[i+1] for i in range(4)],0) #4,6,L\n        x2_ =  x2_[:,[0,1,4,5]].reshape(4*4,-1)\n        x2 = butter_lowpass_filter(x2_.T, cutoff_freq=40).T\n        x2 = np.nan_to_num(x2, nan=0)\n        x2 = np.clip(x2, -1024,1024)\n        imgs = line2img(x2)\n        imgs = (imgs-0.5)/0.5\n        return imgs\n    return _eeg2img\n\ndef get_eegimg_frame():\n    class TrainingDatasetEEG(Dataset):\n        def __init__(self, metadata):\n            self.metadata = metadata\n            self.eeg_ids = np.array(sorted(metadata[\"eeg_id\"].unique()), dtype=np.int64).squeeze()\n            self.func = eeg2img()\n            print()\n            \n        def __len__(self):\n            return self.eeg_ids.shape[0]\n\n        def __getitem__(self, idx):\n            imgs = self.func(idx, self.metadata, self.eeg_ids)\n            imgs = tc.from_numpy(np.ascontiguousarray(imgs)).type(tc.uint8)\n            return imgs\n    \n    bs=4\n    train_dataset = TrainingDatasetEEG(df)\n    train_dataloader = DataLoader(train_dataset, batch_size=bs, shuffle=False, num_workers=4)\n    print(len(train_dataset))\n    print(len(train_dataloader))\n    all_imgs = np.zeros((df[\"eeg_id\"].unique().shape[0],16,64,512), dtype=np.uint8)\n    j = 0\n    for i, imgs in enumerate(tqdm(train_dataloader), 0):\n        for img in imgs:\n            all_imgs[j] = img\n            j+=1\n    return all_imgs\n        \nif read_all_eegs:\n    if domain == 'train':\n        all_eegs, all_is_corrupted_list = get_eeg_frame()\n\nif read_all_specs:\n    if domain == 'train':\n        all_specs = get_spec_frame()\n        \nif read_all_l2i:\n    if domain == 'train':\n        all_imgs = get_eegimg_frame()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TrainingDatasetEEG(Dataset):\n    def __init__(self, metadata, train=True, is_test_random=False):\n        self.train = train\n        self.metadata = metadata\n        \n        neeg_ids = np.array(sorted(metadata[\"eeg_id\"].unique()), dtype=np.int64).squeeze()\n        if read_all_eegs:\n            eeg_ids = np.array([eeg_id for eeg_id,is_corrupted in zip(neeg_ids,all_is_corrupted_list)]).astype(np.int64)\n        else:\n            eeg_ids = neeg_ids\n        \n        num_eeg_ids = eeg_ids.shape[0]\n        train_size_percentage = 85 if not is_test_random else 100\n        train_set_size = int(train_size_percentage/100 * num_eeg_ids)\n        self.train_set_size = train_set_size\n        eeg_ids_train = eeg_ids[:train_set_size]\n        eeg_ids_valid = eeg_ids[train_set_size:] if not is_test_random else np.random.choice(eeg_ids_train, 500)\n        self.eeg_ids = eeg_ids_train if train else eeg_ids_valid\n        \n        names = ['Fp1','F3','C3','P3','F7','T3','T5','O1','Fz','Cz','Pz','Fp2','F4','C4','P4','F8','T4','T6','O2','EKG']\n\n        pairs = list(zip(range(20),names))\n        pairs = {k.lower():v for v,k in pairs}\n        cen_locs = ['fp1,fp1,fp1,fp2,fp2,fp2'.split(\",\"),'f7,f3,fz,fz,f4,f8'.split(\",\"), 't3,c3,cz,cz,c4,t4'.split(\",\"), 't5,p3,pz,pz,p4,t6'.split(\",\"),'o1,o1,o1,o2,o2,o2'.split(\",\")]\n        cen_ilocs = [pairs[i] for i in sum(cen_locs,[])]\n        cen_ilocs = np.array(cen_ilocs).reshape(5,6)\n        cen_ilocs.astype(np.int32)\n        self.cen_ilocs = cen_ilocs\n        \n        self.transform = A.Compose([\n                            A.Blur(blur_limit=3, p=0.2),\n                            A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.2),\n                            A.GaussNoise(var_limit=(10.0, 50.0), p=0.2),\n                            A.OneOf([\n                                A.CoarseDropout(max_holes=12, max_height=60, max_width=64, p=0.05),\n                                A.GridDropout(p=0.05),\n                            ], p=0.1),\n                            \n                            A.ElasticTransform(p=0.05),\n                            A.GridDistortion(p=0.05),\n                            A.OpticalDistortion(p=0.05)\n                        ], p=0.2)\n        \n    def __len__(self):\n        return self.eeg_ids.shape[0]\n        \n    def __getitem__(self, idx):\n        eeg_id = self.eeg_ids[idx]\n        idx = idx if self.train else idx+self.train_set_size\n        occurances_df = self.metadata.loc[self.metadata[\"eeg_id\"]==eeg_id,:]\n        \n        if read_all_eegs:\n            eeg = all_eegs[idx].copy()\n        else:\n            eeg = get_eeg_from_parquet(idx-self.train_set_size, self.metadata, self.eeg_ids)\n            \n            for i in range(eeg.shape[0]):\n                nans = np.isnan(eeg[i]).sum()\n                if nans/1e4 > 0.7:\n                    nidx = idx+1\n                    if nidx > self.__len__():\n                        nidx = 0\n                    return self.__getitem__(nidx)\n            \n            eeg = np.nan_to_num(eeg, nan=0)\n        \n        cutoff_freq = 40\n        gn = 0.0\n        direction = 1\n        polarity = 1.\n        \n        if is_augument:\n            if self.train:\n                #butter pass\n                if np.random.rand()>0.8:\n                    cutoff_freq = np.random.choice(range(16,30,4), [])\n\n                #Random noise\n                if np.random.rand()>0.8:\n                    gnmx = eeg.mean(1)[:,None]*np.random.choice((0.025,0.05,0.1), [])\n                    polarity = np.random.choice((-1.,1.), eeg[0].shape)[None, :]\n                    gn = polarity * np.random.random(eeg[0].shape)[None,:] * gnmx\n\n                #Random reversal\n                if np.random.rand()>0.8:\n                    direction = -1\n\n                #Random polarity reversal\n                if np.random.rand()>0.8:\n                    polarity = np.random.choice((-1.,1.), [])\n        \n        eeg = eeg[:,::direction] + gn\n        eeg *= polarity\n        \n        x1 = butter_lowpass_filter(eeg.T, cutoff_freq=cutoff_freq).T\n        x1 = np.clip(x1, -1024,1024)\n        \n        x2_ = eeg[self.cen_ilocs,:] #5,6,L\n        x2_ = np.split(x2_, 5, axis=0) #list of 5 x (1,6,L)\n        x2_ = np.concatenate([x2_[i]-x2_[i+1] for i in range(4)],0) #4,6,L\n        x2_ =  x2_.reshape(4*6,-1)\n        \n        x2 = butter_lowpass_filter(x2_.T, cutoff_freq=cutoff_freq).T\n        x2 = np.clip(x2, -1024,1024)\n        \n        x3_ = eeg[self.cen_ilocs[[0,2,4],:][:,[0,1,4,5]],:] #3,4,L\n        x3_ = np.split(x3_, 3, axis=0) #list of 3 x (1,4,L)\n        x3_ = np.concatenate([x3_[i]-x3_[i+1] for i in range(2)],0) #2,4,L\n        x3_ =  x3_.reshape(2*4,-1)\n        \n        x3 = butter_lowpass_filter(x3_.T, cutoff_freq=cutoff_freq).T\n        x3 = np.clip(x3, -1024,1024)\n\n        if read_all_specs:\n            x4 = all_specs[idx]\n        else:\n            x4 = get_spec_from_parquet(idx-self.train_set_size, self.metadata, self.eeg_ids)\n            x4 = process_spec(x4)\n            \n        if read_all_l2i:\n            x5 = all_imgs[idx]\n        else:\n            x5 = x2[:,4000:6000].reshape(4,6,-1)[:,[0,1,4,5]].reshape(4*4,-1)\n            x5 = line2img(x5).astype(np.float32)\n            \n        if is_augument:\n            if self.train:\n                for i in range(16):\n                    augmented = self.transform(image=x5[i])\n                    x5[i] = augmented['image']\n                    \n        x5 = (x5-0.5)/0.5\n        \n        targets = occurances_df.iloc[:,-6:].values.squeeze()\n        \n        x1 = tc.from_numpy(np.ascontiguousarray(x1)).type(tc.float32)/128.\n        x2 = tc.from_numpy(np.ascontiguousarray(x2)).type(tc.float32)/32.\n        x3 = tc.from_numpy(np.ascontiguousarray(x3)).type(tc.float32)\n        x4 = tc.from_numpy(np.ascontiguousarray(x4)).type(tc.float32)\n        x5 = tc.from_numpy(np.ascontiguousarray(x5)).type(tc.float32)\n        targets = tc.from_numpy(targets).type(tc.float32)\n        return (x1,x2,x3,x4,x5), targets.squeeze()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DatasetEEG(Dataset):\n    def __init__(self, metadata):\n        self.metadata = metadata\n        names = ['Fp1','F3','C3','P3','F7','T3','T5','O1','Fz','Cz','Pz','Fp2','F4','C4','P4','F8','T4','T6','O2','EKG']\n\n        pairs = list(zip(range(20),names))\n        pairs = {k.lower():v for v,k in pairs}\n        cen_locs = ['fp1,fp1,fp1,fp2,fp2,fp2'.split(\",\"),'f7,f3,fz,fz,f4,f8'.split(\",\"), 't3,c3,cz,cz,c4,t4'.split(\",\"), 't5,p3,pz,pz,p4,t6'.split(\",\"),'o1,o1,o1,o2,o2,o2'.split(\",\")]\n        cen_ilocs = [pairs[i] for i in sum(cen_locs,[])]\n        cen_ilocs = np.array(cen_ilocs).reshape(5,6)\n        cen_ilocs.astype(np.int32)\n        self.cen_ilocs = cen_ilocs\n        \n    def __len__(self):\n        return self.metadata.shape[0]\n        \n    def __getitem__(self, idx):\n\n        eeg_id = self.metadata.loc[idx, \"eeg_id\"]\n        eeg_path = path + \"/test_eegs/\" + str(eeg_id) + \".parquet\"\n        eeg = pd.read_parquet(eeg_path).iloc[:10_000,:]\n        eeg = eeg.values.T\n        eeg = np.nan_to_num(eeg, nan=0)\n        \n        cutoff_freq = 40\n        \n        x1 = butter_lowpass_filter(eeg.T, cutoff_freq=cutoff_freq).T\n        x1 = np.clip(x1, -1024,1024)\n        \n        x2_ = eeg[self.cen_ilocs,:] #5,6,L\n        x2_ = np.split(x2_, 5, axis=0) #list of 5 x (1,6,L)\n        x2_ = np.concatenate([x2_[i]-x2_[i+1] for i in range(4)],0) #4,6,L\n        x2_ =  x2_.reshape(4*6,-1)\n        \n        x2 = butter_lowpass_filter(x2_.T, cutoff_freq=cutoff_freq).T\n        x2 = np.clip(x2, -1024,1024)\n        \n        x3_ = eeg[self.cen_ilocs[[0,2,4],:][:,[0,1,4,5]],:] #3,4,L\n        x3_ = np.split(x3_, 3, axis=0) #list of 3 x (1,4,L)\n        x3_ = np.concatenate([x3_[i]-x3_[i+1] for i in range(2)],0) #2,4,L\n        x3_ =  x3_.reshape(2*4,-1)\n        \n        x3 = butter_lowpass_filter(x3_.T, cutoff_freq=cutoff_freq).T\n        x3 = np.clip(x3, -1024,1024)\n\n        spec_id = self.metadata.loc[idx, \"spectrogram_id\"]\n        spec_path = path + \"/test_spectrograms/\" + str(spec_id) + \".parquet\"\n        x4 = pd.read_parquet(spec_path)\n        x4 = x4.values[:300,1:].reshape(300,4,100).transpose((1,2,0))\n        x4 = process_spec(x4)\n            \n\n        x5 = x2[:,4000:6000].reshape(4,6,-1)[:,[0,1,4,5]].reshape(4*4,-1)\n        x5 = line2img(x5).astype(np.float32)                    \n        x5 = (x5-0.5)/0.5\n        \n        x1 = tc.from_numpy(np.ascontiguousarray(x1)).type(tc.float32)/128.\n        x2 = tc.from_numpy(np.ascontiguousarray(x2)).type(tc.float32)/32.\n        x3 = tc.from_numpy(np.ascontiguousarray(x3)).type(tc.float32)\n        x4 = tc.from_numpy(np.ascontiguousarray(x4)).type(tc.float32)\n        x5 = tc.from_numpy(np.ascontiguousarray(x5)).type(tc.float32)\n        return x1,x2,x3,x4,x5,eeg_id","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"import torch\np_dropout = 0.3 if domain==\"train\" else 0\n\nclass BBlock(nn.Module):\n    def __init__(self,o,g,s=2,h=None, *args, **kwargs) -> None:\n        super().__init__(*args, **kwargs)\n        \n        self.c = nn.Sequential(nn.Conv1d(o,o*s,3,padding=1,groups=g),nn.BatchNorm1d(o*s),nn.SiLU(),\n                                nn.Conv1d(o*s,o*s,3,padding=1,groups=g),nn.BatchNorm1d(o*s),nn.SiLU(),\n                                nn.Conv1d(o*s,o*s,3,padding=1,groups=g))#1000\n        \n        self.c_res = nn.Conv1d(o,o*s,1)\n        self.c_add_norm = nn.Sequential(nn.BatchNorm1d(o*s),nn.SiLU())\n        o_ = o if h is None else h\n        self.c_final = nn.Sequential(nn.Conv1d(o*s,o_,3,padding=1,groups=g),nn.BatchNorm1d(o_),nn.SiLU())\n    \n    def forward(self, x):\n        c = self.c(x) #NCS\n        c_res = self.c_res(x)\n        c_norm = self.c_add_norm(c+c_res)\n        c_final = self.c_final(c_norm)\n        return c_final\n    \nclass BBlock2(nn.Module):\n    def __init__(self,o,g,s=2,h=None, is_strided=False, stride=None, *args, **kwargs) -> None:\n        super().__init__(*args, **kwargs)\n        \n        self.c = nn.Sequential(nn.Conv2d(o,o*s,3,padding=1,groups=g),nn.BatchNorm2d(o*s),nn.SiLU(),\n                                nn.Conv2d(o*s,o*s,3,padding=1,groups=g),nn.BatchNorm2d(o*s),nn.SiLU(),\n                                nn.Conv2d(o*s,o*s,3,padding=1,groups=g))\n        \n        self.c_res = nn.Conv2d(o,o*s,1,groups=g)\n        self.c_add_norm = nn.Sequential(nn.BatchNorm2d(o*s),nn.SiLU())\n        o_ = o if h is None else h\n        self.c_final = nn.Sequential(nn.Conv2d(o*s,o_,3,stride if is_strided else 1, padding=1,groups=g),nn.BatchNorm2d(o_),nn.SiLU())\n    \n    def forward(self, x):\n        c = self.c(x) #NCS\n        c_res = self.c_res(x)\n        c_norm = self.c_add_norm(c+c_res)\n        c_final = self.c_final(c_norm)\n        return c_final\n    \nclass TBlock(nn.Module):\n    def __init__(self,o,g,s=2, *args, **kwargs) -> None:\n        super().__init__(*args, **kwargs)\n        \n        self.c = nn.Sequential(nn.Conv1d(o,o*s,3,padding=1,groups=g),nn.Tanh(),\n                                nn.Conv1d(o*s,o*s,3,padding=1,groups=g),nn.Tanh(),\n                                nn.Conv1d(o*s,o*s,3,padding=1,groups=g))#1000\n        \n        self.c_res = nn.Conv1d(o,o*s,1, groups=g)\n        self.c_add_norm = nn.Tanh()\n        self.c_final = nn.Sequential(nn.Conv1d(o*s,o,3,padding=1,groups=g), nn.Tanh())\n    \n    def forward(self, x):\n        c = self.c(x) #NCS\n        c_res = self.c_res(x)\n        c_norm = self.c_add_norm(c+c_res)\n        c_final = self.c_final(c_norm)\n        return c_final\n\nclass EEGFeatureExtractor(nn.Module):\n    def __init__(self, *args, **kwargs) -> None:\n        super().__init__(*args, **kwargs)\n        \n        self.c1 = TBlock(20,20,s=2)\n        \n        self.c2 = nn.Sequential(\n            nn.Conv1d(190,190,1,groups=190),nn.Tanh(),nn.AvgPool1d(5),\n            nn.Conv1d(190,32,3,padding=1,groups=1),nn.BatchNorm1d(32),nn.SiLU())\n        \n        self.c3 = nn.Sequential(\n            BBlock(32,1,s=2, h=16),\n            BBlock(16,1,s=2, h=24),\n            nn.AvgPool1d(5),\n            \n            BBlock(24,1,s=2, h=32),\n            BBlock(32,1,s=2, h=48),\n            nn.AvgPool1d(5),\n            \n            BBlock(48,1,s=2, h=64),\n            BBlock(64,1,s=2, h=72),\n            nn.AvgPool1d(5),\n            \n            BBlock(72,1,s=2, h=96)\n        )\n        \n        self.d1 = nn.Sequential(\n            BBlock(1,1,s=4,h=4),\n            nn.AvgPool1d(5),\n            BBlock(4,1,s=2,h=8),\n            nn.AvgPool1d(5),\n            BBlock(8,1,s=2,h=12),\n            nn.AvgPool1d(5),\n            BBlock(12,1,s=2,h=16),\n            nn.AvgPool1d(5),\n            BBlock(16,1,s=2,h=24),\n        )\n        \n        self.d2 = nn.Sequential(nn.Dropout(p_dropout), nn.Linear(480,128), nn.SiLU())\n        \n        self.e1 = nn.Sequential(\n            BBlock(1,1,s=4,h=4),\n            nn.AvgPool1d(5),\n            BBlock(4,1,s=2,h=8),\n            nn.AvgPool1d(5),\n            BBlock(8,1,s=2,h=12),\n            nn.AvgPool1d(5),\n            BBlock(12,1,s=2,h=16),\n            nn.AvgPool1d(5),\n            BBlock(16,1,s=2,h=24),\n        )\n        self.e2 = nn.Sequential(nn.Dropout(p_dropout), nn.Linear(576,128), nn.SiLU())\n        \n        \n        self.f1 = nn.Sequential(\n            BBlock(4,1,s=4,h=8),\n            nn.AvgPool1d(5),\n            BBlock(8,1,s=2,h=12),\n            nn.AvgPool1d(5),\n            BBlock(12,1,s=2,h=16),\n            nn.AvgPool1d(5),\n            BBlock(16,1,s=2,h=24),\n            nn.AvgPool1d(5),\n            BBlock(24,1,s=2,h=32),\n        )\n        self.f2 = nn.Sequential(nn.Dropout(p_dropout), nn.Linear(192,96), nn.SiLU())\n        \n        self.g1 = nn.Sequential(\n            BBlock(1,1,s=4,h=8),\n            nn.AvgPool1d(2),\n            BBlock(8,1,s=2,h=16),\n            nn.AvgPool1d(2),\n            BBlock(16,1,s=2,h=32),\n            nn.AvgPool1d(2),\n            BBlock(32,1,s=2,h=64),\n            nn.AvgPool1d(2),\n            BBlock(64,1,s=2,h=128),\n            nn.AvgPool1d(2),\n            BBlock(128,1,s=1,h=256),\n            BBlock(256,1,s=1,h=324),\n            BBlock(324,1,s=1,h=512),\n        )\n        self.g2 = nn.Sequential(nn.Dropout(p_dropout), nn.Linear(512*8,512), nn.SiLU())\n        \n        self.h1 = nn.Sequential(\n            BBlock2(1,1,s=2,h=8),nn.AvgPool2d(2),\n            BBlock2(8,1,s=2,h=12),nn.AvgPool2d(2),\n            BBlock2(12,1,s=2,h=24),nn.AvgPool2d(2),\n            BBlock2(24,1,s=2,h=32),nn.AvgPool2d(2),\n            BBlock2(32,1,s=2,h=48),nn.AvgPool2d(2),\n            BBlock2(48,1,s=2,h=64),\n        )\n        self.h2 = nn.Sequential(nn.Dropout(p_dropout), nn.Linear(64*4,256), nn.SiLU())\n        \n        self.i1 = nn.Sequential(\n            BBlock2(1,1,s=2,h=8, is_strided=True, stride=(1,2)),\n            BBlock2(8,1,s=2,h=12, is_strided=True, stride=(1,2)),\n            BBlock2(12,1,s=2,h=24),nn.AvgPool2d(2),\n            BBlock2(24,1,s=2,h=32),nn.AvgPool2d(2),\n            BBlock2(32,1,s=2,h=48),nn.AvgPool2d(2),\n            BBlock2(48,1,s=2,h=64),\n        )\n        self.i2 = nn.Sequential(nn.Dropout(p_dropout), nn.Linear(64*16,512), nn.SiLU())\n        \n        self.j1 = timm.create_model('tf_efficientnet_b0', pretrained=False, in_chans=4)\n        self.j1.classifier = nn.Linear(self.j1.classifier.in_features, 512)\n        self.j2 = nn.Sequential(nn.Dropout(p_dropout), nn.Linear(512*4,512), nn.SiLU())\n        \n    def ind_feature_extractor(self, x):\n        d1 = tc.cat([self.d1(x[:,i,:].unsqueeze(1)) for i in range(x.shape[1])], 1)\n        d2 = self.d2(d1.mean(-1))\n        return d2\n    \n    def total_feature_extractor(self, x):\n        c1 = self.c1(x) #NCS        \n        c2 = []\n        for i in range(20):\n            X = c1[:,i,:]\n            for j in range(i+1,20):\n                x = X-c1[:,j,:]\n                c2.append(x)\n\n        c2 = self.c2(F.tanh(tc.stack(c2,1)))\n        c3 = self.c3(c2)\n        return c3.mean(-1)\n    \n    def montage_features_extractor(self, x):\n        e1 = tc.cat([self.e1(x[:,i,:].unsqueeze(1)) for i in range(x.shape[1])],1)\n        e2 =  self.e2(e1.mean(-1))\n        return e2\n    \n    def montage_features_extractor2(self, x):\n        with tc.no_grad():\n            x = self.norm(x, 2)\n            \n        g1 = tc.cat([self.g1(x[:,i,:].unsqueeze(1)) for i in range(x.shape[1])],1)\n        g2 =  self.g2(g1.mean(-1))\n        return g2\n    \n    \n    def norm(self, x, axis=1):\n        mx = x.max(axis, keepdims=True)[0]\n        mn = x.min(axis, keepdims=True)[0]\n        n = (x-mn)/(mx-mn+1e-5)\n        return n\n    \n    def gr_montage_features_extractor(self, x):            \n        N,S,L = x.shape\n        x =  x.view(N,4,6,L)\n        \n        f1 = tc.cat([self.f1(x[:,:,i,:]).mean(-1) for i in range(6)],1)\n        f2 =  self.f2(f1)\n        return f2\n    \n    def kagg_spec_feature_extractor(self, x):\n        if is_augument:\n            with tc.no_grad():\n                if self.train:\n                    if tc.rand([])>0.85:\n                        x = tc.flip(x, dims=[-2])\n\n                    if tc.rand([])>0.85:\n                        x = tc.flip(x, dims=[-1])\n                        \n        h1 = tc.cat([self.h1(x[:,i,...].unsqueeze(1)).mean((-2,-1)) for i in range(4)],1)\n        h2 = self.h2(h1)\n        return h2\n    \n    def eeg_img_feature_extractor(self, x):\n        if is_augument:\n            with tc.no_grad():\n                if self.train:\n                    if tc.rand([])>0.85:\n                        x = tc.flip(x, dims=[-2])\n\n                    if tc.rand([])>0.85:\n                        x = tc.flip(x, dims=[-1])\n                        \n        i1 = tc.cat([self.i1(x[:,i,...].unsqueeze(1)).mean((-2,-1)) for i in range(16)],1)\n        i2 = self.i2(i1)\n        return i2\n    \n    def eeg_img_feature_extractor2(self, x):\n        with tc.no_grad():\n            N,C,H,W = x.shape\n            if is_augument:\n                if self.train:\n                    if tc.rand([])>0.85:\n                        x = tc.flip(x, dims=[-2])\n\n                    if tc.rand([])>0.85:\n                        x = tc.flip(x, dims=[-1])\n                    \n            x = x.view(N,4,4,H,W)\n            \n        j1 = tc.cat([self.j1(x[:,:,i,...]) for i in range(4)],1)\n        j2 = self.j2(j1)\n        return j2\n    \n    def train_branch(self,func,x,randomize=False):\n        if randomize:\n            if tc.rand([])>0.5:\n                with tc.no_grad():\n                    return func(x)\n            else:\n                return func(x)\n        else:\n            return func(x)\n            \n    \n    def forward(self, x1,x2,x3,x4,x5):\n        #N,C,S\n        h1 = self.train_branch(self.total_feature_extractor,x1, True if self.train else False)\n        h2 = self.train_branch(self.ind_feature_extractor,x1, True if self.train else False)\n        h3 = self.train_branch(self.montage_features_extractor,x2, True if self.train else False)\n        h4 = self.train_branch(self.gr_montage_features_extractor,x2, True if self.train else False)\n        h5 = self.train_branch(self.montage_features_extractor2,x3, True if self.train else False)\n        h6 = self.train_branch(self.kagg_spec_feature_extractor,x4, True if self.train else False)\n        h7 = self.train_branch(self.eeg_img_feature_extractor,x5, True if self.train else False)\n        h8 = self.train_branch(self.eeg_img_feature_extractor2,x5, True if self.train else False)\n        \n        \n#         h2 = self.ind_feature_extractor(x1)\n#         h3 = self.montage_features_extractor(x2)\n#         h4 = self.gr_montage_features_extractor(x2)\n#         h5 = self.montage_features_extractor2(x3)\n#         h6 = self.kagg_spec_feature_extractor(x4)\n#         h7 = self.eeg_img_feature_extractor(x5)\n#         h8 = self.eeg_img_feature_extractor2(x5)\n        h = tc.cat([h1,h2,h3,h4,h5,h6,h7,h8],1)\n        return h\n        \nclass EEGBasedClassifier(nn.Module):\n    def __init__(self, load_weights=True, load_out_weights=True, *args, **kwargs) -> None:\n        super().__init__(*args, **kwargs)\n        self.feature_extractor = EEGFeatureExtractor()\n        self.out = nn.Sequential(\n            nn.Dropout(p_dropout),\n            \n            nn.Linear(96+128+128+96+512+256+512+512, 1024),\n            nn.BatchNorm1d(1024),\n            nn.SiLU(),\n            nn.Dropout(p_dropout),\n            \n            nn.Linear(1024, 768),\n            nn.BatchNorm1d(768),\n            nn.SiLU(),\n            nn.Dropout(p_dropout),\n            \n            nn.Linear(768, 512),\n            nn.BatchNorm1d(512),\n            nn.SiLU(),\n            nn.Dropout(p_dropout),\n            \n            nn.Linear(512, 512),\n            nn.BatchNorm1d(512),\n            nn.SiLU(),\n            nn.Dropout(p_dropout),\n            \n            nn.Linear(512, 6)\n        )\n        \n        if load_weights:\n            path = \"/kaggle/input/hms-dataset/HMS_EEG_MODEL0_1.pth\"\n            checkpoint = tc.load(path, map_location=tc.device(device))[\"model\"]\n            new_weights = self.feature_extractor.state_dict()\n            \n            for key in checkpoint.keys():\n                if \"out\" in key:\n                    continue\n                nkey = \".\".join(key.split(\".\")[1:])\n                new_weights[nkey] = checkpoint[key]\n            self.feature_extractor.load_state_dict(new_weights)\n            \n            if load_out_weights:\n                new_weights = self.out.state_dict()\n                for key in checkpoint.keys():\n                    if \"out\" in key:\n                        nkey = \".\".join(key.split(\".\")[1:])\n                        new_weights[nkey] = checkpoint[key]\n                self.out.load_state_dict(new_weights)\n        \n        self.load_out_weights = load_out_weights\n    \n    def forward(self, x1,x2,x3,x4,x5):\n        if self.load_out_weights:\n            h = self.feature_extractor(x1, x2, x3, x4, x5)\n        else:\n            with tc.no_grad():\n                h = self.feature_extractor(x1, x2, x3, x4, x5)\n#         return h  \n        out = self.out(h)\n        return out\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nmodel = EEGBasedClassifier(load_weights=True, load_out_weights=True).to(device)\nsum([p.numel() for p in model.parameters() if p.requires_grad])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\n\ndef test(model, criterion, ret=False):\n    test_dataset = TrainingDatasetEEG(df, train=False, is_test_random=is_test_random)\n    test_dataloader = DataLoader(test_dataset, batch_size=max(batch_size, 128), shuffle=False, num_workers=min(batch_size,4))\n    print('Starting Testing')\n    model.eval()\n    running_loss = 0.0\n    ya = []\n    yp = []\n    for i, data in enumerate(test_dataloader, 0):\n        (x1,x2,x3,x4,x5), targets = data\n            \n        ya.append(targets.squeeze())\n        \n        # forward + backward + optimize\n        with tc.cuda.amp.autocast():\n            with tc.no_grad():\n                outputs = model(x1.to(device), x2.to(device), x3.to(device), x4.to(device), x5.to(device))\n                outputs = F.softmax(outputs, -1)\n                outputs = tc.log(outputs+1e-5)\n                loss = criterion(outputs, targets.to(device))\n                \n                yp.append(outputs.detach().cpu().numpy().squeeze())\n        \n        # print statistics\n        running_loss += loss.item()\n    print(f'Test Loss: {running_loss / i:.6f}')\n    del loss,outputs\n    \n    print('Finished Testing')\n    print(\"\")\n    if ret:\n        return np.concatenate(ya), np.concatenate(yp)\n    \ndef train(model, optimizer, criterion, scalar, scheduler, epochs):    \n    for epoch in range(epochs):\n        model.train()\n        running_loss = 0.0\n        loss_count=0\n\n        for i, data in enumerate(train_dataloader, 0):\n            (x1,x2,x3,x4,x5), targets = data\n            \n            # forward + backward + optimize\n            with tc.cuda.amp.autocast():\n                outputs = model(x1.to(device), x2.to(device), x3.to(device), x4.to(device), x5.to(device))\n                outputs = F.log_softmax(outputs,-1)\n                \n                loss = criterion(outputs, targets.to(device))\n                \n                scaler.scale(loss).backward()\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n                \n            running_loss += loss.item()\n            # print statistics\n            loss_count += 1\n#             \n            if i % 100 == 99:    # print every 10 mini-batches\n                print(f'[{epoch + 1}, {i + 1:5d}] loss: {running_loss / loss_count:.6f}')\n                running_loss = 0.0\n                loss_count = 0\n            del loss,outputs\n            \n#         scheduler.step()\n        test(model, criterion)\n            \n    print('Finished Training')\n\ndef submit(model):\n    dataset = DatasetEEG(df)\n    dataloader = DataLoader(dataset, batch_size=32, shuffle=False, num_workers=min(batch_size,4))\n    print('Starting')\n    model.eval()\n    submission = []\n\n    for x1,x2,x3,x4,x5,eeg_ids in dataloader:\n\n        with tc.cuda.amp.autocast():\n            with tc.no_grad():\n                outputs = F.softmax(model(x1.to(device), x2.to(device), x3.to(device), x4.to(device), x5.to(device)),-1)\n                \n        for output, eeg_id in zip(outputs.detach().cpu().numpy(), eeg_ids.numpy()):\n            d = {\n                \"eeg_id\": eeg_id,\n                \"seizure_vote\":output[0],\n                \"lpd_vote\":output[1], \n                \"gpd_vote\":output[2],\n                \"lrda_vote\":output[3],\n                \"grda_vote\":output[4],\n                \"other_vote\":output[5]\n            }\n            submission.append(d)\n\n    print(\"completed\")\n    return pd.DataFrame(submission)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# del model\n# import gc\n# gc.collect()\n# tc.cuda.empty_cache()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if domain == \"train\":\n    lr = 6e-6\n    epochs = 100\n\n    train_dataset = TrainingDatasetEEG(df, train=True, is_test_random=is_test_random)\n    train_dataloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=min(batch_size,4))\n\n    criterion = nn.KLDivLoss(reduction=\"batchmean\")\n    optimizer = tc.optim.AdamW(model.parameters(),lr=lr)\n    scaler = tc.cuda.amp.GradScaler()\n    scheduler = tc.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs, eta_min=1e-5)\n    train(model, optimizer, criterion, scaler, scheduler, epochs)\nelse:\n    submission = submit(model)\n    submission.to_csv(\"submission.csv\", index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# pd.read_csv(\"submission.csv\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# d = test(model, criterion, ret=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# checkpoint = {'model': model.state_dict()}\n# tc.save(checkpoint, 'HMS_EEG_MODEL0_1.pth')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}