[MISC] Consolidate cleanup() and refactor offline_inference_with_prefix.py (#9510)
This commit is contained in:
@@ -20,6 +20,7 @@ If you only need to use the distributed environment without model/pipeline
|
||||
steps.
|
||||
"""
|
||||
import contextlib
|
||||
import gc
|
||||
import pickle
|
||||
import weakref
|
||||
from collections import namedtuple
|
||||
@@ -36,7 +37,7 @@ from torch.distributed import Backend, ProcessGroup
|
||||
import vllm.envs as envs
|
||||
from vllm.logger import init_logger
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.utils import supports_custom_op
|
||||
from vllm.utils import is_cpu, supports_custom_op
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -1129,6 +1130,19 @@ def destroy_distributed_environment():
|
||||
torch.distributed.destroy_process_group()
|
||||
|
||||
|
||||
def cleanup_dist_env_and_memory(shutdown_ray: bool = False):
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
with contextlib.suppress(AssertionError):
|
||||
torch.distributed.destroy_process_group()
|
||||
if shutdown_ray:
|
||||
import ray # Lazy import Ray
|
||||
ray.shutdown()
|
||||
gc.collect()
|
||||
if not is_cpu():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
def in_the_same_node_as(pg: ProcessGroup, source_rank: int = 0) -> List[bool]:
|
||||
"""
|
||||
This is a collective operation that returns if each rank is in the same node
|
||||
|
||||
Reference in New Issue
Block a user