Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions slime/ray/rollout.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
10 changes: 10 additions & 0 deletions slime/utils/logging_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
3 changes: 2 additions & 1 deletion train.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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__":
Expand Down
3 changes: 2 additions & 1 deletion train_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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__":
Expand Down