初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD Source: Original Platform
This commit is contained in:
94
slime/utils/misc.py
Normal file
94
slime/utils/misc.py
Normal file
@@ -0,0 +1,94 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import importlib
|
||||
import subprocess
|
||||
|
||||
import ray
|
||||
|
||||
from slime.utils.http_utils import is_port_available
|
||||
|
||||
|
||||
def load_function(path):
|
||||
"""
|
||||
Load a function from a module.
|
||||
:param path: The path to the function, e.g. "module.submodule.function".
|
||||
:return: The function object.
|
||||
"""
|
||||
module_path, _, attr = path.rpartition(".")
|
||||
module = importlib.import_module(module_path)
|
||||
return getattr(module, attr)
|
||||
|
||||
|
||||
class SingletonMeta(type):
|
||||
"""
|
||||
A metaclass for creating singleton classes.
|
||||
"""
|
||||
|
||||
_instances = {}
|
||||
|
||||
def __call__(cls, *args, **kwargs):
|
||||
if cls not in cls._instances:
|
||||
instance = super().__call__(*args, **kwargs)
|
||||
cls._instances[cls] = instance
|
||||
return cls._instances[cls]
|
||||
|
||||
|
||||
def exec_command(cmd: str, capture_output: bool = False) -> str | None:
|
||||
print(f"EXEC: {cmd}", flush=True)
|
||||
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["bash", "-c", cmd],
|
||||
shell=False,
|
||||
check=True,
|
||||
capture_output=capture_output,
|
||||
**(dict(text=True) if capture_output else {}),
|
||||
)
|
||||
except subprocess.CalledProcessError as e:
|
||||
if capture_output:
|
||||
print(f"{e.stdout=} {e.stderr=}")
|
||||
raise
|
||||
|
||||
if capture_output:
|
||||
print(f"Captured stdout={result.stdout} stderr={result.stderr}")
|
||||
return result.stdout
|
||||
|
||||
|
||||
def get_current_node_ip():
|
||||
address = ray._private.services.get_node_ip_address()
|
||||
# strip ipv6 address
|
||||
address = address.strip("[]")
|
||||
return address
|
||||
|
||||
|
||||
def get_free_port(start_port=10000, consecutive=1):
|
||||
# find the port where port, port + 1, port + 2, ... port + consecutive - 1 are all available
|
||||
port = start_port
|
||||
while not all(is_port_available(port + i) for i in range(consecutive)):
|
||||
port += 1
|
||||
return port
|
||||
|
||||
|
||||
def should_run_periodic_action(
|
||||
rollout_id: int,
|
||||
interval: int | None,
|
||||
num_rollout_per_epoch: int | None = None,
|
||||
num_rollout: int | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Return True when a periodic action (eval/save/checkpoint) should run.
|
||||
|
||||
Args:
|
||||
rollout_id: The current rollout index (0-based).
|
||||
interval: Desired cadence; disables checks when None.
|
||||
num_rollout_per_epoch: Optional epoch boundary to treat as a trigger.
|
||||
"""
|
||||
if interval is None:
|
||||
return False
|
||||
|
||||
if num_rollout is not None and rollout_id == num_rollout - 1:
|
||||
return True
|
||||
|
||||
step = rollout_id + 1
|
||||
return (step % interval == 0) or (num_rollout_per_epoch is not None and step % num_rollout_per_epoch == 0)
|
||||
Reference in New Issue
Block a user