{"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":"### This is an example of custom objective metric for LightGBM","metadata":{}},{"cell_type":"markdown","source":"In this competition we need to optimize recall@20, but there's no such metric in LightGBM. We can implement it ourselves.\n\n","metadata":{}},{"cell_type":"markdown","source":"### Simple version","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport lightgbm as lgb\n\nfrom numba import njit\n\n\ndef recall20(preds, targets, groups):\n    total = 0\n    nonempty = 0\n    group_starts = np.cumsum(groups)\n\n    for group_id in range(len(groups)):\n        group_end = group_starts[group_id]\n        group_start = group_end - groups[group_id]\n        ranks = np.argsort(preds[group_start:group_end])[::-1]\n        hits = 0\n        for i in range(min(len(ranks), 20)):\n            hits += targets[group_start + ranks[i]]\n\n        actual = min(20, targets[group_start:group_end].sum())\n        if actual > 0:\n            total += hits / actual\n            nonempty += 1\n\n    return total / nonempty\n\n# custom metric for LightGBM should return \n# \"metric name\", \"metric value\" and \"greater is better\" flag\ndef lgb_recall(preds, lgb_dataset):\n    metric = recall20(preds, lgb_dataset.label, lgb_dataset.group)\n    return 'recall@20', metric, True","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-01-23T16:13:50.678821Z","iopub.execute_input":"2023-01-23T16:13:50.679423Z","iopub.status.idle":"2023-01-23T16:13:51.810378Z","shell.execute_reply.started":"2023-01-23T16:13:50.679363Z","shell.execute_reply":"2023-01-23T16:13:51.809400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Numba version ","metadata":{}},{"cell_type":"code","source":"@njit() # the only difference from previous version\ndef numba_recall20(preds, targets, groups):\n    total = 0\n    nonempty = 0\n    group_starts = np.cumsum(groups)\n\n    for group_id in range(len(groups)):\n        group_end = group_starts[group_id]\n        group_start = group_end - groups[group_id]\n        ranks = np.argsort(preds[group_start:group_end])[::-1]\n        hits = 0\n        for i in range(min(len(ranks), 20)):\n            hits += targets[group_start + ranks[i]]\n\n        actual = min(20, targets[group_start:group_end].sum())\n        if actual > 0:\n            total += hits / actual\n            nonempty += 1\n\n    return total / nonempty\n\n\ndef lgb_numba_recall(preds, lgb_dataset):\n    metric = numba_recall20(preds, lgb_dataset.label, lgb_dataset.group)\n    return 'numba_recall@20', metric, True\n\n","metadata":{"execution":{"iopub.status.busy":"2023-01-23T16:13:51.815104Z","iopub.execute_input":"2023-01-23T16:13:51.815384Z","iopub.status.idle":"2023-01-23T16:13:51.850714Z","shell.execute_reply.started":"2023-01-23T16:13:51.815357Z","shell.execute_reply":"2023-01-23T16:13:51.849842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Synthetic data","metadata":{}},{"cell_type":"code","source":"def create_dataset(size: int) -> lgb.Dataset:\n    data = np.random.normal(size=(size, 10))\n    target = np.random.randint(0, 2, size)\n    groups = np.ones(size // 40, dtype=np.int32) * 40 # groups of 40\n    return lgb.Dataset(data, target, group=groups)","metadata":{"execution":{"iopub.status.busy":"2023-01-23T16:13:51.854569Z","iopub.execute_input":"2023-01-23T16:13:51.854902Z","iopub.status.idle":"2023-01-23T16:13:51.861078Z","shell.execute_reply.started":"2023-01-23T16:13:51.854870Z","shell.execute_reply":"2023-01-23T16:13:51.859991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"size = 1_000_000","metadata":{"execution":{"iopub.status.busy":"2023-01-23T16:13:51.863389Z","iopub.execute_input":"2023-01-23T16:13:51.863786Z","iopub.status.idle":"2023-01-23T16:13:51.873073Z","shell.execute_reply.started":"2023-01-23T16:13:51.863755Z","shell.execute_reply":"2023-01-23T16:13:51.872270Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = create_dataset(size)\neval_dataset = create_dataset(size)","metadata":{"execution":{"iopub.status.busy":"2023-01-23T16:13:51.874206Z","iopub.execute_input":"2023-01-23T16:13:51.874598Z","iopub.status.idle":"2023-01-23T16:13:52.547642Z","shell.execute_reply.started":"2023-01-23T16:13:51.874538Z","shell.execute_reply":"2023-01-23T16:13:52.546609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Built-in map@20 metric","metadata":{}},{"cell_type":"code","source":"%%time\nparams = {'objective': 'lambdarank', 'metric': 'map', 'eval_at': [20]}\nmodel = lgb.train(\n    params,\n    train_dataset,\n    valid_sets=eval_dataset,\n    callbacks=[lgb.log_evaluation(5)],\n)","metadata":{"execution":{"iopub.status.busy":"2023-01-23T16:13:52.549063Z","iopub.execute_input":"2023-01-23T16:13:52.549414Z","iopub.status.idle":"2023-01-23T16:14:08.584163Z","shell.execute_reply.started":"2023-01-23T16:13:52.549383Z","shell.execute_reply":"2023-01-23T16:14:08.583319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### recall@20","metadata":{}},{"cell_type":"code","source":"%%time\nparams = {'objective': 'lambdarank', 'metric': '\"None\"'}\nmodel = lgb.train(\n    params,\n    train_dataset,\n    valid_sets=eval_dataset,\n    feval=lgb_recall,\n    callbacks=[lgb.log_evaluation(5)],\n)","metadata":{"execution":{"iopub.status.busy":"2023-01-23T16:14:08.585670Z","iopub.execute_input":"2023-01-23T16:14:08.586310Z","iopub.status.idle":"2023-01-23T16:15:24.093190Z","shell.execute_reply.started":"2023-01-23T16:14:08.586267Z","shell.execute_reply":"2023-01-23T16:15:24.092194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### numba recall@20","metadata":{}},{"cell_type":"code","source":"%%time\n# call function before benchmarking lightgbm, because we need numba to compile it\nlgb_numba_recall(np.zeros(size), eval_dataset)","metadata":{"execution":{"iopub.status.busy":"2023-01-23T16:15:24.097407Z","iopub.execute_input":"2023-01-23T16:15:24.098344Z","iopub.status.idle":"2023-01-23T16:15:26.491611Z","shell.execute_reply.started":"2023-01-23T16:15:24.098303Z","shell.execute_reply":"2023-01-23T16:15:26.490398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nparams = {'objective': 'lambdarank', 'metric': '\"None\"'}\nmodel = lgb.train(\n    params,\n    train_dataset,\n    valid_sets=eval_dataset,\n    feval=lgb_numba_recall,\n    callbacks=[lgb.log_evaluation(5)],\n)","metadata":{"execution":{"iopub.status.busy":"2023-01-23T16:15:26.493370Z","iopub.execute_input":"2023-01-23T16:15:26.494076Z","iopub.status.idle":"2023-01-23T16:15:48.395709Z","shell.execute_reply.started":"2023-01-23T16:15:26.494028Z","shell.execute_reply":"2023-01-23T16:15:48.394825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Summary","metadata":{}},{"cell_type":"markdown","source":"Custom metric function slows training time significantly (from **16s** to **1min 15s**). With numba we can speed it up to **22s** (which is also 50% slower than built-in, but we can afford it).\n\nAlso, training time impact heavily depends on the size of the training/eval set. (e.g. if training one iteration takes 9s, and eval takes 1s, changing eval function to 2s will change training time only by 10%).\n\nOn multi-core machines you can make further speed-up by parallelization.","metadata":{}},{"cell_type":"code","source":"\nfrom numba import prange\n\n@njit(parallel=True) # added parallel flag\ndef numba_parallel_recall20(preds, targets, groups):\n    total = 0\n    nonempty = 0\n    group_starts = np.cumsum(groups)\n\n    for group_id in prange(len(groups)): # changed range to prange\n        group_end = group_starts[group_id]\n        group_start = group_end - groups[group_id]\n        ranks = np.argsort(preds[group_start:group_end])[::-1]\n        hits = 0\n        for i in range(min(len(ranks), 20)):\n            hits += targets[group_start + ranks[i]]\n\n        actual = min(20, targets[group_start:group_end].sum())\n        if actual > 0:\n            total += hits / actual\n            nonempty += 1\n\n    return total / nonempty","metadata":{"execution":{"iopub.status.busy":"2023-01-23T16:15:48.399610Z","iopub.execute_input":"2023-01-23T16:15:48.400429Z","iopub.status.idle":"2023-01-23T16:15:48.409062Z","shell.execute_reply.started":"2023-01-23T16:15:48.400385Z","shell.execute_reply":"2023-01-23T16:15:48.407731Z"},"trusted":true},"execution_count":null,"outputs":[]}]}