{"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":"# **Main ideas is as follows：**\n\n1.Basic transformer：  Portugese-English translation . I consider contend_id as Portugese，answered_correctly as English, so SAINT is contend_id -answered_correctly translation.\n\n\n2.modify create_masks function\n\n3.Trained on 12M data.\n","metadata":{}},{"cell_type":"markdown","source":"The reference of this notebook is as follows :\n\n1.https://tensorflow.google.cn/tutorials/text/transformer\n\n2.our host's papers (https://arxiv.org/abs/2002.07033, https://arxiv.org/abs/2010.12042)","metadata":{}},{"cell_type":"markdown","source":"# **SANIT++**\n1.Encoder add features:\n\npart\\tags1\\tags2\n\nuser_lecture_lv\n\n2 Decoder add features:\n\nlagtime\\lagtime2\\lagtime3\n\nelapsed time\\had explation\n","metadata":{}},{"cell_type":"code","source":"!pip install ../input/python-datatable/datatable-0.11.0-cp37-cp37m-manylinux2010_x86_64.whl > /dev/null 2>&1\nimport datatable as dt","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:07:14.392915Z","iopub.execute_input":"2021-06-06T13:07:14.393730Z","iopub.status.idle":"2021-06-06T13:07:37.149806Z","shell.execute_reply.started":"2021-06-06T13:07:14.393686Z","shell.execute_reply":"2021-06-06T13:07:37.148900Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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)\nimport tensorflow as tf\nfrom sklearn.metrics import roc_auc_score\nfrom collections import defaultdict\nimport gc\n\nimport time\nimport matplotlib.pyplot as plt\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\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","execution":{"iopub.status.busy":"2021-06-06T13:07:37.151670Z","iopub.execute_input":"2021-06-06T13:07:37.151951Z","iopub.status.idle":"2021-06-06T13:07:43.746720Z","shell.execute_reply.started":"2021-06-06T13:07:37.151915Z","shell.execute_reply":"2021-06-06T13:07:43.745592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Data**","metadata":{}},{"cell_type":"code","source":"content_count=13523\ntarget_count=2\nuser_lecture_count=500\nlagtime_count=20002\nelapsed_time_count=500\n\nMAX_LENGTH=50\n\nBUFFER_SIZE = 2000\nBATCH_SIZE = 128\n\n#train parameter\n#The values used in the base model of transformer were:\n#num_layers=6, d_model = 512, dff = 2048. See the paper for all the other versions of the transformer.\n\n#SAINT:The window size, dropout rate, and batch size are set to 100, 0.1, and 128 respectively.\n#The best performing model of SAINT has 4 layers and a latent space dimension of 512.\nnum_layers = 2\nd_model = 128\ndff = 512\nnum_heads = 8\n\n# input_vocab_size = content_count + 1\n# target_vocab_size = target_count + 1\n\n#content_id have 0，so add 1\ninput_vocab_size = content_count+1\ntarget_vocab_size = target_count+1\ndropout_rate = 0.1\n\n#\nEPOCHS = 10\ntrain_flag=False","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:07:43.750337Z","iopub.execute_input":"2021-06-06T13:07:43.750693Z","iopub.status.idle":"2021-06-06T13:07:43.757641Z","shell.execute_reply.started":"2021-06-06T13:07:43.750659Z","shell.execute_reply":"2021-06-06T13:07:43.756485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target = 'answered_correctly'","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:07:43.758891Z","iopub.execute_input":"2021-06-06T13:07:43.759285Z","iopub.status.idle":"2021-06-06T13:07:43.768851Z","shell.execute_reply.started":"2021-06-06T13:07:43.759254Z","shell.execute_reply":"2021-06-06T13:07:43.767881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ndata_types_dict = {\n    'timestamp': 'int64',\n    'user_id': 'int32', \n    'content_id': 'int16', \n    'content_type_id':'int8', \n    #'task_container_id': 'int16',\n    #'user_answer': 'int8',\n    'answered_correctly': 'int8', \n    'prior_question_elapsed_time': 'float32', \n    'prior_question_had_explanation': 'bool'\n}\n\nprint('start read train data...')\ntrain_df = dt.fread('../input/riiid-test-answer-prediction/train.csv', columns=set(data_types_dict.keys())).to_pandas()","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:07:43.772120Z","iopub.execute_input":"2021-06-06T13:07:43.772535Z","iopub.status.idle":"2021-06-06T13:09:09.737792Z","shell.execute_reply.started":"2021-06-06T13:07:43.772493Z","shell.execute_reply":"2021-06-06T13:09:09.736898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"user_lecture_agg = train_df.groupby('user_id')['content_type_id'].agg(['sum', 'count'])\nuser_lecture_agg=user_lecture_agg.astype('int16')\n\nuser_lecture_sum_dict = user_lecture_agg['sum'].astype('int16').to_dict(defaultdict(int))\nuser_lecture_count_dict = user_lecture_agg['count'].astype('int16').to_dict(defaultdict(int))\n\ndel user_lecture_agg","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:09:09.740769Z","iopub.execute_input":"2021-06-06T13:09:09.741078Z","iopub.status.idle":"2021-06-06T13:09:13.998819Z","shell.execute_reply.started":"2021-06-06T13:09:09.741047Z","shell.execute_reply":"2021-06-06T13:09:13.997836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cum = train_df.groupby('user_id')['content_type_id'].agg(['cumsum', 'cumcount'])\ncum['cumcount']=cum['cumcount']+1\n\n#train_df['user_lecture_lv'] = cum['cumsum'] / cum['cumcount']\ntrain_df['user_lecture_lv'] = cum['cumsum'] \ntrain_df.user_lecture_lv=train_df.user_lecture_lv.astype('int16')\ndel cum","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:09:14.000004Z","iopub.execute_input":"2021-06-06T13:09:14.000281Z","iopub.status.idle":"2021-06-06T13:09:23.623203Z","shell.execute_reply.started":"2021-06-06T13:09:14.000253Z","shell.execute_reply":"2021-06-06T13:09:23.622239Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['user_lecture_lv']","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:09:23.624402Z","iopub.execute_input":"2021-06-06T13:09:23.624662Z","iopub.status.idle":"2021-06-06T13:09:23.635001Z","shell.execute_reply.started":"2021-06-06T13:09:23.624635Z","shell.execute_reply":"2021-06-06T13:09:23.633874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['prior_question_had_explanation'].fillna(2, inplace=True)\ntrain_df = train_df.astype(data_types_dict)\ntrain_df = train_df[train_df[target] != -1].reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:09:23.636277Z","iopub.execute_input":"2021-06-06T13:09:23.636814Z","iopub.status.idle":"2021-06-06T13:09:45.544504Z","shell.execute_reply.started":"2021-06-06T13:09:23.636776Z","shell.execute_reply":"2021-06-06T13:09:45.543622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"max_timestamp_u = train_df[['user_id','timestamp']].groupby(['user_id']).agg(['max']).reset_index()\nmax_timestamp_u.columns = ['user_id', 'max_time_stamp']\nmax_timestamp_u.user_id=max_timestamp_u.user_id.astype('int32')\n\ntrain_df['lagtime'] = train_df.groupby('user_id')['timestamp'].shift()\n\nmax_timestamp_u2 = train_df[['user_id','lagtime']].groupby(['user_id']).agg(['max']).reset_index()\nmax_timestamp_u2.columns = ['user_id', 'max_time_stamp2']\nmax_timestamp_u2.user_id=max_timestamp_u2.user_id.astype('int32')\n\ntrain_df['lagtime']=train_df['timestamp']-train_df['lagtime']\n","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:09:45.546709Z","iopub.execute_input":"2021-06-06T13:09:45.547114Z","iopub.status.idle":"2021-06-06T13:09:56.287160Z","shell.execute_reply.started":"2021-06-06T13:09:45.547068Z","shell.execute_reply":"2021-06-06T13:09:56.286199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['lagtime']=train_df['lagtime']/(10000)\n","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:09:56.288761Z","iopub.execute_input":"2021-06-06T13:09:56.289159Z","iopub.status.idle":"2021-06-06T13:09:56.901870Z","shell.execute_reply.started":"2021-06-06T13:09:56.289119Z","shell.execute_reply":"2021-06-06T13:09:56.900761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.lagtime[train_df.lagtime>20000]=20000","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:09:56.903138Z","iopub.execute_input":"2021-06-06T13:09:56.903433Z","iopub.status.idle":"2021-06-06T13:09:57.477423Z","shell.execute_reply.started":"2021-06-06T13:09:56.903403Z","shell.execute_reply":"2021-06-06T13:09:57.474976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['lagtime']","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:09:57.479143Z","iopub.execute_input":"2021-06-06T13:09:57.479548Z","iopub.status.idle":"2021-06-06T13:09:57.489898Z","shell.execute_reply.started":"2021-06-06T13:09:57.479507Z","shell.execute_reply":"2021-06-06T13:09:57.488529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['lagtime2'] = train_df.groupby('user_id')['timestamp'].shift(2)\nmax_timestamp_u3 = train_df[['user_id','lagtime2']].groupby(['user_id']).agg(['max']).reset_index()\nmax_timestamp_u3.columns = ['user_id', 'max_time_stamp3']\nmax_timestamp_u3.user_id=max_timestamp_u3.user_id.astype('int32')\n\ntrain_df['lagtime2']=train_df['timestamp']-train_df['lagtime2']\ntrain_df['lagtime2']=train_df['lagtime2']/(10000)\n#train_df.lagtime2=train_df.lagtime2.astype('float32')","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:09:57.491468Z","iopub.execute_input":"2021-06-06T13:09:57.491859Z","iopub.status.idle":"2021-06-06T13:10:07.796050Z","shell.execute_reply.started":"2021-06-06T13:09:57.491817Z","shell.execute_reply":"2021-06-06T13:10:07.795027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.lagtime2[train_df.lagtime2>20000]=20000","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:07.797321Z","iopub.execute_input":"2021-06-06T13:10:07.797647Z","iopub.status.idle":"2021-06-06T13:10:08.141480Z","shell.execute_reply.started":"2021-06-06T13:10:07.797616Z","shell.execute_reply":"2021-06-06T13:10:08.140473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['lagtime2'].max()","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:08.142711Z","iopub.execute_input":"2021-06-06T13:10:08.142987Z","iopub.status.idle":"2021-06-06T13:10:08.266359Z","shell.execute_reply.started":"2021-06-06T13:10:08.142959Z","shell.execute_reply":"2021-06-06T13:10:08.265464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['lagtime3'] = train_df.groupby('user_id')['timestamp'].shift(3)\ntrain_df['lagtime3']=train_df['lagtime3']/(10000)\n#train_df.lagtime3=train_df.lagtime3.astype('float32')","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:08.267612Z","iopub.execute_input":"2021-06-06T13:10:08.267854Z","iopub.status.idle":"2021-06-06T13:10:12.800150Z","shell.execute_reply.started":"2021-06-06T13:10:08.267830Z","shell.execute_reply":"2021-06-06T13:10:12.799277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.lagtime3[train_df.lagtime3>20000]=20000","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:12.801473Z","iopub.execute_input":"2021-06-06T13:10:12.801748Z","iopub.status.idle":"2021-06-06T13:10:16.368700Z","shell.execute_reply.started":"2021-06-06T13:10:12.801721Z","shell.execute_reply":"2021-06-06T13:10:16.367768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"max_timestamp_u_dict=max_timestamp_u.set_index('user_id').to_dict()\nmax_timestamp_u_dict2=max_timestamp_u2.set_index('user_id').to_dict()\nmax_timestamp_u_dict3=max_timestamp_u3.set_index('user_id').to_dict()\n\ndel max_timestamp_u\ndel max_timestamp_u2\ndel max_timestamp_u3\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:16.369887Z","iopub.execute_input":"2021-06-06T13:10:16.370134Z","iopub.status.idle":"2021-06-06T13:10:17.125692Z","shell.execute_reply.started":"2021-06-06T13:10:16.370110Z","shell.execute_reply":"2021-06-06T13:10:17.124636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"prior_question_elapsed_time_mean=train_df['prior_question_elapsed_time'].mean()\ntrain_df['prior_question_elapsed_time'].fillna(prior_question_elapsed_time_mean, inplace=True)\nlagtime_mean=train_df['lagtime'].mean()\ntrain_df['lagtime'].fillna(0, inplace=True)\ntrain_df.lagtime=train_df.lagtime.astype('int16')\nlagtime_mean2=train_df['lagtime2'].mean()\ntrain_df['lagtime2'].fillna(0, inplace=True)\ntrain_df.lagtime2=train_df.lagtime2.astype('int16')\nlagtime_mean3=train_df['lagtime3'].mean()\ntrain_df['lagtime3'].fillna(0, inplace=True)\ntrain_df.lagtime3=train_df.lagtime3.astype('int16')","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:17.127143Z","iopub.execute_input":"2021-06-06T13:10:17.127433Z","iopub.status.idle":"2021-06-06T13:10:20.366960Z","shell.execute_reply.started":"2021-06-06T13:10:17.127406Z","shell.execute_reply":"2021-06-06T13:10:20.365980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['prior_question_elapsed_time']=train_df['prior_question_elapsed_time']/(1000)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:20.368601Z","iopub.execute_input":"2021-06-06T13:10:20.369046Z","iopub.status.idle":"2021-06-06T13:10:20.712649Z","shell.execute_reply.started":"2021-06-06T13:10:20.368974Z","shell.execute_reply":"2021-06-06T13:10:20.711797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.prior_question_elapsed_time=train_df.prior_question_elapsed_time.astype('int16')","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:20.713912Z","iopub.execute_input":"2021-06-06T13:10:20.714168Z","iopub.status.idle":"2021-06-06T13:10:20.953495Z","shell.execute_reply.started":"2021-06-06T13:10:20.714143Z","shell.execute_reply":"2021-06-06T13:10:20.952480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.prior_question_had_explanation=train_df.prior_question_had_explanation.astype('int8')","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:20.955148Z","iopub.execute_input":"2021-06-06T13:10:20.955458Z","iopub.status.idle":"2021-06-06T13:10:21.029786Z","shell.execute_reply.started":"2021-06-06T13:10:20.955427Z","shell.execute_reply":"2021-06-06T13:10:21.028828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"questions_df = pd.read_csv(\n    '../input/riiid-test-answer-prediction/questions.csv', \n    usecols=[0, 1,3,4],\n    dtype={'question_id': 'int16','bundle_id': 'int16', 'part': 'int8','tags': 'str'}\n)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:21.031067Z","iopub.execute_input":"2021-06-06T13:10:21.031337Z","iopub.status.idle":"2021-06-06T13:10:21.057321Z","shell.execute_reply.started":"2021-06-06T13:10:21.031311Z","shell.execute_reply":"2021-06-06T13:10:21.056593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tag = questions_df[\"tags\"].str.split(\" \", n = 10, expand = True)\ntag.columns = ['tags1','tags2','tags3','tags4','tags5','tags6']\n#\n#有的tag没值，赋0值\ntag.fillna(0, inplace=True)\ntag = tag.astype('int16')\nquestions_df =  pd.concat([questions_df,tag],axis=1).drop(['tags'],axis=1)\n\nquestions_df.rename(columns={'question_id':'content_id'}, inplace=True)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:21.058437Z","iopub.execute_input":"2021-06-06T13:10:21.058850Z","iopub.status.idle":"2021-06-06T13:10:21.114399Z","shell.execute_reply.started":"2021-06-06T13:10:21.058808Z","shell.execute_reply":"2021-06-06T13:10:21.113529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"questions_df['tags2'].max()","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:21.115710Z","iopub.execute_input":"2021-06-06T13:10:21.116264Z","iopub.status.idle":"2021-06-06T13:10:21.122448Z","shell.execute_reply.started":"2021-06-06T13:10:21.116219Z","shell.execute_reply":"2021-06-06T13:10:21.121601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#train_df=train_df[0:3300*10000]","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:21.123552Z","iopub.execute_input":"2021-06-06T13:10:21.123800Z","iopub.status.idle":"2021-06-06T13:10:21.133581Z","shell.execute_reply.started":"2021-06-06T13:10:21.123774Z","shell.execute_reply":"2021-06-06T13:10:21.132603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.merge(train_df, questions_df, on='content_id', how='left',right_index=True)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:21.134825Z","iopub.execute_input":"2021-06-06T13:10:21.135307Z","iopub.status.idle":"2021-06-06T13:10:35.864744Z","shell.execute_reply.started":"2021-06-06T13:10:21.135278Z","shell.execute_reply":"2021-06-06T13:10:35.863827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['answer_shift'] = train_df.groupby('user_id')['answered_correctly'].shift()\ntrain_df['answer_shift'].fillna(2, inplace=True)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:35.866410Z","iopub.execute_input":"2021-06-06T13:10:35.866681Z","iopub.status.idle":"2021-06-06T13:10:39.894761Z","shell.execute_reply.started":"2021-06-06T13:10:35.866655Z","shell.execute_reply":"2021-06-06T13:10:39.894065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.answer_shift=train_df.answer_shift.astype('int8')","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:39.895757Z","iopub.execute_input":"2021-06-06T13:10:39.896184Z","iopub.status.idle":"2021-06-06T13:10:40.146837Z","shell.execute_reply.started":"2021-06-06T13:10:39.896144Z","shell.execute_reply":"2021-06-06T13:10:40.146058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.dtypes","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:40.148041Z","iopub.execute_input":"2021-06-06T13:10:40.148475Z","iopub.status.idle":"2021-06-06T13:10:40.158191Z","shell.execute_reply.started":"2021-06-06T13:10:40.148432Z","shell.execute_reply":"2021-06-06T13:10:40.156833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#for inference\\train\ntrain_df=train_df.groupby('user_id').tail(50)\nlen(train_df)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:40.165778Z","iopub.execute_input":"2021-06-06T13:10:40.166107Z","iopub.status.idle":"2021-06-06T13:10:49.883813Z","shell.execute_reply.started":"2021-06-06T13:10:40.166074Z","shell.execute_reply":"2021-06-06T13:10:49.882720Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"##content_id have 0，so add 1\ntrain_df['content_id'] += 1\n# train_df['answered_correctly'] += 1\n\n#train_df = train_df[train_df.content_type_id == False]","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:49.885698Z","iopub.execute_input":"2021-06-06T13:10:49.886126Z","iopub.status.idle":"2021-06-06T13:10:49.925302Z","shell.execute_reply.started":"2021-06-06T13:10:49.886065Z","shell.execute_reply":"2021-06-06T13:10:49.924510Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train_flag:\n    valid_df=train_df[1300*10000:]\n    train_df=train_df[0:1200*10000]#less 100\nelse:\n    valid_df=train_df[1200*10000:1300*10000]\n    test_df=train_df[1200*10000:1201*10000]","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:49.926439Z","iopub.execute_input":"2021-06-06T13:10:49.926826Z","iopub.status.idle":"2021-06-06T13:10:49.932356Z","shell.execute_reply.started":"2021-06-06T13:10:49.926798Z","shell.execute_reply":"2021-06-06T13:10:49.931195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = train_df.groupby('user_id').apply(lambda r:[\n            r['content_id'].values.tolist(),\n            r['answered_correctly'].values.tolist(),\n            r['prior_question_elapsed_time'].values.tolist(),\n            r['prior_question_had_explanation'].values.tolist(),\n            r['user_lecture_lv'].values.tolist(),\n            r['lagtime'].values.tolist(),\n            r['lagtime2'].values.tolist(),\n            r['answer_shift'].values.tolist(),\n            r['part'].values.tolist(),\n            r['tags1'].values.tolist(),\n            r['tags2'].values.tolist()])\n\nvalid_dataset = valid_df.groupby('user_id').apply(lambda r:[\n            r['content_id'].values.tolist(),\n            r['answered_correctly'].values.tolist(),\n            r['prior_question_elapsed_time'].values.tolist(),\n            r['prior_question_had_explanation'].values.tolist(),\n            r['user_lecture_lv'].values.tolist(),\n            r['lagtime'].values.tolist(),\n            r['lagtime2'].values.tolist(),\n            r['answer_shift'].values.tolist(),\n            r['part'].values.tolist(),\n            r['tags1'].values.tolist(),\n            r['tags2'].values.tolist()])","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:10:49.933786Z","iopub.execute_input":"2021-06-06T13:10:49.934224Z","iopub.status.idle":"2021-06-06T13:12:46.053145Z","shell.execute_reply.started":"2021-06-06T13:10:49.934195Z","shell.execute_reply":"2021-06-06T13:12:46.052207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del train_df\ndel valid_df","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:12:46.054612Z","iopub.execute_input":"2021-06-06T13:12:46.055011Z","iopub.status.idle":"2021-06-06T13:12:46.060507Z","shell.execute_reply.started":"2021-06-06T13:12:46.054972Z","shell.execute_reply":"2021-06-06T13:12:46.059260Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('len(train_dataset):',len(train_dataset))","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:12:46.061756Z","iopub.execute_input":"2021-06-06T13:12:46.062060Z","iopub.status.idle":"2021-06-06T13:12:46.078863Z","shell.execute_reply.started":"2021-06-06T13:12:46.062031Z","shell.execute_reply":"2021-06-06T13:12:46.077803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def getlimitdata(x):    \n    seq_len = len(x)\n    r = np.zeros(MAX_LENGTH, dtype=np.float32)\n\n    if seq_len >= MAX_LENGTH:          \n        r[:] = x[-MAX_LENGTH:]\n       \n    else:        \n        #r[-seq_len:] = x\n        r[0:seq_len] = x\n            \n    return r.tolist()\n\ndef getlimitdata2(x,start_token):#response\n    seq_len = len(x)\n    r = np.zeros(MAX_LENGTH, dtype=np.float32)\n\n    if seq_len >= MAX_LENGTH:          \n        r[:] = x[-MAX_LENGTH:]\n        r[0:1]=start_token #start token\n       \n    else:        \n        #r[-seq_len:] = x\n        r[0:seq_len] = x\n        r[0:1]=start_token #start token\n            \n    return r.tolist()","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:12:46.080241Z","iopub.execute_input":"2021-06-06T13:12:46.080816Z","iopub.status.idle":"2021-06-06T13:12:46.095185Z","shell.execute_reply.started":"2021-06-06T13:12:46.080773Z","shell.execute_reply":"2021-06-06T13:12:46.094391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#MAX_LENGTH=50\ngetlimitdata2([1,3,3,4,5],8)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:12:46.096844Z","iopub.execute_input":"2021-06-06T13:12:46.097391Z","iopub.status.idle":"2021-06-06T13:12:46.112770Z","shell.execute_reply.started":"2021-06-06T13:12:46.097337Z","shell.execute_reply":"2021-06-06T13:12:46.112021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def getslice(dataset):\n    content_ids=[]# r['content_id'].values.tolist(),\n    answered_correctlys=[]# r['answered_correctly'].values.tolist(),\n    prior_question_elapsed_times=[]# r['prior_question_elapsed_time'].values.tolist(),\n    prior_question_had_explanations=[]# r['prior_question_had_explanation'].values.tolist(),\n    user_lecture_lvs=[]# r['user_lecture_lv'].values.tolist(),\n    lagtimes=[]# r['lagtime'].values.tolist(),\n    lagtime2s=[]# r['lagtime2'].values.tolist(),\n    answer_shifts=[]# r['answer_shift'].values.tolist(),\n    parts=[]# r['part'].values.tolist(),\n    tags1s=[]# r['tags1'].values.tolist(),\n    tags2s=[]# r['tags2'].values.tolist()])\n\n    i=0\n    for v in dataset.values:\n        #print(v[0])\n        content_ids.append(getlimitdata(v[0]))\n        answered_correctlys.append(getlimitdata(v[1]))\n        prior_question_elapsed_times.append(getlimitdata2(v[2],elapsed_time_count-1))#response decoder\n        prior_question_had_explanations.append(getlimitdata2(v[3],2))#response decoder\n        user_lecture_lvs.append(getlimitdata(v[4]))\n        lagtimes.append(getlimitdata2(v[5],lagtime_count-1))#response decoder\n        lagtime2s.append(getlimitdata2(v[6],lagtime_count-1))#response decoder\n        answer_shifts.append(getlimitdata2(v[7],2))#response decoder\n        parts.append(getlimitdata(v[8]))\n        tags1s.append(getlimitdata(v[9]))\n        tags2s.append(getlimitdata(v[10]))\n    #     if i==1:\n    #         break\n    #     i=i+1\n    \n    return content_ids,answered_correctlys,prior_question_elapsed_times,prior_question_had_explanations,user_lecture_lvs,lagtimes,lagtime2s,answer_shifts,parts,tags1s,tags2s","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:12:46.113932Z","iopub.execute_input":"2021-06-06T13:12:46.114393Z","iopub.status.idle":"2021-06-06T13:12:46.126532Z","shell.execute_reply.started":"2021-06-06T13:12:46.114352Z","shell.execute_reply":"2021-06-06T13:12:46.125345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nif train_flag:\n    train_dataset = tf.data.Dataset.from_tensor_slices((getslice(train_dataset)))\n\nvalid_dataset = tf.data.Dataset.from_tensor_slices((getslice(valid_dataset)))","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:12:46.128280Z","iopub.execute_input":"2021-06-06T13:12:46.128630Z","iopub.status.idle":"2021-06-06T13:13:23.263909Z","shell.execute_reply.started":"2021-06-06T13:12:46.128599Z","shell.execute_reply":"2021-06-06T13:13:23.263211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# del content_ids\n# del answered_correctlys\n# del prior_question_elapsed_times\n# del prior_question_had_explanations\n# del user_lecture_lvs\n# del lagtimes\n# del lagtime2s\n# del answer_shifts\n# del parts\n# del tags1s\n# del tags2s","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.265065Z","iopub.execute_input":"2021-06-06T13:13:23.265314Z","iopub.status.idle":"2021-06-06T13:13:23.268585Z","shell.execute_reply.started":"2021-06-06T13:13:23.265290Z","shell.execute_reply":"2021-06-06T13:13:23.267897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train_flag:\n    train_dataset = train_dataset.cache()\n    train_dataset = train_dataset.shuffle(BUFFER_SIZE).padded_batch(BATCH_SIZE)\n    train_dataset = train_dataset.prefetch(tf.data.experimental.AUTOTUNE)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.269471Z","iopub.execute_input":"2021-06-06T13:13:23.269800Z","iopub.status.idle":"2021-06-06T13:13:23.288832Z","shell.execute_reply.started":"2021-06-06T13:13:23.269775Z","shell.execute_reply":"2021-06-06T13:13:23.287583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_dataset = valid_dataset.padded_batch(2000)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.290460Z","iopub.execute_input":"2021-06-06T13:13:23.290808Z","iopub.status.idle":"2021-06-06T13:13:23.306469Z","shell.execute_reply.started":"2021-06-06T13:13:23.290775Z","shell.execute_reply":"2021-06-06T13:13:23.305525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"content_ids,answered_correctlys,prior_question_elapsed_times,prior_question_had_explanations,user_lecture_lvs,lagtimes,lagtime2s,answer_shifts,parts,tags1s,tags2s =next(iter(valid_dataset))","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.308005Z","iopub.execute_input":"2021-06-06T13:13:23.308421Z","iopub.status.idle":"2021-06-06T13:13:23.449560Z","shell.execute_reply.started":"2021-06-06T13:13:23.308358Z","shell.execute_reply":"2021-06-06T13:13:23.448774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"content_ids","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.451895Z","iopub.execute_input":"2021-06-06T13:13:23.452417Z","iopub.status.idle":"2021-06-06T13:13:23.458680Z","shell.execute_reply.started":"2021-06-06T13:13:23.452385Z","shell.execute_reply":"2021-06-06T13:13:23.457630Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"user_lecture_lvs","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.460176Z","iopub.execute_input":"2021-06-06T13:13:23.460573Z","iopub.status.idle":"2021-06-06T13:13:23.474537Z","shell.execute_reply.started":"2021-06-06T13:13:23.460534Z","shell.execute_reply":"2021-06-06T13:13:23.473730Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"user_lecture_lv_embeddings= tf.keras.layers.Dense(d_model, use_bias=False)(user_lecture_lvs)\nuser_lecture_lv_embeddings","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.476048Z","iopub.execute_input":"2021-06-06T13:13:23.476390Z","iopub.status.idle":"2021-06-06T13:13:23.578779Z","shell.execute_reply.started":"2021-06-06T13:13:23.476295Z","shell.execute_reply":"2021-06-06T13:13:23.577989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"content_embeddings = tf.keras.layers.Embedding(input_vocab_size, d_model)(content_ids)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.579927Z","iopub.execute_input":"2021-06-06T13:13:23.580189Z","iopub.status.idle":"2021-06-06T13:13:23.618597Z","shell.execute_reply.started":"2021-06-06T13:13:23.580163Z","shell.execute_reply":"2021-06-06T13:13:23.617604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Model**","metadata":{}},{"cell_type":"code","source":"def get_angles(pos, i, d_model):\n    angle_rates = 1 / np.power(10000, (2 * (i//2)) / np.float32(d_model))\n    return pos * angle_rates\n\ndef positional_encoding(position, d_model):\n    angle_rads = get_angles(np.arange(position)[:, np.newaxis],\n                          np.arange(d_model)[np.newaxis, :],\n                          d_model)\n\n    # apply sin to even indices in the array; 2i\n    angle_rads[:, 0::2] = np.sin(angle_rads[:, 0::2])\n\n    # apply cos to odd indices in the array; 2i+1\n    angle_rads[:, 1::2] = np.cos(angle_rads[:, 1::2])\n    #print(angle_rads)\n    pos_encoding = angle_rads[np.newaxis, ...]\n    \n    #pos_encoding=np.repeat(pos_encoding,BATCH_SIZE,0)#在0 这个维度扩展batchsize遍\n    \n    return tf.cast(pos_encoding, dtype=tf.float32)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.619661Z","iopub.execute_input":"2021-06-06T13:13:23.619899Z","iopub.status.idle":"2021-06-06T13:13:23.626895Z","shell.execute_reply.started":"2021-06-06T13:13:23.619875Z","shell.execute_reply":"2021-06-06T13:13:23.626106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x=positional_encoding(50,d_model)\nx","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.627800Z","iopub.execute_input":"2021-06-06T13:13:23.628065Z","iopub.status.idle":"2021-06-06T13:13:23.645824Z","shell.execute_reply.started":"2021-06-06T13:13:23.628039Z","shell.execute_reply":"2021-06-06T13:13:23.645203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x[:28,:,]","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.646741Z","iopub.execute_input":"2021-06-06T13:13:23.647092Z","iopub.status.idle":"2021-06-06T13:13:23.660255Z","shell.execute_reply.started":"2021-06-06T13:13:23.647057Z","shell.execute_reply":"2021-06-06T13:13:23.659465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch\n# from torch import nn\n# seq = torch.arange(50).unsqueeze(0)\n# pos=nn.Embedding(50,128)(seq)\n# pos","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.663421Z","iopub.execute_input":"2021-06-06T13:13:23.663787Z","iopub.status.idle":"2021-06-06T13:13:23.666858Z","shell.execute_reply.started":"2021-06-06T13:13:23.663757Z","shell.execute_reply":"2021-06-06T13:13:23.665955Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Mask**","metadata":{}},{"cell_type":"code","source":"def create_padding_mask(seq):\n    seq = tf.cast(tf.math.equal(seq, 0), tf.float32)#\n\n    # add extra dimensions to add the padding\n    # to the attention logits.\n    return seq[:, tf.newaxis, tf.newaxis, :]  # (batch_size, 1, 1, seq_len)\n\ndef create_look_ahead_mask(size):\n    mask = 1 - tf.linalg.band_part(tf.ones((size, size)), -1, 0)\n    return mask  # (seq_len, seq_len)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.668471Z","iopub.execute_input":"2021-06-06T13:13:23.669093Z","iopub.status.idle":"2021-06-06T13:13:23.678267Z","shell.execute_reply.started":"2021-06-06T13:13:23.669052Z","shell.execute_reply":"2021-06-06T13:13:23.677616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"create_look_ahead_mask(5)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.679619Z","iopub.execute_input":"2021-06-06T13:13:23.680226Z","iopub.status.idle":"2021-06-06T13:13:23.699187Z","shell.execute_reply.started":"2021-06-06T13:13:23.680186Z","shell.execute_reply":"2021-06-06T13:13:23.698419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch\n# seq_len=5\n# torch.triu(torch.ones(seq_len,seq_len),diagonal=1).to(dtype=torch.bool)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.700229Z","iopub.execute_input":"2021-06-06T13:13:23.700554Z","iopub.status.idle":"2021-06-06T13:13:23.705223Z","shell.execute_reply.started":"2021-06-06T13:13:23.700528Z","shell.execute_reply":"2021-06-06T13:13:23.704133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Scaled dot product attention**","metadata":{}},{"cell_type":"code","source":"def scaled_dot_product_attention(q, k, v, mask):\n#     \"\"\"Calculate the attention weights.\n#       q, k, v must have matching leading dimensions.\n#       k, v must have matching penultimate dimension, i.e.: seq_len_k = seq_len_v.\n#       The mask has different shapes depending on its type(padding or look ahead) \n#       but it must be broadcastable for addition.\n\n#       Args:\n#         q: query shape == (..., seq_len_q, depth)\n#         k: key shape == (..., seq_len_k, depth)\n#         v: value shape == (..., seq_len_v, depth_v)\n#         mask: Float tensor with shape broadcastable \n#               to (..., seq_len_q, seq_len_k). Defaults to None.\n\n#       Returns:\n#         output, attention_weights\n#     \"\"\"\n\n    matmul_qk = tf.matmul(q, k, transpose_b=True)  # (..., seq_len_q, seq_len_k)\n\n    # scale matmul_qk\n    dk = tf.cast(tf.shape(k)[-1], tf.float32)\n    scaled_attention_logits = matmul_qk / tf.math.sqrt(dk)\n\n    # add the mask to the scaled tensor.\n    if mask is not None:\n        scaled_attention_logits += (mask * -1e9)  \n\n    # softmax is normalized on the last axis (seq_len_k) so that the scores\n    # add up to 1.\n    attention_weights = tf.nn.softmax(scaled_attention_logits, axis=-1)  # (..., seq_len_q, seq_len_k)\n\n    output = tf.matmul(attention_weights, v)  # (..., seq_len_q, depth_v)\n\n    return output, attention_weights","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.706253Z","iopub.execute_input":"2021-06-06T13:13:23.706615Z","iopub.status.idle":"2021-06-06T13:13:23.716060Z","shell.execute_reply.started":"2021-06-06T13:13:23.706588Z","shell.execute_reply":"2021-06-06T13:13:23.715290Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Multi-head attention**","metadata":{}},{"cell_type":"code","source":"class MultiHeadAttention(tf.keras.layers.Layer):\n    def __init__(self, d_model, num_heads):\n        super(MultiHeadAttention, self).__init__()\n        self.num_heads = num_heads\n        self.d_model = d_model\n\n        assert d_model % self.num_heads == 0\n\n        self.depth = d_model // self.num_heads\n\n        self.wq = tf.keras.layers.Dense(d_model)\n        self.wk = tf.keras.layers.Dense(d_model)\n        self.wv = tf.keras.layers.Dense(d_model)\n\n        self.dense = tf.keras.layers.Dense(d_model)\n\n    def split_heads(self, x, batch_size):\n       \n        x = tf.reshape(x, (batch_size, -1, self.num_heads, self.depth))\n        return tf.transpose(x, perm=[0, 2, 1, 3])\n\n    def call(self, v, k, q, mask):\n        batch_size = tf.shape(q)[0]\n\n        q = self.wq(q)  # (batch_size, seq_len, d_model)\n        k = self.wk(k)  # (batch_size, seq_len, d_model)\n        v = self.wv(v)  # (batch_size, seq_len, d_model)\n\n        q = self.split_heads(q, batch_size)  # (batch_size, num_heads, seq_len_q, depth)\n        k = self.split_heads(k, batch_size)  # (batch_size, num_heads, seq_len_k, depth)\n        v = self.split_heads(v, batch_size)  # (batch_size, num_heads, seq_len_v, depth)\n\n        # scaled_attention.shape == (batch_size, num_heads, seq_len_q, depth)\n        # attention_weights.shape == (batch_size, num_heads, seq_len_q, seq_len_k)\n        scaled_attention, attention_weights = scaled_dot_product_attention(\n            q, k, v, mask)\n\n        scaled_attention = tf.transpose(scaled_attention, perm=[0, 2, 1, 3])  # (batch_size, seq_len_q, num_heads, depth)\n\n        concat_attention = tf.reshape(scaled_attention, \n                                      (batch_size, -1, self.d_model))  # (batch_size, seq_len_q, d_model)\n\n        output = self.dense(concat_attention)  # (batch_size, seq_len_q, d_model)\n\n        return output, attention_weights","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.717513Z","iopub.execute_input":"2021-06-06T13:13:23.717894Z","iopub.status.idle":"2021-06-06T13:13:23.731254Z","shell.execute_reply.started":"2021-06-06T13:13:23.717854Z","shell.execute_reply":"2021-06-06T13:13:23.730223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Point wise feed forward network**","metadata":{}},{"cell_type":"code","source":"def point_wise_feed_forward_network(d_model, dff):\n    return tf.keras.Sequential([\n      tf.keras.layers.Dense(dff, activation='relu'),  # (batch_size, seq_len, dff)\n      tf.keras.layers.Dense(d_model)  # (batch_size, seq_len, d_model)\n    ])","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.732633Z","iopub.execute_input":"2021-06-06T13:13:23.733178Z","iopub.status.idle":"2021-06-06T13:13:23.746986Z","shell.execute_reply.started":"2021-06-06T13:13:23.733134Z","shell.execute_reply":"2021-06-06T13:13:23.746210Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Encoder and decoder**","metadata":{}},{"cell_type":"code","source":"class EncoderLayer(tf.keras.layers.Layer):\n    def __init__(self, d_model, num_heads, dff, rate=0.1):\n        super(EncoderLayer, self).__init__()\n\n        self.mha = MultiHeadAttention(d_model, num_heads)\n        self.ffn = point_wise_feed_forward_network(d_model, dff)\n\n        self.layernorm1 = tf.keras.layers.LayerNormalization(epsilon=1e-6)\n        self.layernorm2 = tf.keras.layers.LayerNormalization(epsilon=1e-6)\n\n        self.dropout1 = tf.keras.layers.Dropout(rate)\n        self.dropout2 = tf.keras.layers.Dropout(rate)\n\n    def call(self, x, training, mask):\n\n        attn_output, _ = self.mha(x, x, x, mask)  # (batch_size, input_seq_len, d_model)\n        attn_output = self.dropout1(attn_output, training=training)\n        out1 = self.layernorm1(x + attn_output)  # (batch_size, input_seq_len, d_model)\n\n        ffn_output = self.ffn(out1)  # (batch_size, input_seq_len, d_model)\n        ffn_output = self.dropout2(ffn_output, training=training)\n        out2 = self.layernorm2(out1 + ffn_output)  # (batch_size, input_seq_len, d_model)\n\n        return out2","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.748070Z","iopub.execute_input":"2021-06-06T13:13:23.748468Z","iopub.status.idle":"2021-06-06T13:13:23.764373Z","shell.execute_reply.started":"2021-06-06T13:13:23.748439Z","shell.execute_reply":"2021-06-06T13:13:23.763474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Encoder(tf.keras.layers.Layer):\n    def __init__(self, num_layers, d_model, num_heads, dff, input_vocab_size,\n               maximum_position_encoding, rate=0.1):\n        super(Encoder, self).__init__()\n\n        self.d_model = d_model\n        self.num_layers = num_layers\n\n        self.embedding1 = tf.keras.layers.Embedding(input_vocab_size, d_model)\n        self.embedding2 =tf.keras.layers.Embedding(user_lecture_count, d_model)\n        self.embedding3 =tf.keras.layers.Embedding(8, d_model)\n        self.embedding4 = tf.keras.layers.Embedding(188, d_model)\n        self.embedding5 = tf.keras.layers.Embedding(188, d_model)\n        self.pos_encoding = positional_encoding(maximum_position_encoding, \n                                                self.d_model)\n\n\n        self.enc_layers = [EncoderLayer(d_model, num_heads, dff, rate) \n                           for _ in range(num_layers)]\n\n        self.dropout = tf.keras.layers.Dropout(rate)\n\n    def call(self, x, training, mask):\n       \n        content_ids=x[0]\n        content_embeddings = self.embedding1(content_ids)\n        #user_lecture_lv_embeddings = tf.keras.layers.Dense(d_model, use_bias=False)(x[1])\n        user_lecture_lv_embeddings = self.embedding2(x[1])\n       \n        part_embeddings = self.embedding3(x[2])\n        tags1_embeddings = self.embedding4(x[3])\n        tags2_embeddings = self.embedding5(x[4])\n        \n        seq_len = tf.shape(content_ids)[1]\n        batch=tf.shape(content_ids)[0]\n        \n        pos= self.pos_encoding[:, :seq_len, :]\n        \n        x=content_embeddings+part_embeddings+tags1_embeddings\n        #x=x+tags2_embeddings+user_lecture_lv_embeddings\n       \n        # Add embeddings\n#         x = tf.keras.layers.Add()([            \n#             content_embeddings,\n#             user_lecture_lv_embeddings,\n#             part_embeddings,\n#             tags1_embeddings,\n#             tags2_embeddings\n#             #pos\n#         ])\n        \n        x *= tf.math.sqrt(tf.cast(self.d_model, tf.float32))\n        x+=pos\n        # \n        #x = self.embedding(x)  # (batch_size, input_seq_len, d_model)\n     \n\n        x = self.dropout(x, training=training)\n\n        for i in range(self.num_layers):\n            x = self.enc_layers[i](x, training, mask)\n\n        return x  # (batch_size, input_seq_len, d_model)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.765552Z","iopub.execute_input":"2021-06-06T13:13:23.765853Z","iopub.status.idle":"2021-06-06T13:13:23.781025Z","shell.execute_reply.started":"2021-06-06T13:13:23.765825Z","shell.execute_reply":"2021-06-06T13:13:23.779953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DecoderLayer(tf.keras.layers.Layer):\n    def __init__(self, d_model, num_heads, dff, rate=0.1):\n        super(DecoderLayer, self).__init__()\n\n        self.mha1 = MultiHeadAttention(d_model, num_heads)\n        self.mha2 = MultiHeadAttention(d_model, num_heads)\n\n        self.ffn = point_wise_feed_forward_network(d_model, dff)\n\n        self.layernorm1 = tf.keras.layers.LayerNormalization(epsilon=1e-6)\n        self.layernorm2 = tf.keras.layers.LayerNormalization(epsilon=1e-6)\n        self.layernorm3 = tf.keras.layers.LayerNormalization(epsilon=1e-6)\n\n        self.dropout1 = tf.keras.layers.Dropout(rate)\n        self.dropout2 = tf.keras.layers.Dropout(rate)\n        self.dropout3 = tf.keras.layers.Dropout(rate)\n\n\n    def call(self, x, enc_output, training, \n           look_ahead_mask, padding_mask):\n    # enc_output.shape == (batch_size, input_seq_len, d_model)\n\n        attn1, attn_weights_block1 = self.mha1(x, x, x, look_ahead_mask)  # (batch_size, target_seq_len, d_model)\n        attn1 = self.dropout1(attn1, training=training)\n        out1 = self.layernorm1(attn1 + x)\n\n        attn2, attn_weights_block2 = self.mha2(\n            enc_output, enc_output, out1, padding_mask)  # (batch_size, target_seq_len, d_model)\n        attn2 = self.dropout2(attn2, training=training)\n        out2 = self.layernorm2(attn2 + out1)  # (batch_size, target_seq_len, d_model)\n\n        ffn_output = self.ffn(out2)  # (batch_size, target_seq_len, d_model)\n        ffn_output = self.dropout3(ffn_output, training=training)\n        out3 = self.layernorm3(ffn_output + out2)  # (batch_size, target_seq_len, d_model)\n\n        return out3, attn_weights_block1, attn_weights_block2","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.782697Z","iopub.execute_input":"2021-06-06T13:13:23.783077Z","iopub.status.idle":"2021-06-06T13:13:23.795752Z","shell.execute_reply.started":"2021-06-06T13:13:23.783038Z","shell.execute_reply":"2021-06-06T13:13:23.794586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Decoder(tf.keras.layers.Layer):\n    def __init__(self, num_layers, d_model, num_heads, dff, target_vocab_size,\n               maximum_position_encoding, rate=0.1):\n        super(Decoder, self).__init__()\n\n        self.d_model = d_model\n        self.num_layers = num_layers\n\n        self.embedding1 = tf.keras.layers.Embedding(3, d_model)\n        self.embedding2 = tf.keras.layers.Embedding(3, d_model)\n        self.embedding3 = tf.keras.layers.Embedding(elapsed_time_count, d_model)\n        self.embedding4 = tf.keras.layers.Embedding(lagtime_count, d_model)\n        self.embedding5 = tf.keras.layers.Embedding(lagtime_count, d_model)\n        \n        self.pos_encoding = positional_encoding(maximum_position_encoding, d_model)\n\n        self.dec_layers = [DecoderLayer(d_model, num_heads, dff, rate) \n                           for _ in range(num_layers)]\n        self.dropout = tf.keras.layers.Dropout(rate)\n\n    def call(self, x, enc_output, training, \n           look_ahead_mask, padding_mask):\n        #encode_inputs=[content_ids,user_lecture_lvs,parts,tags1s,tags2s]\n        #decode_inputs=[answered_correctlys,prior_question_elapsed_times,prior_question_had_explanations,lagtimes,lagtime2s,answer_shifts]\n     \n        prior_question_elapsed_times=x[1]     \n        prior_question_had_explanations=x[2]   \n        lagtimes=x[3]       \n        lagtime2s=x[4]       \n        answer_shifts=x[5]\n        \n        \n        answer_shifts_embeddings = self.embedding1(answer_shifts)\n        had_explanations_embeddings = self.embedding2(prior_question_had_explanations)\n#         elapsed_time_embeddings = tf.keras.layers.Dense(d_model, use_bias=False)(prior_question_elapsed_times)\n#         lagtimes_embeddings = tf.keras.layers.Dense(d_model, use_bias=False)(lagtimes)\n#         lagtime2s_embeddings = tf.keras.layers.Dense(d_model, use_bias=False)(lagtime2s)\n#         lagtime_count=20002\n#         elapsed_time_count=500\n        elapsed_time_embeddings = self.embedding3(prior_question_elapsed_times)\n        lagtimes_embeddings = self.embedding4(lagtimes)\n        lagtime2s_embeddings = self.embedding5(lagtime2s)\n        #lagtime3s_embeddings = tf.keras.layers.Dense(d_model, use_bias=False)(lagtime3s)\n        \n#         elapsed_time_embeddings=prior_question_elapsed_times\n#         lagtimes_embeddings=lagtimes\n#         lagtime2s_embeddings=lagtime2s\n        \n        seq_len = tf.shape(answer_shifts)[1]\n        batch=tf.shape(answer_shifts)[0]\n        attention_weights = {}\n        \n        pos = self.pos_encoding[:, :seq_len, :]\n        \n        x=answer_shifts_embeddings+had_explanations_embeddings\n        #x=x+elapsed_time_embeddings+lagtimes_embeddings+lagtime2s_embeddings\n        \n#         x = tf.keras.layers.Add()([            \n#             answer_shifts_embeddings,\n#             had_explanations_embeddings,\n#             elapsed_time_embeddings,\n#             lagtimes_embeddings,\n#             lagtime2s_embeddings\n#             #lagtime3s_embeddings,\n#             #pos\n#         ])\n        \n        x *= tf.math.sqrt(tf.cast(self.d_model, tf.float32))\n        x+=pos   \n        #x = self.embedding(x)  # (batch_size, target_seq_len, d_model)\n    \n\n        x = self.dropout(x, training=training)\n\n        for i in range(self.num_layers):\n            x, block1, block2 = self.dec_layers[i](x, enc_output, training,\n                                                 look_ahead_mask, padding_mask)\n\n            attention_weights['decoder_layer{}_block1'.format(i+1)] = block1\n            attention_weights['decoder_layer{}_block2'.format(i+1)] = block2\n\n        # x.shape == (batch_size, target_seq_len, d_model)\n        return x, attention_weights","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.797793Z","iopub.execute_input":"2021-06-06T13:13:23.798206Z","iopub.status.idle":"2021-06-06T13:13:23.815082Z","shell.execute_reply.started":"2021-06-06T13:13:23.798166Z","shell.execute_reply":"2021-06-06T13:13:23.814067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ** Transformer**","metadata":{}},{"cell_type":"code","source":"class Transformer(tf.keras.Model):\n    def __init__(self, num_layers, d_model, num_heads, dff, input_vocab_size, \n               target_vocab_size, pe_input, pe_target, rate=0.1):\n        super(Transformer, self).__init__()\n\n        self.encoder = Encoder(num_layers, d_model, num_heads, dff, \n                               input_vocab_size, pe_input, rate)\n\n        self.decoder = Decoder(num_layers, d_model, num_heads, dff, \n                               target_vocab_size, pe_target, rate)\n\n        self.final_layer = tf.keras.layers.Dense(target_vocab_size,activation='softmax')#\n\n    def call(self, encode_inputs, decode_inputs, training, enc_padding_mask, \n           look_ahead_mask, dec_padding_mask):\n       \n        enc_output = self.encoder(encode_inputs, training, enc_padding_mask)  # (batch_size, inp_seq_len, d_model)\n\n        # dec_output.shape == (batch_size, tar_seq_len, d_model)\n        dec_output, attention_weights = self.decoder(\n            decode_inputs, enc_output, training, look_ahead_mask, dec_padding_mask)\n\n        final_output = self.final_layer(dec_output)  # (batch_size, tar_seq_len, target_vocab_size)\n        \n\n        return final_output, attention_weights","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","execution":{"iopub.status.busy":"2021-06-06T13:13:23.816579Z","iopub.execute_input":"2021-06-06T13:13:23.816940Z","iopub.status.idle":"2021-06-06T13:13:23.829807Z","shell.execute_reply.started":"2021-06-06T13:13:23.816902Z","shell.execute_reply":"2021-06-06T13:13:23.828946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Optimizer**","metadata":{}},{"cell_type":"code","source":"class CustomSchedule(tf.keras.optimizers.schedules.LearningRateSchedule):\n    def __init__(self, d_model, warmup_steps=4000):\n        super(CustomSchedule, self).__init__()\n\n        self.d_model = d_model\n        self.d_model = tf.cast(self.d_model, tf.float32)\n\n        self.warmup_steps = warmup_steps\n\n    def __call__(self, step):\n        arg1 = tf.math.rsqrt(step)\n        arg2 = step * (self.warmup_steps ** -1.5)\n\n        return tf.math.rsqrt(self.d_model) * tf.math.minimum(arg1, arg2)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.830989Z","iopub.execute_input":"2021-06-06T13:13:23.831241Z","iopub.status.idle":"2021-06-06T13:13:23.843895Z","shell.execute_reply.started":"2021-06-06T13:13:23.831217Z","shell.execute_reply":"2021-06-06T13:13:23.842817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learning_rate = CustomSchedule(d_model)\n#SAINT:We use the Adam optimizer  with lr =0:001;b1 = 0:9;b2 = 0:999 and epsilon = 1e􀀀8.\noptimizer = tf.keras.optimizers.Adam(learning_rate, beta_1=0.9, beta_2=0.98, \n                                     epsilon=1e-9)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.847184Z","iopub.execute_input":"2021-06-06T13:13:23.847552Z","iopub.status.idle":"2021-06-06T13:13:23.856136Z","shell.execute_reply.started":"2021-06-06T13:13:23.847521Z","shell.execute_reply":"2021-06-06T13:13:23.855217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Loss and metrics**","metadata":{}},{"cell_type":"code","source":"# from_logits=False，output为经过softmax输出的概率值。\n# from_logits=True，output为经过网络直接输出的 logits张量。\nloss_object = tf.keras.losses.SparseCategoricalCrossentropy(\n    from_logits=False, reduction='none')\n\ndef loss_function(mask,real, pred):\n    #mask = tf.math.logical_not(tf.math.equal(real, 0))\n    loss_ = loss_object(real, pred)\n\n    mask = tf.cast(mask, dtype=loss_.dtype)\n    loss_ *= mask\n\n    return tf.reduce_sum(loss_)/tf.reduce_sum(mask)\n\n\ndef accuracy_function(mask,real, pred):\n    real = tf.cast(real, dtype=tf.int64)\n    accuracies = tf.equal(real, tf.argmax(pred, axis=2))\n\n    #mask = tf.math.logical_not(tf.math.equal(real, 0))\n    accuracies = tf.math.logical_and(mask, accuracies)\n\n    accuracies = tf.cast(accuracies, dtype=tf.float32)\n    mask = tf.cast(mask, dtype=tf.float32)\n    return tf.reduce_sum(accuracies)/tf.reduce_sum(mask)\n\ntrain_loss = tf.keras.metrics.Mean(name='train_loss')\ntrain_accuracy = tf.keras.metrics.Mean(name='train_accuracy')\nval_loss = tf.keras.metrics.Mean(name='val_loss')\nval_accuracy = tf.keras.metrics.Mean(name='val_accuracy')","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.857378Z","iopub.execute_input":"2021-06-06T13:13:23.857632Z","iopub.status.idle":"2021-06-06T13:13:23.893844Z","shell.execute_reply.started":"2021-06-06T13:13:23.857607Z","shell.execute_reply":"2021-06-06T13:13:23.892897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Training and checkpointing**","metadata":{}},{"cell_type":"code","source":"transformer = Transformer(num_layers, d_model, num_heads, dff,\n                          input_vocab_size, target_vocab_size, \n                          pe_input=50, \n                          pe_target=50,\n                          rate=dropout_rate)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.894935Z","iopub.execute_input":"2021-06-06T13:13:23.895188Z","iopub.status.idle":"2021-06-06T13:13:23.961135Z","shell.execute_reply.started":"2021-06-06T13:13:23.895164Z","shell.execute_reply":"2021-06-06T13:13:23.960173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def create_masks(inp, tar):\n#     # \n#     enc_padding_mask = create_look_ahead_mask(tf.shape(inp)[1]) \n#     dec_padding_mask = create_look_ahead_mask(tf.shape(inp)[1])\n#     #dec_padding_mask=dec_padding_mask[1:,]\n#     combined_mask = create_look_ahead_mask(tf.shape(tar)[1])\n\n#     return enc_padding_mask, combined_mask, dec_padding_mask","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.962669Z","iopub.execute_input":"2021-06-06T13:13:23.963048Z","iopub.status.idle":"2021-06-06T13:13:23.967809Z","shell.execute_reply.started":"2021-06-06T13:13:23.963008Z","shell.execute_reply":"2021-06-06T13:13:23.966672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_masks(inp, tar):\n  # Encoder padding mask\n    enc_padding_mask = create_padding_mask(inp)\n    look_ahead_mask = create_look_ahead_mask(tf.shape(inp)[1])\n    enc_padding_mask = tf.maximum(enc_padding_mask, look_ahead_mask)\n\n    # Used in the 2nd attention block in the decoder.\n    # This padding mask is used to mask the encoder outputs.\n    dec_padding_mask = create_padding_mask(inp)\n    look_ahead_mask = create_look_ahead_mask(tf.shape(inp)[1])\n    dec_padding_mask = tf.maximum(dec_padding_mask, look_ahead_mask)\n    \n    # Used in the 1st attention block in the decoder.\n    # It is used to pad and mask future tokens in the input received by \n    # the decoder.\n    look_ahead_mask = create_look_ahead_mask(tf.shape(inp)[1])\n    dec_target_padding_mask = create_padding_mask(inp)\n    combined_mask = tf.maximum(dec_target_padding_mask, look_ahead_mask)\n\n    return enc_padding_mask, combined_mask, dec_padding_mask","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.969508Z","iopub.execute_input":"2021-06-06T13:13:23.969887Z","iopub.status.idle":"2021-06-06T13:13:23.978150Z","shell.execute_reply.started":"2021-06-06T13:13:23.969848Z","shell.execute_reply":"2021-06-06T13:13:23.977423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint_path = \"../input/riiid-saint-transformer-6-1/checkpoints/train\"\n\nckpt = tf.train.Checkpoint(transformer=transformer,\n                           optimizer=optimizer)\n\nckpt_manager = tf.train.CheckpointManager(ckpt, checkpoint_path, max_to_keep=5)\n\n# \nif ckpt_manager.latest_checkpoint:\n    ckpt.restore(ckpt_manager.latest_checkpoint)\n    print ('Latest checkpoint restored!!')\n    \n    \n#save path\ncheckpoint_path = \"./checkpoints/train\"\n\nckpt_manager = tf.train.CheckpointManager(ckpt, checkpoint_path, max_to_keep=5)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:23.979386Z","iopub.execute_input":"2021-06-06T13:13:23.979738Z","iopub.status.idle":"2021-06-06T13:13:24.104072Z","shell.execute_reply.started":"2021-06-06T13:13:23.979700Z","shell.execute_reply":"2021-06-06T13:13:24.103312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# The @tf.function trace-compiles train_step into a TF graph for faster\n# execution. The function specializes to the precise shape of the argument\n# tensors. To avoid re-tracing due to the variable sequence lengths or variable\n# batch sizes (the last batch is smaller), use input_signature to specify\n# more generic shapes.\n\n# train_step_signature = [\n#     tf.TensorSpec(shape=(None, None), dtype=tf.int64),\n#     tf.TensorSpec(shape=(None, None), dtype=tf.int64),\n# ]\n\n# @tf.function(input_signature=train_step_signature)\ndef train_step(content_ids,answered_correctlys,prior_question_elapsed_times,prior_question_had_explanations,user_lecture_lvs,lagtimes,lagtime2s,answer_shifts,parts,tags1s,tags2s):\n#     tar_inp = answered_correctlys[:, :-1]\n#     tar_real = answered_correctlys[:, 1:]    \n    tar_inp = answer_shifts\n    tar_real = answered_correctlys\n    #sigmoid,just last tar\n    #tar_real=tar[: ,-1]\n    encode_inputs=[content_ids,user_lecture_lvs,parts,tags1s,tags2s]\n    decode_inputs=[answered_correctlys,prior_question_elapsed_times,prior_question_had_explanations,lagtimes,lagtime2s,answer_shifts]\n                                      \n     \n    mask = tf.math.logical_not(tf.math.equal(content_ids, 0))#content_ids[:, 1:]\n\n    enc_padding_mask, combined_mask, dec_padding_mask = create_masks(content_ids, tar_inp)\n\n    with tf.GradientTape() as tape:\n        predictions, _ = transformer(encode_inputs, decode_inputs, \n                                     True, \n                                     enc_padding_mask, \n                                     combined_mask, \n                                     dec_padding_mask)\n        loss = loss_function(mask,tar_real, predictions)\n\n    gradients = tape.gradient(loss, transformer.trainable_variables)    \n    optimizer.apply_gradients(zip(gradients, transformer.trainable_variables))\n\n    train_loss(loss)\n    train_accuracy(accuracy_function(mask,tar_real, predictions))\n    \ndef val_step(content_ids,answered_correctlys,prior_question_elapsed_times,prior_question_had_explanations,user_lecture_lvs,lagtimes,lagtime2s,answer_shifts,parts,tags1s,tags2s):\n#     tar_inp = answered_correctlys[:, :-1]\n#     tar_real = answered_correctlys[:, 1:]    \n    tar_inp = answer_shifts\n    tar_real = answered_correctlys\n    #sigmoid,just last tar\n    #tar_real=tar[: ,-1]\n    encode_inputs=[content_ids,user_lecture_lvs,parts,tags1s,tags2s]\n    decode_inputs=[answered_correctlys,prior_question_elapsed_times,prior_question_had_explanations,lagtimes,lagtime2s,answer_shifts]\n                                      \n     \n    mask = tf.math.logical_not(tf.math.equal(content_ids, 0))#content_ids[:, 1:]\n\n    enc_padding_mask, combined_mask, dec_padding_mask = create_masks(content_ids, tar_inp)\n\n    \n    predictions, _ = transformer(encode_inputs, decode_inputs, \n                                 False, \n                                 enc_padding_mask, \n                                 combined_mask, \n                                 dec_padding_mask)\n    loss = loss_function(mask,tar_real, predictions)   \n\n    val_loss(loss)\n    val_accuracy(accuracy_function(mask,tar_real, predictions))","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:24.105246Z","iopub.execute_input":"2021-06-06T13:13:24.105501Z","iopub.status.idle":"2021-06-06T13:13:24.117631Z","shell.execute_reply.started":"2021-06-06T13:13:24.105477Z","shell.execute_reply":"2021-06-06T13:13:24.116404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('start train....')","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:24.119287Z","iopub.execute_input":"2021-06-06T13:13:24.119709Z","iopub.status.idle":"2021-06-06T13:13:24.133824Z","shell.execute_reply.started":"2021-06-06T13:13:24.119666Z","shell.execute_reply":"2021-06-06T13:13:24.132721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_val_accuracy=0\npersistent=0\n\nif train_flag:\n    for epoch in range(EPOCHS):\n        start = time.time()\n\n        train_loss.reset_states()\n        train_accuracy.reset_states()\n        val_loss.reset_states()\n        val_accuracy.reset_states()\n\n        # inp -> portuguese, tar -> english\n        #\n        for (batch, (content_ids,answered_correctlys,prior_question_elapsed_times,prior_question_had_explanations,user_lecture_lvs,lagtimes,lagtime2s,answer_shifts,parts,tags1s,tags2s)) in enumerate(train_dataset):\n            train_step(content_ids,answered_correctlys,prior_question_elapsed_times,prior_question_had_explanations,user_lecture_lvs,lagtimes,lagtime2s,answer_shifts,parts,tags1s,tags2s)\n\n            if batch % 5 == 0:\n                print ('Epoch {} Batch {} train_Loss {:.4f} train_accuracy {:.4f}'.format(\n                  epoch + 1, batch, train_loss.result(), train_accuracy.result()))\n\n        if (epoch + 1) % 5 == 0:\n            ckpt_save_path = ckpt_manager.save()\n            print ('Saving checkpoint for epoch {} at {}'.format(epoch+1,\n                                                                 ckpt_save_path))\n\n        for (batch,(v_content_ids,v_answered_correctlys,v_prior_question_elapsed_times,v_prior_question_had_explanations,v_user_lecture_lvs,v_lagtimes,v_lagtime2s,v_answer_shifts,v_parts,v_tags1s,v_tags2s)) in enumerate(valid_dataset):\n            val_step(v_content_ids,v_answered_correctlys,v_prior_question_elapsed_times,v_prior_question_had_explanations,v_user_lecture_lvs,v_lagtimes,v_lagtime2s,v_answer_shifts,v_parts,v_tags1s,v_tags2s)\n            if batch % 5 == 0:\n                print ('Epoch {} Batch {}  val_Loss {:.4f} val_accuracy {:.4f}'.format(\n                  epoch + 1, batch,val_loss.result(), val_accuracy.result()))\n\n        print ('Epoch {} train_Loss {:.4f} train_accuracy {:.4f} val_Loss {:.4f} val_accuracy {:.4f}'.format(epoch + 1, \n                                                    train_loss.result(), \n                                                    train_accuracy.result(),\n                                                    val_loss.result(), \n                                                    val_accuracy.result()))\n        if val_accuracy.result()>best_val_accuracy:\n            best_val_accuracy=val_accuracy.result()\n            persistent=0\n            print('best val_accuracy:',best_val_accuracy)\n        else:\n            persistent+=1\n            print('val_accuracy not improve:',persistent)\n        if persistent>=4:\n            print('earling stop...................')\n            break\n\n        print ('Time taken for 1 epoch: {} secs\\n'.format(time.time() - start))","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:24.135376Z","iopub.execute_input":"2021-06-06T13:13:24.135748Z","iopub.status.idle":"2021-06-06T13:13:24.149924Z","shell.execute_reply.started":"2021-06-06T13:13:24.135708Z","shell.execute_reply":"2021-06-06T13:13:24.149039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Inference**","metadata":{}},{"cell_type":"code","source":"auc_sum=0\nfor (batch,(content_ids,answered_correctlys,prior_question_elapsed_times,prior_question_had_explanations,user_lecture_lvs,lagtimes,lagtime2s,answer_shifts,parts,tags1s,tags2s)) in enumerate(valid_dataset):\n    print(batch)\n    label=answered_correctlys[: ,-1]\n    label\n    #input\n    tar_inp = answer_shifts\n    encode_inputs=[content_ids,user_lecture_lvs,parts,tags1s,tags2s]\n    decode_inputs=[answered_correctlys,prior_question_elapsed_times,prior_question_had_explanations,lagtimes,lagtime2s,answer_shifts]\n    enc_padding_mask, combined_mask, dec_padding_mask = create_masks(\n            content_ids, tar_inp)\n    predictions, attention_weights = transformer(encode_inputs, \n                                                     decode_inputs,\n                                                     False,\n                                                     enc_padding_mask,\n                                                     combined_mask,\n                                                     dec_padding_mask)\n    predictions = predictions[: ,-1:, :]  # (batch_size, 1, vocab_size)\n    predictions=predictions[:,:,-2]\n    auc=roc_auc_score(label, predictions)\n    auc_sum+=auc\n    print('valid-auc:',auc,'auc_avg:',auc_sum/(batch+1) )\n    if batch ==10:\n        break","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:13:24.151130Z","iopub.execute_input":"2021-06-06T13:13:24.151603Z","iopub.status.idle":"2021-06-06T13:14:47.118773Z","shell.execute_reply.started":"2021-06-06T13:13:24.151570Z","shell.execute_reply":"2021-06-06T13:14:47.117959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_dataset = train_df.groupby('user_id').apply(lambda r:[\n#             r['content_id'].values.tolist(),\n#             r['answered_correctly'].values.tolist(),\n#             r['prior_question_elapsed_time'].values.tolist(),\n#             r['prior_question_had_explanation'].values.tolist(),\n#             r['user_lecture_lv'].values.tolist(),\n#             r['lagtime'].values.tolist(),\n#             r['lagtime2'].values.tolist(),\n#             r['answer_shift'].values.tolist(),\n#             r['part'].values.tolist(),\n#             r['tags1'].values.tolist(),\n#             r['tags2'].values.tolist()])","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:14:47.119741Z","iopub.execute_input":"2021-06-06T13:14:47.120015Z","iopub.status.idle":"2021-06-06T13:14:47.124198Z","shell.execute_reply.started":"2021-06-06T13:14:47.119989Z","shell.execute_reply":"2021-06-06T13:14:47.123083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def getpredatas(test_df):   \n    content_ids=[]\n    answered_correctlys=[]    \n    prior_question_elapsed_times=[]\n    prior_question_had_explanations=[]\n    user_lecture_lvs=[]\n    lagtimes=[]\n    lagtime2s=[]\n    answer_shifts=[]\n    parts=[]\n    tags1s=[]\n    tags2s=[]\n    \n    test_user_ids = test_df['user_id'].values\n    test_content_ids = test_df['content_id'].values\n    test_elapsed_times = test_df['prior_question_elapsed_time'].values\n    test_had_explanations = test_df['prior_question_had_explanation'].values\n    test_user_lecture_lvs = test_df['user_lecture_lv'].values\n    test_lagtimes = test_df['lagtime'].values\n    test_lagtime2s = test_df['lagtime2'].values\n    #test_lagtime3s = test_df['lagtime3'].values\n    test_parts = test_df['part'].values\n    test_tags1s = test_df['tags1'].values\n    test_tags2s = test_df['tags2'].values\n    for user_id, content_id,prior_question_elapsed_time,prior_question_had_explanation,lecture_lv,lagtime,lagtime2,part,tags1,tags2 in zip(test_user_ids,test_content_ids,test_elapsed_times,test_had_explanations,test_user_lecture_lvs,test_lagtimes,test_lagtime2s,test_parts,test_tags1s,test_tags2s):\n        if user_id in train_dataset.keys():\n            test_content=train_dataset[user_id][0]\n            test_content=np.concatenate((test_content,[content_id+1]))            \n            \n            #answered_correctly has no future.\n            test_answered_correctly=train_dataset[user_id][1]\n            test_answered_correctly=np.concatenate((test_answered_correctly,[2]))#no use          \n            \n            test_elapsed_time=train_dataset[user_id][2]\n            test_elapsed_time=np.concatenate((test_elapsed_time,[prior_question_elapsed_time]))\n            #\n            \n            test_had_explanation=train_dataset[user_id][3]\n            test_had_explanation=np.concatenate((test_had_explanation,[prior_question_had_explanation]))           \n            \n            test_user_lecture_lv=train_dataset[user_id][4]\n            test_user_lecture_lv=np.concatenate((test_user_lecture_lv,[lecture_lv]))           \n            \n            test_lagtime=train_dataset[user_id][5]         \n            test_lagtime=np.concatenate((test_lagtime,[lagtime]))          \n            \n            test_lagtime2=train_dataset[user_id][6]         \n            test_lagtime2=np.concatenate((test_lagtime2,[lagtime2]))           \n            \n            test_answer_shifts=train_dataset[user_id][7]# when prior_test_df, update                    \n            \n            #,part,tags1,tags2\n            test_part=train_dataset[user_id][8]         \n            test_part=np.concatenate((test_part,[part]))            \n            \n            test_tags1=train_dataset[user_id][9]         \n            test_tags1=np.concatenate((test_tags1,[tags1]))           \n            \n            test_tags2=train_dataset[user_id][10]         \n            test_tags2=np.concatenate((test_tags2,[tags2]))            \n        else:\n            #test_content = np.zeros(MAX_LENGTH-1, dtype=int)\n            test_content=np.array([content_id+1])\n            \n            #test_answered_correctly = np.zeros(MAX_LENGTH-1, dtype=int)\n            test_answered_correctly=np.array([2])#no use\n            \n            #test_elapsed_time = np.zeros(MAX_LENGTH-1, dtype=int)\n            test_elapsed_time=np.array([prior_question_elapsed_time])\n            \n            #test_had_explanation = np.zeros(MAX_LENGTH-1, dtype=int)\n            test_had_explanation=np.array([prior_question_had_explanation])            \n            \n            #test_user_lecture_lv = np.zeros(MAX_LENGTH-1, dtype=int)\n            test_user_lecture_lv=np.array([lecture_lv])\n            \n            #test_lagtime = np.zeros(MAX_LENGTH-1, dtype=int)\n            test_lagtime=np.array([lagtime])\n            \n            #test_lagtime2 = np.zeros(MAX_LENGTH-1, dtype=int)\n            test_lagtime2=np.array([lagtime2])\n            \n            \n            test_answer_shifts=np.array([2])\n            \n            #test_part = np.zeros(MAX_LENGTH-1, dtype=int)\n            test_part=np.array([part])\n            \n            #test_tags1 = np.zeros(MAX_LENGTH-1, dtype=int)\n            test_tags1=np.array([tags1])\n            \n            #test_tags2 = np.zeros(MAX_LENGTH-1, dtype=int)\n            test_tags2=np.array([tags2])\n        \n        train_dataset[user_id]=[test_content,test_answered_correctly,test_elapsed_time,test_had_explanation,test_user_lecture_lv,test_lagtime,test_lagtime2,test_answer_shifts,test_part,test_tags1,test_tags2]\n        \n        test_content=getlimitdata(test_content)        \n        content_ids.append(test_content)      \n        \n        #answered_correctly has no future.so max-1  dimension\n        test_answered_correctly=getlimitdata(test_answered_correctly)\n        answered_correctlys.append(test_answered_correctly)\n        \n        test_elapsed_time=getlimitdata2(test_elapsed_time,elapsed_time_count-1)        \n        prior_question_elapsed_times.append(test_elapsed_time)\n        \n        test_had_explanation=getlimitdata2(test_had_explanation,2)\n        prior_question_had_explanations.append(test_had_explanation)\n        \n        test_user_lecture_lv=getlimitdata(test_user_lecture_lv)\n        user_lecture_lvs.append(test_user_lecture_lv)\n        \n        test_lagtime=getlimitdata2(test_lagtime,lagtime_count-1)\n        lagtimes.append(test_lagtime)\n            \n        test_lagtime2=getlimitdata2(test_lagtime2,lagtime_count-1)\n        lagtime2s.append(test_lagtime2)\n            \n        test_answer_shifts=getlimitdata2(test_answer_shifts,2)\n        answer_shifts.append(test_answer_shifts)\n            \n        test_part=getlimitdata(test_part)\n        parts.append(test_part)\n            \n        test_tags1=getlimitdata(test_tags1)\n        tags1s.append(test_tags1)\n            \n        test_tags2=getlimitdata(test_tags2)\n        tags2s.append(test_tags2)\n            \n        \n    \n    content_ids=tf.convert_to_tensor(content_ids)\n    user_lecture_lvs=tf.convert_to_tensor(user_lecture_lvs)\n    parts=tf.convert_to_tensor(parts)\n    tags1s=tf.convert_to_tensor(tags1s)\n    tags2s=tf.convert_to_tensor(tags2s)\n    answered_correctlys=tf.convert_to_tensor(answered_correctlys)\n    prior_question_elapsed_times=tf.convert_to_tensor(prior_question_elapsed_times)\n    prior_question_had_explanations=tf.convert_to_tensor(prior_question_had_explanations)\n    lagtimes=tf.convert_to_tensor(lagtimes)\n    lagtime2s=tf.convert_to_tensor(lagtime2s)\n    answer_shifts=tf.convert_to_tensor(answer_shifts)\n\n    encode_inputs=[content_ids,user_lecture_lvs,parts,tags1s,tags2s]\n    decode_inputs=[answered_correctlys,prior_question_elapsed_times,prior_question_had_explanations,lagtimes,lagtime2s,answer_shifts]\n    return encode_inputs,decode_inputs","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:14:47.125954Z","iopub.execute_input":"2021-06-06T13:14:47.126328Z","iopub.status.idle":"2021-06-06T13:14:47.154658Z","shell.execute_reply.started":"2021-06-06T13:14:47.126287Z","shell.execute_reply":"2021-06-06T13:14:47.153592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encode_inputs,decode_inputs =getpredatas(test_df[0:10])\ncontent_ids=encode_inputs[0]\ntar_inp=decode_inputs[5]\nenc_padding_mask, combined_mask, dec_padding_mask = create_masks(\n        content_ids, tar_inp)\npredictions, attention_weights = transformer(encode_inputs, \n                                                 decode_inputs,\n                                                 False,\n                                                 enc_padding_mask,\n                                                 combined_mask,\n                                                 dec_padding_mask)","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:14:47.156139Z","iopub.execute_input":"2021-06-06T13:14:47.156429Z","iopub.status.idle":"2021-06-06T13:14:47.415874Z","shell.execute_reply.started":"2021-06-06T13:14:47.156402Z","shell.execute_reply":"2021-06-06T13:14:47.414892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = predictions[: ,-1:, :]  # (batch_size, 1, vocab_size)\npredictions=predictions[:,:,-2]\nnp.squeeze(predictions.numpy())","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:14:47.417354Z","iopub.execute_input":"2021-06-06T13:14:47.417746Z","iopub.status.idle":"2021-06-06T13:14:47.425998Z","shell.execute_reply.started":"2021-06-06T13:14:47.417707Z","shell.execute_reply":"2021-06-06T13:14:47.425209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import riiideducation\nenv = riiideducation.make_env()\niter_test = env.iter_test()\nprior_test_df = None","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:14:47.427092Z","iopub.execute_input":"2021-06-06T13:14:47.427377Z","iopub.status.idle":"2021-06-06T13:14:47.465952Z","shell.execute_reply.started":"2021-06-06T13:14:47.427335Z","shell.execute_reply":"2021-06-06T13:14:47.465051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ntarget = 'answered_correctly'\nfor (test_df, sample_prediction_df) in iter_test: \n    if prior_test_df is not None:\n        prior_test_df[target] = eval(test_df['prior_group_answers_correct'].iloc[0])\n        prior_test_df = prior_test_df[prior_test_df[target] != -1].reset_index(drop=True)\n        user_ids = prior_test_df['user_id'].values\n        content_ids = prior_test_df['content_id'].values\n        targets = prior_test_df[target].values \n        for user_id, answered_correctly,content_id in zip(user_ids,targets,content_ids):\n            #when train ,content_id add 1\n            if user_id in train_dataset.keys():\n                #train_dataset[user_id][0]=np.concatenate((train_dataset[user_id][0],[content_id+1]))\n                train_dataset[user_id][1]=np.concatenate((train_dataset[user_id][1],[answered_correctly]))\n#             else:\n# #                 train_dataset[user_id][0]=[content_id+1]\n# #                 train_dataset[user_id][1]=[answered_correctly]\n#                 train_df[user_id]=[np.array([content_id+1]),np.array([answered_correctly])]\n             \n            \n        \n    prior_test_df = test_df.copy()\n    \n    question_len=len( test_df[test_df['content_type_id'] == 0])\n    \n    user_lecture_sum = np.zeros(question_len, dtype=np.int16)\n    user_lecture_count = np.zeros(question_len, dtype=np.int16)\n    lagtime = np.zeros(question_len, dtype=np.float32)\n    lagtime2 = np.zeros(question_len, dtype=np.float32)\n    lagtime3 = np.zeros(question_len, dtype=np.float32)\n    \n    i=0\n    for j, (user_id,content_type_id,timestamp) in enumerate(zip(test_df['user_id'].values,test_df['content_type_id'].values,test_df['timestamp'].values)):\n        user_lecture_sum_dict[user_id] += content_type_id\n        user_lecture_count_dict[user_id] += 1\n        if(content_type_id==0):#question    \n            user_lecture_sum[i] = user_lecture_sum_dict[user_id]\n            user_lecture_count[i] = user_lecture_count_dict[user_id]\n            \n            if user_id in max_timestamp_u_dict['max_time_stamp'].keys():\n                lagtime[i]=timestamp-max_timestamp_u_dict['max_time_stamp'][user_id]\n                if(max_timestamp_u_dict2['max_time_stamp2'][user_id]==lagtime_mean2):#第二次也要赋值平均值\n                    lagtime2[i]=lagtime_mean2\n                    lagtime3[i]=lagtime_mean3\n                    #max_timestamp_u_dict3['max_time_stamp3'].update({user_id:lagtime_mean3})\n                else:\n                    lagtime2[i]=timestamp-max_timestamp_u_dict2['max_time_stamp2'][user_id]\n                    if(max_timestamp_u_dict3['max_time_stamp3'][user_id]==lagtime_mean3):\n                        lagtime3[i]=lagtime_mean3 #lagtime_mean3第3次也要赋值平均值\n                    else:\n                        lagtime3[i]=timestamp-max_timestamp_u_dict3['max_time_stamp3'][user_id]\n                    \n                    max_timestamp_u_dict3['max_time_stamp3'][user_id]=max_timestamp_u_dict2['max_time_stamp2'][user_id]\n                        \n                max_timestamp_u_dict2['max_time_stamp2'][user_id]=max_timestamp_u_dict['max_time_stamp'][user_id]\n                max_timestamp_u_dict['max_time_stamp'][user_id]=timestamp\n#                 lagtime_means[i]=(lagtime_mean_dict[user_id]+lagtime[i])/2\n#                 lagtime_mean_dict[user_id]=lagtime_means[i]\n            else:\n                lagtime[i]=lagtime_mean\n                max_timestamp_u_dict['max_time_stamp'].update({user_id:timestamp})\n                lagtime2[i]=lagtime_mean2#第一次赋值平均值\n                max_timestamp_u_dict2['max_time_stamp2'].update({user_id:lagtime_mean2})\n                lagtime3[i]=lagtime_mean3#第一次赋值平均值\n                max_timestamp_u_dict3['max_time_stamp3'].update({user_id:lagtime_mean3})\n            i=i+1 \n    \n    \n    test_df = test_df[test_df['content_type_id'] == 0].reset_index(drop=True)\n    test_df['prior_question_had_explanation'].fillna(False, inplace=True)\n    test_df.prior_question_had_explanation=test_df.prior_question_had_explanation.astype('int8')\n    \n    test_df['user_lecture_lv'] = user_lecture_sum\n    test_df[\"lagtime\"]=lagtime\n    test_df[\"lagtime2\"]=lagtime2\n    test_df[\"lagtime3\"]=lagtime3\n    \n    test_df['lagtime']=test_df['lagtime']/(10000)\n    test_df.lagtime[test_df.lagtime>20000]=20000\n    test_df['lagtime'].fillna(0, inplace=True)\n    test_df.lagtime=test_df.lagtime.astype('int16')\n\n    test_df['lagtime2']=test_df['lagtime2']/(10000)\n    test_df.lagtime2[test_df.lagtime2>20000]=20000\n    test_df['lagtime2'].fillna(0, inplace=True)\n    test_df.lagtime2=test_df.lagtime2.astype('int16')\n\n    test_df['prior_question_elapsed_time'].fillna(prior_question_elapsed_time_mean, inplace=True)\n    test_df['prior_question_elapsed_time']=test_df['prior_question_elapsed_time']/(1000)\n    test_df.prior_question_elapsed_time=test_df.prior_question_elapsed_time.astype('int16')\n    \n#     user_ids = test_df['user_id'].values\n#     content_ids = test_df['content_id'].values\n    \n    test_df=test_df.merge(questions_df.loc[questions_df.index.isin(test_df['content_id'])],\n                  how='left', on='content_id', right_index=True)\n \n    encode_inputs,decode_inputs=getpredatas(test_df)\n    content_ids=encode_inputs[0]\n    tar_inp=decode_inputs[5]\n\n    enc_padding_mask, combined_mask, dec_padding_mask = create_masks(\n        content_ids, tar_inp)\n    predictions, attention_weights = transformer(encode_inputs, \n                                                 decode_inputs,\n                                                 False,\n                                                 enc_padding_mask,\n                                                 combined_mask,\n                                                 dec_padding_mask)\n    predictions = predictions[: ,-1:, :]  # (batch_size, 1, vocab_size)\n    predictions=predictions[:,:,-2]\n    \n    test_df[target]=np.squeeze(predictions.numpy())\n       \n    #    \n    env.predict(test_df[['row_id', target]])\n    ","metadata":{"execution":{"iopub.status.busy":"2021-06-06T13:14:47.467641Z","iopub.execute_input":"2021-06-06T13:14:47.468069Z","iopub.status.idle":"2021-06-06T13:14:48.645224Z","shell.execute_reply.started":"2021-06-06T13:14:47.468026Z","shell.execute_reply":"2021-06-06T13:14:48.644355Z"},"trusted":true},"execution_count":null,"outputs":[]}]}