{"cells":[{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"import gc\nimport pickle\nimport random\nimport math\nimport numpy as np\nimport pandas as pd\nimport os\nfrom tqdm import tqdm\nfrom sklearn.metrics import roc_auc_score\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\nimport torch.nn.utils.rnn as rnn_utils\nfrom torch.autograd import Variable\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as F\nfrom torch.nn import TransformerEncoder, TransformerEncoderLayer","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def seed_everything(s):\n    random.seed(s)\n    os.environ['PYTHONHASHSEED'] = str(s)\n    np.random.seed(s)\n    # Torch\n    torch.manual_seed(s)\n    torch.cuda.manual_seed(s)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(s)\n\nseed = 2020\nseed_everything(seed)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#HDKIM\nMAX_SEQ = 256\n#HDKIMHDKIM","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Load data"},{"metadata":{"trusted":true},"cell_type":"code","source":"%%time\ndtype = {'timestamp':'int64', \n         'user_id':'int32' ,\n         'content_id':'int16',\n         'content_type_id':'int8',\n         'answered_correctly':'int8'}\nfeld_needed = ['timestamp', 'user_id', 'content_id', 'content_type_id', 'answered_correctly']\n\ntrain_df = pd.read_pickle('../input/riiid-cross-validation-files/cv1_train.pickle')[feld_needed]\ntrain_df.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"%%time\nvalid_df = pd.read_pickle('../input/riiid-cross-validation-files/cv1_valid.pickle')[feld_needed]\nvalid_df.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_df = train_df[train_df.content_type_id == False]\n\n#arrange by timestamp\ntrain_df = train_df.sort_values(['timestamp'], ascending=True).reset_index(drop = True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"valid_df = valid_df[valid_df.content_type_id == False]\n\n#arrange by timestamp\nvalid_df = valid_df.sort_values(['timestamp'], ascending=True).reset_index(drop = True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"questions_df = pd.read_csv('../input/riiid-test-answer-prediction/questions.csv')\n\ndef get_tags_id(df):\n    tags_dict = {}\n    for i,v in enumerate(df['tags'].unique()):\n        tags_dict[v] = i\n    return tags_dict\ntags_dict_labels = get_tags_id(questions_df)\ndef label_tags(df, tags_dict):\n    df['tag_label'] = None\n    for i,v in enumerate(df['tags']):\n        df['tag_label'].iloc[i] = tags_dict[v]\n    return df\nquestions_df = label_tags(questions_df, tags_dict_labels)\nquestions_df['tag_label'] = questions_df['tag_label'].astype(int)\n\n# train_df = pd.merge(train_df, questions_df[['question_id', 'tag_label', 'part']], left_on = 'content_id', right_on = 'question_id', how = 'left')\n# valid_df = pd.merge(valid_df, questions_df[['question_id', 'tag_label', 'part']], left_on = 'content_id', right_on = 'question_id', how = 'left')","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Preprocess"},{"metadata":{"trusted":true},"cell_type":"code","source":"parts = questions_df[\"part\"].unique()\nn_part = len(parts)\nprint(\"number parts\", len(parts))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"tags = questions_df[\"tag_label\"].unique()\nn_tag = len(tags)\nprint(\"number tags\", len(tags))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"skills = train_df[\"content_id\"].unique()\nn_skill = len(skills)\nprint(\"number skills\", len(skills))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"gc.collect()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_group = pickle.load(open(\"../input/group-features-seq-160/train_group.pkl\", \"rb\"))\nvalid_group = pickle.load(open(\"../input/group-features-seq-160/valid_group.pkl\", \"rb\"))\n# train_group = train_df[['user_id', 'content_id', 'tag_label', 'part', 'answered_correctly']].groupby('user_id').apply(lambda r: (\n#             r['content_id'].values,\n#             r['tag_label'].values,\n#             r['part'].values,\n#             r['answered_correctly'].values))\n\n# valid_group = valid_df[['user_id', 'content_id', 'tag_label', 'part', 'answered_correctly']].groupby('user_id').apply(lambda r: (\n#             r['content_id'].values,\n#             r['tag_label'].values,\n#             r['part'].values,\n#             r['answered_correctly'].values))\n\ndel train_df, valid_df, questions_df\ngc.collect()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class SAKTDataset(Dataset):\n    def __init__(self, group, n_skill, n_tag, n_part, max_seq=MAX_SEQ): #HDKIM 100\n        super(SAKTDataset, self).__init__()\n        self.max_seq = max_seq\n        self.n_skill = n_skill\n        self.n_tag = n_tag\n        self.n_part = n_part\n        self.samples = group\n        \n        self.user_ids = []\n        for user_id in group.index:\n            q,tag, part, qa = group[user_id]\n            if len(q) < 2: #HDKIM 10\n                continue\n            self.user_ids.append(user_id)\n            \n\n    def __len__(self):\n        return len(self.user_ids)\n\n    def __getitem__(self, index):\n        user_id = self.user_ids[index]\n        q_, tag_, part_, qa_ = self.samples[user_id]\n        seq_len = len(q_)\n\n        q = np.zeros(self.max_seq, dtype=int)\n        tag = np.zeros(self.max_seq, dtype=int)\n        part = np.zeros(self.max_seq, dtype=int)\n        qa = np.zeros(self.max_seq, dtype=int)\n        \n        if seq_len >= self.max_seq:\n            #HDKIM\n            if random.random()>0.1:\n                start = random.randint(0,(seq_len-self.max_seq))\n                end = start + self.max_seq\n                q[:] = q_[start:end]\n                tag[:] = tag_[start:end]\n                part[:] = part_[start:end]\n                qa[:] = qa_[start:end]\n            else:\n                #HDKIMHDKIM\n                q[:] = q_[-self.max_seq:]\n                tag[:] = tag_[-self.max_seq:]\n                part[:] = part_[-self.max_seq:]\n                qa[:] = qa_[-self.max_seq:]\n        else:\n            #HDKIM\n            if random.random()>0.1:\n                #HDKIMHDKIM\n                start = 0\n                end = random.randint(2,seq_len)\n                seq_len = end - start\n                q[-seq_len:] = q_[0:seq_len]\n                tag[-seq_len:] = tag_[0:seq_len]\n                part[-seq_len:] = part_[0:seq_len]\n                qa[-seq_len:] = qa_[0:seq_len]\n            else:\n                #HDKIMHDKIM\n                q[-seq_len:] = q_\n                tag[-seq_len:] = tag_\n                part[-seq_len:] = part_\n                qa[-seq_len:] = qa_\n\n        \n        target_id = q[1:]\n        tag_id = tag[1:]\n        part_id = part[1:]\n        label = qa[1:]\n\n        x = np.zeros(self.max_seq-1, dtype=int)\n        x = q[:-1].copy()\n        x += (qa[:-1] == 1) * self.n_skill\n        \n        x_tag = np.zeros(self.max_seq-1, dtype=int)\n        x_tag = tag[:-1].copy()\n        x_tag += (qa[:-1] == 1) * self.n_tag\n        \n        x_part = np.zeros(self.max_seq-1, dtype=int)\n        x_part = part[:-1].copy()\n        x_part += (qa[:-1] == 1) * self.n_part\n\n        return x, x_tag, x_part, target_id, tag_id, part_id, label","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_dataset = SAKTDataset(train_group, n_skill, n_tag, n_part)\ntrain_dataloader = DataLoader(train_dataset, batch_size=256, shuffle=True, num_workers=8)\n\n# item = train_dataset.__getitem__(5)\n# print(item)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"valid_dataset = SAKTDataset(valid_group, n_skill, n_tag, n_part)\nvalid_dataloader = DataLoader(valid_dataset, batch_size=256, shuffle=True, num_workers=8)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import gc\ndel valid_dataset, train_dataset, train_group, valid_group\ngc.collect()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Define model"},{"metadata":{"trusted":true},"cell_type":"code","source":"class TransformerModel(nn.Module):\n\n    def __init__(self, ninp:int=32, nhead:int=2, nhid:int=64, nlayers:int=2, dropout:float=0.3):\n        '''\n        nhead -> number of heads in the transformer multi attention thing.\n        nhid -> the number of hidden dimension neurons in the model.\n        nlayers -> how many layers we want to stack.\n        '''\n        super(TransformerModel, self).__init__()\n        self.src_mask = None\n        encoder_layers = TransformerEncoderLayer(d_model=ninp, nhead=nhead, dim_feedforward=nhid, dropout=dropout, activation='relu')\n        self.transformer_encoder = TransformerEncoder(encoder_layer=encoder_layers, num_layers=nlayers)\n        self.exercise_embeddings = nn.Embedding(num_embeddings=n_skill+1, embedding_dim=ninp) # exercise_id\n        self.pos_embedding = nn.Embedding(ninp, ninp) # positional embeddings\n        self.part_embeddings = nn.Embedding(num_embeddings=n_part+1, embedding_dim=ninp) # part_id_embeddings\n        self.tag_label_embedding = nn.Embedding(num_embeddings=n_tag+1, embedding_dim=ninp) # prior_question_elapsed_time\n        self.ninp = ninp\n        self.decoder = nn.Linear(ninp, 1)\n        self.init_weights()\n\n    def init_weights(self):\n        initrange = 0.1\n        # init embeddings\n        self.exercise_embeddings.weight.data.uniform_(-initrange, initrange)\n        self.part_embeddings.weight.data.uniform_(-initrange, initrange)\n        self.tag_label_embedding.weight.data.uniform_(-initrange, initrange)\n        self.decoder.bias.data.zero_()\n        self.decoder.weight.data.uniform_(-initrange, initrange)\n    \n    def generate_square_subsequent_mask(self, sz):\n        mask = (torch.triu(torch.ones(sz, sz)) == 1).transpose(0, 1)\n        mask = mask.float().masked_fill(mask == 0, float('-inf')).masked_fill(mask == 1, float(0.0))\n        return mask\n    \n    def forward(self, content_id, tag_id, part_id, mask_src=None):\n        '''\n        S is the sequence length, N the batch size and E the Embedding Dimension (number of features).\n        src: (S, N, E)\n        src_mask: (S, S)\n        src_key_padding_mask: (N, S)\n        padding mask is (N, S) with boolean True/False.\n        SRC_MASK is (S, S) with float(’-inf’) and float(0.0).\n        '''\n        device = content_id.device  \n        q_em = self.exercise_embeddings(content_id)\n        part_em = self.part_embeddings(part_id)\n        tag_em = self.tag_label_embedding(tag_id)\n        \n\n        embedded_src = q_em + part_em + tag_em # (N, S, E)\n        embedded_src = embedded_src.transpose(0, 1) # (S, N, E)\n        \n        _src = embedded_src * np.sqrt(self.ninp)\n        \n        mask_src = self.generate_square_subsequent_mask(_src.shape[0]).to(device)\n        \n        output = self.transformer_encoder(src=_src, mask=mask_src)\n        output = self.decoder(output)\n        output = output.transpose(1, 0)\n        return output.squeeze(-1)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class FFN(nn.Module):\n    def __init__(self, state_size=200):\n        super(FFN, self).__init__()\n        self.state_size = state_size\n\n        self.lr1 = nn.Linear(state_size, state_size)\n        self.relu = nn.ReLU()\n        self.lr2 = nn.Linear(state_size+state_size, state_size)\n        self.lr3 = nn.Linear(state_size, state_size)\n        self.lr4 = nn.Linear(state_size + state_size, state_size)\n        self.dropout = nn.Dropout(0.2)\n    \n    def forward(self, x):\n        x1 = self.lr1(x)\n        x1 = self.relu(x1)\n        \n        x2 = torch.cat((x, x1), 2)\n        x2 = self.lr2(x2)\n        x2 = self.relu(x2)\n        \n        x3 = self.lr3(x2)\n        x3 = self.relu(x3)\n        \n        x4 = torch.cat((x3, x2), 2)\n        \n        x4 = self.lr4(x4)\n        return self.dropout(x4)\n\ndef future_mask(seq_length):\n    future_mask = np.triu(np.ones((seq_length, seq_length)), k=1).astype('bool')\n    return torch.from_numpy(future_mask)\n\n\nclass SAKTModel(nn.Module):\n    def __init__(self, n_skill, n_tag, n_part, max_seq=MAX_SEQ, embed_dim=128): #HDKIM 100->MAX_SEQ\n        super(SAKTModel, self).__init__()\n        self.n_skill = n_skill\n        self.n_tag = n_tag\n        self.n_part = n_part\n        self.embed_dim = embed_dim\n        \n\n        # question related models\n        self.q_embedding = nn.Embedding(2*n_skill+1, embed_dim)\n        self.q_pos_embedding = nn.Embedding(max_seq-1, embed_dim)\n        self.q_e_embedding = nn.Embedding(n_skill+1, embed_dim)\n        self.q_multi_att = nn.MultiheadAttention(embed_dim=embed_dim, num_heads=16, dropout=0.3)\n        self.layer_normal_q = nn.LayerNorm(embed_dim)\n        \n        # tags related models\n        self.t_embedding = nn.Embedding(2*n_tag+1, embed_dim)\n        self.t_pos_embedding = nn.Embedding(max_seq-1, embed_dim)\n        self.t_e_embedding = nn.Embedding(n_tag+1, embed_dim)\n        self.t_multi_att = nn.MultiheadAttention(embed_dim=embed_dim, num_heads=16, dropout=0.3)\n        self.layer_normal_tag = nn.LayerNorm(embed_dim)\n        \n        # parts related models\n        self.p_embedding = nn.Embedding(2*n_part+1, embed_dim)\n        self.p_pos_embedding = nn.Embedding(max_seq-1, embed_dim)\n        self.p_e_embedding = nn.Embedding(n_part+1, embed_dim)\n        self.p_multi_att = nn.MultiheadAttention(embed_dim=embed_dim, num_heads=16, dropout=0.3)\n        self.layer_normal_part = nn.LayerNorm(embed_dim)\n\n        # Common models and layers\n        self.dropout = nn.Dropout(0.25)\n        self.layer_normal = nn.LayerNorm(embed_dim*3)\n        self.ffn = FFN(embed_dim*3)\n        self.pred = nn.Linear(embed_dim*3, 1)\n    \n    def forward(self, x, x_tag, x_part, question_ids, tag_ids, part_ids):\n        ###########################\n        # question related features\n        ###########################\n        device = x.device        \n        x = self.q_embedding(x)\n        pos_id = torch.arange(x.size(1)).unsqueeze(0).to(device)\n        \n        pos_x = self.q_pos_embedding(pos_id)\n        x = x + pos_x\n        e = self.q_e_embedding(question_ids)\n        \n        x = x.permute(1, 0, 2) # x: [bs, s_len, embed] => [s_len, bs, embed]\n        e = e.permute(1, 0, 2)\n        att_mask_q = future_mask(x.size(0)).to(device)\n        att_output_q, att_weight_q = self.q_multi_att(e, x, x, attn_mask=att_mask_q)\n        att_output_q = self.layer_normal_q(att_output_q + e)\n        att_output_q = att_output_q.permute(1, 0, 2) # att_output: [s_len, bs, embed] => [bs, s_len, embed]\n        \n        ###########################\n        # tags related features\n        ###########################\n        device = x_tag.device        \n        x_tag = self.t_embedding(x_tag)\n        pos_tag_id = torch.arange(x_tag.size(1)).unsqueeze(0).to(device)\n\n        pos_tag_x = self.t_pos_embedding(pos_tag_id)\n        x_tag = x_tag + pos_tag_x\n        t = self.t_e_embedding(tag_ids)\n\n        x_tag = x_tag.permute(1, 0, 2) # x: [bs, s_len, embed] => [s_len, bs, embed]\n        t = t.permute(1, 0, 2)\n        att_mask_tag = future_mask(x_tag.size(0)).to(device)\n        att_output_tag, att_weight_tag = self.t_multi_att(t, x_tag, x_tag, attn_mask=att_mask_tag)\n        att_output_tag = self.layer_normal_tag(att_output_tag + t)\n        att_output_tag = att_output_tag.permute(1, 0, 2) # att_output: [s_len, bs, embed] => [bs, s_len, embed]\n        \n        ###########################\n        # part related features\n        ###########################\n        device = x_part.device        \n        x_part = self.p_embedding(x_part)\n        pos_part_id = torch.arange(x_part.size(1)).unsqueeze(0).to(device)\n\n        pos_part_x = self.p_pos_embedding(pos_part_id)\n        x_part = x_part + pos_part_x\n        p = self.p_e_embedding(part_ids)\n\n        x_part = x_part.permute(1, 0, 2) # x: [bs, s_len, embed] => [s_len, bs, embed]\n        p = p.permute(1, 0, 2)\n        att_mask_part = future_mask(x_part.size(0)).to(device)\n        att_output_part, att_weight_part = self.p_multi_att(p, x_part, x_part, attn_mask=att_mask_part)\n        att_output_part = self.layer_normal_part(att_output_part + p)\n        att_output_part = att_output_part.permute(1, 0, 2) # att_output: [s_len, bs, embed] => [bs, s_len, embed]\n        \n        ###########################\n        # combine all features\n        ###########################\n        att_output = torch.cat((att_output_q, att_output_tag, att_output_part), 2)\n        #print(att_output.shape)\n        \n        x = self.ffn(att_output)\n        #print(x.shape)\n        x = self.layer_normal(x + att_output)\n        x = self.pred(x)\n\n        return x.squeeze(-1), att_weight_q, att_weight_tag, att_weight_part","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = TransformerModel(ninp=MAX_SEQ, nhead=8, nhid=512, nlayers=4, dropout=0.3)\n\n# try:\n#     print('Loading on GPU')\n#     model.load_state_dict(torch.load('../input/transformer-encoder-model/transformer_model.pt'))\n# except:\n#     model.load_state_dict(torch.load('../input/transformer-encoder-model/transformer_model.pt', map_location='cpu'))\n\n# optimizer = torch.optim.SGD(model.parameters(), lr=1e-3, momentum=0.99, weight_decay=0.005)\noptimizer = torch.optim.Adam(model.parameters(), lr=5e-4)\ncriterion = nn.BCEWithLogitsLoss()\n\nmodel.to(device)\ncriterion.to(device)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def train_epoch(model, train_iterator, optim, criterion, 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(train_iterator)\n    for item in tbar:\n        # x, x_tag, x_part, target_id, tag_id, part_id, label\n        x = item[0].to(device).long()\n        x_tag = item[1].to(device).long()\n        x_part = item[2].to(device).long()\n        target_id = item[3].to(device).long()\n        tag_id = item[4].to(device).long()\n        part_id = item[5].to(device).long()\n        label = item[6].to(device).float()\n\n        optim.zero_grad()\n        # preds, att_weight_q, att_weight_t = net(x, x_tag, question_ids, tag_ids)\n        output = model(target_id, tag_id, part_id)\n        loss = criterion(output, label)\n        loss.backward()\n        optim.step()\n        train_loss.append(loss.item())\n\n        output = output[:, -1]\n        label = label[:, -1] \n        pred = (torch.sigmoid(output) >= 0.5).long()\n        \n        num_corrects += (pred == label).sum().item()\n        num_total += len(label)\n\n        labels.extend(label.view(-1).data.cpu().numpy())\n        outs.extend(output.view(-1).data.cpu().numpy())\n\n        tbar.set_description('loss - {:.4f}'.format(loss))\n\n    acc = num_corrects / num_total\n    auc = roc_auc_score(labels, outs)\n    loss = np.mean(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        # x, x_tag, x_part, target_id, tag_id, part_id, label\n        x = item[0].to(device).long()\n        x_tag = item[1].to(device).long()\n        x_part = item[2].to(device).long()\n        target_id = item[3].to(device).long()\n        tag_id = item[4].to(device).long()\n        part_id = item[5].to(device).long()\n        label = item[6].to(device).float()\n        target_mask = (target_id != 0)\n\n        with torch.no_grad():\n            output = model(target_id, tag_id, part_id)\n        \n        output = torch.masked_select(output, target_mask)\n        label = torch.masked_select(label, target_mask)\n\n        loss = criterion(output, label)\n        train_loss.append(loss.item())\n\n        pred = (torch.sigmoid(output) >= 0.5).long()\n        \n        num_corrects += (pred == label).sum().item()\n        num_total += len(label)\n\n        labels.extend(label.view(-1).data.cpu().numpy())\n        outs.extend(output.view(-1).data.cpu().numpy())\n\n        tbar.set_description('loss - {:.4f}'.format(loss))\n\n    acc = num_corrects / num_total\n    auc = roc_auc_score(labels, outs)\n    loss = np.average(train_loss)\n\n    return loss, acc, auc","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Train"},{"metadata":{"trusted":true},"cell_type":"code","source":"epochs = 35 #HDKIM 20\nlr_scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer,\n                                                          mode='max',\n                                                          factor=0.5,\n                                                          patience=2,\n                                                          threshold=0.0001, \n                                                          threshold_mode='rel',\n                                                          cooldown=0, min_lr=1e-6,\n                                                          eps=1e-08, verbose=True)\n\nover_fit = 0\nlast_auc = 0\nfor epoch in range(epochs):\n    \n    random.seed(epoch)\n    \n    train_loss, train_acc, train_auc = train_epoch(model, train_dataloader, optimizer, criterion, device)\n    print(\"epoch - {} train_loss - {:.2f} acc - {:.3f} auc - {:.3f}\".format(epoch, train_loss, train_acc, train_auc))\n    \n    val_loss, avl_acc, val_auc = val_epoch(model, valid_dataloader, criterion, device)\n    lr_scheduler.step(val_auc)\n    print(\"epoch - {} val_loss - {:.2f} acc - {:.3f} auc - {:.3f}\".format(epoch, val_loss, avl_acc, val_auc))\n    \n    if val_auc > last_auc:\n        last_auc = val_auc\n        over_fit = 0\n    else:\n        over_fit += 1\n        \n    if over_fit >= 4:\n        print(\"early stop epoch \", epoch)\n        break","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"torch.save(model.state_dict(), \"transformer_model.pt\")","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}