{"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_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Preliminaries","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nimport math\nimport pickle\nimport psutil\nimport random\n\nimport json\nimport numpy as np\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport pandas as pd\n\nimport riiideducation","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-08-12T18:38:43.718075Z","iopub.execute_input":"2021-08-12T18:38:43.718572Z","iopub.status.idle":"2021-08-12T18:38:44.146405Z","shell.execute_reply.started":"2021-08-12T18:38:43.718480Z","shell.execute_reply":"2021-08-12T18:38:44.145564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed = 0\nrandom.seed(seed)\ntorch.random.manual_seed(seed)\n\nn_workers = os.cpu_count()\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\ncfg_path = '/kaggle/input/riiid-mydata/cfg.json'\ntrain_path = '/kaggle/input/riiid-mydata/train.pkl'\ntag_path = '/kaggle/input/riiid-mydata/tags.csv'\nstates_path = '/kaggle/input/riiid-mydata/states.pickle'\nmodel_path = '/kaggle/input/riiid-mydata/aPFA_08-12_18-18.pt'\n\nB = 512\nMAX_LEN = 128\nSEQ_LEN = 128\nMAX_LAG = 30 * 7 * 24 * 60\nN_LAYERS = 0\nN_HEADS = 1\nD_MODEL = 256\nIS_LITE = True","metadata":{"execution":{"iopub.status.busy":"2021-08-12T18:38:44.153175Z","iopub.execute_input":"2021-08-12T18:38:44.153539Z","iopub.status.idle":"2021-08-12T18:38:44.180964Z","shell.execute_reply.started":"2021-08-12T18:38:44.153494Z","shell.execute_reply":"2021-08-12T18:38:44.180181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## aPFA Model","metadata":{}},{"cell_type":"code","source":"class FFN(nn.Module): \n    def __init__(self, d_model, dropout=0.0): \n        super().__init__()\n        self.lr1 = nn.Linear(d_model, d_model)\n        self.relu = nn.ReLU()\n        self.lr2 = nn.Linear(d_model, d_model)\n        self.dropout = nn.Dropout(dropout)\n\n    def forward(self, x): \n        x = self.lr1(x)\n        x = self.relu(x)\n        x = self.dropout(x)\n        x = self.lr2(x)\n        return x\n\n\nclass AIKTMultiheadAttention(nn.Module): \n    def __init__(self, d_model, n_heads=8, d_qkv=None, use_proj=True, dropout=0.1): \n        super().__init__()\n        if d_qkv is None or not use_proj: \n            assert d_model % n_heads == 0\n            d_qkv = d_model // n_heads\n        d_inner = d_qkv * n_heads\n        self.scale = d_model ** (-0.5)\n        if use_proj: \n            self.Q = nn.Linear(d_model, d_inner)\n            self.K = nn.Linear(d_model, d_inner)\n            self.V = nn.Linear(d_model, d_inner)\n            self.out_proj = nn.Linear(d_inner, d_model)\n        else: \n            self.Q = self.K = self.V = self.out_proj = lambda x: x\n        # (len, B, d_inner) -> (len, B, n_heads, d_qkv)\n        self.reshape_for_attn = lambda x: x.reshape(*x.shape[:-1], n_heads, d_qkv)\n        # (len, B, n_heads, d_qkv) -> (len, B, d_inner)\n        self.recover_from_attn = lambda x: x.reshape(*x.shape[:-2], d_inner)\n        self.dropout = nn.Dropout(dropout)\n\n    def forward(self, query, key, value, attn_mask=None, rel_pos_embd=None): \n        # in-sample projection\n        query, key, value = (\n            self.reshape_for_attn(query), \n            self.reshape_for_attn(key), \n            self.reshape_for_attn(value)\n        )\n        # self-attention mechanism\n        attn_scores = torch.einsum('ibnd,jbnd->ijbn', (query, key)) * self.scale\n        if rel_pos_embd is not None: \n            attn_scores += rel_pos_embd\n        if attn_mask is not None: \n            assert attn_mask.dtype == torch.bool, 'Only bool type is supported for masks.'\n            assert attn_mask.ndim == 2, 'Only 2D attention mask is supported'\n            assert attn_mask.shape == attn_scores.shape[:2], 'Incorrect mask shape: {}. Expect: {}'.format(attn_mask.shape, attn_scores.shape[:2])\n            mask = torch.zeros_like(attn_mask, dtype=torch.float)\n            mask.masked_fill_(attn_mask, float('-inf'))\n            mask = mask.view(*mask.shape, 1, 1)\n            attn_scores += mask\n        attn_weights = torch.softmax(attn_scores, dim=1)\n        attn_weights = self.dropout(attn_weights)\n        out = torch.einsum('ijbn,jbnd->ibnd', (attn_weights, value))\n        out = self.recover_from_attn(out)\n        # output layer\n        out = self.out_proj(out)\n        attn_weights = attn_weights.mean(-1).permute(-1, 0, 1)\n        return out, attn_weights\n\n\nclass InducedMultiheadAttention(nn.Module): \n    def __init__(self, d_model, n_heads, d_qkv=None, use_proj=False, dropout=0.1): \n        super().__init__()\n        if d_qkv is None or not use_proj: \n            assert d_model % n_heads == 0\n            d_qkv = d_model // n_heads\n        d_inner = d_qkv * n_heads\n        self.scale = d_model ** (-0.5)\n        self.Q = nn.Parameter(torch.randn(n_heads, d_qkv))\n        if use_proj: \n            self.K = nn.Linear(d_model, d_inner)\n            self.V = nn.Linear(d_model, d_inner)\n            self.out_proj = nn.Linear(d_inner, d_model)\n        else: \n            self.K = self.V = self.out_proj = lambda x: x\n        # (len, B, d_inner) -> (len, B, n_heads, d_qkv)\n        self.reshape_for_attn = lambda x: x.reshape(*x.shape[:-1], n_heads, d_qkv)\n        # (len, B, n_heads, d_qkv) -> (len, B, d_inner)\n        self.recover_from_attn = lambda x: x.reshape(*x.shape[:-2], d_inner)\n        self.dropout = nn.Dropout(dropout)\n\n    def forward(self, key, value, attn_mask=None, rel_pos_embd=None): \n        # in-sample projection\n        key, value = (\n            self.reshape_for_attn(self.K(key)), \n            self.reshape_for_attn(self.V(value))\n        )\n        # induced attention mechanism\n        attn_scores = torch.einsum('nd,jbnd->jbn', (self.Q, key)) * self.scale\n        attn_scores = attn_scores.unsqueeze(0).repeat(attn_scores.shape[0], 1, 1, 1)\n        if rel_pos_embd is not None: \n            attn_scores += rel_pos_embd\n        if attn_mask is not None: \n            assert attn_mask.dtype == torch.bool, 'Only bool type is supported for masks.'\n            assert attn_mask.ndim == 2, 'Only 2D attention mask is supported'\n            assert attn_mask.shape == attn_scores.shape[:2], 'Incorrect mask shape: {}. Expect: {}'.format(attn_mask.shape, attn_scores.shape[:2])\n            mask = torch.zeros_like(attn_mask, dtype=torch.float)\n            mask.masked_fill_(attn_mask, float('-inf'))\n            mask = mask.view(*mask.shape, 1, 1)\n            attn_scores += mask\n        attn_weights = torch.softmax(attn_scores, dim=1)\n        attn_weights = self.dropout(attn_weights)\n        out = torch.einsum('ijbn,jbnd->ibnd', (attn_weights, value))\n        out = self.recover_from_attn(out)\n        # output layer\n        out = self.out_proj(out)\n        attn_weights = attn_weights.mean(-1).permute(-1, 0, 1)\n        return out, attn_weights\n\n\nclass APFAModel(nn.Module): \n    def __init__(self, n_exercises, max_lag, n_layers, d_model, n_heads=8, is_lite=False, dropout=0.1): \n        super().__init__()\n        self.exercise_embd = nn.Embedding(n_exercises, d_model)\n        self.correct_embd = nn.Embedding(2, d_model)\n        self.max_lag = max_lag\n        self.n_lag_buckets = 2 * math.ceil(math.log(max_lag))\n        self.rel_lag_embd = nn.Embedding(self.n_lag_buckets, n_heads)\n\n        self.is_lite = is_lite\n        self.enc = nn.ModuleList([\n            nn.ModuleList([\n                AIKTMultiheadAttention(d_model, n_heads, use_proj=not is_lite, dropout=dropout), \n                None if is_lite else nn.LayerNorm(d_model), \n                None if is_lite else FFN(d_model, dropout=dropout), \n                None if is_lite else nn.LayerNorm(d_model)\n            ]) for _ in range(n_layers)\n        ])\n        self.init_mem = nn.Parameter(torch.randn(d_model))\n        self.dec = InducedMultiheadAttention(d_model, n_heads, use_proj=not is_lite, dropout=dropout)\n        self.ln1 = None if is_lite else nn.LayerNorm(d_model)\n        self.ffn = None if is_lite else FFN(d_model, dropout=dropout)\n        self.ln2 = None if is_lite else nn.LayerNorm(d_model)\n\n        self.predict = nn.Linear(d_model, n_exercises)\n        self.predict.weight = self.exercise_embd.weight # tie weights\n        self.dropout = nn.Dropout(dropout)\n\n    def lag_to_bucket(self, lag_time): \n        n_exact = self.n_lag_buckets // 2\n        acc_lag_time = torch.cumsum(lag_time, dim=-1).unsqueeze(-1)\n        rel_lag_time = torch.clamp(\n            acc_lag_time - acc_lag_time.transpose(-1, -2), min=0, max=self.max_lag\n        )\n        rel_lag_time = torch.cat(\n            [rel_lag_time[:, :, :1], rel_lag_time[:, :, :-1]], dim=-1\n        ) # right shift by 1 along the dimension of k_len\n        buckets_for_long_lag = n_exact - 1 + torch.ceil(\n            torch.log(rel_lag_time / n_exact) / math.log(self.max_lag / n_exact) * (self.n_lag_buckets - n_exact)\n        )\n        buckets = torch.where(rel_lag_time < n_exact, rel_lag_time, buckets_for_long_lag.long())\n        return buckets.permute(1, 2, 0) # (q_len, k_len, B)\n\n    def forward(self, e, c, lt, mem=None, attn_mask=None): \n        # encoder\n        src = self.exercise_embd(e)\n        src = src.transpose(0, 1) # (B, L, d) -> (L, B, d)\n        enc_attn_weights = []\n        for self_attn, ln1, ffn, ln2 in self.enc: \n            out, attn_weights = self_attn(src, src, src, attn_mask=attn_mask)\n            if self.is_lite: \n                src = self.dropout(out)\n            else: \n                src = ln1(src + self.dropout(out))\n                out = ffn(src)\n                src = ln2(src + self.dropout(out))\n            enc_attn_weights.append(attn_weights)\n        src = src.transpose(0, 1) # (L, B, d) -> (B, L, d)\n        # decoder\n        src = src + self.correct_embd(c)\n        if mem is None: \n            mem = self.init_mem.view(1, 1, -1).repeat(src.shape[0], 1, 1) # (B, 1, d)\n        src = torch.cat([mem, src[:, :-1, :]], dim=1)\n        src = src.transpose(0, 1) # (B, L, d) -> (L, B, d)\n        rel_pos_embd = self.rel_lag_embd(self.lag_to_bucket(lt))\n        out, dec_attn_weights = self.dec(src, src, attn_mask=attn_mask, rel_pos_embd=rel_pos_embd)\n        tgt = self.dropout(out)\n        if not self.is_lite: \n            tgt = self.ln1(tgt)\n            out = self.ffn(tgt)\n            tgt = self.ln2(tgt + self.dropout(out))\n        tgt = tgt.transpose(0, 1) # (L, B, d) -> (B, L, d)\n        mem = tgt[:, -1:, :]\n        # prediction\n        weight = self.predict.weight[e]\n        bias = self.predict.bias[e]\n        logits = torch.einsum('bid,bid->bi', (tgt, weight)) + bias\n        return logits, mem, (enc_attn_weights, dec_attn_weights)","metadata":{"execution":{"iopub.status.busy":"2021-08-12T18:38:44.182442Z","iopub.execute_input":"2021-08-12T18:38:44.182801Z","iopub.status.idle":"2021-08-12T18:38:44.226881Z","shell.execute_reply.started":"2021-08-12T18:38:44.182745Z","shell.execute_reply":"2021-08-12T18:38:44.226041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Testing Phase","metadata":{}},{"cell_type":"code","source":"class KaggleOnlineDataset(Dataset): \n    def __init__(self, train_path, tag_path, states_path, n_exercises, cols, max_len): \n        super().__init__()\n        self.df = pd.read_pickle(train_path)\n        self.test_df = None\n        tag_df = pd.read_csv(tag_path, usecols=['exercise_id', 'bundle_id', 'part', 'correct_rate', 'frequency'])\n        assert np.all(tag_df['exercise_id'].values == np.arange(n_exercises))\n        self.parts, self.correct_rate, self.frequency = tag_df[['part', 'correct_rate', 'frequency']].values.T\n        self.lag_info = pickle.load(open(states_path, 'rb'))\n        self.n_exercises = n_exercises\n        self.cols = cols\n        self.max_len = max_len\n\n    def __len__(self): \n        assert self.test_df is not None, 'Please call update() first'\n        return len(self.test_df)\n\n    def __getitem__(self, idx): \n        new_observation = self.test_df.iloc[idx]\n        # 'correct' is set to 0 temporarily\n        user_id = new_observation['user_id']\n        new_data = {col: np.array([new_observation.get(col, 0)]) for col in self.cols}\n        # retrieve old observations\n        if user_id in self.df.index: \n            old_items = self.df[user_id]\n            old_len = min(len(old_items[0]), self.max_len - 1)\n            data = {key: np.append(old_item[-old_len:], new_data[key]) for key, old_item in zip(self.cols, old_items)}\n        else: \n            old_len = 0\n            data = new_data\n        seq_len = old_len + 1\n        # retrieve addtional features\n        data['part'] = self.parts[data['exercise_id']]\n        data['correct_rate'] = self.correct_rate[data['exercise_id']]\n        # pad to max_len and set dtype\n        dtype_map = {key: int for key in self.cols + ['part']}\n        dtype_map['correct_rate'] = float\n        data = KaggleOnlineDataset._postpad_and_asdtype(data, self.max_len - seq_len, dtype_map)\n        data['valid_len'] = np.array([seq_len], dtype=int)\n        return data\n    \n    @staticmethod\n    def _postpad_and_asdtype(data, pad, dtype_map): \n        return {\n            key: np.pad(item, [[0, pad]]).astype(dtype_map[key]) for key, item in data.items()\n        }\n    \n    def update(self, test_df): \n        if self.test_df is not None and psutil.virtual_memory().percent < 90: \n            # update df according to previous labels\n            prev_df = self.test_df\n            prev_df['correct'] = np.array(eval(test_df.iloc[0]['prior_group_answers_correct']))[self.was_exercise]\n            user_df = prev_df.groupby('user_id').apply(lambda udf: tuple(udf[col].values for col in self.cols))\n            for user_id, new_items in user_df.iteritems(): \n                if user_id in self.df.index: \n                    self.df[user_id] = tuple(map(\n                        lambda old_item, new_item: np.append(old_item, new_item)[-min(self.max_len, len(old_item) + 1):], \n                        self.df[user_id], \n                        new_items\n                    )) # truncate at max_len to prevent OOM\n                else: \n                    self.df[user_id] = tuple(new_item for new_item in new_items) # create a new row\n            # update correct rate\n            # self._update_correct_rate(prev_df)\n        # process test_df\n        is_exercise = (test_df['content_type_id'] == 0)\n        test_df = test_df[is_exercise]\n        test_df = test_df.rename(columns={'content_id': 'exercise_id', 'prior_question_elapsed_time': 'prior_elapsed'})\n        # compute lag and convert ms -> min\n        test_df['prior_elapsed'] = test_df['prior_elapsed'].fillna(0).astype(int)\n        lag = self._compute_new_lag(test_df)\n        test_df['lag'] = np.where(\n            np.logical_and(0 < lag, lag < 60 * 1000), 1, np.round(lag / (1000 * 60))\n        ).astype(int)\n        # as for prior_elapsed, convert ms -> s\n        prior_elapsed = test_df['prior_elapsed'].values\n        test_df['prior_elapsed'] = np.where(\n            np.logical_and(0 < prior_elapsed, prior_elapsed < 1000), 1, np.round(prior_elapsed / 1000)\n        ).astype(int)\n        test_df.reset_index(drop=True, inplace=True)\n        # save relevent information\n        self.test_df = test_df\n        self.was_exercise = is_exercise.values\n        return test_df\n        \n    def _update_correct_rate(self, prev_df): \n        exercise_df = prev_df.groupby('exercise_id').aggregate({'correct': [sum, len]})['correct']\n        n_correct = np.arange(self.n_exercises)\n        np.put(n_correct, exercise_df.index, exercise_df['sum'].values)\n        n_correct = n_correct + np.round(self.correct_rate * self.frequency)\n        more_frequency = np.arange(self.n_exercises)\n        np.put(more_frequency, exercise_df.index, exercise_df['len'].values)\n        self.frequency += more_frequency\n        correct_rate = n_correct / self.frequency\n        self.correct_rate = np.where(np.isfinite(correct_rate), correct_rate, 0.5)\n        \n    def _compute_new_lag(self, df): \n        last_states, exercise_id_to_bundle, bundle_id_to_size = self.lag_info\n        # compute_lag from the original implementation\n        lag = np.zeros(len(df))\n        for i, (user_id, curr_timestamp, curr_exercise_id, prior_elapsed) in enumerate(\n            df[['user_id', 'timestamp', 'exercise_id', 'prior_elapsed']].values\n        ): \n            curr_bundle_id = exercise_id_to_bundle[curr_exercise_id]\n            last_state = last_states.get(user_id, None)\n            if last_state is None: \n                last_states[user_id] = (curr_timestamp, curr_bundle_id)\n                lag[i] = 0\n            else: \n                last_timestamp, last_bundle_id = last_state\n                if curr_bundle_id == last_bundle_id: \n                    # same bundle, do not update last_states\n                    lag[i] = 0\n                else: \n                    last_states[user_id] = (curr_timestamp, curr_bundle_id)\n                    elapsed_offset = bundle_id_to_size[last_bundle_id] * prior_elapsed\n                    lag[i] = curr_timestamp - last_timestamp - elapsed_offset\n        lag = np.clip(lag, a_min=0, a_max=None)\n        self.lag_info = (last_states, exercise_id_to_bundle, bundle_id_to_size)\n        return lag","metadata":{"execution":{"iopub.status.busy":"2021-08-12T18:38:44.228080Z","iopub.execute_input":"2021-08-12T18:38:44.228468Z","iopub.status.idle":"2021-08-12T18:38:44.258390Z","shell.execute_reply.started":"2021-08-12T18:38:44.228432Z","shell.execute_reply":"2021-08-12T18:38:44.257398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def truncate_and_prepare_masks(items, valid_len, need_pad_mask=True, need_attn_mask=True): \n    max_len = valid_len.max()\n    device = max_len.device\n    # truncate at the max_len for each sample\n    out = [None if item is None else item[:, :max_len] for item in items]\n    # pad to the same length for batch-ification\n    pad_mask = torch.arange(max_len, device=device) >= valid_len if need_pad_mask else None\n    # assume q_len = k_len in attention\n    attn_mask = torch.triu(torch.ones(max_len, max_len), diagonal=1).to(device, torch.bool) if need_attn_mask else None\n    return out, pad_mask, attn_mask","metadata":{"execution":{"iopub.status.busy":"2021-08-12T18:38:44.259663Z","iopub.execute_input":"2021-08-12T18:38:44.260155Z","iopub.status.idle":"2021-08-12T18:38:44.270925Z","shell.execute_reply.started":"2021-08-12T18:38:44.260118Z","shell.execute_reply":"2021-08-12T18:38:44.270132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg = json.load(open(cfg_path, 'r'))\nmodel = APFAModel(cfg['n_exercises'], MAX_LAG, N_LAYERS, D_MODEL, n_heads=N_HEADS, is_lite=IS_LITE).to(device)\nmodel.load_state_dict(torch.load(model_path))\nmodel.eval()\ntestset = KaggleOnlineDataset(train_path, tag_path, states_path, cfg['n_exercises'], cfg['cols'], MAX_LEN)\n\nenv = riiideducation.make_env()\niter_test = env.iter_test()\nfor test_df, _ in iter_test: \n    test_df = testset.update(test_df)\n    testloader = DataLoader(testset, batch_size=B, shuffle=False, num_workers=n_workers, drop_last=False)\n    outs = np.array([], dtype='float32')\n    for data in testloader: \n        valid_len = data['valid_len'].to(device, torch.long)\n        (*inputs, labels), _, attn_mask = truncate_and_prepare_masks(\n            [\n                data['exercise_id'].to(device, torch.long), \n                data['correct'].to(device, torch.long), \n                data['lag'].to(device, torch.long), \n                data['correct'].to(device, torch.float), \n            ], \n            valid_len, \n            need_pad_mask=False\n        )\n        \n        max_len = valid_len.max().item()\n        mem = None\n        out = torch.empty_like(labels)\n        for start in range(0, max_len, SEQ_LEN): \n            end = min(start + SEQ_LEN, max_len)\n            inputs_i = tuple(item[:, start:end] for item in inputs)\n            attn_mask_i = attn_mask[start:end, start:end]\n            out_i, mem, _ = model(*inputs_i, mem=mem, attn_mask=attn_mask_i)\n            out[:, start:end] = out_i.detach()\n        out = torch.gather(out, 1, valid_len - 1).squeeze(-1)\n        outs = np.append(outs, torch.sigmoid(out).detach().cpu().numpy())\n    test_df['answered_correctly'] = outs\n    env.predict(test_df[['row_id', 'answered_correctly']])","metadata":{"execution":{"iopub.status.busy":"2021-08-12T18:38:44.272205Z","iopub.execute_input":"2021-08-12T18:38:44.272564Z","iopub.status.idle":"2021-08-12T18:38:54.124632Z","shell.execute_reply.started":"2021-08-12T18:38:44.272529Z","shell.execute_reply":"2021-08-12T18:38:54.123621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}