{"cells":[{"metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2020-12-24T15:33:00.553425Z","iopub.status.busy":"2020-12-24T15:33:00.552764Z","iopub.status.idle":"2020-12-24T15:33:00.571829Z","shell.execute_reply":"2020-12-24T15:33:00.572549Z"},"papermill":{"duration":0.052563,"end_time":"2020-12-24T15:33:00.572742","exception":false,"start_time":"2020-12-24T15:33:00.520179","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# installation without internet\n!pip install ../input/python-datatable/datatable-0.11.0-cp37-cp37m-manylinux2010_x86_64.whl","execution_count":null,"outputs":[]},{"metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","execution":{"iopub.execute_input":"2020-12-24T15:33:00.63175Z","iopub.status.busy":"2020-12-24T15:33:00.631079Z","iopub.status.idle":"2020-12-24T15:33:03.243918Z","shell.execute_reply":"2020-12-24T15:33:03.245147Z"},"papermill":{"duration":2.646243,"end_time":"2020-12-24T15:33:03.245371","exception":false,"start_time":"2020-12-24T15:33:00.599128","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"import gc\nimport random\nfrom tqdm.notebook 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\nimport joblib\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\nfrom pathlib import Path\nimport datatable as dt","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.041609,"end_time":"2020-12-24T15:33:03.342403","exception":false,"start_time":"2020-12-24T15:33:03.300794","status":"completed"},"tags":[]},"cell_type":"markdown","source":"### Load Data"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:33:03.43719Z","iopub.status.busy":"2020-12-24T15:33:03.435507Z","iopub.status.idle":"2020-12-24T15:33:03.438056Z","shell.execute_reply":"2020-12-24T15:33:03.436483Z"},"papermill":{"duration":0.051424,"end_time":"2020-12-24T15:33:03.43822","exception":false,"start_time":"2020-12-24T15:33:03.386796","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"path = Path('/kaggle/input')\nassert path.exists()","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:33:03.52318Z","iopub.status.busy":"2020-12-24T15:33:03.522234Z","iopub.status.idle":"2020-12-24T15:35:53.39115Z","shell.execute_reply":"2020-12-24T15:35:53.392123Z"},"papermill":{"duration":169.915909,"end_time":"2020-12-24T15:35:53.392379","exception":false,"start_time":"2020-12-24T15:33:03.47647","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"%%time\n\ndata_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}\ntarget = 'answered_correctly'\ntrain_df = dt.fread(path/'riiid-test-answer-prediction/train.csv', columns=set(data_types_dict.keys())).to_pandas()","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:35:53.475924Z","iopub.status.busy":"2020-12-24T15:35:53.473553Z","iopub.status.idle":"2020-12-24T15:35:53.492309Z","shell.execute_reply":"2020-12-24T15:35:53.491721Z"},"papermill":{"duration":0.069476,"end_time":"2020-12-24T15:35:53.492433","exception":false,"start_time":"2020-12-24T15:35:53.422957","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"train_df.info()","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:35:53.556749Z","iopub.status.busy":"2020-12-24T15:35:53.555886Z","iopub.status.idle":"2020-12-24T15:36:31.397924Z","shell.execute_reply":"2020-12-24T15:36:31.398558Z"},"papermill":{"duration":37.876948,"end_time":"2020-12-24T15:36:31.398711","exception":false,"start_time":"2020-12-24T15:35:53.521763","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"%%time\n\ntrain_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":{"execution":{"iopub.execute_input":"2020-12-24T15:36:31.496506Z","iopub.status.busy":"2020-12-24T15:36:31.489106Z","iopub.status.idle":"2020-12-24T15:36:31.499367Z","shell.execute_reply":"2020-12-24T15:36:31.498842Z"},"papermill":{"duration":0.072132,"end_time":"2020-12-24T15:36:31.499493","exception":false,"start_time":"2020-12-24T15:36:31.427361","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"del train_df['timestamp']\ndel train_df['content_type_id']","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.028366,"end_time":"2020-12-24T15:36:31.557454","exception":false,"start_time":"2020-12-24T15:36:31.529088","status":"completed"},"tags":[]},"cell_type":"markdown","source":"### Pre-process"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:36:31.619539Z","iopub.status.busy":"2020-12-24T15:36:31.618228Z","iopub.status.idle":"2020-12-24T15:36:32.44074Z","shell.execute_reply":"2020-12-24T15:36:32.440227Z"},"papermill":{"duration":0.855613,"end_time":"2020-12-24T15:36:32.440839","exception":false,"start_time":"2020-12-24T15:36:31.585226","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"n_skill = train_df[\"content_id\"].nunique()\nprint(\"number skills\", n_skill)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"joblib.dump(n_skill, \"skills.pkl.zip\")","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:36:32.935299Z","iopub.status.busy":"2020-12-24T15:36:32.933956Z","iopub.status.idle":"2020-12-24T15:37:15.531703Z","shell.execute_reply":"2020-12-24T15:37:15.53253Z"},"papermill":{"duration":43.062908,"end_time":"2020-12-24T15:37:15.532698","exception":false,"start_time":"2020-12-24T15:36:32.46979","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"%%time\n\ngroup = train_df[['user_id', 'content_id', 'answered_correctly']].groupby('user_id').apply(lambda r: (r['content_id'].values, r['answered_correctly'].values))\n\ndel train_df","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"joblib.dump(group, \"group.pkl.zip\")","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.029026,"end_time":"2020-12-24T15:37:15.591296","exception":false,"start_time":"2020-12-24T15:37:15.56227","status":"completed"},"tags":[]},"cell_type":"markdown","source":"### Data Loaders"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:15.656906Z","iopub.status.busy":"2020-12-24T15:37:15.65603Z","iopub.status.idle":"2020-12-24T15:37:15.658882Z","shell.execute_reply":"2020-12-24T15:37:15.659388Z"},"papermill":{"duration":0.038192,"end_time":"2020-12-24T15:37:15.659515","exception":false,"start_time":"2020-12-24T15:37:15.621323","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# 210 2 256 96 0.1","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"MAX_SEQ = 240 # 210\nACCEPTED_USER_CONTENT_SIZE = 2 # 2\nEMBED_SIZE = 256 # 256\nBATCH_SIZE = 64+32 # 96\nDROPOUT = 0.1 # 0.1\n\nprint(f\"{MAX_SEQ} {ACCEPTED_USER_CONTENT_SIZE} {EMBED_SIZE} {BATCH_SIZE} {DROPOUT}\")","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:15.742724Z","iopub.status.busy":"2020-12-24T15:37:15.741799Z","iopub.status.idle":"2020-12-24T15:37:15.744677Z","shell.execute_reply":"2020-12-24T15:37:15.744071Z"},"papermill":{"duration":0.054375,"end_time":"2020-12-24T15:37:15.74477","exception":false,"start_time":"2020-12-24T15:37:15.690395","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"class SAKTDataset(Dataset):\n    def __init__(self, group, n_skill, max_seq=100):\n        super(SAKTDataset, 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, answered_correctly = 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], answered_correctly[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:], answered_correctly[end:])\n                else:\n                    index = f'{user_id}'\n                    self.user_ids.append(index)\n                    self.samples[index] = (content_id, answered_correctly)\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        content_id, answered_correctly = self.samples[user_id]\n        seq_len = len(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        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        target_id = content_id_seq[1:]\n        label = answered_correctly_seq[1:]\n        \n        x = content_id_seq[:-1].copy()\n        x += (answered_correctly_seq[:-1] == 1) * self.n_skill\n        \n        return x, target_id, label","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:15.812583Z","iopub.status.busy":"2020-12-24T15:37:15.811521Z","iopub.status.idle":"2020-12-24T15:37:15.870207Z","shell.execute_reply":"2020-12-24T15:37:15.869469Z"},"papermill":{"duration":0.09523,"end_time":"2020-12-24T15:37:15.870318","exception":false,"start_time":"2020-12-24T15:37:15.775088","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"TEST_SIZE = 0.05\ntrain, val = train_test_split(group, test_size = TEST_SIZE, random_state=42)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:15.936442Z","iopub.status.busy":"2020-12-24T15:37:15.935526Z","iopub.status.idle":"2020-12-24T15:37:21.051549Z","shell.execute_reply":"2020-12-24T15:37:21.050446Z"},"papermill":{"duration":5.151062,"end_time":"2020-12-24T15:37:21.051684","exception":false,"start_time":"2020-12-24T15:37:15.900622","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"train_dataset = SAKTDataset(train, n_skill, max_seq=MAX_SEQ)\ntrain_dataloader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=8)\ndel train","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:21.133569Z","iopub.status.busy":"2020-12-24T15:37:21.132607Z","iopub.status.idle":"2020-12-24T15:37:21.63529Z","shell.execute_reply":"2020-12-24T15:37:21.634602Z"},"papermill":{"duration":0.546084,"end_time":"2020-12-24T15:37:21.635418","exception":false,"start_time":"2020-12-24T15:37:21.089334","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"val_dataset = SAKTDataset(val, n_skill, max_seq=MAX_SEQ)\nval_dataloader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=8)\ndel val","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:21.720854Z","iopub.status.busy":"2020-12-24T15:37:21.720055Z","iopub.status.idle":"2020-12-24T15:37:22.211323Z","shell.execute_reply":"2020-12-24T15:37:22.214941Z"},"papermill":{"duration":0.542483,"end_time":"2020-12-24T15:37:22.215191","exception":false,"start_time":"2020-12-24T15:37:21.672708","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"sample_batch = next(iter(train_dataloader))\nsample_batch[0].shape, sample_batch[1].shape, sample_batch[2].shape","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.052656,"end_time":"2020-12-24T15:37:22.323314","exception":false,"start_time":"2020-12-24T15:37:22.270658","status":"completed"},"tags":[]},"cell_type":"markdown","source":"### Define model"},{"metadata":{"trusted":true},"cell_type":"code","source":"# class PositionalEncoding(nn.Module):\n\n#     def __init__(self, d_model, dropout=0.1, max_len=5000):\n#         super(PositionalEncoding, self).__init__()\n#         self.dropout = nn.Dropout(p=dropout)\n\n#         pe = torch.zeros(max_len, d_model)\n#         position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)\n#         div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))\n#         pe[:, 0::2] = torch.sin(position * div_term)\n#         pe[:, 1::2] = torch.cos(position * div_term)\n#         pe = pe.unsqueeze(0).transpose(0, 1)\n#         self.register_buffer('pe', pe)\n\n#     def forward(self, x):\n#         #print(x.shape,self.pe.shape)\n#         x = x + self.pe.permute(1, 0, 2)\n#         return self.dropout(x)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:22.445106Z","iopub.status.busy":"2020-12-24T15:37:22.44418Z","iopub.status.idle":"2020-12-24T15:37:22.469271Z","shell.execute_reply":"2020-12-24T15:37:22.470143Z"},"papermill":{"duration":0.09352,"end_time":"2020-12-24T15:37:22.47032","exception":false,"start_time":"2020-12-24T15:37:22.3768","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"class FFN(nn.Module):\n    def __init__(self, state_size = 200, forward_expansion = 1, bn_size = MAX_SEQ - 1, dropout=0.2):\n        super(FFN, self).__init__()\n        self.state_size = state_size\n        \n        self.lr1 = nn.Linear(state_size, forward_expansion * state_size)\n        self.relu = nn.ReLU()\n        self.bn = nn.BatchNorm1d(bn_size)\n        self.lr2 = nn.Linear(forward_expansion * state_size, state_size)\n        self.dropout = nn.Dropout(dropout)\n        \n    def forward(self, x):\n        x = self.relu(self.lr1(x))\n        x = self.bn(x)\n        x = self.lr2(x)\n        return self.dropout(x)\n    \nclass FFN0(nn.Module):\n    def __init__(self, state_size = 200, forward_expansion = 1, bn_size = MAX_SEQ - 1, dropout=0.2):\n        super(FFN0, self).__init__()\n        self.state_size = state_size\n\n        self.lr1 = nn.Linear(state_size, forward_expansion * state_size)\n        self.relu = nn.ReLU()\n        self.lr2 = nn.Linear(forward_expansion * state_size, state_size)\n        self.layer_normal = nn.LayerNorm(state_size) \n        self.dropout = nn.Dropout(0.2)\n    \n    def forward(self, x):\n        x = self.lr1(x)\n        x = self.relu(x)\n        x = self.lr2(x)\n        x=self.layer_normal(x)\n        return self.dropout(x)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:22.604366Z","iopub.status.busy":"2020-12-24T15:37:22.602273Z","iopub.status.idle":"2020-12-24T15:37:22.613657Z","shell.execute_reply":"2020-12-24T15:37:22.614502Z"},"papermill":{"duration":0.083322,"end_time":"2020-12-24T15:37:22.614715","exception":false,"start_time":"2020-12-24T15:37:22.531393","status":"completed"},"tags":[],"trusted":true},"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\nfuture_mask(5)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:22.739539Z","iopub.status.busy":"2020-12-24T15:37:22.738664Z","iopub.status.idle":"2020-12-24T15:37:22.769956Z","shell.execute_reply":"2020-12-24T15:37:22.77098Z"},"papermill":{"duration":0.09698,"end_time":"2020-12-24T15:37:22.77116","exception":false,"start_time":"2020-12-24T15:37:22.67418","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"class TransformerBlock(nn.Module):\n    def __init__(self, embed_dim, heads = 8, dropout = DROPOUT, forward_expansion = 1):\n        super(TransformerBlock, self).__init__()\n        self.multi_att = nn.MultiheadAttention(embed_dim=embed_dim, num_heads=heads, dropout=dropout)\n        self.dropout = nn.Dropout(dropout)\n        self.layer_normal = nn.LayerNorm(embed_dim)\n        self.ffn = FFN(embed_dim, forward_expansion = forward_expansion, dropout=dropout)\n        self.ffn0  = FFN0(embed_dim, forward_expansion = forward_expansion, dropout=dropout)\n        self.layer_normal_2 = nn.LayerNorm(embed_dim)\n\n    def forward(self, value, key, query, att_mask):\n        att_output, att_weight = self.multi_att(value, key, query, attn_mask=att_mask)\n        att_output = self.dropout(self.layer_normal(att_output + value))\n        att_output = att_output.permute(1, 0, 2) # att_output: [s_len, bs, embed] => [bs, s_len, embed]\n        x = self.ffn(att_output)\n        x1 = self.ffn0(att_output)\n        x = self.dropout(self.layer_normal_2(x + x1 + att_output))\n        return x.squeeze(-1), att_weight\n    \nclass Encoder(nn.Module):\n    def __init__(self, n_skill, max_seq=100, embed_dim=128, dropout = DROPOUT, forward_expansion = 1, num_layers=1, heads = 8):\n        super(Encoder, self).__init__()\n        self.n_skill, self.embed_dim = n_skill, embed_dim\n        self.embedding = nn.Embedding(2 * n_skill + 1, embed_dim)\n        self.pos_embedding = nn.Embedding(max_seq - 1, embed_dim)\n        self.e_embedding = nn.Embedding(n_skill+1, embed_dim)\n        self.layers = nn.ModuleList([TransformerBlock(embed_dim, forward_expansion = forward_expansion) for _ in range(num_layers)])\n        self.dropout = nn.Dropout(dropout)\n        \n    def forward(self, x, question_ids):\n        device = x.device\n        x = self.embedding(x)\n        pos_id = torch.arange(x.size(1)).unsqueeze(0).to(device)\n        pos_x = self.pos_embedding(pos_id)\n        x = self.dropout(x + pos_x)\n        x = x.permute(1, 0, 2) # x: [bs, s_len, embed] => [s_len, bs, embed]\n        e = self.e_embedding(question_ids)\n        e = e.permute(1, 0, 2)\n        for layer in self.layers:\n            att_mask = future_mask(e.size(0)).to(device)\n            x, att_weight = layer(e, x, x, att_mask=att_mask)\n            x = x.permute(1, 0, 2)\n        x = x.permute(1, 0, 2)\n        return x, att_weight\n\nclass SAKTModel(nn.Module):\n    def __init__(self, n_skill, max_seq=100, embed_dim=128, dropout = DROPOUT, forward_expansion = 1, enc_layers=1, heads = 8):\n        super(SAKTModel, self).__init__()\n        self.encoder = Encoder(n_skill, max_seq, embed_dim, dropout, forward_expansion, num_layers=enc_layers)\n        self.pred = nn.Linear(embed_dim, 1)\n        \n    def forward(self, x, question_ids):\n        x, att_weight = self.encoder(x, question_ids)\n        x = self.pred(x)\n        return x.squeeze(-1), att_weight","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:23.379505Z","iopub.status.busy":"2020-12-24T15:37:23.378363Z","iopub.status.idle":"2020-12-24T15:37:23.385858Z","shell.execute_reply":"2020-12-24T15:37:23.387021Z"},"papermill":{"duration":0.556963,"end_time":"2020-12-24T15:37:23.387209","exception":false,"start_time":"2020-12-24T15:37:22.830246","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:23.477599Z","iopub.status.busy":"2020-12-24T15:37:23.476896Z","iopub.status.idle":"2020-12-24T15:37:23.544605Z","shell.execute_reply":"2020-12-24T15:37:23.545406Z"},"papermill":{"duration":0.111857,"end_time":"2020-12-24T15:37:23.545556","exception":false,"start_time":"2020-12-24T15:37:23.433699","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"# Main changes are possibility of forward expansion and stacking of encoding layers\ndef create_model():\n    return SAKTModel(n_skill, max_seq=MAX_SEQ, embed_dim=EMBED_SIZE, forward_expansion=1, enc_layers=1, heads=4, dropout=0.1)\n\nmodel = create_model()\nmodel","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:23.627762Z","iopub.status.busy":"2020-12-24T15:37:23.627184Z","iopub.status.idle":"2020-12-24T15:37:24.309097Z","shell.execute_reply":"2020-12-24T15:37:24.308483Z"},"papermill":{"duration":0.724855,"end_time":"2020-12-24T15:37:24.309213","exception":false,"start_time":"2020-12-24T15:37:23.584358","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"model(sample_batch[0], sample_batch[1])[0]","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.042188,"end_time":"2020-12-24T15:37:24.392921","exception":false,"start_time":"2020-12-24T15:37:24.350733","status":"completed"},"tags":[]},"cell_type":"markdown","source":"### Training"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:24.477815Z","iopub.status.busy":"2020-12-24T15:37:24.477209Z","iopub.status.idle":"2020-12-24T15:37:24.481376Z","shell.execute_reply":"2020-12-24T15:37:24.480751Z"},"papermill":{"duration":0.047738,"end_time":"2020-12-24T15:37:24.481495","exception":false,"start_time":"2020-12-24T15:37:24.433757","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"MODEL_PATH = '/kaggle/working/sakt_model.pt'","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:24.585567Z","iopub.status.busy":"2020-12-24T15:37:24.584675Z","iopub.status.idle":"2020-12-24T15:37:24.587668Z","shell.execute_reply":"2020-12-24T15:37:24.588186Z"},"papermill":{"duration":0.066412,"end_time":"2020-12-24T15:37:24.588313","exception":false,"start_time":"2020-12-24T15:37:24.521901","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"def load_from_item(item):\n    x = item[0].to(device).long()\n    target_id = item[1].to(device).long()\n    label = item[2].to(device).float()\n    target_mask = (target_id != 0)\n    return x, target_id, label, target_mask\n\ndef 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\ndef 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        x, target_id, label, target_mask = load_from_item(item)\n        \n        optim.zero_grad()\n        output, _ = model(x, target_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        loss.backward()\n        optim.step()\n        scheduler.step()\n        \n        tbar.set_description('loss - {:.4f}'.format(loss))\n\ndef 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, target_id, label, target_mask = load_from_item(item)\n\n        with torch.no_grad():\n            output, atten_weight = model(x, target_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        \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 = roc_auc_score(labels, outs)\n    loss = np.average(train_loss)\n\n    return loss, acc, auc\n","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T15:37:24.679762Z","iopub.status.busy":"2020-12-24T15:37:24.67886Z","iopub.status.idle":"2020-12-24T15:37:24.681812Z","shell.execute_reply":"2020-12-24T15:37:24.681321Z"},"papermill":{"duration":0.052799,"end_time":"2020-12-24T15:37:24.68191","exception":false,"start_time":"2020-12-24T15:37:24.629111","status":"completed"},"tags":[],"trusted":true},"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    \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} acc - {avl_acc:.3f} 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":{"trusted":true},"cell_type":"code","source":"# bbbbbb","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"LR = 2e-3\nEPOCHS = 11\nprint(f'{LR}, {EPOCHS}')\n\ndo_train()","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T16:27:33.445171Z","iopub.status.busy":"2020-12-24T16:27:33.444251Z","iopub.status.idle":"2020-12-24T16:42:58.826186Z","shell.execute_reply":"2020-12-24T16:42:58.826856Z"},"papermill":{"duration":925.452792,"end_time":"2020-12-24T16:42:58.827044","exception":false,"start_time":"2020-12-24T16:27:33.374252","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"LR = 2e-4\nEPOCHS = 3\n\nprint(f'{LR}, {EPOCHS}')\n\ndo_train()","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.07017,"end_time":"2020-12-24T16:42:58.969132","exception":false,"start_time":"2020-12-24T16:42:58.898962","status":"completed"},"tags":[]},"cell_type":"markdown","source":"### Predict"},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T16:42:59.118417Z","iopub.status.busy":"2020-12-24T16:42:59.117416Z","iopub.status.idle":"2020-12-24T16:42:59.205726Z","shell.execute_reply":"2020-12-24T16:42:59.206289Z"},"papermill":{"duration":0.165814,"end_time":"2020-12-24T16:42:59.206443","exception":false,"start_time":"2020-12-24T16:42:59.040629","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"model = create_model()\nmodel.load_state_dict(torch.load(MODEL_PATH))\nmodel.to(device)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T16:42:59.365849Z","iopub.status.busy":"2020-12-24T16:42:59.36408Z","iopub.status.idle":"2020-12-24T16:42:59.3666Z","shell.execute_reply":"2020-12-24T16:42:59.36718Z"},"papermill":{"duration":0.088891,"end_time":"2020-12-24T16:42:59.367336","exception":false,"start_time":"2020-12-24T16:42:59.278445","status":"completed"},"tags":[],"trusted":true},"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":{"execution":{"iopub.execute_input":"2020-12-24T16:42:59.512756Z","iopub.status.busy":"2020-12-24T16:42:59.512101Z","iopub.status.idle":"2020-12-24T16:42:59.544357Z","shell.execute_reply":"2020-12-24T16:42:59.544846Z"},"papermill":{"duration":0.1065,"end_time":"2020-12-24T16:42:59.544972","exception":false,"start_time":"2020-12-24T16:42:59.438472","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"import riiideducation\n\nenv = riiideducation.make_env()\niter_test = env.iter_test()","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2020-12-24T16:42:59.710491Z","iopub.status.busy":"2020-12-24T16:42:59.70252Z","iopub.status.idle":"2020-12-24T16:43:00.528934Z","shell.execute_reply":"2020-12-24T16:43:00.528448Z"},"papermill":{"duration":0.912138,"end_time":"2020-12-24T16:43:00.52908","exception":false,"start_time":"2020-12-24T16:42:59.616942","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"import psutil\n\nmodel.eval()\n\nprev_test_df = None\n\nfor (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":{"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}