{"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":"code","source":"import pandas as pd\n\ntrain = pd.read_parquet('../input/otto-full-optimized-memory-footprint/train.parquet')\ntest = pd.read_parquet('../input/otto-full-optimized-memory-footprint/test.parquet')\n\n!pip install pickle5\nimport pickle5 as pickle\n\nwith open('../input/otto-full-optimized-memory-footprint/id2type.pkl', \"rb\") as fh:\n    id2type = pickle.load(fh)\nwith open('../input/otto-full-optimized-memory-footprint/type2id.pkl', \"rb\") as fh:\n    type2id = pickle.load(fh)\n    \nsample_sub = pd.read_csv('../input/otto-recommender-system/sample_submission.csv')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-12-15T09:47:33.685368Z","iopub.execute_input":"2022-12-15T09:47:33.685869Z","iopub.status.idle":"2022-12-15T09:47:55.393602Z","shell.execute_reply.started":"2022-12-15T09:47:33.685833Z","shell.execute_reply":"2022-12-15T09:47:55.391996Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare Data","metadata":{}},{"cell_type":"code","source":"num_events = train.shape[0]\nprint(num_events)\nquarter_train = train[0:1000]\nprint(quarter_train.shape)","metadata":{"execution":{"iopub.status.busy":"2022-12-15T09:47:55.396906Z","iopub.execute_input":"2022-12-15T09:47:55.397460Z","iopub.status.idle":"2022-12-15T09:47:55.405224Z","shell.execute_reply.started":"2022-12-15T09:47:55.397389Z","shell.execute_reply":"2022-12-15T09:47:55.404177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\n# Concatenate 'type' column with 'aid' column\nquarter_train['type + aid'] = quarter_train['type'].astype(str) + quarter_train['aid'].astype(str)","metadata":{"execution":{"iopub.status.busy":"2022-12-15T09:47:55.406698Z","iopub.execute_input":"2022-12-15T09:47:55.407365Z","iopub.status.idle":"2022-12-15T09:47:55.429617Z","shell.execute_reply.started":"2022-12-15T09:47:55.407320Z","shell.execute_reply":"2022-12-15T09:47:55.428283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-15T09:47:55.433102Z","iopub.execute_input":"2022-12-15T09:47:55.433570Z","iopub.status.idle":"2022-12-15T09:47:55.448270Z","shell.execute_reply.started":"2022-12-15T09:47:55.433525Z","shell.execute_reply":"2022-12-15T09:47:55.447076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"session_aids_untruncated = quarter_train.groupby('session')['type + aid'].apply(lambda x: list(x)[-26:])","metadata":{"execution":{"iopub.status.busy":"2022-12-15T09:47:55.449643Z","iopub.execute_input":"2022-12-15T09:47:55.450675Z","iopub.status.idle":"2022-12-15T09:47:55.462414Z","shell.execute_reply.started":"2022-12-15T09:47:55.450627Z","shell.execute_reply":"2022-12-15T09:47:55.461030Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"session_types = ['clicks', 'carts', 'orders']","metadata":{"execution":{"iopub.status.busy":"2022-12-15T09:47:55.464519Z","iopub.execute_input":"2022-12-15T09:47:55.465046Z","iopub.status.idle":"2022-12-15T09:47:55.473380Z","shell.execute_reply.started":"2022-12-15T09:47:55.464998Z","shell.execute_reply":"2022-12-15T09:47:55.471809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"session_aids_untruncated.head()","metadata":{"execution":{"iopub.status.busy":"2022-12-15T09:47:55.474551Z","iopub.execute_input":"2022-12-15T09:47:55.474938Z","iopub.status.idle":"2022-12-15T09:47:55.491283Z","shell.execute_reply.started":"2022-12-15T09:47:55.474902Z","shell.execute_reply":"2022-12-15T09:47:55.489682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Separating into inputs and labels\nbeginning_sentence = []\nend_sentence = []\n\n# looping over each aid of each session\nfor session, aids in session_aids_untruncated.iteritems():\n    n_aids = len(aids) # number of aids in the session\n    \n    # In the case we have only one, \n    # we keep it as input and copy it as label\n    if n_aids == 1: \n        beginning_sentence.append(aids)\n        end_sentence.append(aids)\n        \n    # In the case we have less than 12, \n    # we keep half as inputs and copy the other half as labels\n    elif n_aids <= 12: \n        if(n_aids % 2 == 0): # case it is even\n            beginning_sentence.append(aids[:n_aids//2])\n            end_sentence.append(aids[n_aids//2:])\n        elif(n_aids % 2 != 0): # case it is odd\n            beginning_sentence.append(aids[:n_aids//2+1])\n            end_sentence.append(aids[n_aids//2+1:])\n    \n    # In the case we have more than 26, \n    # we only look at the last 26: we keep 6 as inputs, copy the last twenty as labels\n    elif n_aids > 26: \n        beginning_sentence.append(aids[-26:-20])\n        end_sentence.append(aids[-20:])\n        \n    # In the case we have more than 12, but less than 26, \n    # we keep 6 as inputs and copy the rest as labels\n    else: \n        beginning_sentence.append(aids[:6])\n        end_sentence.append(aids[6:])               ","metadata":{"execution":{"iopub.status.busy":"2022-12-15T09:47:55.492553Z","iopub.execute_input":"2022-12-15T09:47:55.492916Z","iopub.status.idle":"2022-12-15T09:47:55.504241Z","shell.execute_reply.started":"2022-12-15T09:47:55.492883Z","shell.execute_reply":"2022-12-15T09:47:55.502671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(session_aids_untruncated[5:10])\nprint(beginning_sentence[5:10])\nprint(end_sentence[5:10])","metadata":{"execution":{"iopub.status.busy":"2022-12-15T09:47:55.506041Z","iopub.execute_input":"2022-12-15T09:47:55.506410Z","iopub.status.idle":"2022-12-15T09:47:55.525192Z","shell.execute_reply.started":"2022-12-15T09:47:55.506378Z","shell.execute_reply":"2022-12-15T09:47:55.523659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.DataFrame(beginning_sentence).to_csv(\"/kaggle/working/begginning_sentence.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-12-15T09:47:55.527871Z","iopub.execute_input":"2022-12-15T09:47:55.528290Z","iopub.status.idle":"2022-12-15T09:47:55.539702Z","shell.execute_reply.started":"2022-12-15T09:47:55.528252Z","shell.execute_reply":"2022-12-15T09:47:55.538557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.DataFrame(end_sentence).to_csv(\"/kaggle/working/end_sentence.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-12-15T09:47:55.540954Z","iopub.execute_input":"2022-12-15T09:47:55.541683Z","iopub.status.idle":"2022-12-15T09:47:55.551612Z","shell.execute_reply.started":"2022-12-15T09:47:55.541647Z","shell.execute_reply":"2022-12-15T09:47:55.550703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.DataFrame(session_aids_untruncated).to_csv(\"/kaggle/working/session_aids_untruncated.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-12-15T09:56:17.474927Z","iopub.execute_input":"2022-12-15T09:56:17.475332Z","iopub.status.idle":"2022-12-15T09:56:17.483480Z","shell.execute_reply.started":"2022-12-15T09:56:17.475300Z","shell.execute_reply":"2022-12-15T09:56:17.482152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Encoding of the position","metadata":{}},{"cell_type":"code","source":"import numpy as np\n\ndef getPositionEncoding(seq_len, d, aid_type, n=10000):\n    factor=0.1 #(aid_type==\"clicks\")\n    if(aid_type==\"carts\"):\n        factor=0.3\n    elif(aid_type==\"orders\"):\n        factor=0.6\n    \n    P = np.zeros((seq_len, d))\n    for k in range(seq_len):\n        for i in np.arange(int(d/2)):\n            denominator = np.power(n, 2*i/d)\n            P[k, 2*i] = factor * np.sin(k/denominator)\n            P[k, 2*i+1] = factor * np.cos(k/denominator)\n    return P\n\nP = getPositionEncoding(seq_len=4, d=4, aid_type=\"carts\", n=100)\nprint(P)","metadata":{"execution":{"iopub.status.busy":"2022-12-15T11:25:38.041323Z","iopub.execute_input":"2022-12-15T11:25:38.041660Z","iopub.status.idle":"2022-12-15T11:25:38.055231Z","shell.execute_reply.started":"2022-12-15T11:25:38.041634Z","shell.execute_reply":"2022-12-15T11:25:38.054165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model definition","metadata":{}},{"cell_type":"code","source":"import torch\n\ntransformer_model = torch.nn.Transformer()","metadata":{"execution":{"iopub.status.busy":"2022-12-15T09:47:55.552860Z","iopub.execute_input":"2022-12-15T09:47:55.553367Z","iopub.status.idle":"2022-12-15T09:47:55.971383Z","shell.execute_reply.started":"2022-12-15T09:47:55.553336Z","shell.execute_reply":"2022-12-15T09:47:55.970240Z"},"trusted":true},"execution_count":null,"outputs":[]}]}