diff --git a/slime/ray/rollout.py b/slime/ray/rollout.py index cea293a508..cd501bd10d 100644 --- a/slime/ray/rollout.py +++ b/slime/ray/rollout.py @@ -414,6 +414,7 @@ def _try_ci_fault_injection(self): def dispose(self): for monitor in self._health_monitors: monitor.stop() + logging_utils.finish_tracking(self.args) @property def server(self) -> RolloutServer | None: diff --git a/slime/utils/logging_utils.py b/slime/utils/logging_utils.py index 5fc0fad357..1fc3d94b23 100644 --- a/slime/utils/logging_utils.py +++ b/slime/utils/logging_utils.py @@ -31,6 +31,16 @@ def init_tracking(args, primary: bool = True, **kwargs): wandb_utils.init_wandb_secondary(args, **kwargs) +def finish_tracking(args): + if not args.use_wandb: + return + try: + if wandb.run is not None: + wandb.finish() + except Exception: + logging.getLogger(__name__).exception("Failed to finish wandb run") + + # TODO further refactor, e.g. put TensorBoard init to the "init" part def log(args, metrics, step_key: str): if args.use_wandb: diff --git a/train.py b/train.py index 15b3ca8b24..554b26b453 100644 --- a/train.py +++ b/train.py @@ -2,7 +2,7 @@ from slime.ray.placement_group import create_placement_groups, create_rollout_manager, create_training_models from slime.utils.arguments import parse_args -from slime.utils.logging_utils import configure_logger, init_tracking +from slime.utils.logging_utils import configure_logger, init_tracking, finish_tracking from slime.utils.misc import should_run_periodic_action @@ -98,6 +98,7 @@ def save(rollout_id): ray.get(rollout_manager.eval.remote(rollout_id)) ray.get(rollout_manager.dispose.remote()) + finish_tracking(args) if __name__ == "__main__": diff --git a/train_async.py b/train_async.py index 182309cb08..17f9cff47d 100644 --- a/train_async.py +++ b/train_async.py @@ -2,7 +2,7 @@ from slime.ray.placement_group import create_placement_groups, create_rollout_manager, create_training_models from slime.utils.arguments import parse_args -from slime.utils.logging_utils import configure_logger, init_tracking +from slime.utils.logging_utils import configure_logger, init_tracking, finish_tracking from slime.utils.misc import should_run_periodic_action @@ -72,6 +72,7 @@ def train(args): ray.get(rollout_manager.eval.remote(rollout_id)) ray.get(rollout_manager.dispose.remote()) + finish_tracking(args) if __name__ == "__main__":