{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":38760,"databundleVersionId":4493939,"sourceType":"competition"},{"sourceId":4499280,"sourceType":"datasetVersion","datasetId":2631076,"isSourceIdPinned":false}],"dockerImageVersionId":30804,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# OTTO – Multi-Objective Recommender System ","metadata":{}},{"cell_type":"markdown","source":"## 1. Setup\n\n### **1. Nhập thư viện**\nChương trình nhập các thư viện quan trọng sau:\n- `os`: Cung cấp các chức năng để tương tác với hệ điều hành.\n- `gc`: Module thu gom rác (garbage collection) giúp quản lý bộ nhớ hiệu quả.\n- `heapq`: Cấu trúc dữ liệu hàng đợi ưu tiên (min-heap).\n- `pickle`: Dùng để lưu trữ và đọc dữ liệu dưới dạng nhị phân.\n- `numba as nb`: Một trình biên dịch Just-In-Time (JIT) giúp tăng tốc tính toán số học.\n- `numpy as np`: Thư viện tính toán số học quan trọng trong Python.\n- `pandas as pd`: Thư viện mạnh mẽ để xử lý và phân tích dữ liệu.\n- `tqdm.auto import tqdm`: Cung cấp thanh tiến trình để theo dõi vòng lặp.\n\n---\n\n### **2. Thiết lập các tham số**\nChương trình định nghĩa các hằng số quan trọng:\n- `tail = 30`: Có thể là số lượng tương tác gần nhất trong một phiên được sử dụng.\n- `parallel = 1024`: Xử lý 1024 phiên (sessions) song song để tăng tốc độ tính toán.\n- `topn = 20`: Xem xét 20 đề xuất hàng đầu cho mỗi phiên.\n- `ops_weights = np.array([1.0, 6.0, 3.0])`: Các trọng số được gán cho các loại tương tác khác nhau.\n- `OP_WEIGHT = 0; TIME_WEIGHT = 1`: Xác định các chỉ số để phân biệt giữa trọng số thao tác và trọng số theo thời gian.\n- `parallel = 1024` (được định nghĩa lại): Đảm bảo giá trị xử lý song song.\n- `test_ops_weights = np.array([1.0, 6.0, 3.0])`: Trọng số của các thao tác được sử dụng trong quá trình kiểm thử.\n\n---","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport heapq\nimport pickle\nimport numba as nb\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\ntail = 30\nparallel = 1024\ntopn = 20\nops_weights = np.array([1.0, 6.0, 3.0])\nOP_WEIGHT = 0; TIME_WEIGHT = 1\nparallel = 1024\ntest_ops_weights = np.array([1.0, 6.0, 3.0])","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":0.946424,"end_time":"2022-11-13T15:03:06.858235","exception":false,"start_time":"2022-11-13T15:03:05.911811","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-02-13T06:07:31.486498Z","iopub.execute_input":"2025-02-13T06:07:31.486879Z","iopub.status.idle":"2025-02-13T06:07:33.739626Z","shell.execute_reply.started":"2025-02-13T06:07:31.486847Z","shell.execute_reply":"2025-02-13T06:07:33.738600Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Dữ liệu\n\n#### DataFrame Meta\n\nCác cột trong `train.csv/test.csv`:\n\n- `session`: ID của các phiên, được sắp xếp theo thứ tự\n- `length`: độ dài của mỗi phiên\n- `start_time`: thời gian bắt đầu phiên (đơn vị = 1 giây)\n\n#### Mảng Dữ Liệu\n\n`train.npz/test.npz`:\n\nTệp npz chứa ba mảng một chiều:\n\n- `aids`: article id\n- `ts`: \"the Unix timestamp of the event\" mảng dấu thời gian (đã trừ đi `start_time` của phiên tương ứng)\n- `ops`: mảng loại sự kiện, trong đó `clicks=0`, `carts=1`, `orders=2`\n\nMỗi mảng là sự kết hợp của các phiên, được sắp xếp theo ID phiên.\n\nVí dụ:\n\n```python\naids = np.concatenate([\n        [1, 2, 3], # aid trong phiên 0 theo thứ tự thời gian\n        [4, 5, 6], # aid trong phiên 1 theo thứ tự thời gian\n        ...])\n```","metadata":{"papermill":{"duration":0.003222,"end_time":"2022-11-13T15:03:06.865386","exception":false,"start_time":"2022-11-13T15:03:06.862164","status":"completed"},"tags":[]}},{"cell_type":"code","source":"df = pd.read_csv(\"../input/otto-data/train.csv\")\ndf_test = pd.read_csv(\"../input/otto-data/test.csv\")\ndf = pd.concat([df, df_test]).reset_index(drop = True)\nnpz = np.load(\"../input/otto-data/train.npz\")\nnpz_test = np.load(\"../input/otto-data/test.npz\")\naids = np.concatenate([npz['aids'], npz_test['aids']])\nts = np.concatenate([npz['ts'], npz_test['ts']])\nops = np.concatenate([npz['ops'], npz_test['ops']])\n\ndf[\"idx\"] = np.cumsum(df.length) - df.length\ndf[\"end_time\"] = df.start_time + ts[df.idx + df.length - 1]","metadata":{"papermill":{"duration":41.941645,"end_time":"2022-11-13T15:03:48.810448","exception":false,"start_time":"2022-11-13T15:03:06.868803","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-02-13T06:07:33.741603Z","iopub.execute_input":"2025-02-13T06:07:33.742102Z","iopub.status.idle":"2025-02-13T06:08:01.065027Z","shell.execute_reply.started":"2025-02-13T06:07:33.742066Z","shell.execute_reply":"2025-02-13T06:08:01.063698Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Các hàm Numba\n\nĐể tận dụng tối đa khả năng của JIT trong Numba và tránh lỗi bộ nhớ, chúng tôi xử lý `parallel=1024` phiên trong một lần gọi hàm JIT.\n\n- Trước tiên, Em đếm các cặp `(aid1, aid2)` trong mỗi phiên một cách song song.\n\n- Sau đó, chúng tôi hợp nhất các cặp này thành một bộ đếm lồng nhau `{aid1: {aid2: n}, ...}`.\n\n- Cuối cùng, Em tìm `aid2` hàng đầu (`top-k`) từ bộ đếm của `aid1`.","metadata":{"papermill":{"duration":0.005089,"end_time":"2022-11-13T15:03:48.823258","exception":false,"start_time":"2022-11-13T15:03:48.818169","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"# get pair dict {(aid1, aid2): weight} for each session\n# The maximum time span between two points is 1 day = 24 * 60 * 60 sec\n@nb.jit(nopython = True, cache = True)\ndef get_single_pairs(pairs, aids, ts, ops, idx, length, start_time, ops_weights, mode):\n    max_idx = idx + length\n    min_idx = max(max_idx - tail, idx)\n    for i in range(min_idx, max_idx):\n        for j in range(i + 1, max_idx):\n            if ts[j] - ts[i] >= 24 * 60 * 60: break\n            if aids[i] == aids[j]: continue\n            if mode == OP_WEIGHT:\n                w1 = ops_weights[ops[j]]\n                w2 = ops_weights[ops[i]]\n            elif mode == TIME_WEIGHT:\n                w1 = 1 + 3 * (ts[i] + start_time - 1659304800) / (1662328791 - 1659304800)\n                w2 = 1 + 3 * (ts[j] + start_time - 1659304800) / (1662328791 - 1659304800)\n            pairs[(aids[i], aids[j])] = w1\n            pairs[(aids[j], aids[i])] = w2\n\n# get pair dict of each session in parallel\n# merge pairs into a nested dict format (cnt)\n@nb.jit(nopython = True, parallel = True, cache = True)\ndef get_pairs(aids, ts, ops, row, cnts, ops_weights, mode):\n    par_n = len(row)\n    pairs = [{(0, 0): 0.0 for _ in range(0)} for _ in range(par_n)]\n    for par_i in nb.prange(par_n):\n        _, idx, length, start_time = row[par_i]\n        get_single_pairs(pairs[par_i], aids, ts, ops, idx, length, start_time, ops_weights, mode)\n    for par_i in range(par_n):\n        for (aid1, aid2), w in pairs[par_i].items():\n            if aid1 not in cnts: cnts[aid1] = {0: 0.0 for _ in range(0)}\n            cnt = cnts[aid1]\n            if aid2 not in cnt: cnt[aid2] = 0.0\n            cnt[aid2] += w\n    \n# util function to get most common keys from a counter dict using min-heap\n# overwrite == 1 means the later item with equal weight is more important\n# otherwise, means the former item with equal weight is more important\n# the result is ordered from higher weight to lower weight\n@nb.jit(nopython = True, cache = True)\ndef heap_topk(cnt, overwrite, cap):\n    q = [(0.0, 0, 0) for _ in range(0)]\n    for i, (k, n) in enumerate(cnt.items()):\n        if overwrite == 1:\n            heapq.heappush(q, (n, i, k))\n        else:\n            heapq.heappush(q, (n, -i, k))\n        if len(q) > cap:\n            heapq.heappop(q)\n    return [heapq.heappop(q)[2] for _ in range(len(q))][::-1]\n   \n# save top-k aid2 for each aid1's cnt\n@nb.jit(nopython = True, cache = True)\ndef get_topk(cnts, topk, k):\n    for aid1, cnt in cnts.items():\n        topk[aid1] = np.array(heap_topk(cnt, 1, k))","metadata":{"papermill":{"duration":0.376871,"end_time":"2022-11-13T15:03:49.206642","exception":false,"start_time":"2022-11-13T15:03:48.829771","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-02-13T06:08:01.066035Z","iopub.execute_input":"2025-02-13T06:08:01.066348Z","iopub.status.idle":"2025-02-13T06:08:01.183099Z","shell.execute_reply.started":"2025-02-13T06:08:01.066306Z","shell.execute_reply":"2025-02-13T06:08:01.181639Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Huấn luyện\n\nỞ đây, chúng tôi sử dụng phương pháp đếm **topk** từ [notebook của Deotte](https://www.kaggle.com/code/cdeotte/candidate-rerank-model-lb-0-573). Tôi đã lược bỏ một số ý tưởng của anh ấy để làm cho toàn bộ quy trình đơn giản hơn nhưng vẫn mạnh mẽ.\n\n- **Mode 0**: Bộ đếm được tính trọng số theo loại thao tác, với `op_weights = [1.0, 6.0, 3.0]`. Chế độ này sẽ được sử dụng cho dự đoán giỏ hàng và đơn hàng.\n\n- **Mode 1**: Bộ đếm được tính trọng số theo thời gian thao tác. Chế độ này sẽ được sử dụng cho dự đoán lượt nhấp.\n\nVới mỗi chế độ, thời gian chạy ước tính khoảng **7~8 phút** trên notebook Kaggle.","metadata":{"papermill":{"duration":0.005349,"end_time":"2022-11-13T15:03:49.217833","exception":false,"start_time":"2022-11-13T15:03:49.212484","status":"completed"},"tags":[]}},{"cell_type":"code","source":"topks = {}\n\n# for two modes\nfor mode in [OP_WEIGHT, TIME_WEIGHT]:\n    # get nested counter\n    cnts = nb.typed.Dict.empty(\n        key_type = nb.types.int64,\n        value_type = nb.typeof(nb.typed.Dict.empty(key_type = nb.types.int64, value_type = nb.types.float64)))\n    max_idx = len(df)\n    for idx in tqdm(range(0, max_idx, parallel)):\n        row = df.iloc[idx:min(idx + parallel, max_idx)][['session', 'idx', 'length', 'start_time']].values\n        get_pairs(aids, ts, ops, row, cnts, ops_weights, mode)\n\n    # get topk from counter\n    topk = nb.typed.Dict.empty(\n            key_type = nb.types.int64,\n            value_type = nb.types.int64[:])\n    get_topk(cnts, topk, topn)\n\n    del cnts; gc.collect()\n    topks[mode] = topk","metadata":{"papermill":{"duration":1100.78763,"end_time":"2022-11-13T15:22:10.011201","exception":false,"start_time":"2022-11-13T15:03:49.223571","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-02-13T06:08:01.185787Z","iopub.execute_input":"2025-02-13T06:08:01.186225Z","iopub.status.idle":"2025-02-13T06:24:59.888495Z","shell.execute_reply.started":"2025-02-13T06:08:01.186171Z","shell.execute_reply":"2025-02-13T06:24:59.887025Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Suy luận Hàm Numba\n\nLogic của phần này là phiên bản đơn giản hơn của [notebook của Deotte](https://www.kaggle.com/code/cdeotte/candidate-rerank-model-lb-0-573).\n\nTương tự như các hàm huấn luyện, chúng tôi xử lý `parallel=1024` phiên trong một lần gọi hàm jit. Đối với mỗi phiên, nếu số lượng `aids` duy nhất nhiều hơn 20, chúng tôi sắp xếp lại chúng theo trọng số suy giảm theo thời gian. Ngược lại, chúng tôi sử dụng `aids` hiện tại để gọi lại các `aids` mới từ danh sách top-k đồng truy cập của chúng, với trọng số dựa trên loại thao tác của truy vấn (`test_ops_weights=[1.0, 6.0, 3.0]`).\n\nSử dụng top-k **mode0** để tạo dự đoán cho giỏ hàng và đơn hàng, và top-k **mode1** để tạo dự đoán cho lượt nhấp.","metadata":{"papermill":{"duration":0.00353,"end_time":"2022-11-13T15:22:10.018903","exception":false,"start_time":"2022-11-13T15:22:10.015373","status":"completed"},"tags":[]}},{"cell_type":"code","source":"@nb.jit(nopython = True, cache = True)\ndef inference_(aids, ops, row, result, topk, test_ops_weights, seq_weight):\n    for session, idx, length in row:\n        unique_aids = nb.typed.Dict.empty(key_type = nb.types.int64, value_type = nb.types.float64)\n        cnt = nb.typed.Dict.empty(key_type = nb.types.int64, value_type = nb.types.float64)\n        \n        candidates = aids[idx:idx + length][::-1]\n        candidates_ops = ops[idx:idx + length][::-1]\n        for a in candidates:\n            unique_aids[a] = 0\n                \n        if len(unique_aids) >= 20:\n            sequence_weight = np.power(2, np.linspace(seq_weight, 1, len(candidates)))[::-1] - 1\n            for a, op, w in zip(candidates, candidates_ops, sequence_weight):\n                if a not in cnt: cnt[a] = 0\n                cnt[a] += w * test_ops_weights[op]\n            result_candidates = heap_topk(cnt, 0, 20)\n        else:\n            result_candidates = list(unique_aids)\n            for a in result_candidates:\n                if a not in topk: continue\n                for b in topk[a]:\n                    if b in unique_aids: continue\n                    if b not in cnt: cnt[b] = 0\n                    cnt[b] += 1\n            result_candidates.extend(heap_topk(cnt, 0, 20 - len(result_candidates)))\n        result[session] = np.array(result_candidates)\n        \n@nb.jit(nopython = True)\ndef inference(aids, ops, row, \n              result_clicks, result_buy,\n              topk_clicks, topk_buy,\n              test_ops_weights):\n    inference_(aids, ops, row, result_clicks, topk_clicks, test_ops_weights, 0.1)\n    inference_(aids, ops, row, result_buy, topk_buy, test_ops_weights, 0.5)","metadata":{"papermill":{"duration":3.252033,"end_time":"2022-11-13T15:22:13.275135","exception":false,"start_time":"2022-11-13T15:22:10.023102","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-02-13T06:24:59.936576Z","iopub.execute_input":"2025-02-13T06:24:59.937209Z","iopub.status.idle":"2025-02-13T06:25:04.712369Z","shell.execute_reply.started":"2025-02-13T06:24:59.937172Z","shell.execute_reply":"2025-02-13T06:25:04.711329Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# result place holder\nresult_clicks = nb.typed.Dict.empty(\n    key_type = nb.types.int64,\n    value_type = nb.types.int64[:])\nresult_buy = nb.typed.Dict.empty(\n    key_type = nb.types.int64,\n    value_type = nb.types.int64[:])\nfor idx in tqdm(range(len(df) - len(df_test), len(df), parallel)):\n    row = df.iloc[idx:min(idx + parallel, len(df))][['session', 'idx', 'length']].values\n    inference(aids, ops, row, result_clicks, result_buy, topks[TIME_WEIGHT], topks[OP_WEIGHT], test_ops_weights)","metadata":{"papermill":{"duration":72.25349,"end_time":"2022-11-13T15:23:25.532644","exception":false,"start_time":"2022-11-13T15:22:13.279154","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-02-13T06:25:04.713879Z","iopub.execute_input":"2025-02-13T06:25:04.714214Z","iopub.status.idle":"2025-02-13T06:25:57.917764Z","shell.execute_reply.started":"2025-02-13T06:25:04.714180Z","shell.execute_reply":"2025-02-13T06:25:57.915142Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"subs = []\nop_names = [\"clicks\", \"carts\", \"orders\"]\nfor result, op in zip([result_clicks, result_buy, result_buy], op_names):\n\n    sub = pd.DataFrame({\"session_type\": result.keys(), \"labels\": result.values()})\n    sub.session_type = sub.session_type.astype(str) + f\"_{op}\"\n    sub.labels = sub.labels.apply(lambda x: \" \".join(x.astype(str)))\n    subs.append(sub)\n    \nsub = pd.concat(subs).reset_index(drop = True)\nsub.to_csv('submission.csv', index = False)\nsub.head()","metadata":{"papermill":{"duration":141.966809,"end_time":"2022-11-13T15:25:47.503801","exception":false,"start_time":"2022-11-13T15:23:25.536992","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-02-13T06:25:57.925712Z","iopub.execute_input":"2025-02-13T06:25:57.926182Z","iopub.status.idle":"2025-02-13T06:27:15.104764Z","shell.execute_reply.started":"2025-02-13T06:25:57.926138Z","shell.execute_reply":"2025-02-13T06:27:15.103609Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}