{"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":"Here is a sample code for gru4rec using recbole. \nI made it based on the following code from a past competition\nhttps://www.kaggle.com/code/astrung/recbole-lstm-sequential-for-recomendation-tutorial\n\nI think you can get a better score if you change the training data or epoch size, etc. Enjoy!","metadata":{}},{"cell_type":"code","source":"!pip install polars","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:34:40.973409Z","iopub.execute_input":"2023-05-11T10:34:40.974118Z","iopub.status.idle":"2023-05-11T10:34:58.396455Z","shell.execute_reply.started":"2023-05-11T10:34:40.974012Z","shell.execute_reply":"2023-05-11T10:34:58.395207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install recbole\n!pip install torch==1.8.0\n# !pip install --upgrade torch","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:34:58.399919Z","iopub.execute_input":"2023-05-11T10:34:58.401036Z","iopub.status.idle":"2023-05-11T10:36:24.160622Z","shell.execute_reply.started":"2023-05-11T10:34:58.400991Z","shell.execute_reply":"2023-05-11T10:36:24.15942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\ntorch.__version__","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:36:28.860955Z","iopub.execute_input":"2023-05-11T10:36:28.861957Z","iopub.status.idle":"2023-05-11T10:36:28.887475Z","shell.execute_reply.started":"2023-05-11T10:36:28.861915Z","shell.execute_reply":"2023-05-11T10:36:28.886452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tqdm\nimport polars as pl\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport random\nimport os \nimport h5py\nimport sys\nimport gc\n\nfrom matplotlib import pyplot as plt\nimport pyarrow.parquet as pq","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:36:33.647431Z","iopub.execute_input":"2023-05-11T10:36:33.648065Z","iopub.status.idle":"2023-05-11T10:36:34.830902Z","shell.execute_reply.started":"2023-05-11T10:36:33.648012Z","shell.execute_reply":"2023-05-11T10:36:34.829815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 1. Create atomic file","metadata":{}},{"cell_type":"code","source":"train = pl.read_parquet('/kaggle/input/otto-train-and-test-data-for-local-validation/test.parquet')\ntest = pl.read_parquet('/kaggle/input/otto-full-optimized-memory-footprint/test.parquet')\n\ndf = pl.concat([train, test])\n#df = pl.read_parquet('../input/otto-train-and-test-data-for-local-validation/test.parquet')\n\ndf = df.sort(['session', 'aid', 'ts'])\ndf = df.with_columns((pl.col('ts') * 1e9).alias('ts'))\ndf = df.rename({'session': 'session:token', 'aid': 'aid:token', 'ts': 'ts:float'})\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:36:35.546023Z","iopub.execute_input":"2023-05-11T10:36:35.546494Z","iopub.status.idle":"2023-05-11T10:36:39.092856Z","shell.execute_reply.started":"2023-05-11T10:36:35.54645Z","shell.execute_reply":"2023-05-11T10:36:39.091824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir /kaggle/working/recbox_data\ndf.select([pl.col('session:token','aid:token','ts:float')]).write_csv('/kaggle/working/recbox_data/recbox_data.inter', separator='\\t')\n\ndel df, train, test\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:36:39.094938Z","iopub.execute_input":"2023-05-11T10:36:39.095334Z","iopub.status.idle":"2023-05-11T10:36:44.666011Z","shell.execute_reply.started":"2023-05-11T10:36:39.095294Z","shell.execute_reply":"2023-05-11T10:36:44.664856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3. Create dataset and train model with Recbole\n\nFor anyone need instruction document, please check this link: https://recbole.io/docs/user_guide/usage/use_modules.html","metadata":{}},{"cell_type":"code","source":"import logging\nfrom logging import getLogger\nfrom recbole.config import Config\nfrom recbole.data import create_dataset, data_preparation\nfrom recbole.model.sequential_recommender import GRU4Rec\nfrom recbole.trainer import Trainer\nfrom recbole.utils import init_seed, init_logger\n\nfrom recbole.utils.case_study import full_sort_topk","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:36:44.668266Z","iopub.execute_input":"2023-05-11T10:36:44.668678Z","iopub.status.idle":"2023-05-11T10:36:45.178597Z","shell.execute_reply.started":"2023-05-11T10:36:44.668634Z","shell.execute_reply":"2023-05-11T10:36:45.177558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#64维\nMAX_ITEM = 20\n\nparameter_dict = {\n    'data_path': '/kaggle/working/',\n    'USER_ID_FIELD': 'session',\n    'ITEM_ID_FIELD': 'aid',\n    'TIME_FIELD': 'ts',\n    'user_inter_num_interval': \"[5,Inf)\",\n    'item_inter_num_interval': \"[5,Inf)\",\n    'load_col': {'inter': ['session', 'aid', 'ts']},\n    'train_neg_sample_args': None,\n    'learning_rate': 0.002,\n    'epochs': 15,\n    'stopping_step':3,\n    'embedding_size':64,\n    'eval_batch_size': 1024,\n    #'train_batch_size': 1024,\n#    'enable_amp':True,\n    'MAX_ITEM_LIST_LENGTH': MAX_ITEM,\n    'eval_args': {\n        'split': {'RS': [9, 1, 0]},\n        'group_by': 'user',\n        'order': 'TO',\n        'mode': 'full'}\n}\n\nconfig = Config(model='GRU4Rec', dataset='recbox_data', config_dict=parameter_dict)\n\n# init random seed\ninit_seed(config['seed'], config['reproducibility'])\n\n# logger initialization\ninit_logger(config)\nlogger = getLogger()\n\n# Create handlers\nc_handler = logging.StreamHandler()\nc_handler.setLevel(logging.INFO)\nlogger.addHandler(c_handler)\n\n# write config info into log\nlogger.info(config)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:11:33.831801Z","iopub.execute_input":"2023-05-11T10:11:33.832495Z","iopub.status.idle":"2023-05-11T10:11:34.278816Z","shell.execute_reply.started":"2023-05-11T10:11:33.832458Z","shell.execute_reply":"2023-05-11T10:11:34.278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#32维\nMAX_ITEM = 20\n\nparameter_dict = {\n    'data_path': '/kaggle/working/',\n    'USER_ID_FIELD': 'session',\n    'ITEM_ID_FIELD': 'aid',\n    'TIME_FIELD': 'ts',\n    'user_inter_num_interval': \"[5,Inf)\",\n    'item_inter_num_interval': \"[5,Inf)\",\n    'load_col': {'inter': ['session', 'aid', 'ts']},\n    'train_neg_sample_args': None,\n    'learning_rate': 0.002,\n    'epochs': 15,\n    'stopping_step':3,\n    'embedding_size':32,\n    'eval_batch_size': 1024,\n    #'train_batch_size': 1024,\n#    'enable_amp':True,\n    'MAX_ITEM_LIST_LENGTH': MAX_ITEM,\n    'eval_args': {\n        'split': {'RS': [9, 1, 0]},\n        'group_by': 'user',\n        'order': 'TO',\n        'mode': 'full'}\n}\n\nconfig = Config(model='GRU4Rec', dataset='recbox_data', config_dict=parameter_dict)\n\n# init random seed\ninit_seed(config['seed'], config['reproducibility'])\n\n# logger initialization\ninit_logger(config)\nlogger = getLogger()\n\n# Create handlers\nc_handler = logging.StreamHandler()\nc_handler.setLevel(logging.INFO)\nlogger.addHandler(c_handler)\n\n# write config info into log\nlogger.info(config)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T02:55:06.124881Z","iopub.execute_input":"2023-05-11T02:55:06.125288Z","iopub.status.idle":"2023-05-11T02:55:06.680169Z","shell.execute_reply.started":"2023-05-11T02:55:06.125252Z","shell.execute_reply":"2023-05-11T02:55:06.679209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#128维\nMAX_ITEM = 20\n\nparameter_dict = {\n    'data_path': '/kaggle/working/',\n    'USER_ID_FIELD': 'session',\n    'ITEM_ID_FIELD': 'aid',\n    'TIME_FIELD': 'ts',\n    'user_inter_num_interval': \"[5,Inf)\",\n    'item_inter_num_interval': \"[5,Inf)\",\n    'load_col': {'inter': ['session', 'aid', 'ts']},\n    'train_neg_sample_args': None,\n    'learning_rate': 0.001,\n    'epochs': 15,\n    'stopping_step': 3,\n    'num_layers': 1,\n    'embedding_size':128,\n    'eval_batch_size': 1024,\n    #'train_batch_size': 1024,\n#    'enable_amp':True,\n    'MAX_ITEM_LIST_LENGTH': MAX_ITEM,\n    'eval_args': {\n        'split': {'RS': [9, 1, 0]},\n        'group_by': 'user',\n        'order': 'TO',\n        'mode': 'full'}\n}\n\nconfig = Config(model='GRU4Rec', dataset='recbox_data', config_dict=parameter_dict)\n\n# init random seed\ninit_seed(config['seed'], config['reproducibility'])\n\n# logger initialization\ninit_logger(config)\nlogger = getLogger()\n\n# Create handlers\nc_handler = logging.StreamHandler()\nc_handler.setLevel(logging.INFO)\nlogger.addHandler(c_handler)\n\n# write config info into log\nlogger.info(config)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T09:37:59.641251Z","iopub.execute_input":"2023-05-11T09:37:59.641664Z","iopub.status.idle":"2023-05-11T09:38:04.543737Z","shell.execute_reply.started":"2023-05-11T09:37:59.641627Z","shell.execute_reply":"2023-05-11T09:38:04.543043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = create_dataset(config)\nlogger.info(dataset)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:37:05.760153Z","iopub.execute_input":"2023-05-11T10:37:05.760609Z","iopub.status.idle":"2023-05-11T10:40:15.799196Z","shell.execute_reply.started":"2023-05-11T10:37:05.760567Z","shell.execute_reply":"2023-05-11T10:40:15.798395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dataset splitting\ntrain_data, valid_data, test_data = data_preparation(config, dataset)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:40:15.800762Z","iopub.execute_input":"2023-05-11T10:40:15.801162Z","iopub.status.idle":"2023-05-11T10:42:38.922558Z","shell.execute_reply.started":"2023-05-11T10:40:15.801121Z","shell.execute_reply":"2023-05-11T10:42:38.921758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # 64维\n# model = GRU4Rec(config, train_data.dataset).to(config['device'])\n# logger.info(model)\n\n# # trainer loading and initialization\n# trainer = Trainer(config, model)\n\n# # model training\n# best_valid_score, best_valid_result = trainer.fit(train_data, valid_data)\n# #best_valid_score, best_valid_result = trainer.fit(train_data)","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-05-10T14:02:05.504842Z","iopub.execute_input":"2023-05-10T14:02:05.505174Z","iopub.status.idle":"2023-05-10T16:07:22.61204Z","shell.execute_reply.started":"2023-05-10T14:02:05.505141Z","shell.execute_reply":"2023-05-10T16:07:22.61106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#32维\nmodel = GRU4Rec(config, train_data.dataset).to(config['device'])\nlogger.info(model)\n\n# trainer loading and initialization\ntrainer = Trainer(config, model)\n\n# model training\nbest_valid_score, best_valid_result = trainer.fit(train_data, valid_data)\n#best_valid_score, best_valid_result = trainer.fit(train_data)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T03:05:32.111877Z","iopub.execute_input":"2023-05-11T03:05:32.112329Z","iopub.status.idle":"2023-05-11T04:53:38.165433Z","shell.execute_reply.started":"2023-05-11T03:05:32.112291Z","shell.execute_reply":"2023-05-11T04:53:38.164227Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 128维\nmodel = GRU4Rec(config, train_data.dataset).to(config['device'])\nlogger.info(model)\n\n# trainer loading and initialization\ntrainer = Trainer(config, model)\n\n# model training\nbest_valid_score, best_valid_result = trainer.fit(train_data, valid_data)\n#best_valid_score, best_valid_result = trainer.fit(train_data)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T05:12:38.625299Z","iopub.execute_input":"2023-05-11T05:12:38.625708Z","iopub.status.idle":"2023-05-11T06:53:54.779694Z","shell.execute_reply.started":"2023-05-11T05:12:38.625652Z","shell.execute_reply":"2023-05-11T06:53:54.778719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import torch\nfrom recbole.model.abstract_recommender import GeneralRecommender\nclass TransformRec(GeneralRecommender):\n    def __init__(self, config, dataset):\n        super(TransformRec, self).__init__(config, dataset)\n        self.embedding_size = config['embedding_size']\n        self.max_seq_length = config['max_seq_length']\n        self.num_heads = config['num_heads']\n        self.num_layers = config['num_layers']\n        self.ff_num_hidden_units = config['ff_num_hidden_units']\n        self.dropout = config['dropout']\n\n        self.embedding_layer = torch.nn.Embedding(315004, self.embedding_size)\n        self.transformer_layer = torch.nn.TransformerEncoder(\n            layer=torch.nn.TransformerEncoderLayer(d_model=self.embedding_size,\n                                                    nhead=self.num_heads,\n                                                    dim_feedforward=self.ff_num_hidden_units,\n                                                    dropout=self.dropout),\n            num_layers=self.num_layers\n        )\n        self.fc_layer = torch.nn.Linear(self.embedding_size, 315004)\n\n    def forward(self, user, item_seq):\n        # item_seq shape: [batch_size, seq_len]\n        input_seq = self.embedding_layer(item_seq)  # shape: [batch_size, seq_len, embedding_size]\n        input_seq = input_seq.permute(1, 0, 2)  # shape: [seq_len, batch_size, embedding_size]\n\n        output_seq = self.transformer_layer(input_seq)  # shape: [seq_len, batch_size, embedding_size]\n\n        # 选择序列中最后一个物品做预测\n        output = output_seq[-1]  # shape: [batch_size, embedding_size]\n\n        output = self.fc_layer(output)  # shape: [batch_size, num_items]\n\n        return output\n","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:42:49.269105Z","iopub.execute_input":"2023-05-11T10:42:49.269772Z","iopub.status.idle":"2023-05-11T10:42:49.355092Z","shell.execute_reply.started":"2023-05-11T10:42:49.269722Z","shell.execute_reply":"2023-05-11T10:42:49.353821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.__version__","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:33:52.196817Z","iopub.execute_input":"2023-05-11T10:33:52.197219Z","iopub.status.idle":"2023-05-11T10:33:52.204494Z","shell.execute_reply.started":"2023-05-11T10:33:52.197183Z","shell.execute_reply":"2023-05-11T10:33:52.203333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config['ff_num_hidden_units']","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:22:16.14629Z","iopub.execute_input":"2023-05-11T10:22:16.146666Z","iopub.status.idle":"2023-05-11T10:22:16.153845Z","shell.execute_reply.started":"2023-05-11T10:22:16.146635Z","shell.execute_reply":"2023-05-11T10:22:16.152882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip uninstall recbole -y\n!pip install git+https://github.com/RUCAIBox/RecBole.git@master","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:20:17.605772Z","iopub.execute_input":"2023-05-11T10:20:17.606155Z","iopub.status.idle":"2023-05-11T10:20:47.128707Z","shell.execute_reply.started":"2023-05-11T10:20:17.60611Z","shell.execute_reply":"2023-05-11T10:20:47.12739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:33:21.008092Z","iopub.execute_input":"2023-05-11T10:33:21.008942Z","iopub.status.idle":"2023-05-11T10:33:32.40899Z","shell.execute_reply.started":"2023-05-11T10:33:21.008902Z","shell.execute_reply":"2023-05-11T10:33:32.407621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset","metadata":{"execution":{"iopub.status.busy":"2023-05-11T09:57:18.704395Z","iopub.execute_input":"2023-05-11T09:57:18.705111Z","iopub.status.idle":"2023-05-11T09:57:24.327039Z","shell.execute_reply.started":"2023-05-11T09:57:18.705071Z","shell.execute_reply":"2023-05-11T09:57:24.325759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ndel trainer, train_data, valid_data, test_data","metadata":{"execution":{"iopub.status.busy":"2023-05-10T16:07:22.61365Z","iopub.execute_input":"2023-05-10T16:07:22.614355Z","iopub.status.idle":"2023-05-10T16:07:22.620344Z","shell.execute_reply.started":"2023-05-10T16:07:22.614308Z","shell.execute_reply":"2023-05-10T16:07:22.619046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-05-10T16:07:22.622245Z","iopub.execute_input":"2023-05-10T16:07:22.62281Z","iopub.status.idle":"2023-05-10T16:07:23.266644Z","shell.execute_reply.started":"2023-05-10T16:07:22.622773Z","shell.execute_reply":"2023-05-10T16:07:23.265465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt \nplt.figure(figsize=(10, 7))\nplt.plot(list(range(1, len(auc798) + 1)), auc798)\nplt.scatter(list(range(1, len(auc798) + 1)), auc798)\nplt.plot(list(range(1, len(auc835) + 1)), auc835)\nplt.scatter(list(range(1, len(auc835) + 1)), auc835)\nplt.plot(list(range(1, len(auc848) + 1)), auc848)\nplt.scatter(list(range(1, len(auc848) + 1)), auc848)\nplt.plot(list(range(1, len(auc827) + 1)), auc827)\nplt.scatter(list(range(1, len(auc827) + 1)), auc827)\nplt.xlabel('Epoch', fontsize=15)\nplt.ylabel('pr auc', fontsize=15)\nplt.legend(['dim 4','','dim 8','','dim 16','','dim 20',''])\n\nplt.title('pr auc', fontsize=20)\nplt.savefig('pr_auc_hist.png')","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:42:57.944848Z","iopub.execute_input":"2023-05-10T17:42:57.94523Z","iopub.status.idle":"2023-05-10T17:42:57.952902Z","shell.execute_reply.started":"2023-05-10T17:42:57.945195Z","shell.execute_reply":"2023-05-10T17:42:57.95163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4. Create recommendation result from trained model\n\nI note document here for any one want to customize it: https://recbole.io/docs/user_guide/usage/case_study.html","metadata":{}},{"cell_type":"code","source":"# # https://qiita.com/fufufukakaka/items/e03df3a7299b2b8f99cf\n# from typing import List, Tuple\n# import numpy as np\n# import torch\n\n# from pydantic import BaseModel\n# from recbole.data import create_dataset\n# from recbole.data.dataset.sequential_dataset import SequentialDataset\n# from recbole.data.interaction import Interaction\n# from recbole.model.sequential_recommender.sine import SINE\n# from recbole.utils import get_model, init_seed\n\n# class ItemHistory(BaseModel):\n#     sequence: List[str]\n#     topk: int\n\n# class RecommendedItems(BaseModel):\n#     score_list: List[float]\n#     item_list: List[str]\n\n\n# def pred_user_to_item(item_history: ItemHistory):\n#     item_history_dict = item_history.dict()\n#     item_sequence = item_history_dict[\"sequence\"]\n#     item_length = len(item_sequence)\n#     pad_length = MAX_ITEM  # pre-defined by recbole\n\n#     padded_item_sequence = torch.nn.functional.pad(\n#         torch.tensor(dataset.token2id(dataset.iid_field, item_sequence)),\n#         (0, pad_length - item_length),\n#         \"constant\",\n#         0,\n#     )\n\n#     input_interaction = Interaction(\n#         {\n#             \"aid_list\": padded_item_sequence.reshape(1, -1),\n#             \"item_length\": torch.tensor([item_length]),\n#         }\n#     )\n#     scores = model.full_sort_predict(input_interaction.to(model.device))\n#     scores = scores.view(-1, dataset.item_num)\n#     scores[:, 0] = -np.inf  # pad item score -> -inf\n#     topk_score, topk_iid_list = torch.topk(scores, item_history_dict[\"topk\"])\n\n#     predicted_score_list = topk_score.tolist()[0]\n#     predicted_item_list = dataset.id2token(\n#         dataset.iid_field, topk_iid_list.tolist()\n#     ).tolist()\n\n#     recommended_items = {\n#         \"score_list\": predicted_score_list,\n#         \"item_list\": predicted_item_list,\n#     }\n#     return recommended_items","metadata":{"execution":{"iopub.status.busy":"2023-05-10T16:07:23.268381Z","iopub.execute_input":"2023-05-10T16:07:23.270156Z","iopub.status.idle":"2023-05-10T16:07:23.391197Z","shell.execute_reply.started":"2023-05-10T16:07:23.270113Z","shell.execute_reply":"2023-05-10T16:07:23.390112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# #test = pl.read_parquet('../input/otto-train-and-test-data-for-local-validation/test.parquet')\n# test = pl.read_parquet('../input/otto-full-optimized-memory-footprint/test.parquet')\n\n# import pandas as pd\n# import numpy as np\n\n# from collections import defaultdict\n\n# #sample_sub = pd.read_csv('../input/otto-recommender-system//sample_submission.csv')\n\n# session_types = ['clicks', 'carts', 'orders']\n# test_session_AIDs = test.to_pandas().reset_index(drop=True).groupby('session')['aid'].apply(list)\n# test_session_types = test.to_pandas().reset_index(drop=True).groupby('session')['type'].apply(list)\n\n# del test\n# gc.collect()\n\n# labels = []\n\n# type_weight_multipliers = {0: 1, 1: 6, 2: 3}\n# for AIDs, types in zip(test_session_AIDs, test_session_types):\n#     if len(AIDs) >= 20:\n#         # if we have enough aids (over equals 20) we don't need to look for candidates! we just use the old logic\n#         weights=np.logspace(0.1,1,len(AIDs),base=2, endpoint=True)-1\n#         aids_temp=defaultdict(lambda: 0)\n#         for aid,w,t in zip(AIDs,weights,types): \n#             aids_temp[aid]+= w * type_weight_multipliers[t]\n            \n#         sorted_aids=[k for k, v in sorted(aids_temp.items(), key=lambda item: -item[1])]\n#         labels.append(sorted_aids[:20])\n#     else:\n#         AIDs = list(dict.fromkeys(AIDs))\n#         item = ItemHistory(sequence=AIDs, topk=20)\n#         try:\n#             nns = [ int(v) for v in pred_user_to_item(item)['item_list']]\n#         except:\n#             nns = []\n\n#         for word in nns:\n#             if len(AIDs) == 20:\n#                 break\n#             if int(word) not in AIDs:\n#                 AIDs.append(word)\n\n#         labels.append(AIDs[:20])","metadata":{"execution":{"iopub.status.busy":"2023-05-10T16:07:23.392838Z","iopub.execute_input":"2023-05-10T16:07:23.39321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# pred_user_to_item(item)['item_list']","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# labels_as_strings = [' '.join([str(l) for l in lls]) for lls in labels]\n# predictions = pd.DataFrame(data={'session_type': test_session_AIDs.index, 'labels': labels_as_strings})","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# labels_as_strings = [' '.join([str(l) for l in lls]) for lls in labels]\n\n# predictions = pd.DataFrame(data={'session_type': test_session_AIDs.index, 'labels': labels_as_strings})\n\n# prediction_dfs = []\n\n# for st in session_types:\n#     modified_predictions = predictions.copy()\n#     modified_predictions.session_type = modified_predictions.session_type.astype('str') + f'_{st}'\n#     prediction_dfs.append(modified_predictions)\n\n# submission = pd.concat(prediction_dfs).reset_index(drop=True)\n# submission.to_csv('submission.csv', index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}