{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":21651,"databundleVersionId":1595136,"sourceType":"competition"}],"dockerImageVersionId":30034,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"###### * Base Source: https://www.kaggle.com/wangsg/a-self-attentive-model-for-knowledge-tracing\n* My First Work: https://www.kaggle.com/leadbest/sakt-self-attentive-knowledge-tracing-submitter\n\n1. Version 1: State Updates -> LB 0.765\n2. Version 3: Random Selection of User Interactions -> LB 0.768\n3. Version 6: Small Optimization -> LB 0.771?","metadata":{}},{"cell_type":"markdown","source":"## Loss: 0.6044 - Acc: 0.6662 - AUC: 0.7285","metadata":{}},{"cell_type":"code","source":"import os\nos.environ['CUDA_LAUNCH_BLOCKING'] = '1'\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:53:18.660276Z","iopub.execute_input":"2025-08-03T19:53:18.660578Z","iopub.status.idle":"2025-08-03T19:53:18.664414Z","shell.execute_reply.started":"2025-08-03T19:53:18.660553Z","shell.execute_reply":"2025-08-03T19:53:18.663466Z"}},"outputs":[],"execution_count":null},{"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","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:53:18.668551Z","iopub.execute_input":"2025-08-03T19:53:18.668793Z","iopub.status.idle":"2025-08-03T19:53:18.686091Z","shell.execute_reply.started":"2025-08-03T19:53:18.668769Z","shell.execute_reply":"2025-08-03T19:53:18.685449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\nimport random\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","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:53:18.687477Z","iopub.execute_input":"2025-08-03T19:53:18.687678Z","iopub.status.idle":"2025-08-03T19:53:19.492949Z","shell.execute_reply.started":"2025-08-03T19:53:18.687658Z","shell.execute_reply":"2025-08-03T19:53:19.492227Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MAX_SEQ = 160","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:53:19.494359Z","iopub.execute_input":"2025-08-03T19:53:19.494592Z","iopub.status.idle":"2025-08-03T19:53:19.497784Z","shell.execute_reply.started":"2025-08-03T19:53:19.494570Z","shell.execute_reply":"2025-08-03T19:53:19.496892Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load data","metadata":{}},{"cell_type":"markdown","source":"## Interaction data","metadata":{}},{"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',\n        'prior_question_elapsed_time':'float32',\n        'prior_question_had_explanation':'int8'\n        }\n\ntrain_df = pd.read_csv('/kaggle/input/riiid-test-answer-prediction/train.csv', usecols=[1, 2, 3, 4,7,8,9], dtype=dtype)\ntrain_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:53:19.499778Z","iopub.execute_input":"2025-08-03T19:53:19.500139Z","iopub.status.idle":"2025-08-03T19:54:54.214406Z","shell.execute_reply.started":"2025-08-03T19:53:19.500098Z","shell.execute_reply":"2025-08-03T19:54:54.213499Z"}},"outputs":[],"execution_count":null},{"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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:54:54.215835Z","iopub.execute_input":"2025-08-03T19:54:54.216213Z","iopub.status.idle":"2025-08-03T19:55:19.008013Z","shell.execute_reply.started":"2025-08-03T19:54:54.216176Z","shell.execute_reply":"2025-08-03T19:55:19.007175Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Lag time","metadata":{}},{"cell_type":"code","source":"# train_df['lag_time'] = train_df.groupby('user_id')['timestamp'].diff()\n# train_df['lag_time'] = np.log1p(train_df['lag_time'].clip(lower=0))\n# train_df['prior_question_elapsed_time'] = np.log1p(train_df['prior_question_elapsed_time'].clip(lower=0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:55:19.009522Z","iopub.execute_input":"2025-08-03T19:55:19.009882Z","iopub.status.idle":"2025-08-03T19:55:19.013245Z","shell.execute_reply.started":"2025-08-03T19:55:19.009845Z","shell.execute_reply":"2025-08-03T19:55:19.012558Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Difficulty Calculation","metadata":{}},{"cell_type":"code","source":"# calculate the accuracy of each question\nquestion_difficulty = train_df.groupby('content_id')['answered_correctly'].mean()\ntrain_df['difficulty'] = train_df['content_id'].map(question_difficulty)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:55:19.014505Z","iopub.execute_input":"2025-08-03T19:55:19.014718Z","iopub.status.idle":"2025-08-03T19:55:22.947320Z","shell.execute_reply.started":"2025-08-03T19:55:19.014698Z","shell.execute_reply":"2025-08-03T19:55:22.946620Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Question Data","metadata":{}},{"cell_type":"code","source":"questions_df = pd.read_csv('/kaggle/input/riiid-test-answer-prediction/questions.csv', usecols=[0,3,4], dtype=dtype)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:55:22.948541Z","iopub.execute_input":"2025-08-03T19:55:22.948755Z","iopub.status.idle":"2025-08-03T19:55:22.967913Z","shell.execute_reply.started":"2025-08-03T19:55:22.948733Z","shell.execute_reply":"2025-08-03T19:55:22.967281Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tagsarray = set()\nfor tags in questions_df['tags'].dropna():\n    tagsarray.update([int(tag) for tag in tags.split()])\n\n# the tag range from 0-187, shift the tag by +1 as 0 will be used for padding","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:55:22.969140Z","iopub.execute_input":"2025-08-03T19:55:22.969492Z","iopub.status.idle":"2025-08-03T19:55:22.988435Z","shell.execute_reply.started":"2025-08-03T19:55:22.969457Z","shell.execute_reply":"2025-08-03T19:55:22.987788Z"},"_kg_hide-output":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"questions_df['tags'] = questions_df['tags'].fillna(\"pad\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:55:22.989631Z","iopub.execute_input":"2025-08-03T19:55:22.989954Z","iopub.status.idle":"2025-08-03T19:55:23.000462Z","shell.execute_reply.started":"2025-08-03T19:55:22.989921Z","shell.execute_reply":"2025-08-03T19:55:22.999727Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def encode_tags_shifted(tag_str):\n    if tag_str == \"pad\":\n        return [0] \n    return [int(t) + 1 for t in tag_str.split()]\nquestions_df['tag_ids'] = questions_df['tags'].apply(encode_tags_shifted)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:58:51.631347Z","iopub.execute_input":"2025-08-03T19:58:51.631685Z","iopub.status.idle":"2025-08-03T19:58:51.653946Z","shell.execute_reply.started":"2025-08-03T19:58:51.631659Z","shell.execute_reply":"2025-08-03T19:58:51.653290Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# shift the content_id and question_id by 1 so that we can use 0 for padding.\ntrain_df['content_id'] += 1\nquestions_df['question_id'] += 1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:58:53.849820Z","iopub.execute_input":"2025-08-03T19:58:53.850147Z","iopub.status.idle":"2025-08-03T19:58:54.065473Z","shell.execute_reply.started":"2025-08-03T19:58:53.850118Z","shell.execute_reply":"2025-08-03T19:58:54.064805Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MAX_TAGS = 4\nPAD_TAG = 0\n\ndef pad_tags(tag_list):\n    return tag_list[:MAX_TAGS] + [PAD_TAG] * (MAX_TAGS - len(tag_list))\nquestions_df['tag_ids_padded'] = questions_df['tag_ids'].apply(pad_tags)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:58:56.010643Z","iopub.execute_input":"2025-08-03T19:58:56.010924Z","iopub.status.idle":"2025-08-03T19:58:56.024887Z","shell.execute_reply.started":"2025-08-03T19:58:56.010900Z","shell.execute_reply":"2025-08-03T19:58:56.024109Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"qid_to_tags = dict(zip(questions_df['question_id'], questions_df['tag_ids_padded']))\ntrain_df['tags'] = train_df['content_id'].map(qid_to_tags)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:58:58.729693Z","iopub.execute_input":"2025-08-03T19:58:58.730002Z","iopub.status.idle":"2025-08-03T19:59:04.157556Z","shell.execute_reply.started":"2025-08-03T19:58:58.729953Z","shell.execute_reply":"2025-08-03T19:59:04.156865Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"qid_to_part = dict(zip(questions_df['question_id'], questions_df['part']))\ntrain_df['part'] = train_df['content_id'].map(qid_to_part)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:59:04.158987Z","iopub.execute_input":"2025-08-03T19:59:04.159233Z","iopub.status.idle":"2025-08-03T19:59:05.650238Z","shell.execute_reply.started":"2025-08-03T19:59:04.159208Z","shell.execute_reply":"2025-08-03T19:59:05.649550Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# train_df['lag_time'] = train_df['lag_time'].astype('float32') \n# train_df['prior_question_had_explanation'] = train_df['prior_question_had_explanation'].astype('int8')\n\ntrain_df['difficulty'] = train_df['difficulty'].astype('float32')\ntrain_df['part'] = train_df['part'].astype('int8')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:59:06.449712Z","iopub.execute_input":"2025-08-03T19:59:06.450006Z","iopub.status.idle":"2025-08-03T19:59:06.841756Z","shell.execute_reply.started":"2025-08-03T19:59:06.449981Z","shell.execute_reply":"2025-08-03T19:59:06.841140Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Fill Na","metadata":{}},{"cell_type":"code","source":"# elapsed_median = train_df['prior_question_elapsed_time'].median()\n# train_df['prior_question_elapsed_time'] = train_df['prior_question_elapsed_time'].fillna(elapsed_median).astype('float32')\n# train_df['lag_time'] = train_df['lag_time'].fillna(0).astype('float32')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:55:23.112099Z","iopub.status.idle":"2025-08-03T19:55:23.112648Z","shell.execute_reply":"2025-08-03T19:55:23.112362Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Preprocess","metadata":{}},{"cell_type":"code","source":"group=train_df.groupby('user_id').apply(lambda x: {\n    'content_id':x['content_id'].values,\n    'answered_correctly': x['answered_correctly'].values,\n    # 'prior_question_elapsed_time': x['prior_question_elapsed_time'].values,\n    # 'prior_question_had_explanation': x['prior_question_had_explanation'].values,\n    'cumulative_accuracy': x['answered_correctly'].shift().expanding().mean().fillna(0).values,\n    'part':x['part'].values,\n    'tags':x['tags'].tolist(),\n    'difficulty':x['difficulty'].values,\n    # 'lag_time':x['lag_time'].values\n})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T19:59:11.874746Z","iopub.execute_input":"2025-08-03T19:59:11.875037Z","iopub.status.idle":"2025-08-03T20:03:47.565195Z","shell.execute_reply.started":"2025-08-03T19:59:11.875011Z","shell.execute_reply":"2025-08-03T20:03:47.564427Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"group[115]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T20:05:11.395483Z","iopub.execute_input":"2025-08-03T20:05:11.395786Z","iopub.status.idle":"2025-08-03T20:05:11.417073Z","shell.execute_reply.started":"2025-08-03T20:05:11.395761Z","shell.execute_reply":"2025-08-03T20:05:11.416305Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del train_df\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T20:05:45.009412Z","iopub.execute_input":"2025-08-03T20:05:45.009741Z","iopub.status.idle":"2025-08-03T20:05:46.601308Z","shell.execute_reply.started":"2025-08-03T20:05:45.009709Z","shell.execute_reply":"2025-08-03T20:05:46.600517Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Define Dataset","metadata":{}},{"cell_type":"code","source":"import random\nrandom.seed(1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T20:05:48.609767Z","iopub.execute_input":"2025-08-03T20:05:48.610077Z","iopub.status.idle":"2025-08-03T20:05:48.613996Z","shell.execute_reply.started":"2025-08-03T20:05:48.610047Z","shell.execute_reply":"2025-08-03T20:05:48.613086Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MVAKTDataset(Dataset):\n    def __init__(self, group, max_seq=100, training=True):\n        self.group = group\n        self.max_seq = max_seq\n        self.training = training\n        self.user_ids = [uid for uid in group.index if len(group[uid]['content_id']) > 2]\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        sample = self.group[user_id]\n        seq_len = len(sample['content_id'])\n\n        # Random truncation if sequence is too long\n        if seq_len > self.max_seq:\n            start = random.randint(0, seq_len - self.max_seq) if self.training else seq_len - self.max_seq\n            for key in sample:\n                sample[key] = sample[key][start : start + self.max_seq]\n\n        seq_len = len(sample['content_id'])\n        pad_len = max(0, self.max_seq - seq_len)\n\n        def pad(seq, pad_value=0):\n            if isinstance(seq[0], list):  # tags\n                return [([pad_value] * len(seq[0]))] * pad_len + seq\n            return [pad_value] * pad_len + list(seq)\n\n        # Padding all fields\n        content_id = pad(sample['content_id'])\n        answered_correctly = pad(sample['answered_correctly'])\n        cumulative_accuracy = pad(sample['cumulative_accuracy'])\n        part = pad(sample['part'])\n        tags = pad(sample['tags'], pad_value=0)\n        difficulty = pad(sample['difficulty'])\n\n        # Transformer input (up to t-1)\n        content_id_input = content_id[:-1]\n        answered_correctly_input = answered_correctly[:-1]\n        cumulative_accuracy_input = cumulative_accuracy[:-1]\n\n        # Prediction target (at t)\n        content_id_target = content_id[1:]\n        part_target = part[1:]\n        tags_target = tags[1:]\n        difficulty_target = difficulty[1:]\n        cumulative_accuracy_target = cumulative_accuracy[1:]\n        label = answered_correctly[1:]\n        label_mask = [1 if cid != 0 else 0 for cid in content_id_target]\n\n        return {\n            'content_id': torch.tensor(content_id_input, dtype=torch.long),\n            'answered_correctly': torch.tensor(answered_correctly_input, dtype=torch.float32),\n            'cumulative_accuracy_input': torch.tensor(cumulative_accuracy_input, dtype=torch.float32),\n\n            'target_id': torch.tensor(content_id_target, dtype=torch.long),\n            'part_target': torch.tensor(part_target, dtype=torch.long),\n            'tags_target': torch.tensor(tags_target, dtype=torch.long),\n            'difficulty_target': torch.tensor(difficulty_target, dtype=torch.float32),\n            'cumulative_accuracy_target': torch.tensor(cumulative_accuracy_target, dtype=torch.float32),\n\n            'label': torch.tensor(label, dtype=torch.float32),\n            'label_mask': torch.tensor(label_mask, dtype=torch.bool)\n        }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T20:05:50.770271Z","iopub.execute_input":"2025-08-03T20:05:50.770595Z","iopub.status.idle":"2025-08-03T20:05:50.785774Z","shell.execute_reply.started":"2025-08-03T20:05:50.770565Z","shell.execute_reply":"2025-08-03T20:05:50.784943Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset=MVAKTDataset(group)\ndataloader = DataLoader(\n    dataset,\n    batch_size=512, \n    shuffle=True,\n    num_workers=8   \n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T20:05:54.149525Z","iopub.execute_input":"2025-08-03T20:05:54.149806Z","iopub.status.idle":"2025-08-03T20:05:55.817940Z","shell.execute_reply.started":"2025-08-03T20:05:54.149782Z","shell.execute_reply":"2025-08-03T20:05:55.817261Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Define model","metadata":{}},{"cell_type":"code","source":"\nimport random\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import roc_auc_score\nfrom tqdm import tqdm\n\n# Constants\nMAX_CONTENT_ID = 15023\nMAX_PART_ID = 7\nMAX_TAG_ID = 188\nMAX_SEQ = 100\nEMBED_DIM = 128\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\n\n# Dataset\nclass MVAKTDataset(Dataset):\n    def __init__(self, group, max_seq=MAX_SEQ, training=True):\n        self.group = group\n        self.max_seq = max_seq\n        self.training = training\n        self.user_ids = [uid for uid in group.index if len(group[uid]['content_id']) > 2]\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        sample = self.group[user_id]\n        seq_len = len(sample['content_id'])\n\n        if seq_len > self.max_seq:\n            start = random.randint(0, seq_len - self.max_seq) if self.training else seq_len - self.max_seq\n            for key in sample:\n                sample[key] = sample[key][start : start + self.max_seq]\n\n        seq_len = len(sample['content_id'])\n        pad_len = max(0, self.max_seq - seq_len)\n\n        def pad(seq, pad_value=0):\n            if isinstance(seq[0], list):\n                return [([pad_value] * len(seq[0]))] * pad_len + seq\n            return [pad_value] * pad_len + list(seq)\n\n        content_id = pad(sample['content_id'])\n        answered_correctly = pad(sample['answered_correctly'])\n        cumulative_accuracy = pad(sample['cumulative_accuracy'])\n        part = pad(sample['part'])\n        tags = pad(sample['tags'], pad_value=0)\n        difficulty = pad(sample['difficulty'])\n\n        content_id_input = content_id[:-1]\n        answered_correctly_input = answered_correctly[:-1]\n        cumulative_accuracy_input = cumulative_accuracy[:-1]\n\n        content_id_target = content_id[1:]\n        part_target = part[1:]\n        tags_target = tags[1:]\n        difficulty_target = difficulty[1:]\n        cumulative_accuracy_target = cumulative_accuracy[1:]\n        label = answered_correctly[1:]\n        label_mask = [1 if cid != 0 else 0 for cid in content_id_target]\n\n        return {\n            'content_id': torch.tensor(content_id_input, dtype=torch.long),\n            'answered_correctly': torch.tensor(answered_correctly_input, dtype=torch.float32),\n            'cumulative_accuracy_input': torch.tensor(cumulative_accuracy_input, dtype=torch.float32),\n            'target_id': torch.tensor(content_id_target, dtype=torch.long),\n            'part_target': torch.tensor(part_target, dtype=torch.long),\n            'tags_target': torch.tensor(tags_target, dtype=torch.long),\n            'difficulty_target': torch.tensor(difficulty_target, dtype=torch.float32),\n            'cumulative_accuracy_target': torch.tensor(cumulative_accuracy_target, dtype=torch.float32),\n            'label': torch.tensor(label, dtype=torch.float32),\n            'label_mask': torch.tensor(label_mask, dtype=torch.bool)\n        }\n\n# Model\ndef future_mask(seq_len):\n    return torch.from_numpy(np.triu(np.ones((seq_len, seq_len)), k=1).astype('bool'))\n\nclass FFN(nn.Module):\n    def __init__(self, in_dim, out_dim):\n        super().__init__()\n        self.linear1 = nn.Linear(in_dim, in_dim)\n        self.relu = nn.ReLU()\n        self.linear2 = nn.Linear(in_dim, out_dim)\n        self.dropout = nn.Dropout(0.2)\n        self.layer_norm = nn.LayerNorm(out_dim)\n\n    def forward(self, x):\n        x = self.linear1(x)\n        x = self.relu(x)\n        x = self.linear2(x)\n        x = self.dropout(x)\n        return self.layer_norm(x)\n\nclass MVAKTModel(nn.Module):\n    def __init__(self, n_questions, n_parts, n_tags, embed_dim=128, n_heads=8, dropout=0.2, n_transformer_layers=1, max_seq=MAX_SEQ):\n        super().__init__()\n        self.embed_dim = embed_dim\n        self.content_embed = nn.Embedding(n_questions + 1, embed_dim, padding_idx=0)\n        self.answered_embed = nn.Linear(1, embed_dim)\n        self.cumacc_embed = nn.Linear(1, embed_dim)\n        self.pos_embed = nn.Embedding(max_seq, embed_dim)\n\n        encoder_layer = nn.TransformerEncoderLayer(d_model=embed_dim, nhead=n_heads, dropout=dropout)\n        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=n_transformer_layers)\n\n        self.part_embed = nn.Embedding(n_parts + 1, embed_dim, padding_idx=0)\n        self.tag_embed = nn.Embedding(n_tags, embed_dim, padding_idx=0)\n        self.agg_tag_proj = nn.Linear(embed_dim, embed_dim)\n        self.diff_proj = nn.Linear(1, embed_dim)\n        self.target_cumacc_proj = nn.Linear(1, embed_dim)\n\n        self.post_proj = FFN(embed_dim * 5, embed_dim)\n        self.fc = nn.Linear(embed_dim, 1)\n\n    def forward(self, content_id, answered_correctly, cumulative_accuracy_input,\n                target_id, part_target, tags_target, difficulty_target, cumulative_accuracy_target):\n\n        B, T = content_id.size()\n        content_emb = self.content_embed(content_id)\n        ans_emb = self.answered_embed(answered_correctly.unsqueeze(-1))\n        cumacc_emb = self.cumacc_embed(cumulative_accuracy_input.unsqueeze(-1))\n        pos_ids = torch.arange(T, device=content_id.device).unsqueeze(0).expand(B, T)\n        pos_emb = self.pos_embed(pos_ids)\n\n        x = content_emb + ans_emb + cumacc_emb + pos_emb\n        x = x.transpose(0, 1)\n        mask = future_mask(T).to(x.device)\n        x = self.transformer(x, mask=mask).transpose(0, 1)\n\n        part_emb = self.part_embed(part_target)\n        tag_mask = (tags_target != 0).unsqueeze(-1)\n        tag_embeds = self.tag_embed(tags_target) * tag_mask\n        tag_sum = tag_embeds.sum(dim=2)\n        tag_count = tag_mask.sum(dim=2).clamp(min=1)\n        tag_mean = tag_sum / tag_count\n        tag_emb = self.agg_tag_proj(tag_mean)\n\n        diff_emb = self.diff_proj(difficulty_target.unsqueeze(-1))\n        cumacc_t_emb = self.target_cumacc_proj(cumulative_accuracy_target.unsqueeze(-1))\n\n        concat = torch.cat([x, part_emb, tag_emb, diff_emb, cumacc_t_emb], dim=-1)\n        out = self.post_proj(concat)\n        return self.fc(out).squeeze(-1)\n\n# Training function\ndef train_epoch(model, train_loader, optimizer, criterion, device='cuda'):\n    model.train()\n    total_loss, total_correct, total_samples = 0.0, 0, 0\n    all_labels, all_preds = [], []\n\n    for batch in tqdm(train_loader):\n        content_id = batch['content_id'].to(device)\n        answered_correctly = batch['answered_correctly'].to(device)\n        cumulative_accuracy_input = batch['cumulative_accuracy_input'].to(device)\n        part_target = batch['part_target'].to(device)\n        tags_target = batch['tags_target'].to(device)\n        difficulty_target = batch['difficulty_target'].to(device)\n        cumulative_accuracy_target = batch['cumulative_accuracy_target'].to(device)\n        target_id = batch['target_id'].to(device)\n        label = batch['label'].to(device)\n        label_mask = batch['label_mask'].to(device)\n\n        optimizer.zero_grad()\n        output = model(content_id, answered_correctly, cumulative_accuracy_input,\n                       target_id, part_target, tags_target, difficulty_target, cumulative_accuracy_target)\n\n        last_output = output[:, -1]\n        last_label = label[:, -1]\n        pred = (torch.sigmoid(last_output) >= 0.5).long()\n        total_correct += (pred == last_label.long()).sum().item()\n        total_samples += last_label.size(0)\n        all_labels.extend(last_label.detach().cpu().numpy())\n        all_preds.extend(torch.sigmoid(last_output).detach().cpu().numpy())\n\n        output_masked = torch.masked_select(output, label_mask)\n        label_masked = torch.masked_select(label, label_mask)\n        loss = criterion(output_masked, label_masked)\n        loss.backward()\n        optimizer.step()\n        total_loss += loss.item()\n\n    avg_loss = total_loss / len(train_loader)\n    acc = total_correct / total_samples\n    auc = roc_auc_score(all_labels, all_preds) if len(set(all_labels)) > 1 else 0.0\n    return avg_loss, acc, auc\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T20:22:49.770645Z","iopub.execute_input":"2025-08-03T20:22:49.770933Z","iopub.status.idle":"2025-08-03T20:22:49.809634Z","shell.execute_reply.started":"2025-08-03T20:22:49.770908Z","shell.execute_reply":"2025-08-03T20:22:49.808881Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\n\n# ----------------------------\n# 1. Set the device\n# ----------------------------\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# ----------------------------\n# 2. Define dataset and dataloader\n# ----------------------------\ndataset = MVAKTDataset(group)  # Assumes `group` is your grouped dictionary\ndataloader = DataLoader(\n    dataset,\n    batch_size=512,\n    shuffle=True,\n    num_workers=4  # Or 0 if using Windows\n)\n\n# ----------------------------\n# 3. Model initialization\n# ----------------------------\nMAX_CONTENT_ID = 15023\nMAX_PART_ID = 7\nMAX_TAG_ID = 188\n\nmodel = MVAKTModel(\n    n_questions=MAX_CONTENT_ID + 3,\n    n_parts=MAX_PART_ID + 3,\n    n_tags=MAX_TAG_ID + 3,\n    embed_dim=128,\n    n_heads=8,\n    dropout=0.2,\n    n_transformer_layers=1,\n).to(device)\n\n# ----------------------------\n# 4. Optimizer and Loss Function\n# ----------------------------\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\nloss_fn = nn.BCEWithLogitsLoss()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T20:26:31.376280Z","iopub.execute_input":"2025-08-03T20:26:31.376634Z","iopub.status.idle":"2025-08-03T20:26:33.220727Z","shell.execute_reply.started":"2025-08-03T20:26:31.376601Z","shell.execute_reply":"2025-08-03T20:26:33.219875Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for epoch in range(1, 36):\n    print(f\"\\n📘 Epoch {epoch} / 5\")\n    train_loss, train_acc, train_auc = train_epoch(\n        model=model,\n        train_loader=dataloader,\n        optimizer=optimizer,\n        criterion=loss_fn,\n        device=device\n    )\n    print(f\"✅ Epoch {epoch}: Loss = {train_loss:.4f}, Acc = {train_acc:.4f}, AUC = {train_auc:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-03T20:48:19.750279Z","iopub.execute_input":"2025-08-03T20:48:19.750621Z","iopub.status.idle":"2025-08-03T21:44:13.755385Z","shell.execute_reply.started":"2025-08-03T20:48:19.750593Z","shell.execute_reply":"2025-08-03T21:44:13.754237Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import riiideducation\nimport pandas as pd\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm import tqdm\n\n# Initialize Riiid environment\nenv = riiideducation.make_env()\niter_test = env.iter_test()\n\n# Load metadata\nquestion_difficulty_dict = dict(question_difficulty)\nqid_to_tags = dict(zip(questions_df['question_id'], questions_df['tag_ids_padded']))\nqid_to_part = dict(zip(questions_df['question_id'], questions_df['part']))\n\n# Initialize user history\nuser_histories = {}\n\n# Dataset class that works for single test batch\nclass OnlineSAKTDataset(Dataset):\n    def __init__(self, test_df, user_histories, max_seq=100):\n        self.samples = []\n        self.max_seq = max_seq\n\n        for _, row in test_df.iterrows():\n            user_id = row['user_id']\n            content_id = row['content_id'] + 1  # Shift for padding\n            \n            # Retrieve or initialize history\n            history = user_histories.get(user_id, {\n                'content_id': [],\n                'answered_correctly': [],\n                'cumulative_accuracy': [],\n                'part': [],\n                'tags': [],\n                'difficulty': [],\n            })\n\n            current_tags = qid_to_tags.get(row['content_id'], [0, 0, 0, 0])\n            current_part = qid_to_part.get(row['content_id'], 1)\n            current_difficulty = question_difficulty_dict.get(row['content_id'], 0.5)\n            last_cumacc = history['cumulative_accuracy'][-1] if history['cumulative_accuracy'] else 0.0\n\n            # Append current to history\n            temp_history = {\n                'content_id': history['content_id'] + [content_id],\n                'answered_correctly': history['answered_correctly'] + [0],\n                'cumulative_accuracy': history['cumulative_accuracy'] + [last_cumacc],\n                'part': history['part'] + [current_part],\n                'tags': history['tags'] + [current_tags],\n                'difficulty': history['difficulty'] + [current_difficulty],\n            }\n\n            # Truncate if longer than max_seq\n            for k in temp_history:\n                temp_history[k] = temp_history[k][-self.max_seq:]\n\n            self.samples.append({\n                'row_id': row['row_id'],\n                'user_id': user_id,\n                'history': temp_history\n            })\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        history = sample['history']\n        row_id = sample['row_id']\n\n        def pad(seq, pad_val=0, dim=1):\n            pad_len = self.max_seq - len(seq)\n            if dim > 1:\n                return [[pad_val]*dim]*pad_len + seq\n            return [pad_val]*pad_len + seq\n\n        return {\n            'content_id': torch.tensor(pad(history['content_id'][:-1]), dtype=torch.long),\n            'answered_correctly': torch.tensor(pad(history['answered_correctly'][:-1]), dtype=torch.float32),\n            'cumulative_accuracy_input': torch.tensor(pad(history['cumulative_accuracy'][:-1]), dtype=torch.float32),\n            'target_id': torch.tensor(pad(history['content_id'][1:]), dtype=torch.long),\n            'part_target': torch.tensor(pad(history['part'][1:]), dtype=torch.long),\n            'tags_target': torch.tensor(pad(history['tags'][1:], pad_val=0, dim=4), dtype=torch.long),\n            'difficulty_target': torch.tensor(pad(history['difficulty'][1:]), dtype=torch.float32),\n            'cumulative_accuracy_target': torch.tensor(pad(history['cumulative_accuracy'][1:]), dtype=torch.float32),\n            'row_id': row_id,\n            'user_id': sample['user_id']\n        }\n\n# Begin loop over test batches\nfor test_df, sample_prediction_df in iter_test:\n    if test_df['content_type_id'].iloc[0] == 1:\n        # Skip lectures\n        env.predict(sample_prediction_df)\n        continue\n\n    # Prepare dataset and loader\n    dataset = OnlineSAKTDataset(test_df, user_histories)\n    loader = DataLoader(dataset, batch_size=64, shuffle=False)\n\n    row_ids, preds = [], []\n\n    with torch.no_grad():\n        for batch in loader:\n            output = model(\n                batch['content_id'].to(device),\n                batch['answered_correctly'].to(device),\n                batch['cumulative_accuracy_input'].to(device),\n                batch['target_id'].to(device),\n                batch['part_target'].to(device),\n                batch['tags_target'].to(device),\n                batch['difficulty_target'].to(device),\n                batch['cumulative_accuracy_target'].to(device),\n            )\n            prob = torch.sigmoid(output[:, -1])\n            preds.extend(prob.cpu().numpy())\n            row_ids.extend(batch['row_id'])\n\n    # Submit predictions\n    sample_prediction_df = pd.DataFrame({\n        'row_id': row_ids,\n        'answered_correctly': preds\n    })\n    env.predict(sample_prediction_df)\n\n    # Update history with correct answers\n    for idx, row in test_df.iterrows():\n        user_id = row['user_id']\n        if row['prior_group_answers_correct'] is np.nan:\n            continue  # Skip if no answers yet (first batch)\n\n        # Initialize if not present\n        if user_id not in user_histories:\n            user_histories[user_id] = {\n                'content_id': [],\n                'answered_correctly': [],\n                'cumulative_accuracy': [],\n                'part': [],\n                'tags': [],\n                'difficulty': []\n            }\n\n        correct = eval(row['prior_group_answers_correct'])\n        responses = eval(row['prior_group_responses'])\n        \n        for content_id, is_correct in zip(responses, correct):\n            cid = content_id + 1\n            user_histories[user_id]['content_id'].append(cid)\n            user_histories[user_id]['answered_correctly'].append(is_correct)\n            cumacc = np.mean(user_histories[user_id]['answered_correctly'])\n            user_histories[user_id]['cumulative_accuracy'].append(cumacc)\n            user_histories[user_id]['part'].append(qid_to_part.get(content_id, 1))\n            user_histories[user_id]['tags'].append(qid_to_tags.get(content_id, [0, 0, 0, 0]))\n            user_histories[user_id]['difficulty'].append(question_difficulty_dict.get(content_id, 0.5))\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}