{
  "id": 338712,
  "title": "For those looking for early_stopping for LGBM (dart)",
  "url": "/competitions/amex-default-prediction/discussion/338712",
  "author_name": "",
  "post_date": "2022-07-21T16:40:26.712380600Z",
  "votes": 5,
  "comment_count": 4,
  "views": 0,
  "content": "<p>You can use this: </p>\n<p>link: <a href=\"https://qiita.com/birdwatcher/items/78e0158957fec2d6c9e8\" target=\"_blank\">https://qiita.com/birdwatcher/items/78e0158957fec2d6c9e8</a></p>\n<p>Example:</p>\n<p>des = DartEarlyStopping(\"valid_1\", \"amex_metric\", 100)<br>\nmodel = lgb.train(<br>\n    params, data,<br>\n    valid_sets=[data, eval_data], <br>\n    num_boost_round=100,<br>\n    callbacks=[des],<br>\n    verbose_eval=1)</p>\n<p>class DartEarlyStopping(object):</p>\n<pre><code>def __init__(self, data_name, monitor_metric, stopping_round):\n    self.data_name = data_name\n    self.monitor_metric = monitor_metric\n    self.stopping_round = stopping_round\n    self.best_score = None\n    self.best_model = None\n    self.best_score_list = []\n    self.best_iter = 0\n\ndef _is_higher_score(self, metric_score, is_higher_better):\n    if self.best_score is None:\n        return True\n    return (self.best_score &lt; metric_score) if is_higher_better else (self.best_score &gt; metric_score)\n\ndef _deepcopy(self, x):\n    # copy.deepcopyではlightgbmのモデルは完全にコピーされないためpickleを使用\n    return pickle.loads(pickle.dumps(x))\n\ndef __call__(self, env):\n    evals = env.evaluation_result_list\n    for data, metric, score, is_higher_better in evals:\n        if data != self.data_name or metric != self.monitor_metric:\n            continue\n        if not self._is_higher_score(score, is_higher_better):\n            if env.iteration - self.best_iter &gt; self.stopping_round:\n                # 終了させる\n                eval_result_str = '\\t'.join([lgb.callback._format_eval_result(x) for x in self.best_score_list])\n                lgb.basic._log_info(f\"Early stopping, best iteration is:\\n[{self.best_iter+1}]\\t{eval_result_str}\") \n                lgb.basic._log_info(f\"You can get best model by \\\"DartEarlyStopping.best_model\\\"\")\n                raise lgb.callback.EarlyStopException(self.best_iter, self.best_score_list)\n            return\n        # dartでは過去の木も更新されてしまうため、deepcopyしておく\n        self.best_model = self._deepcopy(env.model)\n        self.best_iter = env.iteration\n        self.best_score_list = evals\n        self.best_score = score\n        return\n    raise ValueError(\"monitoring metric not found\")\n</code></pre>",
  "messages": [
    {
      "id": "1865234",
      "postDate": "07/21/2022 16:40:26",
      "content": "<p>You can use this: </p>\n<p>link: <a href=\"https://qiita.com/birdwatcher/items/78e0158957fec2d6c9e8\" target=\"_blank\">https://qiita.com/birdwatcher/items/78e0158957fec2d6c9e8</a></p>\n<p>Example:</p>\n<p>des = DartEarlyStopping(\"valid_1\", \"amex_metric\", 100)<br>\nmodel = lgb.train(<br>\n    params, data,<br>\n    valid_sets=[data, eval_data], <br>\n    num_boost_round=100,<br>\n    callbacks=[des],<br>\n    verbose_eval=1)</p>\n<p>class DartEarlyStopping(object):</p>\n<pre><code>def __init__(self, data_name, monitor_metric, stopping_round):\n    self.data_name = data_name\n    self.monitor_metric = monitor_metric\n    self.stopping_round = stopping_round\n    self.best_score = None\n    self.best_model = None\n    self.best_score_list = []\n    self.best_iter = 0\n\ndef _is_higher_score(self, metric_score, is_higher_better):\n    if self.best_score is None:\n        return True\n    return (self.best_score &lt; metric_score) if is_higher_better else (self.best_score &gt; metric_score)\n\ndef _deepcopy(self, x):\n    # copy.deepcopyではlightgbmのモデルは完全にコピーされないためpickleを使用\n    return pickle.loads(pickle.dumps(x))\n\ndef __call__(self, env):\n    evals = env.evaluation_result_list\n    for data, metric, score, is_higher_better in evals:\n        if data != self.data_name or metric != self.monitor_metric:\n            continue\n        if not self._is_higher_score(score, is_higher_better):\n            if env.iteration - self.best_iter &gt; self.stopping_round:\n                # 終了させる\n                eval_result_str = '\\t'.join([lgb.callback._format_eval_result(x) for x in self.best_score_list])\n                lgb.basic._log_info(f\"Early stopping, best iteration is:\\n[{self.best_iter+1}]\\t{eval_result_str}\") \n                lgb.basic._log_info(f\"You can get best model by \\\"DartEarlyStopping.best_model\\\"\")\n                raise lgb.callback.EarlyStopException(self.best_iter, self.best_score_list)\n            return\n        # dartでは過去の木も更新されてしまうため、deepcopyしておく\n        self.best_model = self._deepcopy(env.model)\n        self.best_iter = env.iteration\n        self.best_score_list = evals\n        self.best_score = score\n        return\n    raise ValueError(\"monitoring metric not found\")\n</code></pre>",
      "rawMarkdown": "You can use this: \n\nlink: https://qiita.com/birdwatcher/items/78e0158957fec2d6c9e8\n\n\nExample:\n    \ndes = DartEarlyStopping(\"valid_1\", \"amex_metric\", 100)\nmodel = lgb.train(\n    params, data,\n    valid_sets=[data, eval_data], \n    num_boost_round=100,\n    callbacks=[des],\n    verbose_eval=1)\n\nclass DartEarlyStopping(object):\n\n\n    def __init__(self, data_name, monitor_metric, stopping_round):\n        self.data_name = data_name\n        self.monitor_metric = monitor_metric\n        self.stopping_round = stopping_round\n        self.best_score = None\n        self.best_model = None\n        self.best_score_list = []\n        self.best_iter = 0\n\n    def _is_higher_score(self, metric_score, is_higher_better):\n        if self.best_score is None:\n            return True\n        return (self.best_score < metric_score) if is_higher_better else (self.best_score > metric_score)\n\n    def _deepcopy(self, x):\n        # copy.deepcopyではlightgbmのモデルは完全にコピーされないためpickleを使用\n        return pickle.loads(pickle.dumps(x))\n\n    def __call__(self, env):\n        evals = env.evaluation_result_list\n        for data, metric, score, is_higher_better in evals:\n            if data != self.data_name or metric != self.monitor_metric:\n                continue\n            if not self._is_higher_score(score, is_higher_better):\n                if env.iteration - self.best_iter > self.stopping_round:\n                    # 終了させる\n                    eval_result_str = '\\t'.join([lgb.callback._format_eval_result(x) for x in self.best_score_list])\n                    lgb.basic._log_info(f\"Early stopping, best iteration is:\\n[{self.best_iter+1}]\\t{eval_result_str}\") \n                    lgb.basic._log_info(f\"You can get best model by \\\"DartEarlyStopping.best_model\\\"\")\n                    raise lgb.callback.EarlyStopException(self.best_iter, self.best_score_list)\n                return\n            # dartでは過去の木も更新されてしまうため、deepcopyしておく\n            self.best_model = self._deepcopy(env.model)\n            self.best_iter = env.iteration\n            self.best_score_list = evals\n            self.best_score = score\n            return\n        raise ValueError(\"monitoring metric not found\")",
      "votes": null
    },
    {
      "id": "1865335",
      "postDate": "07/21/2022 17:56:31",
      "content": "<p>Thank you <a href=\"https://www.kaggle.com/mbburabak\" target=\"_blank\">@mbburabak</a> </p>",
      "rawMarkdown": "Thank you @mbburabak",
      "votes": null
    },
    {
      "id": "1871279",
      "postDate": "07/26/2022 07:37:50",
      "content": "<p>How is it different then simply using the classic early stopping?</p>",
      "rawMarkdown": "How is it different then simply using the classic early stopping?",
      "votes": null
    },
    {
      "id": "1893815",
      "postDate": "08/11/2022 05:17:56",
      "content": "<p>I have the same question. Hope anyone could explain.</p>",
      "rawMarkdown": "I have the same question. Hope anyone could explain.",
      "votes": null
    },
    {
      "id": "1893920",
      "postDate": "08/11/2022 06:51:29",
      "content": "<p>Early stopping and dart cannot work well together:<br>\nwhen using dart,the previous trees will be updated at each iteration. Thus best iteration does not contain best trees as dart  updated the previous trees. Higher is stopping_rounds parameter more different is the best model at the end of training</p>",
      "rawMarkdown": "Early stopping and dart cannot work well together:\nwhen using dart,the previous trees will be updated at each iteration. Thus best iteration does not contain best trees as dart  updated the previous trees. Higher is stopping_rounds parameter more different is the best model at the end of training",
      "votes": null
    }
  ],
  "comments": [
    {
      "id": 1865335,
      "author_name": "saberghaderi",
      "author_url": "",
      "post_date": "07/21/2022 17:56:31",
      "content": "<p>Thank you <a href=\"https://www.kaggle.com/mbburabak\" target=\"_blank\">@mbburabak</a> </p>",
      "votes": null,
      "replies": []
    },
    {
      "id": 1871279,
      "author_name": "thedevastator",
      "author_url": "",
      "post_date": "07/26/2022 07:37:50",
      "content": "<p>How is it different then simply using the classic early stopping?</p>",
      "votes": null,
      "replies": [
        {
          "id": 1893815,
          "author_name": "kimberlynie",
          "author_url": "",
          "post_date": "08/11/2022 05:17:56",
          "content": "<p>I have the same question. Hope anyone could explain.</p>",
          "votes": null,
          "replies": []
        },
        {
          "id": 1893920,
          "author_name": "steubk",
          "author_url": "",
          "post_date": "08/11/2022 06:51:29",
          "content": "<p>Early stopping and dart cannot work well together:<br>\nwhen using dart,the previous trees will be updated at each iteration. Thus best iteration does not contain best trees as dart  updated the previous trees. Higher is stopping_rounds parameter more different is the best model at the end of training</p>",
          "votes": null,
          "replies": []
        }
      ]
    }
  ],
  "raw_markdown_by_id": {
    "1865234": "You can use this: \n\nlink: https://qiita.com/birdwatcher/items/78e0158957fec2d6c9e8\n\n\nExample:\n    \ndes = DartEarlyStopping(\"valid_1\", \"amex_metric\", 100)\nmodel = lgb.train(\n    params, data,\n    valid_sets=[data, eval_data], \n    num_boost_round=100,\n    callbacks=[des],\n    verbose_eval=1)\n\nclass DartEarlyStopping(object):\n\n\n    def __init__(self, data_name, monitor_metric, stopping_round):\n        self.data_name = data_name\n        self.monitor_metric = monitor_metric\n        self.stopping_round = stopping_round\n        self.best_score = None\n        self.best_model = None\n        self.best_score_list = []\n        self.best_iter = 0\n\n    def _is_higher_score(self, metric_score, is_higher_better):\n        if self.best_score is None:\n            return True\n        return (self.best_score < metric_score) if is_higher_better else (self.best_score > metric_score)\n\n    def _deepcopy(self, x):\n        # copy.deepcopyではlightgbmのモデルは完全にコピーされないためpickleを使用\n        return pickle.loads(pickle.dumps(x))\n\n    def __call__(self, env):\n        evals = env.evaluation_result_list\n        for data, metric, score, is_higher_better in evals:\n            if data != self.data_name or metric != self.monitor_metric:\n                continue\n            if not self._is_higher_score(score, is_higher_better):\n                if env.iteration - self.best_iter > self.stopping_round:\n                    # 終了させる\n                    eval_result_str = '\\t'.join([lgb.callback._format_eval_result(x) for x in self.best_score_list])\n                    lgb.basic._log_info(f\"Early stopping, best iteration is:\\n[{self.best_iter+1}]\\t{eval_result_str}\") \n                    lgb.basic._log_info(f\"You can get best model by \\\"DartEarlyStopping.best_model\\\"\")\n                    raise lgb.callback.EarlyStopException(self.best_iter, self.best_score_list)\n                return\n            # dartでは過去の木も更新されてしまうため、deepcopyしておく\n            self.best_model = self._deepcopy(env.model)\n            self.best_iter = env.iteration\n            self.best_score_list = evals\n            self.best_score = score\n            return\n        raise ValueError(\"monitoring metric not found\")",
    "1865335": "Thank you @mbburabak",
    "1871279": "How is it different then simply using the classic early stopping?",
    "1893815": "I have the same question. Hope anyone could explain.",
    "1893920": "Early stopping and dart cannot work well together:\nwhen using dart,the previous trees will be updated at each iteration. Thus best iteration does not contain best trees as dart  updated the previous trees. Higher is stopping_rounds parameter more different is the best model at the end of training"
  },
  "source": "meta"
}