{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.7.12"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":38760,"databundleVersionId":4493939,"sourceType":"competition"},{"sourceId":4461402,"sourceType":"datasetVersion","datasetId":2611514},{"sourceId":4474043,"sourceType":"datasetVersion","datasetId":2601572}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":17967.851697,"end_time":"2024-12-31T05:17:10.878262","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-12-31T00:17:43.026565","version":"2.3.4"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"70a7dc8d","cell_type":"code","source":"# 1. Cài đặt các thư viện Data Science cơ bản\n!pip install pandas numpy matplotlib seaborn polars pyarrow tqdm h5py pydantic\n\n# 2. Cài đặt PyTorch với CUDA (Giả sử bạn cần GPU)\n!pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118\n\n# 3. Cài đặt RecBole\n!pip install recbole polars -q","metadata":{"papermill":{"duration":10.833471,"end_time":"2024-12-31T00:18:01.127656","exception":false,"start_time":"2024-12-31T00:17:50.294185","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"8f05d372","cell_type":"code","source":"!pip install protobuf==3.20.0","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"81afec0b","cell_type":"code","source":"# ====== IMPORT THƯ VIỆN CƠ BẢN ======\n# Thư viện xử lý dữ liệu, số liệu và hệ thống\nimport os\nimport sys\nimport gc\nimport random\n\nimport numpy as np\nimport pandas as pd\nimport polars as pl\n\n# Thư viện vẽ biểu đồ (nếu cần phân tích)\nfrom matplotlib import pyplot as plt\nimport seaborn as sns\n\n# IO & định dạng dữ liệu\nimport h5py\nimport pyarrow.parquet as pq\n\n# Tiến trình hiển thị thanh progress\nfrom tqdm.auto import tqdm","metadata":{"papermill":{"duration":1.069532,"end_time":"2024-12-31T00:18:11.798443","exception":false,"start_time":"2024-12-31T00:18:10.728911","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"3f85e709","cell_type":"markdown","source":"# 1. Create atomic file","metadata":{"papermill":{"duration":0.005091,"end_time":"2024-12-31T00:18:11.809315","exception":false,"start_time":"2024-12-31T00:18:11.804224","status":"completed"},"tags":[]}},{"id":"c3fa2c3b","cell_type":"code","source":"# Đọc dữ liệu train/test của OTTO và ghép lại thành một bảng duy nhất\ntrain_df = pl.read_parquet('/kaggle/input/otto-train-and-test-data-for-local-validation/test.parquet')\ntest_df = pl.read_parquet('/kaggle/input/otto-full-optimized-memory-footprint/test.parquet')\n\ninter_df = pl.concat([train_df, test_df])\n\n# Sắp xếp theo session, aid và timestamp để đảm bảo thứ tự\ninter_df = inter_df.sort(['session', 'aid', 'ts'])\n\n# RecBole yêu cầu TIME_FIELD là số (float) và thường dùng ns, ở đây giữ nguyên *1e9 như bản gốc\ninter_df = inter_df.with_columns((pl.col('ts') * 1e9).alias('ts'))\n\n# Đổi tên cột theo đúng format RecBole\ninter_df = inter_df.rename({'session': 'session:token', 'aid': 'aid:token', 'ts': 'ts:float'})","metadata":{"papermill":{"duration":1.978699,"end_time":"2024-12-31T00:18:13.793382","exception":false,"start_time":"2024-12-31T00:18:11.814683","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"df19f053","cell_type":"code","source":"import os\n\n# Đường dẫn thư mục\ndirectory = \"/kaggle/working/recbox_data\"\n\n# Kiểm tra và tạo thư mục\nif not os.path.exists(directory):\n    os.makedirs(directory)\n    print(f\"Đã tạo thư mục: {directory}\")\nelse:\n    print(f\"Thư mục đã tồn tại: {directory}\")\n","metadata":{"papermill":{"duration":0.013766,"end_time":"2024-12-31T00:18:13.814208","exception":false,"start_time":"2024-12-31T00:18:13.800442","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"913f8deb","cell_type":"code","source":"# Xem lại các cột sau khi chuẩn hóa\nprint(inter_df.columns)\n\n# Chuyển sang Pandas chỉ với các cột cần thiết cho RecBole\npandas_df = inter_df[['session:token', 'aid:token', 'ts:float']].to_pandas()\n\n# Ghi ra file .inter (TSV) cho RecBole đọc\npandas_df.to_csv(\n    '/kaggle/working/recbox_data/recbox_data.inter',\n    sep='\\t',\n    index=False\n)\n\n# Giải phóng bộ nhớ\ndel inter_df, pandas_df\ngc.collect()","metadata":{"papermill":{"duration":30.399991,"end_time":"2024-12-31T00:18:44.219546","exception":false,"start_time":"2024-12-31T00:18:13.819555","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"d9f15b5e","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":{"papermill":{"duration":0.005208,"end_time":"2024-12-31T00:18:44.230150","exception":false,"start_time":"2024-12-31T00:18:44.224942","status":"completed"},"tags":[]}},{"id":"aa3c9392","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":{"papermill":{"duration":1.827884,"end_time":"2024-12-31T00:18:46.063470","exception":false,"start_time":"2024-12-31T00:18:44.235586","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"2d84cfec","cell_type":"code","source":"# Độ dài tối đa chuỗi lịch sử per session (số item gần nhất dùng cho GRU4Rec)\nMAX_ITEM = 20\n\n# Cấu hình RecBole cho mô hình GRU4Rec\nrecbole_config = {\n    # Đường dẫn chứa file .inter đã chuẩn hóa từ train/test OTTO\n    'data_path': '/kaggle/working/',\n\n    # Định nghĩa các field chính trong tập dữ liệu\n    'USER_ID_FIELD': 'session',   # user = session\n    'ITEM_ID_FIELD': 'aid',       # item = aid\n    'TIME_FIELD': 'ts',\n\n    # Lọc bớt user/item quá ít tương tác để giảm kích thước embedding\n    'user_inter_num_interval': \"[5,Inf)\",\n    'item_inter_num_interval': \"[5,Inf)\",\n\n    # Các cột sẽ load từ file .inter\n    'load_col': {'inter': ['session', 'aid', 'ts']},\n\n    # Không dùng negative sampling riêng (GRU4Rec với CE loss sẽ tự xử lý)\n    'train_neg_sample_args': None,\n\n    # Thiết lập huấn luyện\n    'epochs': 1,             # giảm xuống 1 epoch để train nhanh, tránh tốn tài nguyên\n    'stopping_step': 3,      # early stopping nếu không cải thiện sau 3 epoch\n    'eval_batch_size': 1024,\n    'train_batch_size': 1024,\n    # 'enable_amp': True,    # có thể bật nếu GPU hỗ trợ để giảm RAM/tăng tốc\n\n    # Chiều dài tối đa chuỗi tương tác\n    'MAX_ITEM_LIST_LENGTH': MAX_ITEM,\n\n    # Thiết lập đánh giá: chia train/valid theo thời gian\n    'eval_args': {\n        'split': {'RS': [9, 1, 0]},  # 90% train, 10% valid theo thời gian\n        'group_by': 'user',          # group theo session\n        'order': 'TO',               # order theo thời gian (Time Order)\n        'mode': 'full',              # full-ranking trên toàn item\n    },\n}\n\n# Tạo đối tượng Config của RecBole\nconfig = Config(model='GRU4Rec', dataset='recbox_data', config_dict=recbole_config)\n\n# Khởi tạo seed để tái lập kết quả\ninit_seed(config['seed'], config['reproducibility'])\n\n# Khởi tạo logger của RecBole\ninit_logger(config)\nlogger = getLogger()\n\n# Thêm handler để log ra stdout (Kaggle cell output)\nconsole_handler = logging.StreamHandler()\nconsole_handler.setLevel(logging.INFO)\nlogger.addHandler(console_handler)\n\n# Ghi cấu hình đầy đủ ra log để dễ debug\nlogger.info(config)","metadata":{"papermill":{"duration":0.553068,"end_time":"2024-12-31T00:18:46.622400","exception":false,"start_time":"2024-12-31T00:18:46.069332","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"15e8545d","cell_type":"code","source":"dataset = create_dataset(config)\nlogger.info(dataset)","metadata":{"papermill":{"duration":135.370534,"end_time":"2024-12-31T00:21:02.013101","exception":false,"start_time":"2024-12-31T00:18:46.642567","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"143502a0","cell_type":"code","source":"# dataset splitting\ntrain_data, valid_data, test_data = data_preparation(config, dataset)","metadata":{"papermill":{"duration":96.036874,"end_time":"2024-12-31T00:22:38.069819","exception":false,"start_time":"2024-12-31T00:21:02.032945","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"ac30481f","cell_type":"code","source":"# model loading and initialization\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":{"papermill":{"duration":9663.378013,"end_time":"2024-12-31T03:03:41.468074","exception":false,"start_time":"2024-12-31T00:22:38.090061","status":"completed"},"scrolled":true,"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"5b2124c1","cell_type":"code","source":"import gc\ndel trainer, train_data, valid_data, test_data","metadata":{"papermill":{"duration":0.038222,"end_time":"2024-12-31T03:03:41.538137","exception":false,"start_time":"2024-12-31T03:03:41.499915","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"414c813c","cell_type":"code","source":"gc.collect()","metadata":{"papermill":{"duration":0.719989,"end_time":"2024-12-31T03:03:42.289999","exception":false,"start_time":"2024-12-31T03:03:41.570010","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"24eed8b1","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":{"papermill":{"duration":0.031247,"end_time":"2024-12-31T03:03:42.353521","exception":false,"start_time":"2024-12-31T03:03:42.322274","status":"completed"},"tags":[]}},{"id":"e29c2d43","cell_type":"code","source":"# https://qiita.com/fufufukakaka/items/e03df3a7299b2b8f99cf\nfrom typing import List, Tuple\nimport numpy as np\nimport torch\n\nfrom pydantic import BaseModel\nfrom recbole.data import create_dataset\nfrom recbole.data.dataset.sequential_dataset import SequentialDataset\nfrom recbole.data.interaction import Interaction\nfrom recbole.model.sequential_recommender.sine import SINE\nfrom recbole.utils import get_model, init_seed\n\nclass ItemHistory(BaseModel):\n    sequence: List[str]\n    topk: int\n\nclass RecommendedItems(BaseModel):\n    score_list: List[float]\n    item_list: List[str]\n\n\ndef 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":{"papermill":{"duration":0.138822,"end_time":"2024-12-31T03:03:42.524494","exception":false,"start_time":"2024-12-31T03:03:42.385672","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"7732425c","cell_type":"code","source":"#test = pl.read_parquet('../input/otto-train-and-test-data-for-local-validation/test.parquet')\ntest = pl.read_parquet('../input/otto-full-optimized-memory-footprint/test.parquet')\n\nimport pandas as pd\nimport numpy as np\n\nfrom collections import defaultdict\n\n#sample_sub = pd.read_csv('../input/otto-recommender-system//sample_submission.csv')\n\nsession_types = ['clicks', 'carts', 'orders']\ntest_session_AIDs = test.to_pandas().reset_index(drop=True).groupby('session')['aid'].apply(list)\ntest_session_types = test.to_pandas().reset_index(drop=True).groupby('session')['type'].apply(list)\n\ndel test\ngc.collect()\n\nlabels = []\n\ntype_weight_multipliers = {0: 1, 1: 6, 2: 3}\nfor 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=[str(x) for x in 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":{"papermill":{"duration":7972.540818,"end_time":"2024-12-31T05:16:35.097680","exception":false,"start_time":"2024-12-31T03:03:42.556862","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"11c6d327","cell_type":"code","source":"pred_user_to_item(item)['item_list']","metadata":{"papermill":{"duration":0.048106,"end_time":"2024-12-31T05:16:35.180800","exception":false,"start_time":"2024-12-31T05:16:35.132694","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"e27a6f8e","cell_type":"code","source":"labels_as_strings = [' '.join([str(l) for l in lls]) for lls in labels]\npredictions = pd.DataFrame(data={'session_type': test_session_AIDs.index, 'labels': labels_as_strings})","metadata":{"papermill":{"duration":5.936263,"end_time":"2024-12-31T05:16:41.149829","exception":false,"start_time":"2024-12-31T05:16:35.213566","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"id":"ced09343","cell_type":"code","source":"labels_as_strings = [' '.join([str(l) for l in lls]) for lls in labels]\n\npredictions = pd.DataFrame(data={'session_type': test_session_AIDs.index, 'labels': labels_as_strings})\n\nprediction_dfs = []\n\nfor 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\nsubmission = pd.concat(prediction_dfs).reset_index(drop=True)\nsubmission.to_csv('submission.csv', index=False)","metadata":{"papermill":{"duration":26.72421,"end_time":"2024-12-31T05:17:07.906522","exception":false,"start_time":"2024-12-31T05:16:41.182312","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null}]}