{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":7392733,"sourceType":"datasetVersion","datasetId":4297749}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nimport numpy as np\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import Dataset\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, OneCycleLR\n\nfrom tqdm.notebook import tqdm\nimport math\nimport time\nimport gc\nfrom torchinfo import summary\nimport warnings\nimport sys\nwarnings.filterwarnings(\"ignore\")\nsys.path.append('/kaggle/input/kaggle-kl-div')\nfrom kaggle_kl_div import score\n\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2024-06-27T23:01:32.736631Z","iopub.execute_input":"2024-06-27T23:01:32.737451Z","iopub.status.idle":"2024-06-27T23:01:32.745554Z","shell.execute_reply.started":"2024-06-27T23:01:32.737415Z","shell.execute_reply":"2024-06-27T23:01:32.744059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"a = pd.read_parquet(\"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/1628180742.parquet\")\na","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(\"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\")\ntrain_df[:15]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Hyper:\n    def __init__(self, **kwargs):\n        for key, value in kwargs.items():\n            setattr(self,key,value)\n    \n    def __repr__(self):\n        out = ''\n        for key in dir(hyper)[26:]: #get the elements \n            out += f'{key}={getattr(self,key)}\\n'\n        return out\n        \nhyper = Hyper(time_step = 1000,\\\n    epochs = 2,\\\n    eval_interval = 200,\\\n    eval_iters = 2,\\\n    batch_size = 32,\\\n    learning_rate = 4e-4,\\\n    d_model = 512,\\\n    n_head = 8,\\\n    n_layer = 5,\\\n    dropout = 0.1,\\\n)\nhyper     ","metadata":{"execution":{"iopub.status.busy":"2024-06-27T23:01:32.746679Z","iopub.execute_input":"2024-06-27T23:01:32.746914Z","iopub.status.idle":"2024-06-27T23:01:32.767099Z","shell.execute_reply.started":"2024-06-27T23:01:32.746893Z","shell.execute_reply":"2024-06-27T23:01:32.766269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader\nfrom torch.utils.data import Dataset\nimport random\nimport os\n\nclass eegData(Dataset):\n    def __init__(self, eeg_dir, meta_data_path, split, train_ratio = 0.95, time_step=hyper.time_step):\n        self.eeg_dir = eeg_dir\n        self.time_step = time_step\n        self.file_list_full = os.listdir(eeg_dir)\n        self.train_ratio = train_ratio\n    \n        if split == 'train':\n            #print(len(self.file_list))\n            self.file_list = self.file_list_full[0:int(len(self.file_list_full) * train_ratio)]\n            #print(len(self.file_list))\n        else:\n            self.file_list = self.file_list_full[int(len(self.file_list_full) * train_ratio):]\n        \n        self.meta = pd.read_csv(meta_data_path)\n        self.split = split\n        \n        self.eeg_id_list = np.array(self.meta['eeg_id']) #contains multiples \n        \n        self.file_list = [i for i in self.file_list if int(i[:-8]) in self.eeg_id_list] #note: file_list order is different from eeg_id_list \n        \n        #self.offset_list = np.array(self.meta['eeg_label_offset_seconds'])\n        self.label_to_int = {'Seizure':0, 'GPD':1, 'LPD':2, 'LRDA':3, 'GRDA':4, 'Other':5}\n        self.eeg_dir_len = len(eeg_dir)\n        \n        \n        self.label_list = self.meta['expert_consensus'] #labels matched to order of eeg_id_list\n        \n        self.label_list = np.array([self.label_to_int[i] for i in self.label_list])\n        \n        self.id_and_label_dict = dict(zip(self.eeg_id_list,self.label_list )) \n        \n    def standardize(self, src):\n        target_mean = np.mean(src)\n        target_std = np.std(src)\n        return (src - target_mean) / target_std\n        \n    def process_eeg(self,idx):\n        \n        file_name = self.file_list[idx] #get a random file \n        \n        \n        file_name_full = os.path.join(eeg_dir, file_name)\n        \n        file_name = int(file_name[:-8])\n       \n        eeg = pd.read_parquet(file_name_full)\n#         eeg = eeg.drop(eeg.index[:4000])\n#         eeg = eeg.drop(eeg.index[-4000:]) #caution: shortest eeg is 10000 long\n        \n        label = torch.tensor(self.id_and_label_dict[file_name])\n        \n       # label = torch.randint(low=0,high=6, size=(1,))\n        \n#         pos =np.where(self.eeg_id_list == int(file_name))[0] #gets the offset list for the current eeg \n#         eeg_offsets = self.offset_list[pos[0]:pos[0]+len(pos)]\n\n#         offset_adjust_total = 0\n#         offset_adjust = 0\n        \n#         for i in range(len(eeg_offsets)-1):\n#             if eeg_offsets[i+1] - eeg_offsets[i] > 10:\n#                 offset_i = eeg_offsets[i] - offset_adjust_total\n#                 offset_i_plus = eeg_offsets[i+1] - offset_adjust_total\n                \n#                 crop_idx_start = int(offset_i * 200 + 2000) #original offset_i * 200 + 6000\n#                 crop_idx_end = int(offset_i_plus * 200) #original ofset_i_plus * 200 + 4000\n                \n#                 eeg = eeg.drop(eeg.index[crop_idx_start:crop_idx_end])\n#                 offset_adjust = (crop_idx_end - crop_idx_start)/200\n                \n#             offset_adjust_total += offset_adjust\n            \n        return eeg, label\n\n\n    def __len__(self):\n        return len(self.file_list)\n\n    def __getitem__(self, idx):\n        eeg,label = self.process_eeg(idx)   \n        \n        eeg = eeg.fillna(0) # always!\n        \n        if len(eeg) > self.time_step:\n            rand_idx = np.random.randint(0, len(eeg)-self.time_step)\n        else:\n            rand_idx = 0\n        chunk = eeg.iloc[rand_idx:rand_idx+self.time_step].to_numpy() #each chunk is size of (time_step, 20) \n        \n        chunk = torch.tensor(self.standardize(chunk))\n        #chunk = torch.tensor(chunk)\n        \n\n\n        return chunk, label","metadata":{"execution":{"iopub.status.busy":"2024-06-27T23:46:30.042872Z","iopub.execute_input":"2024-06-27T23:46:30.043271Z","iopub.status.idle":"2024-06-27T23:46:30.060045Z","shell.execute_reply.started":"2024-06-27T23:46:30.04324Z","shell.execute_reply":"2024-06-27T23:46:30.059235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"eeg_dir = \"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs\"\nmeta_path = \"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\"\n\ntrain_dataset = eegData(eeg_dir, meta_path,split='train', time_step = hyper.time_step)\ntrain_loader = DataLoader(train_dataset, batch_size = hyper.batch_size, shuffle=False)\n\nval_dataset = eegData(eeg_dir, meta_path,split='val', time_step = hyper.time_step)\nval_loader = DataLoader(val_dataset, batch_size = hyper.batch_size, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2024-06-27T23:46:33.350876Z","iopub.execute_input":"2024-06-27T23:46:33.351264Z","iopub.status.idle":"2024-06-27T23:46:34.676578Z","shell.execute_reply.started":"2024-06-27T23:46:33.351232Z","shell.execute_reply":"2024-06-27T23:46:34.675795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import nn, Tensor\nfrom torch.nn import TransformerEncoder, TransformerEncoderLayer\n\n\n\nclass eegEncoder(nn.Module):\n\n    def __init__(self,d_model, nhead,  nlayers, nclasses, d_hid, dropout=0.1):\n        super(eegEncoder, self).__init__()\n        self.d_model = d_model\n\n        #self.pos_encoder = PositionalEncoding(d_model, dropout)\n        self.input_projection = nn.Linear(20, d_model) # (n_chhannels, d_model)  \n        #self.position_embedding_table = nn.Embedding(hyper.time_step, d_model)\n      \n        \n        self.linear = nn.Linear(self.d_model, nclasses)\n        self.pre_linear = nn.Linear(hyper.time_step, 1)\n        \n        encoder_layers = TransformerEncoderLayer(self.d_model, nhead, d_hid, dropout, batch_first=True, norm_first =False) # set batch_first = True b/c input is of shape (b, T, C)!!!!!!!\n        self.transformer_encoder = TransformerEncoder(encoder_layers, nlayers)\n    \n    def pos_encode(self, encoding_1):\n        pe = torch.ones_like(encoding_1[0])\n        position = torch.arange(0, hyper.time_step).unsqueeze(-1)\n        temp = torch.Tensor(range(0, self.d_model, 2))\n        temp = temp * -(math.log(10000) / self.d_model)\n        temp = torch.exp(temp).unsqueeze(0)\n        temp = torch.matmul(position.float(), temp)  # shape:[input, d_model/2]\n        pe[:, 0::2] = torch.sin(temp)\n        pe[:, 1::2] = torch.cos(temp)\n\n    \n        encoding_1 = encoding_1 + pe\n        return encoding_1\n\n    def forward(self, src):\n        \n        B, T, C = src.shape #where B = batch_size, T = time_step, C = n_classes\n        \n        input_proj = self.input_projection(src) #output is shape (B, T, d_model)\n        input_pos = self.pos_encode(input_proj)#self.position_embedding_table(torch.arange(T, device=src.device))  # (T, n_embd)\n        \n        output_1 = self.transformer_encoder(input_pos) #output_1 is shape of (B,T, )\n    \n        #print(output_1.shape)\n        output_1 = torch.permute(output_1, (0,2,1))\n        output_2 = self.pre_linear(output_1)\n        \n        output_2 = torch.squeeze(output_2, dim=-1)\n        #output_2 = torch.mean(output_1, axis=1) #average over sequence length dimension. Output is shape (B,)\n        output_2 = self.linear(output_2)\n                \n        return output_2","metadata":{"execution":{"iopub.status.busy":"2024-06-27T23:01:36.868481Z","iopub.execute_input":"2024-06-27T23:01:36.868775Z","iopub.status.idle":"2024-06-27T23:01:36.880666Z","shell.execute_reply.started":"2024-06-27T23:01:36.868745Z","shell.execute_reply":"2024-06-27T23:01:36.879698Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def __init__(self,d_model, nhead, nlayers, nclasses,d_hid, dropout=0.1):\n# d_hid is the dim of the feed forward in the transformer blocks \n# d_model has to be divisable by n_head\nmodel = eegEncoder(d_model=720, nhead=8, nlayers=8, d_hid=2048, nclasses=6 )\nmodel.to(device)\nprint()","metadata":{"execution":{"iopub.status.busy":"2024-06-27T23:46:37.298427Z","iopub.execute_input":"2024-06-27T23:46:37.299072Z","iopub.status.idle":"2024-06-27T23:46:37.456264Z","shell.execute_reply.started":"2024-06-27T23:46:37.299042Z","shell.execute_reply":"2024-06-27T23:46:37.45527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"summary(model)","metadata":{"execution":{"iopub.status.busy":"2024-06-27T23:11:07.394589Z","iopub.execute_input":"2024-06-27T23:11:07.395256Z","iopub.status.idle":"2024-06-27T23:11:07.41096Z","shell.execute_reply.started":"2024-06-27T23:11:07.395224Z","shell.execute_reply":"2024-06-27T23:11:07.410129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for b in train_loader:\n    test, test_y = b\n    break\n","metadata":{"execution":{"iopub.status.busy":"2024-06-27T22:06:37.798613Z","iopub.execute_input":"2024-06-27T22:06:37.798966Z","iopub.status.idle":"2024-06-27T22:06:37.926931Z","shell.execute_reply.started":"2024-06-27T22:06:37.798939Z","shell.execute_reply":"2024-06-27T22:06:37.926044Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test = pd.read_parquet(\"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/1000913311.parquet\")\ntest = test[:1000]\ntest = torch.tensor(test.to_numpy()).to(torch.float32)\n\ntest = torch.unsqueeze(test, dim=0)\ntest_y = torch.tensor([1])","metadata":{"execution":{"iopub.status.busy":"2024-06-27T21:47:51.185412Z","iopub.execute_input":"2024-06-27T21:47:51.18624Z","iopub.status.idle":"2024-06-27T21:47:51.21237Z","shell.execute_reply.started":"2024-06-27T21:47:51.186206Z","shell.execute_reply":"2024-06-27T21:47:51.211636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del xb\ndel yb\n\ndel loss\ndel optimizer\ndel scaler\ndel model\ndel probs\n\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2024-06-27T22:15:14.707816Z","iopub.execute_input":"2024-06-27T22:15:14.708479Z","iopub.status.idle":"2024-06-27T22:15:15.096478Z","shell.execute_reply.started":"2024-06-27T22:15:14.708443Z","shell.execute_reply":"2024-06-27T22:15:15.095512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"KLDiv = nn.KLDivLoss(reduction = \"batchmean\")\noptimizer = torch.optim.Adam(model.parameters(), lr=0.0001)\nscaler = torch.cuda.amp.GradScaler()\n\nlosses = []\n\nn = 0\nfor b in train_loader:\n    if n == 1000:\n        break\n    \n    xb , yb = b\n    \n    xb = xb.to(device)\n    yb = yb.to(device)\n#     xb = test.to(device)\n#     yb = test_y.to(device)\n    yb = F.one_hot(yb, num_classes=6)\n    yb = torch.squeeze(yb, dim=0).to(torch.float32)\n\n    \n    outputs = model(xb)\n    probs = F.log_softmax(outputs, dim=1)\n    loss = KLDiv(probs, yb)\n    \n    print(f'iter {n}, loss: {loss}')\n    \n    scaler.scale(loss).backward()\n\n    scaler.step(optimizer)\n    scaler.update()\n    optimizer.zero_grad()\n\n    losses.append(loss.detach().cpu().numpy())\n    \n    n+= 1\n    \n","metadata":{"execution":{"iopub.status.busy":"2024-06-27T23:47:14.472563Z","iopub.execute_input":"2024-06-27T23:47:14.472976Z","iopub.status.idle":"2024-06-28T00:09:02.046558Z","shell.execute_reply.started":"2024-06-27T23:47:14.472944Z","shell.execute_reply":"2024-06-28T00:09:02.045769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"losses1= []\n\nfor i in losses:\n    a = i.detach().cpu().numpy()\n    losses1.append(a)","metadata":{"execution":{"iopub.status.busy":"2024-06-27T23:42:27.864735Z","iopub.execute_input":"2024-06-27T23:42:27.865482Z","iopub.status.idle":"2024-06-27T23:42:27.88246Z","shell.execute_reply.started":"2024-06-27T23:42:27.865446Z","shell.execute_reply":"2024-06-27T23:42:27.881517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(losses1)","metadata":{"execution":{"iopub.status.busy":"2024-06-27T23:42:31.224797Z","iopub.execute_input":"2024-06-27T23:42:31.225433Z","iopub.status.idle":"2024-06-27T23:42:31.422938Z","shell.execute_reply.started":"2024-06-27T23:42:31.225403Z","shell.execute_reply":"2024-06-27T23:42:31.422126Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loop(testing=True, print_interval=1, break_at_nan=False)","metadata":{"execution":{"iopub.status.busy":"2024-06-27T21:58:02.3671Z","iopub.execute_input":"2024-06-27T21:58:02.367939Z","iopub.status.idle":"2024-06-27T21:59:50.169574Z","shell.execute_reply.started":"2024-06-27T21:58:02.367905Z","shell.execute_reply":"2024-06-27T21:59:50.168138Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"KLDiv = nn.KLDivLoss(reduction = \"batchmean\")\noptimizer = torch.optim.Adam(model.parameters(), lr=hyper.learning_rate)\nscaler = torch.cuda.amp.GradScaler()\n#scheduler = OneCycleLR(optimizer, max_lr=hyper.learning_rate, epochs=2, steps_per_epoch=len(train_loader) )\n#scheduler = CosineAnnealingLR(optimizer, T_max = (hyper.epochs * len(train_loader)))","metadata":{"execution":{"iopub.status.busy":"2024-06-27T21:42:37.724676Z","iopub.execute_input":"2024-06-27T21:42:37.725617Z","iopub.status.idle":"2024-06-27T21:42:37.731243Z","shell.execute_reply.started":"2024-06-27T21:42:37.725581Z","shell.execute_reply":"2024-06-27T21:42:37.730338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_loop(testing=False, print_interval=2, break_at_nan=False):\n    loss = None\n    \n    model.train()\n    for epoch in range(hyper.epochs):    \n        with tqdm(train_loader, unit=\"batch\") as tepoch:\n            for batch_iter, batch in enumerate(tepoch):\n\n\n                tepoch.set_description(f\"Epoch {epoch}\")\n\n\n                if n_iter % hyper.eval_interval == 0: \n                    val_loss = eval_loop(testing=testing)\n\n                    if loss is not None:\n                        print(f\"step {n_iter}: train loss {loss:.4f}, val loss {val_loss:.4f}\")\n                    else:\n                        print(f\"step {n_iter}: val loss {val_loss:.4f}\")\n\n                xb, yb = batch\n\n\n                xb = xb.to(device)\n                yb = yb.to(device)\n\n                yb = F.one_hot(yb, num_classes=6)\n                \n                \n                yb = torch.squeeze(yb, dim=0)\n\n                yb = yb.to(torch.float32)\n\n\n                if testing:\n                    \n                    \n                    _, _, _, outputs = model(xb)\n                    probs = F.log_softmax(outputs, dim=1) #(B,6)\n                    loss = KLDiv(probs, yb)\n#                     loss_pointwise = yb * (yb - probs)\n#                     loss = loss_pointwise.mean()\n                    \n                    if n_iter % print_interval == 0:\n                        print('============================iter===========================:', n_iter)\n                        #print(\"src after input projection:\", input_proj[0])\n\n                        #print(\"src after pos encode:\", input_pos[0])\n                   \n                        #print(\"output after multiheaded attention:\", output_1[0])\n                        #print('final output:', outputs[0])\n                \n                \n                scaler.scale(loss).backward()\n#                 loss.backward()\n                scaler.step(optimizer)\n                scaler.update()\n#                 optimizer.step()\n               \n                \n                \n                \n                optimizer.zero_grad()\n                #scheduler.step()\n\n   \n                losses.append(loss.item())\n                if testing:\n                    if n_iter % print_interval == 0:\n                        print_p = probs.detach()\n                        print(\"probs:\", torch.round(print_p[0].exp(), decimals=4))\n                        print('labels:', yb[0])\n                        print(\"loss:\", loss.item())\n\n                        \n                if break_at_nan and torch.isnan(loss):\n                    return        \n                        \n                tepoch.set_postfix(loss=loss.item())        \n                n_iter += 1\n\n    #return losses","metadata":{"execution":{"iopub.status.busy":"2024-06-27T21:57:57.987367Z","iopub.execute_input":"2024-06-27T21:57:57.988052Z","iopub.status.idle":"2024-06-27T21:57:58.00441Z","shell.execute_reply.started":"2024-06-27T21:57:57.988023Z","shell.execute_reply":"2024-06-27T21:57:58.003323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef eval_loop(testing=False):\n    model.eval()\n    \n    losses = []\n    for i, b in enumerate(val_loader):\n        if i == hyper.eval_iters:\n            break\n        \n        xb, yb = b\n            \n    \n        xb = xb.to(device)\n        yb = yb.to(device)\n\n        yb = F.one_hot(yb, num_classes=6)\n        yb = yb.to(torch.float32)\n        yb = torch.squeeze(yb, dim=0)\n\n        \n        if testing:\n            _, src, output_1, outputs = model(xb)\n        else:\n            outputs= model(xb)\n                \n        probs = F.log_softmax(outputs, dim=-1) #(1,6)\n\n        loss = KLDiv(probs, yb)\n        \n        losses.append(loss.item())\n    \n    return np.mean(losses)","metadata":{"execution":{"iopub.status.busy":"2024-06-27T21:57:58.777806Z","iopub.execute_input":"2024-06-27T21:57:58.778162Z","iopub.status.idle":"2024-06-27T21:57:58.78611Z","shell.execute_reply.started":"2024-06-27T21:57:58.778131Z","shell.execute_reply":"2024-06-27T21:57:58.78522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"probs = F.softmax(output, dim=-1) #(1,6)\n\nloss = KLDiv(probs, yb)\nprint(loss)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#look at this for how to account stepping for gradient accumulation: https://github.com/huggingface/accelerate/issues/963 \n\n\n\n\n#scaler = torch.cuda.amp.GradScaler()\noptimizer = torch.optim.AdamW(model.parameters(), lr=hyper.learning_rate)\nscheduler = CosineAnnealingLR(optimizer, T_max = (hyper.epochs * len(train_loader) - hyper.warm_up_iters)/ hyper.num_accum_steps)\n\nclass warmup_cosine_scheduler:\n    def __init__(self,optimizer, scheduler, n_warm_up_iters, init_lr):\n        self.optimizer = optimizer \n        self.scheduler = scheduler \n\n        self.n_warm_up_iters = n_warm_up_iters\n        self.init_lr = init_lr \n\n    def step(self, n_iter):\n        '''\n        n_warm_up_iters = batch_size * len(train_loader) \n        '''\n        if n_iter < self.n_warm_up_iters:  \n            self.optimizer.param_groups[0]['lr'] = n_iter / self.n_warm_up_iters * self.init_lr #linearly increase lr\n\n\n        else:\n            if ((n_iter + 1) % hyper.num_accum_steps == 0) or (n_iter + 1 == len(train_loader)):\n                self.scheduler.step()\n\nscheduler = warmup_cosine_scheduler(optimizer, scheduler, hyper.warm_up_iters, hyper.learning_rate)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef submit_reformat(one_hot_y, probs):\n    model.eval()\n    '''\n    make sure that probs have at least 3 dims. use unsqueeze if necessary \n    '''\n    B, T, C = probs.shape\n    \n    scores = []\n    \n    \n    for i in range(B):\n        #indices = torch.argmax(probs[i], dim=-1)\n        x = probs[i]\n        x = x.mean(dim=-2)\n        \n        x = pd.DataFrame(x)\n        x = x.T\n        \n        \n        \n        t = pd.DataFrame(one_hot_y[i])\n        t = t.T\n        \n        x['id'] = 0\n        t['id'] = 0\n        \n        \n        print(x)\n        print(t)\n        \n        score_b = score(t, x, row_id_column_name='id')\n    \n        score_b.append(score_b)\n        \n    return scores_b.mean()    ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''\nloss calculation:\n- avg after softmax\n- avg before softmax \n- linear \n- expand targets (avg during inference)\n\n'''","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Fancy training loop ","metadata":{}},{"cell_type":"code","source":"KLDiv = nn.KLDivLoss(log_target=True)\n\nmodel.train()\nscaler = torch.cuda.amp.GradScaler()\n\n\nn_iter = 0\nloss = None\n\nfor epoch in range(hyper.epochs):    \n    with tqdm(train_loader, unit=\"batch\") as tepoch:\n        for batch_iter, batch in enumerate(tepoch):\n            n_iter += batch_iter\n            tepoch.set_description(f\"Epoch {epoch}\")\n\n            \n            if n_iter % hyper.eval_interval == 0: \n                val_loss = eval_loop()\n                \n                if loss is not None:\n                    print(f\"step {n_iter}: train loss {loss:.4f}, val loss {val_loss:.4f}\")\n                else:\n                    print(f\"step {n_iter}: val loss {val_loss:.4f}\")\n        \n            xb, yb = batch\n            \n            yb = torch.unsqueeze(yb, dim=0)\n            \n            xb = xb.to(device)\n            yb = yb.to(device)\n\n            yb = F.one_hot(yb, num_classes=6)\n            yb = yb.to(torch.float32)\n            \n\n            # Automatic Tensor Casting\n            with torch.cuda.amp.autocast():                    \n                                \n                outputs = model(xb)\n                \n                probs = F.softmax(outputs, dim=-1) #(1,6)\n                \n                loss = KLDiv(probs, yb)\n                \n\n\n            scaler.scale(loss).backward() # Automatic Gradient Scaling\n\n            # Normalize the Gradients\n            loss = loss / hyper.num_accum_steps\n            \n            tepoch.set_postfix(loss=loss.item())\n\n            # Gradient Accumulation\n            if ((batch_iter + 1) % hyper.num_accum_steps == 0) or (batch_iter + 1 == len(train_loader)):     \n                scheduler.step(n_iter)\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n\n    \n        # Garbage Collection\n        \n            torch.cuda.empty_cache()\n            _ = gc.collect()\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# testing training loop ","metadata":{}},{"cell_type":"code","source":"torch.cuda.empty_cache()\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"KLDiv = nn.KLDivLoss(log_target=True)\n\nmodel.train()\noptimizer = torch.optim.AdamW(model.parameters(), lr=hyper.learning_rate)\n\nn_iter = 0\nloss = None\n\nfor epoch in range(hyper.epochs):    \n    with tqdm(train_loader, unit=\"batch\") as tepoch:\n        for batch_iter, batch in enumerate(tepoch):\n           \n        \n            tepoch.set_description(f\"Epoch {epoch}\")\n\n            \n            if n_iter % hyper.eval_interval == 0: \n                val_loss = eval_loop()\n                \n                if loss is not None:\n                    print(f\"step {n_iter}: train loss {loss:.4f}, val loss {val_loss:.4f}\")\n                else:\n                    print(f\"step {n_iter}: val loss {val_loss:.4f}\")\n        \n            xb, yb = batch\n            \n            yb = torch.unsqueeze(yb, dim=0)\n            \n            xb = xb.to(device)\n            yb = yb.to(device)\n\n            yb = F.one_hot(yb, num_classes=6)\n            yb = yb.to(torch.float32)\n            \n                \n                                \n            outputs = model(xb)\n\n            probs = F.softmax(outputs, dim=-1) #(1,6)\n\n            loss = KLDiv(probs, yb)\n                \n\n\n            loss.backward()\n\n            # Normalize the Gradients\n            \n            tepoch.set_postfix(loss=loss.item())\n\n            \n            optimizer.step()\n            optimizer.zero_grad()\n\n            torch.cuda.empty_cache()\n            \n            n_iter += 1\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef eval_loop(val_loader, eval_iters):\n    model.eval()\n    preds = pd.DataFrame(np.zeros((hyper.batch_size * eval_iters, 6)))\n    true = pd.DataFrame(np.zeros((hyper.batch_size * eval_iters, 6)))\n\n    for i, batch in enumerate(val_loader):\n        if i == eval_iters:\n            break\n            \n        xb, yb = batch\n            \n        yb = torch.unsqueeze(yb, dim=-1)\n        yb = yb.expand(hyper.batch_size, hyper.time_step) # this is currently done to enable loss calculation \n\n\n        xb = xb.to(device)\n        yb = yb.to(device)\n\n        yb = F.one_hot(yb, num_classes=6)\n        yb = yb.to(torch.float32)\n        \n        #print(yb)\n    \n        output = model(xb)\n        \n  \n            \n        probs = F.log_softmax(output, dim=2) #(1, 2000, 6)\n      \n        \n        preds.iloc[i*hyper.batch_size:i*hyper.batch_size + hyper.batch_size] = aggregated_log_probs.detach().numpy()\n        true.iloc[i*hyper.batch_size:i*hyper.batch_size + hyper.batch_size] = yb.detach().numpy()\n\n        \n    \n    preds['id'] = np.arange(hyper.batch_size * eval_iters)\n    true['id'] = np.arange(hyper.batch_size * eval_iters)\n\n        \n    return preds, true","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"score(targets, pred)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\n\npreds, true = eval_loop(val_loader,eval_iters=2)\n\n\nprint(preds)\nprint(true)\n\nscore(preds, true,row_id_column_name='id')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"true[-30:]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}