{"cells":[{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:18:43.78769Z","start_time":"2021-01-07T03:18:43.785621Z"},"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5"},"cell_type":"markdown","source":"# UniNet "},{"metadata":{},"cell_type":"markdown","source":"### UniNet is a SAINT+ like Model. Compared with SAINT+, UniNet puts skill into decoder.\n### P.S. Here we also calculate the lag time roughly. (Lag time contributed by https://www.kaggle.com/zlhaaaph)"},{"metadata":{},"cell_type":"markdown","source":"# Some package import"},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:18:44.819295Z","start_time":"2021-01-07T03:18:43.790117Z"},"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","trusted":false},"cell_type":"code","source":"# import gc\n# import random\n# from tqdm import tqdm\n# from sklearn.metrics import roc_auc_score\n# from sklearn.model_selection import train_test_split\n\n# import numpy as np \n# import pandas as pd\n\n# import seaborn as sns\n# import matplotlib.pyplot as plt\n\n# import torch\n# import torch.nn as nn\n# import torch.nn.utils.rnn as rnn_utils\n# from torch.autograd import Variable\n# from torch.utils.data import Dataset, DataLoader\n# from pathlib import Path\n# import datatable as dt\n# import os\n# SEED = 9999\n# os.environ[\"PYTHONHASHSEED\"] = str(SEED)\n# os.environ[\"CUDA_DEVICE_ORDER\"] = 'PCI_BUS_ID'\n# os.environ[\"CUDA_VISIBLE_DEVICES\"] = '0'","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:18:44.823523Z","start_time":"2021-01-07T03:18:44.820707Z"},"trusted":false},"cell_type":"code","source":"# MAX_SEQ = 180\n# ACCEPTED_USER_CONTENT_SIZE = 4\n# EMBED_SIZE = 64\n# BATCH_SIZE = 32\n# DROPOUT = 0.1\n\n# LR = 2e-3\n# EPOCHS = 10\n# MODEL_PATH = 'sakt.pth'","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Data pre-process"},{"metadata":{},"cell_type":"markdown","source":"## Load Data"},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:18:44.83309Z","start_time":"2021-01-07T03:18:44.824583Z"},"trusted":false},"cell_type":"code","source":"# path = Path('/kaggle/input')\n# assert path.exists()","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:19:44.750189Z","start_time":"2021-01-07T03:18:44.834324Z"},"trusted":false},"cell_type":"code","source":"# data_types_dict = {\n#     'content_type_id': 'bool',\n#     'timestamp': 'int64',\n#     'user_id': 'int32', \n#     'content_id': 'int16', \n#     'answered_correctly': 'int8', \n#     'prior_question_elapsed_time': 'float32', \n#     'prior_question_had_explanation': 'bool',\n#     'task_container_id':'int8',\n#     'row_id':'int64'\n# }\n# target = 'answered_correctly'\n# # train_df = dt.fread(path/'riiid-test-answer-prediction/train.csv', columns=set(data_types_dict.keys())).to_pandas()\n# train_df = dt.fread('../data/train.csv', columns=set(data_types_dict.keys())).to_pandas()\n# # train_df[['prior_question_had_explanation']] = train_df[['prior_question_had_explanation']].shift(-1)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Pre-process"},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:20:27.499936Z","start_time":"2021-01-07T03:19:44.752102Z"},"trusted":false},"cell_type":"code","source":"# buf = train_df.drop_duplicates(subset=['user_id','task_container_id'],keep='first',inplace=False).sort_values(by='task_container_id')[['user_id','task_container_id','prior_question_elapsed_time']]\n# buf['question_elapsed_time'] = buf.groupby('user_id')['prior_question_elapsed_time'].shift(-1)","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:20:28.865332Z","start_time":"2021-01-07T03:20:27.502507Z"},"trusted":false},"cell_type":"code","source":"# buf = buf.drop(columns=['prior_question_elapsed_time'])","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:21:28.369418Z","start_time":"2021-01-07T03:20:28.868758Z"},"trusted":false},"cell_type":"code","source":"# train_df = pd.merge(train_df,buf,on=['user_id','task_container_id'],how='left')\n# train_df['question_elapsed_time'] = train_df['question_elapsed_time'].fillna(0).astype(int)","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:22:54.142213Z","start_time":"2021-01-07T03:21:28.371255Z"},"trusted":false},"cell_type":"code","source":"# # 推算lag_time\n# train_df['start_time'] = train_df.groupby('user_id')['prior_question_elapsed_time'].shift(-1)\n# train_df['start_time'] = train_df['timestamp'] - train_df['start_time']# 开始做题时间\n# train_df['lag_time'] = train_df.groupby('user_id')['start_time'].diff()\n# train_df['lag_time'] = train_df['lag_time'].fillna(train_df['lag_time'].quantile(0.5))# decoder\n# train_df['lag_time'] = (train_df['lag_time']/60000).astype(int)\n# train_df.loc[train_df['lag_time']>60*24*7,'lag_time']=60*24*7\n# train_df['lag_time'] = train_df['lag_time']-train_df['lag_time'].min()+1# 保证大于0\n# train_df['content_id'] = train_df['content_id']+1\n# train_df['question_elapsed_time'] = train_df['question_elapsed_time']+1","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:23:29.803295Z","start_time":"2021-01-07T03:22:54.144432Z"},"trusted":false},"cell_type":"code","source":"# train_df = train_df[train_df.content_type_id == False]\n# #arrange by timestamp\n# train_df = train_df.sort_values(['timestamp'], ascending=True).reset_index(drop = True)","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:23:29.813245Z","start_time":"2021-01-07T03:23:29.80535Z"},"trusted":false},"cell_type":"code","source":"# train_df.info()","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:23:30.751339Z","start_time":"2021-01-07T03:23:29.815338Z"},"trusted":false},"cell_type":"code","source":"# n_content_id = train_df[\"content_id\"].nunique()\n# print(\"number content_id\", n_content_id)\n# n_elapsed_time = int(train_df[\"question_elapsed_time\"].max())\n# print(\"number elapsed_time\", n_elapsed_time)\n# n_lag_time = int(train_df[\"lag_time\"].max())\n# print(\"number lag_time\", n_lag_time)\n\n# num_content_id_item = n_content_id+1\n# num_label_item = 3\n# num_skill_item = 2*n_content_id+1\n# num_elapsed_time_item = n_elapsed_time+1\n# num_lag_time_item = n_lag_time+1","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:27.239197Z","start_time":"2021-01-07T03:23:30.752892Z"},"trusted":false},"cell_type":"code","source":"# group = train_df[['user_id', 'content_id', 'question_elapsed_time','answered_correctly','lag_time']].groupby('user_id').apply(lambda r: (r['content_id'].values, r['question_elapsed_time'].values, r['answered_correctly'].values, r['lag_time'].values))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Train"},{"metadata":{},"cell_type":"markdown","source":"## Data Loaders"},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:27.256922Z","start_time":"2021-01-07T03:24:27.241018Z"},"trusted":false},"cell_type":"code","source":"# class UniDataset(Dataset):\n#     def __init__(self, group, n_skill, max_seq=100):\n#         super(UniDataset, self).__init__()\n#         self.samples, self.n_skill, self.max_seq = {}, n_skill, max_seq\n        \n#         self.user_ids = []\n#         for i, user_id in enumerate(group.index):\n#             if(i % 10000 == 0):\n#                 print(f'Processed {i} users')\n#             content_id, question_elapsed_time, answered_correctly, lag_time = group[user_id]\n#             if len(content_id) >= ACCEPTED_USER_CONTENT_SIZE:\n#                 if len(content_id) > self.max_seq:\n#                     total_questions = len(content_id)\n#                     last_pos = total_questions // self.max_seq\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] = (content_id[start:end], question_elapsed_time[start:end], answered_correctly[start:end],lag_time[start:end])\n#                     if len(content_id[end:]) >= ACCEPTED_USER_CONTENT_SIZE:\n#                         index = f\"{user_id}_{last_pos + 1}\"\n#                         self.user_ids.append(index)\n#                         self.samples[index] = (content_id[end:], question_elapsed_time[end:], answered_correctly[end:],lag_time[end:])\n#                 else:\n#                     index = f'{user_id}'\n#                     self.user_ids.append(index)\n#                     self.samples[index] = (content_id, question_elapsed_time, answered_correctly,lag_time)\n#     def __len__(self):\n#         return len(self.user_ids)\n\n#     def __getitem__(self, index):\n#         user_id = self.user_ids[index]\n#         content_id, question_elapsed_time, answered_correctly, lag_time = self.samples[user_id]\n#         seq_len = len(content_id)\n        \n#         content_id_seq = np.zeros(self.max_seq, dtype=int)\n#         question_elapsed_time_seq = np.zeros(self.max_seq, dtype=int)\n#         answered_correctly_seq = np.zeros(self.max_seq, dtype=int)\n#         lag_time_seq = np.zeros(self.max_seq, dtype=int)\n#         label_seq = np.zeros(self.max_seq, dtype=int)\n#         if seq_len >= self.max_seq:\n#             content_id_seq[:] = content_id[-self.max_seq:]\n#             question_elapsed_time_seq[:] = question_elapsed_time[-self.max_seq:]\n#             answered_correctly_seq[:] = answered_correctly[-self.max_seq:]+1# 改了\n#             label_seq[:] = answered_correctly[-self.max_seq:]\n#             lag_time_seq[:] = lag_time[-self.max_seq:]\n#         else:\n#             content_id_seq[-seq_len:] = content_id\n#             question_elapsed_time_seq[-seq_len:] = question_elapsed_time\n#             answered_correctly_seq[-seq_len:] = answered_correctly+1# 改了\n#             label_seq[-seq_len:] = answered_correctly\n#             lag_time_seq[-seq_len:] = lag_time\n            \n#         content_seq = content_id_seq[1:]\n#         elapsed_time_seq = question_elapsed_time_seq[1:]\n#         label_seq = label_seq[1:]\n#         lag_time_seq = lag_time_seq[1:]\n        \n#         skill_seq = content_id_seq[:-1].copy()\n#         skill_seq += (answered_correctly_seq[:-1] == 1) * self.n_skill\n        \n#         return skill_seq,elapsed_time_seq,content_seq,label_seq,lag_time_seq","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:27.323474Z","start_time":"2021-01-07T03:24:27.258215Z"},"trusted":false},"cell_type":"code","source":"# TEST_SIZE = 0.1\n\n# train, val = train_test_split(group, test_size = TEST_SIZE)","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:34.035624Z","start_time":"2021-01-07T03:24:27.32489Z"},"trusted":false},"cell_type":"code","source":"# train_dataset = UniDataset(train, n_content_id, max_seq=MAX_SEQ)\n# train_dataloader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=0)\n# del train","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:34.838173Z","start_time":"2021-01-07T03:24:34.037103Z"},"trusted":false},"cell_type":"code","source":"# val_dataset = UniDataset(val, n_content_id, max_seq=MAX_SEQ)\n# val_dataloader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=0)\n# del val","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:34.921447Z","start_time":"2021-01-07T03:24:34.839978Z"},"trusted":false},"cell_type":"code","source":"# sample_batch = next(iter(train_dataloader))\n# sample_batch[0].shape, sample_batch[1].shape, sample_batch[2].shape, sample_batch[3].shape, sample_batch[4].shape","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Define model"},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:34.928037Z","start_time":"2021-01-07T03:24:34.922698Z"},"trusted":false},"cell_type":"code","source":"# def 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# future_mask(5)","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:34.943691Z","start_time":"2021-01-07T03:24:34.929176Z"},"trusted":false},"cell_type":"code","source":"# class UniModel(nn.Module):\n#     def __init__(self,emb_dim,seq_len,num_content_id_item,num_label_item,num_skill_item,num_elapsed_time_item,num_lag_time_item):\n#         super(UniModel,self).__init__() \n        \n#         self.enc_seq_content_id_emb = nn.Embedding(num_content_id_item, emb_dim)\n        \n#         self.dec_seq_label_emb = nn.Embedding(num_label_item, emb_dim)\n#         self.dec_seq_skill_emb = nn.Embedding(num_skill_item, emb_dim)\n#         self.dec_seq_elapsed_time_emb = nn.Embedding(num_elapsed_time_item, emb_dim)\n#         self.dec_seq_lag_time_emb = nn.Embedding(num_lag_time_item, emb_dim)\n\n#         self.pos_emb = nn.Embedding(seq_len, emb_dim)\n        \n#         self.transformer_model = nn.Transformer(d_model=emb_dim, nhead=16, num_encoder_layers=1)\n#         self.pred = nn.Linear(emb_dim, 1)\n\n        \n#     def forward(self, skill_seq, elapsed_time_seq, content_seq, label_seq, lag_time_seq):\n            \n#             start_token = torch.Tensor(label_seq.size(0),1).fill_(0)\n#             dec_label_seq = torch.cat((start_token.to(device),label_seq[:,:-1].to(device)),-1).to(device)\n            \n#             dec_label_seq = self.dec_seq_label_emb(dec_label_seq.long().to(device))\n#             dec_skill_seq = self.dec_seq_skill_emb(skill_seq.long().to(device))\n#             dec_elapsed_time_seq = self.dec_seq_elapsed_time_emb(elapsed_time_seq.long().to(device))\n#             dec_lag_time_seq = self.dec_seq_lag_time_emb(lag_time_seq.long().to(device))\n#             dec_seq = dec_label_seq+dec_skill_seq+dec_elapsed_time_seq+dec_lag_time_seq\n#             dec_seq = dec_seq.permute(1,0,2)\n            \n#             enc_seq = self.enc_seq_content_id_emb(content_seq.long()).to(device)\n#             pos_id = torch.arange(enc_seq.size(1)).unsqueeze(0).long().to(device)\n#             pos = self.pos_emb(pos_id.long())\n#             enc_seq = enc_seq+pos\n#             enc_seq = enc_seq.permute(1,0,2)\n            \n#             trans_mask = future_mask(enc_seq.size(0)).to(device)\n#             out = self.transformer_model(enc_seq, dec_seq, src_mask=trans_mask, tgt_mask=trans_mask)\n#             out = out.permute(1, 0, 2)\n#             out = self.pred(out)\n            \n#             return out.squeeze(-1)","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:35.043773Z","start_time":"2021-01-07T03:24:34.950236Z"},"trusted":false},"cell_type":"code","source":"# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:35.054638Z","start_time":"2021-01-07T03:24:35.045311Z"},"trusted":false},"cell_type":"code","source":"# device","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:38.190261Z","start_time":"2021-01-07T03:24:35.056119Z"},"scrolled":true,"trusted":false},"cell_type":"code","source":"# def create_model(emb_dim,seq_len,num_content_id_item,num_label_item,num_skill_item,num_elapsed_time_item,num_lag_time_item):\n#     return UniModel(emb_dim,seq_len-1,num_content_id_item,num_label_item,num_skill_item,num_elapsed_time_item,num_lag_time_item).to(device)\n# model = create_model(EMBED_SIZE,MAX_SEQ,num_content_id_item,num_label_item,num_skill_item,num_elapsed_time_item,num_lag_time_item)\n# model","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:38.194359Z","start_time":"2021-01-07T03:24:38.191941Z"},"trusted":false},"cell_type":"code","source":"# skill_seq,elapsed_time_seq,content_seq,label_seq,lag_time_seq\n","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:38.213971Z","start_time":"2021-01-07T03:24:38.195443Z"},"trusted":false},"cell_type":"code","source":"# pd.Series(sample_batch[0].numpy().reshape(-1,)).nunique(),pd.Series(sample_batch[0].numpy().reshape(-1,)).describe()","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:38.222855Z","start_time":"2021-01-07T03:24:38.215803Z"},"trusted":false},"cell_type":"code","source":"# pd.Series(sample_batch[1].numpy().reshape(-1,)).nunique(),pd.Series(sample_batch[1].numpy().reshape(-1,)).describe()","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:38.233441Z","start_time":"2021-01-07T03:24:38.223891Z"},"trusted":false},"cell_type":"code","source":"# pd.Series(sample_batch[2].numpy().reshape(-1,)).nunique(),pd.Series(sample_batch[2].numpy().reshape(-1,)).describe()","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:38.242761Z","start_time":"2021-01-07T03:24:38.234566Z"},"trusted":false},"cell_type":"code","source":"# pd.Series(sample_batch[3].numpy().reshape(-1,)).nunique(),pd.Series(sample_batch[3].numpy().reshape(-1,)).describe()","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:38.253757Z","start_time":"2021-01-07T03:24:38.243911Z"},"trusted":false},"cell_type":"code","source":"# pd.Series(sample_batch[4].numpy().reshape(-1,)).nunique(),pd.Series(sample_batch[4].numpy().reshape(-1,)).describe()","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:38.307033Z","start_time":"2021-01-07T03:24:38.254932Z"},"scrolled":true,"trusted":false},"cell_type":"code","source":"# model(sample_batch[0].to(device), sample_batch[1].to(device), sample_batch[2].to(device),sample_batch[3].to(device),sample_batch[4].to(device))[0]","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Training"},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:38.322848Z","start_time":"2021-01-07T03:24:38.30811Z"},"trusted":false},"cell_type":"code","source":"# def load_from_item(item):\n#     skill_seq = item[0].to(device).long()\n#     elapsed_time_seq = item[1].to(device).long()\n#     content_seq = item[2].to(device).long()\n#     label_seq = item[3].to(device).float()\n#     lag_time_seq = item[4].to(device).long()\n#     content_mask = (content_seq != 0)\n#     return skill_seq, elapsed_time_seq, content_seq, label_seq, lag_time_seq, content_mask\n\n# 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).data.cpu().numpy())\n#     outs.extend(output.view(-1).data.cpu().numpy())\n#     tbar.set_description('loss - {:.4f}'.format(loss))\n#     return num_corrects, num_total\n\n# 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#         skill_seq, elapsed_time_seq, content_seq, label_seq, lag_time_seq, content_mask = load_from_item(item)\n        \n#         optim.zero_grad()\n#         output = model(skill_seq, elapsed_time_seq, content_seq, label_seq, lag_time_seq)\n        \n        \n#         output = torch.masked_select(output, content_mask)\n#         label_seq = torch.masked_select(label_seq, content_mask)\n        \n#         loss = criterion(output, label_seq)\n#         loss.backward()\n#         optim.step()\n#         scheduler.step()\n        \n#         tbar.set_description('loss - {:.4f}'.format(loss))\n\n# 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#         skill_seq, elapsed_time_seq, content_seq, label_seq, lag_time_seq, content_mask = load_from_item(item)\n\n#         with torch.no_grad():\n#             output = model(skill_seq, elapsed_time_seq, content_seq, label_seq, lag_time_seq)\n        \n#         output = torch.masked_select(output, content_mask)\n#         label_seq = torch.masked_select(label_seq, content_mask)\n\n#         loss = criterion(output, label_seq)\n        \n#         num_corrects, num_total = update_stats(tbar, train_loss, loss, output, label_seq, num_corrects, num_total, labels, outs)\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\n","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"end_time":"2021-01-07T03:24:38.33363Z","start_time":"2021-01-07T03:24:38.32376Z"},"trusted":false},"cell_type":"code","source":"# def do_train():\n#     optimizer = torch.optim.Adam(model.parameters(), lr=LR)\n#     criterion = nn.BCEWithLogitsLoss()\n#     scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer, max_lr=LR, \n#                                                     steps_per_epoch=len(train_dataloader), epochs=EPOCHS)\n#     model.to(device)\n#     criterion.to(device)\n#     best_auc = 0.0\n#     for epoch in range(EPOCHS):\n#         train_epoch(model, train_dataloader, optimizer, criterion, scheduler, device)\n#         val_loss, avl_acc, val_auc = val_epoch(model, val_dataloader, criterion, device)\n#         print(f\"epoch - {epoch + 1} val_loss - {val_loss:.3f} val_acc - {avl_acc:.3f} val_auc - {val_auc:.3f}\")\n#         if best_auc < val_auc:\n#             print(f'epoch - {epoch + 1} best model with val auc: {val_auc}')\n#             best_auc = val_auc\n#         torch.save(model.state_dict(), MODEL_PATH)","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"start_time":"2021-01-07T03:18:42.965Z"},"scrolled":true,"trusted":false},"cell_type":"code","source":"# do_train()","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"start_time":"2021-01-07T03:18:42.968Z"},"trusted":false},"cell_type":"code","source":"# LR = 2e-4\n# EPOCHS = 3\n\n# do_train()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Predict (Notice! This predict code is not fix to adapt UniNet)"},{"metadata":{"ExecuteTime":{"start_time":"2021-01-07T03:18:42.97Z"},"trusted":false},"cell_type":"code","source":"# model = create_model()\n# model.load_state_dict(torch.load(MODEL_PATH))\n# model.to(device)","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"start_time":"2021-01-07T03:18:42.972Z"},"trusted":false},"cell_type":"code","source":"# class TestDataset(Dataset):\n#     def __init__(self, samples, test_df, n_skill, max_seq=100):\n#         super(TestDataset, self).__init__()\n#         self.samples, self.user_ids, self.test_df = samples, [x for x in test_df[\"user_id\"].unique()], test_df\n#         self.n_skill, self.max_seq = n_skill, max_seq\n\n#     def __len__(self):\n#         return self.test_df.shape[0]\n    \n#     def __getitem__(self, index):\n#         test_info = self.test_df.iloc[index]\n        \n#         user_id = test_info['user_id']\n#         target_id = test_info['content_id']\n        \n#         content_id_seq = np.zeros(self.max_seq, dtype=int)\n#         answered_correctly_seq = np.zeros(self.max_seq, dtype=int)\n        \n#         if user_id in self.samples.index:\n#             content_id, answered_correctly = self.samples[user_id]\n            \n#             seq_len = len(content_id)\n            \n#             if seq_len >= self.max_seq:\n#                 content_id_seq = content_id[-self.max_seq:]\n#                 answered_correctly_seq = answered_correctly[-self.max_seq:]\n#             else:\n#                 content_id_seq[-seq_len:] = content_id\n#                 answered_correctly_seq[-seq_len:] = answered_correctly\n                \n#         x = content_id_seq[1:].copy()\n#         x += (answered_correctly_seq[1:] == 1) * self.n_skill\n        \n#         questions = np.append(content_id_seq[2:], [target_id])\n        \n#         return x, questions","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"start_time":"2021-01-07T03:18:42.975Z"},"trusted":false},"cell_type":"code","source":"# import riiideducation\n\n# env = riiideducation.make_env()\n# iter_test = env.iter_test()","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"start_time":"2021-01-07T03:18:42.98Z"},"trusted":false},"cell_type":"code","source":"# import psutil\n\n# model.eval()\n\n# prev_test_df = None\n\n# for (test_df, sample_prediction_df) in tqdm(iter_test):\n    \n#     if (prev_test_df is not None) & (psutil.virtual_memory().percent<90):\n#         print(psutil.virtual_memory().percent)\n#         prev_test_df['answered_correctly'] = eval(test_df['prior_group_answers_correct'].iloc[0])\n#         prev_test_df = prev_test_df[prev_test_df.content_type_id == False]\n#         prev_group = prev_test_df[['user_id', 'content_id', 'answered_correctly']].groupby('user_id').apply(lambda r: (\n#             r['content_id'].values,\n#             r['answered_correctly'].values))\n#         for prev_user_id in prev_group.index:\n#             prev_group_content = prev_group[prev_user_id][0]\n#             prev_group_answered_correctly = prev_group[prev_user_id][1]\n#             if prev_user_id in group.index:\n#                 group[prev_user_id] = (np.append(group[prev_user_id][0], prev_group_content), \n#                                        np.append(group[prev_user_id][1], prev_group_answered_correctly))\n#             else:\n#                 group[prev_user_id] = (prev_group_content, prev_group_answered_correctly)\n            \n#             if len(group[prev_user_id][0]) > MAX_SEQ:\n#                 new_group_content = group[prev_user_id][0][-MAX_SEQ:]\n#                 new_group_answered_correctly = group[prev_user_id][1][-MAX_SEQ:]\n#                 group[prev_user_id] = (new_group_content, new_group_answered_correctly)\n                \n#     prev_test_df = test_df.copy()\n#     test_df = test_df[test_df.content_type_id == False]\n    \n#     test_dataset = TestDataset(group, test_df, n_skill, max_seq=MAX_SEQ)\n#     test_dataloader = DataLoader(test_dataset, batch_size=len(test_df), shuffle=False)\n    \n#     item = next(iter(test_dataloader))\n#     x = item[0].to(device).long()\n#     target_id = item[1].to(device).long()\n    \n#     with torch.no_grad():\n#         output, _ = model(x, target_id)\n        \n#     output = torch.sigmoid(output)\n#     output = output[:, -1]\n#     test_df['answered_correctly'] = output.cpu().numpy()\n#     env.predict(test_df.loc[test_df['content_type_id'] == 0, ['row_id', 'answered_correctly']])","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"start_time":"2021-01-07T03:18:42.983Z"},"trusted":false},"cell_type":"code","source":"# test_df","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"start_time":"2021-01-07T03:18:42.985Z"},"trusted":false},"cell_type":"code","source":"# test_dataset = TestDataset(group, test_df, n_skill, max_seq=MAX_SEQ)","execution_count":null,"outputs":[]},{"metadata":{"ExecuteTime":{"start_time":"2021-01-07T03:18:42.988Z"},"trusted":false},"cell_type":"code","source":"# # Save to pickle to usage in other notebooks\n# group.to_pickle('/kaggle/working/group.pkl')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import riiideducation\nenv = riiideducation.make_env()\niter_test = env.iter_test()\nfor (test_df, sample_prediction_df) in iter_test:\n    test_df['answered_correctly'] = 0.5\n    env.predict(test_df.loc[test_df['content_type_id'] == 0, ['row_id', 'answered_correctly']])","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}