{"cells":[{"metadata":{},"cell_type":"markdown","source":"# Riiid! SAINT+ Training\n\n## Versions\n\n### Version 4\n- Everything correct. "},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import gc\nimport random\nimport joblib\nfrom collections import defaultdict\nimport _pickle as cPickle\nimport numpy as np\nimport pandas as pd\nfrom numba import njit\nfrom scipy.stats import rankdata\nfrom tqdm.notebook import tqdm\nfrom sklearn.model_selection import train_test_split\n\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom pathlib import Path\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"@njit\ndef _auc(actual, pred_ranks):\n    actual = np.asarray(actual)\n    pred_ranks = np.asarray(pred_ranks)\n    n_pos = np.sum(actual)\n    n_neg = len(actual) - n_pos\n    return (np.sum(pred_ranks[actual==1]) - n_pos*(n_pos+1)/2) / (n_pos*n_neg)\n\ndef auc_score(actual, predicted):\n    pred_ranks = rankdata(predicted)\n    return _auc(actual, pred_ranks)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"MAX_SEQ = 200\nTRAIN_ACCEPTED_USER_CONTENT_SIZE = 4\nVALID_ACCEPTED_USER_CONTENT_SIZE = 2\nNUM_ENCODER_LAYERS = 3\nNUM_DECODER_LAYERS = 3\nEMBED_SIZE = 256\nD_MODEL = 256\nBATCH_SIZE = 128\nDROPOUT = 0.1\nTEST_SIZE = 0.1\n\nFIRST_ROUND_LR = 2e-3\nFIRST_ROUND_EPOCHS = 10\nFIRST_ROUND_VERBOSE = 5\n\nSECOND_ROUND_LR = 2e-4\nSECOND_ROUND_EPOCHS = 3\nSECOND_ROUND_VERBOSE = 4\n\nMAX_STEPS = 3\nMODEL_PATH = 'saint.pth'","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"n_skill = 13524\nn_parts = 8\nn_lsi_tags = 21\nn_correct_answers = 5","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"```python\ngroup = df.groupby('user_id').apply(\n    lambda r: (\n        # mask\n        r['task_container_id'].to_numpy().astype('int16'),\n\n        # encoder\n        r['content_id'].to_numpy().astype('int16'),\n        r['part'].to_numpy().astype('int8'), \n        r['correct_answer'].to_numpy().astype('int8'),\n        r['lsi_tag'].to_numpy().astype('int8'),\n        r['question_accuracy'].to_numpy().astype('float32'),\n        r['question_elapsed'].to_numpy().astype('float32'),\n        r['question_lagtime'].to_numpy().astype('float32'),\n        r['part_accuracy'].to_numpy().astype('float32'),\n        r['lsi_accuracy'].to_numpy().astype('float32'),\n\n        # decoder\n        r['answered_correctly'].to_numpy().astype('int8'), \n        r['user_answer'].to_numpy().astype('int8'),\n        r['prior_question_had_explanation'].to_numpy().astype('bool'),\n        r['prior_question_elapsed_time'].to_numpy().astype('float32'),\n        r['lagtime'].to_numpy().astype('float32'),\n    )\n)\n```"},{"metadata":{"trusted":true},"cell_type":"code","source":"%%time\n\nwith open(\"../input/riiid-saint-final-training-dataset/group_train.pickle\", \"rb\") as f:\n    group = cPickle.load(f)\ngc.collect()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class SAINTDataset:\n    def __init__(self, group, keys, max_seq=128, min_seqlen=1, dicts_folder=\"\"):\n        super(SAINTDataset, self).__init__()\n        \n        # self.load_dicts(dicts_folder)\n        self.max_seq = max_seq\n        \n        self.user_ids = []\n        self.samples = {}\n        for user_id in tqdm(keys, total=len(keys)):\n            (task_container_id, content_id, part, correct_answer, lsi_tag, question_accuracy, \n             question_elapsed, question_lagtime, part_accuracy, lsi_accuracy,\n             answered_correctly, user_answer, prior_qn_explanation, elapsed_time, lagtime) = group[user_id]\n            \n            if len(content_id) > max_seq:\n                total_questions = len(content_id)\n                last_pos = total_questions // max_seq\n\n                for seq in range(last_pos):\n                    index = f\"{user_id}_{seq}\"\n                    self.user_ids.append(index)\n                    start = seq * self.max_seq\n                    end = (seq + 1) * self.max_seq\n                    self.samples[index] = (\n                        task_container_id[start:end], \n                        content_id[start:end], part[start:end], correct_answer[start:end], \n                        lsi_tag[start:end], question_accuracy[start:end], question_elapsed[start:end], \n                        question_lagtime[start:end], part_accuracy[start:end], lsi_accuracy[start:end],\n                        answered_correctly[start:end], user_answer[start:end], prior_qn_explanation[start:end], elapsed_time[start:end], lagtime[start:end]\n                    )\n\n                if len(content_id[end:]) > min_seqlen:\n                    index = f\"{user_id}_{last_pos + 1}\"\n                    self.user_ids.append(index)\n                    self.samples[index] = (\n                        task_container_id[end:], content_id[end:], part[end:], correct_answer[end:], \n                        lsi_tag[end:], question_accuracy[end:], question_elapsed[end:], \n                        question_lagtime[end:], part_accuracy[end:], lsi_accuracy[end:],\n                        answered_correctly[end:], user_answer[end:], prior_qn_explanation[end:], \n                        elapsed_time[end:], lagtime[end:]\n                    )\n            \n            else:\n                index = f'{user_id}'\n                self.user_ids.append(index)\n                self.samples[index] = (\n                    task_container_id, content_id, part, correct_answer, lsi_tag, question_accuracy, \n                    question_elapsed, question_lagtime, part_accuracy, lsi_accuracy,\n                    answered_correctly, user_answer, prior_qn_explanation, elapsed_time, lagtime\n                )\n        \n    def __len__(self):\n        return len(self.user_ids)\n    \n    def __getitem__(self, index):\n        \n        user_id = self.user_ids[index]\n        \n        # load data\n        (task_container_id, \n         content_id, part, correct_answer, lsi_tag, question_accuracy, \n         question_elapsed, question_lagtime, part_accuracy, lsi_accuracy,\n         answered_correctly, user_answer, prior_qn_explanation, elapsed_time, lagtime) = self.samples[user_id]\n        \n        seq_len = len(content_id)\n        \n        ## create variables\n        # others\n        task_container_id_seq = np.zeros(self.max_seq, dtype=int)\n        label_seq = np.zeros(self.max_seq)\n        # encoder\n        content_id_seq = np.zeros(self.max_seq, dtype=int)\n        part_id_seq = np.zeros(self.max_seq, dtype=int)\n        correct_answer_seq = np.zeros(self.max_seq, dtype=int)\n        lsi_tag_seq = np.zeros(self.max_seq, dtype=int)\n        question_accuracy_seq = np.zeros(self.max_seq)\n        question_elapsed_seq = np.zeros(self.max_seq)\n        question_lagtime_seq = np.zeros(self.max_seq)\n        part_accuracy_seq = np.zeros(self.max_seq)\n        lsi_accuracy_seq = np.zeros(self.max_seq)\n        # decoder\n        response_seq = np.zeros(self.max_seq, dtype=int)\n        user_answer_seq = np.zeros(self.max_seq, dtype=int)\n        prior_qn_explanation_seq = np.zeros(self.max_seq, dtype=int)\n        elapsed_time_seq = np.zeros(self.max_seq)\n        lagtime_seq = np.zeros(self.max_seq)\n        \n        if seq_len >= self.max_seq:\n            # others\n            task_container_id_seq = task_container_id[-self.max_seq:]\n            label_seq = answered_correctly[-self.max_seq:]\n            # encoder\n            content_id_seq = content_id[-self.max_seq:]\n            part_id_seq = part[-self.max_seq:]\n            correct_answer_seq = correct_answer[-self.max_seq:]\n            lsi_tag_seq = lsi_tag[-self.max_seq:]\n            question_accuracy_seq = question_accuracy[-self.max_seq:]\n            question_elapsed_seq = question_elapsed[-self.max_seq:]\n            question_lagtime_seq = question_lagtime[-self.max_seq:]\n            part_accuracy_seq = part_accuracy[-self.max_seq:]\n            lsi_accuracy_seq = lsi_accuracy[-self.max_seq:]\n            # decoder\n            user_answer_seq = user_answer[-self.max_seq:]\n            prior_qn_explanation_seq = prior_qn_explanation[-self.max_seq:]\n            elapsed_time_seq = elapsed_time[-self.max_seq:]\n            lagtime_seq = lagtime[-self.max_seq:]\n            \n        else:\n            # others\n            task_container_id_seq[-seq_len:] = task_container_id\n            label_seq[-seq_len:] = answered_correctly\n            # encoder\n            content_id_seq[-seq_len:] = content_id\n            part_id_seq[-seq_len:] = part\n            correct_answer_seq[-seq_len:] = correct_answer\n            lsi_tag_seq[-seq_len:] = lsi_tag\n            question_accuracy_seq[-seq_len:] = question_accuracy\n            question_elapsed_seq[-seq_len:] = question_elapsed\n            question_lagtime_seq[-seq_len:] = question_lagtime\n            part_accuracy_seq[-seq_len:] = part_accuracy\n            lsi_accuracy_seq[-seq_len:] = lsi_accuracy\n            # decoder\n            user_answer_seq[-seq_len:] = user_answer\n            prior_qn_explanation_seq[-seq_len:] = prior_qn_explanation\n            elapsed_time_seq[-seq_len:] = elapsed_time\n            lagtime_seq[-seq_len:] = lagtime\n            \n        response_seq = np.append([0], label_seq[:-1]).astype('float')\n        user_answer_seq = np.append([0], user_answer_seq[:-1]).astype('int')\n        \n        # Create mask\n        mask = self.create_task_mask(task_container_id_seq)\n        \n        return (\n            content_id_seq, part_id_seq, correct_answer_seq, lsi_tag_seq, question_accuracy_seq,  # encoder\n            question_elapsed_seq, question_lagtime_seq, part_accuracy_seq, lsi_accuracy_seq,  # encoder\n            response_seq, user_answer_seq, prior_qn_explanation_seq, elapsed_time_seq, lagtime_seq,  # decoder\n            label_seq, mask  # others\n        )\n        \n    \n    def create_task_mask(self, tasks):\n        seq_length = len(tasks)\n        future_mask = np.triu(np.ones((seq_length, seq_length)), k=1).astype('bool')\n        container_mask = np.ones((seq_length, seq_length))\n        container_mask = (container_mask * tasks.reshape(1,-1)) == (container_mask * tasks.reshape(-1,1))\n        future_mask = future_mask + container_mask\n        np.fill_diagonal(future_mask, 0)\n        return future_mask\n    \n    def load_dicts(self, dicts_folder):\n        with open(f\"{dicts_folder}/questions_label_mean_dict.pickle\", \"rb\") as f:\n            self.questions_label_mean_dict = cPickle.load(f)\n        with open(f\"{dicts_folder}/questions_elapsed_mean_dict.pickle\", \"rb\") as f:\n            self.questions_elapsed_mean_dict = cPickle.load(f)\n        with open(f\"{dicts_folder}/questions_lag_mean_dict.pickle\", \"rb\") as f:\n            self.questions_lag_mean_dict = cPickle.load(f)\n        \n        with open(f\"{dicts_folder}/part_label_mean_dict.pickle\", \"rb\") as f:\n            self.part_label_mean_dict = cPickle.load(f)\n            \n        with open(f\"{dicts_folder}/lsi_label_mean_dict.pickle\", \"rb\") as f:\n            self.lsi_label_mean_dict = cPickle.load(f)\n            \n        questions_info = pd.read_feather(f\"{dicts_folder}/questions_info.feather\")\n        questions_info.set_index('question_id', inplace=True)\n        self.part_mapping = questions_info['part'].to_dict(defaultdict(int))\n        self.correct_answer_mapping = questions_info['correct_answer'].to_dict(defaultdict(int))\n        self.lsi_mapping = questions_info['lsi_tag'].to_dict(defaultdict(int))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"%%time\n\ntrain, val = train_test_split(list(group.keys()), test_size=TEST_SIZE, random_state=42)\ngc.collect()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# %%time\n# train, val = train_test_split(group, test_size=TEST_SIZE)\n# del group","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"%%time\n\ntrain_dataset = SAINTDataset(group, train, max_seq=MAX_SEQ, min_seqlen=TRAIN_ACCEPTED_USER_CONTENT_SIZE)\ntrain_dataloader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)\ndel train","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"item = next(iter(train_dataloader))\nprint(\n    item[0].shape, item[1].shape, item[2].shape, item[3].shape, item[4].shape, item[5].shape,\n    item[6].shape, item[7].shape, item[8].shape, item[9].shape, item[10].shape, item[11].shape,\n    item[12].shape, item[13].shape, item[14].shape, item[15].shape\n)\ndel item\ngc.collect()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"%%time\n\nval_dataset = SAINTDataset(group, val, max_seq=MAX_SEQ, min_seqlen=VALID_ACCEPTED_USER_CONTENT_SIZE)\nval_dataloader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False)\nnp.save(\"val.npy\", val)\ndel val","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"del group\ngc.collect()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class EncoderEmbed(nn.Module):\n    def __init__(self, n_skill, n_parts, n_correct_answers, n_lsi_tags, max_seq=128, embed_dim=128):\n        super(EncoderEmbed, self).__init__()\n        self.content_id = nn.Embedding(n_skill, embed_dim, padding_idx=0)\n        self.part_id = nn.Embedding(n_parts, embed_dim, padding_idx=0)\n        self.correct_answer = nn.Embedding(n_correct_answers, embed_dim, padding_idx=0)\n        self.lsi_id = nn.Embedding(n_lsi_tags, embed_dim, padding_idx=0)\n        self.question_accuracy = nn.Linear(1, embed_dim)\n        self.question_elapsed = nn.Linear(1, embed_dim)\n        self.question_lagtime = nn.Linear(1, embed_dim)\n        self.part_accuracy = nn.Linear(1, embed_dim)\n        self.lsi_accuracy = nn.Linear(1, embed_dim)\n        self.pos_embed = nn.Embedding(max_seq, embed_dim)\n        \n        self.pos_ids = torch.arange(max_seq).unsqueeze(0)\n        \n    def forward(self, \n                content_id, part_id, correct_answer, lsi_id, \n                question_accuracy, question_elapsed, question_lagtime, \n                part_accuracy, lsi_accuracy):\n        x = self.content_id(content_id) + self.part_id(part_id)\n        x += self.correct_answer(correct_answer) + self.lsi_id(lsi_id)\n        x += self.question_accuracy(question_accuracy.unsqueeze(-1)) \n        x += self.question_elapsed(question_elapsed.unsqueeze(-1)) + self.question_lagtime(question_lagtime.unsqueeze(-1))\n        x += self.part_accuracy(part_accuracy.unsqueeze(-1)) + self.lsi_accuracy(lsi_accuracy.unsqueeze(-1))\n        x += self.pos_embed(self.pos_ids.to(device))\n        \n        return x","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class DecoderEmbed(nn.Module):\n    def __init__(self, n_correct_answers, max_seq=128, embed_dim=128):\n        super(DecoderEmbed, self).__init__()\n        self.responses = nn.Embedding(2, embed_dim)\n        self.user_answer = nn.Embedding(n_correct_answers, embed_dim, padding_idx=0)\n        self.prior_qn_explanation = nn.Embedding(2, embed_dim)\n        self.elapsed_time = nn.Linear(1, embed_dim)\n        self.lagtime = nn.Linear(1, embed_dim)\n        self.pos_embed = nn.Embedding(max_seq, embed_dim)\n        \n        self.pos_ids = torch.arange(max_seq).unsqueeze(0)\n        \n    def forward(self, responses, user_answer, prior_qn_explanation, elapsed_time, lagtime):\n        x = self.responses(responses) + self.user_answer(user_answer)\n        x += self.prior_qn_explanation(prior_qn_explanation)\n        x += self.elapsed_time(elapsed_time.unsqueeze(-1)) + self.lagtime(lagtime.unsqueeze(-1))\n        x += self.pos_embed(self.pos_ids.to(device))\n        return x","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class SAINTModel(nn.Module):\n    def __init__(self, \n                 n_skill, n_parts, n_correct_answers, n_lsi_tags, \n                 max_seq=128, embed_dim=128, d_model=128, dropout=0.0, \n                 forward_expansion=1, enc_layers=1, dec_layers=1, heads=8\n                ):\n        super(SAINTModel, self).__init__()\n        \n        self.enc_layers = enc_layers\n        self.dec_layers = dec_layers\n        self.heads = heads\n        \n        self.enc_embed = EncoderEmbed(n_skill, n_parts, n_correct_answers, n_lsi_tags, max_seq, embed_dim)\n        self.dec_embed = DecoderEmbed(n_correct_answers,  max_seq, embed_dim)\n        self.transformer = torch.nn.Transformer(\n            d_model, heads, enc_layers, dec_layers, embed_dim * forward_expansion, dropout)\n        self.pred = nn.Linear(embed_dim, 1)\n        \n    def forward(\n        self, \n        content_id, part_id, correct_answer, lsi_id,  # encoder inputs\n        question_accuracy, question_elapsed, question_lagtime, part_accuracy, lsi_accuracy,  # encoder inputs\n        responses, user_answer, prior_qn_explanation, elapsed_time, lagtime,  # decoder inputs\n        mask  # mask\n    ):\n        \n        enc_embed = self.enc_embed(\n            content_id, part_id, correct_answer, lsi_id, \n            question_accuracy, question_elapsed, question_lagtime, part_accuracy, lsi_accuracy)\n        dec_embed = self.dec_embed(responses, user_answer, prior_qn_explanation, elapsed_time, lagtime)\n        enc_embed = enc_embed.permute(1, 0, 2)\n        dec_embed = dec_embed.permute(1, 0, 2)\n\n        x = self.transformer(\n            enc_embed, \n            dec_embed,\n            src_mask=mask.repeat(self.heads, 1, 1), \n            tgt_mask=mask.repeat(self.heads, 1, 1), \n            memory_mask=mask.repeat(self.heads, 1, 1)\n        )\n        \n        x = x.permute(1, 0, 2)\n        x = self.pred(x)\n        return x.squeeze(-1)\n    \n    def create_future_mask(self, seqlen1, seqlen2):\n        future_mask = np.triu(np.ones((seqlen1, seqlen2)), k=1).astype('bool')\n        return torch.from_numpy(future_mask)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Main changes are possibility of forward expansion and stacking of encoding layers\ndef create_model():\n    return SAINTModel(\n        n_skill, n_parts, n_correct_answers, n_lsi_tags,\n        max_seq=MAX_SEQ, embed_dim=EMBED_SIZE, d_model=D_MODEL, forward_expansion=1, \n        enc_layers=NUM_ENCODER_LAYERS, dec_layers=NUM_DECODER_LAYERS, heads=8, dropout=DROPOUT\n    )\n\nmodel = create_model()\nmodel.to(device)\nprint(model)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def load_from_item(item):\n    # content_id, part_id, correct_answer, lsi_id,  # encoder inputs\n    # question_accuracy, question_elapsed, question_lagtime, part_accuracy, lsi_accuracy,  # encoder inputs\n    # responses, user_answer, prior_qn_explanation, elapsed_time, lagtime,  # decoder inputs\n    # label_seq, mask  # others\n    \n    content_id = item[0].to(device, non_blocking=True).long()\n    part_id = item[1].to(device, non_blocking=True).long()\n    correct_answer = item[2].to(device, non_blocking=True).long()\n    lsi_id = item[3].to(device, non_blocking=True).long()\n    question_accuracy = item[4].to(device, non_blocking=True).float()\n    question_elapsed = item[5].to(device, non_blocking=True).float()\n    question_lagtime = item[6].to(device, non_blocking=True).float()\n    part_accuracy = item[7].to(device, non_blocking=True).float()\n    lsi_accuracy = item[8].to(device, non_blocking=True).float()\n    \n    responses = item[9].to(device, non_blocking=True).long()\n    user_answer = item[10].to(device, non_blocking=True).long()\n    prior_qn_explanation = item[11].to(device, non_blocking=True).long()\n    elapsed_time = item[12].to(device, non_blocking=True).float()\n    lagtime = item[13].to(device, non_blocking=True).float()\n    \n    label_seq = item[14].to(device, non_blocking=True).float()\n    mask = item[15].to(device, non_blocking=True).bool()\n    target_mask = (content_id != 0)\n    \n    return (\n        content_id, part_id, correct_answer, lsi_id,  # encoder inputs\n        question_accuracy, question_elapsed, question_lagtime, part_accuracy, lsi_accuracy,  # encoder inputs\n        responses, user_answer, prior_qn_explanation, elapsed_time, lagtime,  # decoder inputs\n        label_seq, mask, target_mask  # others\n    )\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def update_stats(tbar, train_loss, loss, output, label, num_corrects, num_total, labels, outs):\n    train_loss.append(loss.item())\n    pred = (torch.sigmoid(output) >= 0.5).long()\n    num_corrects += (pred == label).sum().item()\n    num_total += len(label)\n    labels.extend(label.view(-1).detach().cpu().numpy())\n    outs.extend(output.view(-1).detach().cpu().numpy())\n    tbar.set_description('loss - {:.4f}'.format(loss.item()))\n    return num_corrects, num_total","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def train_epoch(model, dataloader, optim, criterion, scheduler, device=\"cpu\"):\n    model.train()\n    \n    train_loss = []\n    num_corrects = 0\n    num_total = 0\n    labels = []\n    outs = []\n    \n    tbar = tqdm(dataloader)\n    for item in tbar:\n        ( \n            content_id, part_id, correct_answer, lsi_id,  # encoder inputs\n            question_accuracy, question_elapsed, question_lagtime, part_accuracy, lsi_accuracy,  # encoder inputs\n            responses, user_answer, prior_qn_explanation, elapsed_time, lagtime,  # decoder inputs\n            label, mask, target_mask  # others\n        ) = load_from_item(item)\n        \n        optim.zero_grad()\n        output = model(\n            content_id, part_id, correct_answer, lsi_id,  # encoder inputs\n            question_accuracy, question_elapsed, question_lagtime, part_accuracy, lsi_accuracy,  # encoder inputs\n            responses, user_answer, prior_qn_explanation, elapsed_time, lagtime,  # decoder inputs\n            mask  # others\n        )\n        \n        output = torch.masked_select(output, target_mask)\n        label = torch.masked_select(label, target_mask)\n        \n        loss = criterion(output, label)\n        loss.backward()\n        optim.step()\n        scheduler.step()\n        \n        tbar.set_description('loss - {:.4f}'.format(loss.item()))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def train_verbose_epoch(model, dataloader, optim, criterion, scheduler, device=\"cpu\"):\n    model.train()\n    \n    train_loss = []\n    num_corrects = 0\n    num_total = 0\n    labels = []\n    outs = []\n    \n    tbar = tqdm(dataloader)\n    for item in tbar:\n        ( \n            content_id, part_id, correct_answer, lsi_id,  # encoder inputs\n            question_accuracy, question_elapsed, question_lagtime, part_accuracy, lsi_accuracy,  # encoder inputs\n            responses, user_answer, prior_qn_explanation, elapsed_time, lagtime,  # decoder inputs\n            label, mask, target_mask  # others\n        ) = load_from_item(item)\n        \n        optim.zero_grad()\n        output = model(\n            content_id, part_id, correct_answer, lsi_id,  # encoder inputs\n            question_accuracy, question_elapsed, question_lagtime, part_accuracy, lsi_accuracy,  # encoder inputs\n            responses, user_answer, prior_qn_explanation, elapsed_time, lagtime,  # decoder inputs\n            mask  # others\n        )\n        \n        output = torch.masked_select(output, target_mask)\n        label = torch.masked_select(label, target_mask)\n        \n        loss = criterion(output, label)\n        loss.backward()\n        optim.step()\n        scheduler.step()\n        \n        num_corrects, num_total = update_stats(tbar, train_loss, loss, output, label, num_corrects, num_total, labels, outs)\n    \n    acc = num_corrects / num_total\n    auc = auc_score(labels, outs)\n    loss = np.average(train_loss)\n\n    return loss, acc, auc","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def val_epoch(model, val_iterator, criterion, device=\"cpu\"):\n    model.eval()\n\n    train_loss = []\n    num_corrects = 0\n    num_total = 0\n    labels = []\n    outs = []\n\n    tbar = tqdm(val_iterator)\n    for item in tbar:\n        ( \n            content_id, part_id, correct_answer, lsi_id,  # encoder inputs\n            question_accuracy, question_elapsed, question_lagtime, part_accuracy, lsi_accuracy,  # encoder inputs\n            responses, user_answer, prior_qn_explanation, elapsed_time, lagtime,  # decoder inputs\n            label, mask, target_mask  # others\n        ) = load_from_item(item)\n        \n        with torch.no_grad():\n            output = model(\n                content_id, part_id, correct_answer, lsi_id,  # encoder inputs\n                question_accuracy, question_elapsed, question_lagtime, part_accuracy, lsi_accuracy,  # encoder inputs\n                responses, user_answer, prior_qn_explanation, elapsed_time, lagtime,  # decoder inputs\n                mask  # others\n            )\n        \n        output = torch.masked_select(output, target_mask)\n        label = torch.masked_select(label, target_mask)\n\n        loss = criterion(output, label)\n        \n        num_corrects, num_total = update_stats(tbar, train_loss, loss, output, label, num_corrects, num_total, labels, outs)\n\n    acc = num_corrects / num_total\n    auc = auc_score(labels, outs)\n    loss = np.average(train_loss)\n\n    return loss, acc, auc","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def do_train(best_auc=0, learning_rate=0.001, epochs=20, max_steps=3, verbose_eval=1):\n    optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate)\n    criterion = nn.BCEWithLogitsLoss()\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer, max_lr=learning_rate, \n                                                    steps_per_epoch=len(train_dataloader), epochs=epochs)\n    model.to(device, non_blocking=True)\n    criterion.to(device, non_blocking=True)\n    best_auc = best_auc\n    cur_step = 0\n    for epoch in range(epochs):\n        train_epoch(model, train_dataloader, optimizer, criterion, scheduler, device)\n        # if (epoch + 1) % verbose_eval == 0:\n        #     train_loss, train_acc, train_auc = val_epoch(model, train_dataloader, criterion, device)\n        #     print(f\"{epoch + 1}/{epochs}: train: loss - {train_loss:.3f} acc - {train_acc:.3f} auc - {train_auc:.4f}\")\n        val_loss, val_acc, val_auc = val_epoch(model, val_dataloader, criterion, device)\n        print(f\"{epoch + 1}/{epochs}: val: loss - {val_loss:.3f} acc - {val_acc:.3f} auc - {val_auc:.4f}\")\n        if best_auc < val_auc:\n            print(f'epoch - {epoch + 1} best model with val auc: {val_auc:4f}')\n            best_auc = val_auc\n            torch.save(model, MODEL_PATH)\n        else:\n            cur_step += 1\n        \n        if cur_step >= max_steps:\n            break\n    \n    return best_auc","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"best_auc = do_train(0, FIRST_ROUND_LR, FIRST_ROUND_EPOCHS, MAX_STEPS, FIRST_ROUND_VERBOSE)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"torch.save(model, \"saint_final.pth\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat":4,"nbformat_minor":4}