{"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":7392775,"sourceType":"datasetVersion","datasetId":4297782},{"sourceId":7447509,"sourceType":"datasetVersion","datasetId":4334995}],"dockerImageVersionId":30648,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-02-17T17:19:00.929919Z","iopub.execute_input":"2024-02-17T17:19:00.930690Z","iopub.status.idle":"2024-02-17T17:19:02.067550Z","shell.execute_reply.started":"2024-02-17T17:19:00.930654Z","shell.execute_reply":"2024-02-17T17:19:02.066691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as A\nimport gc\nimport matplotlib.pyplot as plt\nimport math\nimport multiprocessing\nimport numpy as np\nimport os\nimport pandas as pd\nimport random\nimport time\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\nfrom albumentations.pytorch import ToTensorV2\nfrom glob import glob\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\nfrom typing import Dict, List\n\nos.environ[\"CUDA_VISIBLE_DEVICES\"] = \"0,1\"\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint('Using', torch.cuda.device_count(), 'GPU(s)')\n","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:22:56.411476Z","iopub.execute_input":"2024-02-17T17:22:56.412363Z","iopub.status.idle":"2024-02-17T17:22:56.420510Z","shell.execute_reply.started":"2024-02-17T17:22:56.412329Z","shell.execute_reply":"2024-02-17T17:22:56.419586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class config:\n    model = \"vit_base_patch16_224\"\n    epoch = 10\n    lr = 1e-3\n    batchsize = 32\n    splits = 5\n    momentum = 0.9\n    MAX_GRAD_NORM = 1e7\n    WEIGHT_DECAY = 0.01\n    device = \"cpu\"\n    FOLDS=5\n    AMP= True\nclass paths:\n    preloadedeeg = \"/kaggle/input/brain-eeg-spectrograms/eeg_specs.npy\"\n    train_eeg_dir = \"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs\"\n    train_spec_dir = \"/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms\"\n    train_csv = \"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\"\n    test_csv = \"/kaggle/input/hms-harmful-brain-activity-classification/test.csv\"\n    test_eeg = \"/kaggle/input/hms-harmful-brain-activity-classification/test_eegs\"\n    test_spec = \"/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms\"\n    out = \"/kaggle/working/\"","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:19:13.391784Z","iopub.execute_input":"2024-02-17T17:19:13.392053Z","iopub.status.idle":"2024-02-17T17:19:13.398164Z","shell.execute_reply.started":"2024-02-17T17:19:13.392030Z","shell.execute_reply":"2024-02-17T17:19:13.397309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(paths.train_csv)\ntargets = df.columns[-6:]\n","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:19:13.400124Z","iopub.execute_input":"2024-02-17T17:19:13.400394Z","iopub.status.idle":"2024-02-17T17:19:13.663150Z","shell.execute_reply.started":"2024-02-17T17:19:13.400371Z","shell.execute_reply":"2024-02-17T17:19:13.662329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"targets","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:19:13.664239Z","iopub.execute_input":"2024-02-17T17:19:13.664549Z","iopub.status.idle":"2024-02-17T17:19:13.671819Z","shell.execute_reply.started":"2024-02-17T17:19:13.664506Z","shell.execute_reply":"2024-02-17T17:19:13.670708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = df.groupby('eeg_id')[['spectrogram_id','spectrogram_label_offset_seconds']].agg({\n    'spectrogram_id':'first',\n    'spectrogram_label_offset_seconds':'min'\n})\ntrain_df.columns = ['spectogram_id','min']\n\naux = df.groupby('eeg_id')[['spectrogram_id','spectrogram_label_offset_seconds']].agg({\n    'spectrogram_label_offset_seconds':'max'\n})\ntrain_df['max'] = aux\n\naux = df.groupby('eeg_id')[['patient_id']].agg('first')\ntrain_df['patient_id'] = aux\n\naux = df.groupby('eeg_id')[targets].agg('sum')\nfor label in targets:\n    train_df[label] = aux[label].values\n    \ny_data = train_df[targets].values\ny_data = y_data / y_data.sum(axis=1,keepdims=True)\ntrain_df[targets] = y_data\n\naux = df.groupby('eeg_id')[['expert_consensus']].agg('first')\ntrain_df['target'] = aux\n\ntrain_df = train_df.reset_index()\nprint('Train non-overlapp eeg_id shape:', train_df.shape )\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:19:13.673080Z","iopub.execute_input":"2024-02-17T17:19:13.673443Z","iopub.status.idle":"2024-02-17T17:19:13.778208Z","shell.execute_reply.started":"2024-02-17T17:19:13.673411Z","shell.execute_reply":"2024-02-17T17:19:13.777259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spectrograms = np.load('/kaggle/input/brain-spectrograms/specs.npy',allow_pickle=True).item()\neegs = np.load(paths.preloadedeeg,allow_pickle = True).item()\n","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:19:13.779305Z","iopub.execute_input":"2024-02-17T17:19:13.779575Z","iopub.status.idle":"2024-02-17T17:21:39.285871Z","shell.execute_reply.started":"2024-02-17T17:19:13.779552Z","shell.execute_reply":"2024-02-17T17:21:39.284976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import KFold, GroupKFold\ngfold = GroupKFold(n_splits = config.splits)\nfor fold,(trainindex,validationindex) in enumerate(gfold.split(train_df,train_df[\"target\"],train_df[\"patient_id\"])):\n    train_df.loc[validationindex,\"fold\"] = int(fold)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:21:39.287093Z","iopub.execute_input":"2024-02-17T17:21:39.287420Z","iopub.status.idle":"2024-02-17T17:21:39.317283Z","shell.execute_reply.started":"2024-02-17T17:21:39.287393Z","shell.execute_reply":"2024-02-17T17:21:39.316587Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:21:39.318426Z","iopub.execute_input":"2024-02-17T17:21:39.318832Z","iopub.status.idle":"2024-02-17T17:21:39.344316Z","shell.execute_reply.started":"2024-02-17T17:21:39.318786Z","shell.execute_reply":"2024-02-17T17:21:39.343370Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomDataset():\n    def __init__(self,traindf,config,mode:str=\"train\",specs:dict[int,np.ndarray]=spectrograms,eegs:dict[int,np.ndarray]=eegs):\n        self.traindf = traindf;\n        self.specs = specs;\n        self.eeg = eegs;\n        self.mode = mode;\n    def __len__(self):\n        return len(self.traindf)\n    def __getitem__(self,idx):\n        X = np.zeros((128, 256, 8), dtype='float32')\n        y = np.zeros(6, dtype='float32')\n        img = np.ones((128,256), dtype='float32')\n        row = self.traindf.iloc[idx]\n        if self.mode=='test': \n            r = 0\n        else: \n            r = int((row['min'] + row['max']) // 4)\n            \n        for region in range(4):\n            img = self.specs[row.spectogram_id][r:r+300, region*100:(region+1)*100].T\n            \n            # Log transform spectogram\n            img = np.clip(img, np.exp(-4), np.exp(8))\n            img = np.log(img)\n\n            # Standarize per image\n            ep = 1e-6\n            mu = np.nanmean(img.flatten())\n            std = np.nanstd(img.flatten())\n            img = (img-mu)/(std+ep)\n            img = np.nan_to_num(img, nan=0.0)\n            X[14:-14, :, region] = img[:, 22:-22] / 2.0\n            img = self.eeg[row.eeg_id]\n            X[:, :, 4:] = img\n        \n        X = torch.tensor(X)\n        spectograms = [X[:, :, i:i+1] for i in range(4)]\n        spectograms = torch.cat(spectograms, dim=0)\n        \n        # === Get EEG spectograms ===\n        eegs = [X[:, :, i:i+1] for i in range(4,8)]\n        eegs = torch.cat(eegs, dim=0)\n        \n        # === Reshape (512,512,3) ===\n        x = torch.cat([spectograms, eegs], dim=1)\n        x = torch.cat([x,x,x], dim=2)\n        x = x.permute(2, 0, 1)\n        if self.mode != 'test':\n            y = row[targets].values.astype(np.float32)\n            \n            \n        return {\"data\":x,\n                \"target\":y}\n        ","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:21:39.349133Z","iopub.execute_input":"2024-02-17T17:21:39.349571Z","iopub.status.idle":"2024-02-17T17:21:39.365021Z","shell.execute_reply.started":"2024-02-17T17:21:39.349516Z","shell.execute_reply":"2024-02-17T17:21:39.364156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"customdataset = CustomDataset(train_df,config)","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:21:39.365964Z","iopub.execute_input":"2024-02-17T17:21:39.366217Z","iopub.status.idle":"2024-02-17T17:21:39.377431Z","shell.execute_reply.started":"2024-02-17T17:21:39.366194Z","shell.execute_reply":"2024-02-17T17:21:39.376673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"customdataset[100][\"data\"]","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:21:39.378789Z","iopub.execute_input":"2024-02-17T17:21:39.379157Z","iopub.status.idle":"2024-02-17T17:21:39.531061Z","shell.execute_reply.started":"2024-02-17T17:21:39.379125Z","shell.execute_reply":"2024-02-17T17:21:39.530098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader\ntrainloader = DataLoader(customdataset,batch_size = config.batchsize)","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:21:39.532547Z","iopub.execute_input":"2024-02-17T17:21:39.533197Z","iopub.status.idle":"2024-02-17T17:21:39.537811Z","shell.execute_reply.started":"2024-02-17T17:21:39.533162Z","shell.execute_reply":"2024-02-17T17:21:39.536923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import timm","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:21:39.539002Z","iopub.execute_input":"2024-02-17T17:21:39.539317Z","iopub.status.idle":"2024-02-17T17:21:39.552245Z","shell.execute_reply.started":"2024-02-17T17:21:39.539285Z","shell.execute_reply":"2024-02-17T17:21:39.551392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision.transforms import transforms\ntras = transforms.Compose([\n    transforms.Resize((224,224))\n])","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:21:39.553504Z","iopub.execute_input":"2024-02-17T17:21:39.554104Z","iopub.status.idle":"2024-02-17T17:21:39.562580Z","shell.execute_reply.started":"2024-02-17T17:21:39.554074Z","shell.execute_reply":"2024-02-17T17:21:39.561567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Custommodel(nn.Module):\n    def __init__(self,config,transform,numclass:int=6):\n        super(Custommodel,self).__init__()\n        self.model = timm.create_model(\n            config.model,\n            pretrained = True,\n            drop_rate = 0.51,\n            #drop_path_rate = 0.2,\n        )\n        #self.features = nn.Sequential(*list(self.model.children())[:-2])\n        #self.customlayer = nn.Sequential(\n        #    nn.AdaptiveAvgPool2d(1),\n        #    nn.Flatten(),\n        #    nn.Linear(self.model.num_features, numclass)\n        #)\n        self.model.head = nn.Linear(self.model.head.in_features, numclass)\n        self.transform = transform\n    def forward(self,x):\n        x = self.transform(x)\n        x = self.model(x)\n        #x = self.customlayer(x)\n        return x\n        ","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:21:39.563651Z","iopub.execute_input":"2024-02-17T17:21:39.563963Z","iopub.status.idle":"2024-02-17T17:21:39.572955Z","shell.execute_reply.started":"2024-02-17T17:21:39.563935Z","shell.execute_reply":"2024-02-17T17:21:39.571950Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = torch.nn.KLDivLoss(reduction=\"mean\")\nfrom torch.optim.lr_scheduler import OneCycleLR\n","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:23:35.186255Z","iopub.execute_input":"2024-02-17T17:23:35.186651Z","iopub.status.idle":"2024-02-17T17:23:35.191837Z","shell.execute_reply.started":"2024-02-17T17:23:35.186622Z","shell.execute_reply":"2024-02-17T17:23:35.190921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def trainer(model,trainloader,optimizer,criterion,device):\n  model.train()\n  iterationloss = 0\n  counter = 0\n  for data in tqdm(trainloader):\n    message = data['data'].to(device)\n    target = data['target'].to(device).squeeze()\n    optimizer.zero_grad()\n    out = model(message)\n    loss = criterion(F.log_softmax(out, dim=1), target)\n    loss.backward()\n    optimizer.step()\n    iterationloss+=loss.item()*message.shape[0]\n    counter+=message.shape[0]\n  return iterationloss/counter","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:23:35.472741Z","iopub.execute_input":"2024-02-17T17:23:35.473156Z","iopub.status.idle":"2024-02-17T17:23:35.480752Z","shell.execute_reply.started":"2024-02-17T17:23:35.473126Z","shell.execute_reply":"2024-02-17T17:23:35.479477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def tester(model,testloader,criterion,device):\n  model.eval()\n  iterationloss = 0\n  counter = 0\n  for data in testloader:\n    message = data['data'].to(device)\n    target = data['target'].to(device).squeeze()\n    with torch.no_grad():\n      out = model(message)\n      loss = criterion(F.log_softmax(out, dim=1), target)\n      iterationloss+=loss*message.shape[0]\n      #print(loss)  \n    counter+=message.shape[0]\n  return float(iterationloss)/float(counter)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:23:35.729683Z","iopub.execute_input":"2024-02-17T17:23:35.730106Z","iopub.status.idle":"2024-02-17T17:23:35.736743Z","shell.execute_reply.started":"2024-02-17T17:23:35.730073Z","shell.execute_reply":"2024-02-17T17:23:35.735806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def trainepoch(fold,train_loss,val_loss):\n    traindf = train_df[train_df[\"fold\"]!=fold].reset_index(drop = True)\n    validdf = train_df[train_df[\"fold\"]==fold].reset_index(drop = True)\n    \n    traindataset = CustomDataset(traindf,config,mode = \"train\")\n    validdataset = CustomDataset(validdf,config,mode = \"train\")\n    \n    trainloader = DataLoader(traindataset,batch_size = config.batchsize)\n    validloader = DataLoader(validdataset,batch_size = config.batchsize)\n    \n    model = Custommodel(config,tras)\n    model.to(device)\n    \n    optimizer = torch.optim.AdamW(model.parameters(), lr=0.1, weight_decay=config.WEIGHT_DECAY)\n    \n    scheduler = OneCycleLR(\n        optimizer,\n        max_lr=1e-3,\n        epochs=config.epoch,\n        steps_per_epoch=len(trainloader),\n        pct_start=0.1,\n        anneal_strategy=\"cos\",\n        final_div_factor=100,\n    )\n    \n    criterion = torch.nn.KLDivLoss(reduction=\"batchmean\")\n    bestloss = np.inf\n    \n    bestloss = np.inf\n    for i in (range(10)):\n        print(\"-\"*120)\n        print(\"Began iteration no.\",i+1)\n        print(\":\"*50,\"=\"*20,\":\"*50)\n        trainloss = trainer(model,trainloader,optimizer,criterion,device)\n        testloss = tester(model,validloader,criterion,device)\n        train_loss.append(trainloss)\n        val_loss.append(testloss)\n        print(\"train loss = \",trainloss)\n        print(\"val loss=\",testloss)\n        print(\"=\"*80,\"\\n\")\n        if testloss<bestloss:\n            bestloss = testloss\n            dic = {\n            'model': model.state_dict()\n            }\n            torch.save(dic,'./Bestmodel.model'+str(fold))\n            print(\"Improved and saved the model\\n\")\n        print(\"=\"*100)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:23:35.972179Z","iopub.execute_input":"2024-02-17T17:23:35.973017Z","iopub.status.idle":"2024-02-17T17:23:35.984306Z","shell.execute_reply.started":"2024-02-17T17:23:35.972984Z","shell.execute_reply":"2024-02-17T17:23:35.983390Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainlist = []\nvalidlist = []\nfor fold in range(config.FOLDS):\n    if fold in [0, 1, 2, 3, 4]:\n        print(\"fold = \",fold)\n        trainepoch(fold,trainlist,validlist)","metadata":{"execution":{"iopub.status.busy":"2024-02-17T17:23:36.206962Z","iopub.execute_input":"2024-02-17T17:23:36.207320Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}