mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
Move new puzzle dist utils from feature/compress to main (#746)
- Move new `modelopt.torch.utils.distributed` from `feature/compress` to `main` branch so they can be used via modelopt in puzzletron gitlab Signed-off-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com>
This commit is contained in:
@@ -20,6 +20,8 @@ import io
|
||||
import os
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from contextlib import suppress
|
||||
from datetime import timedelta
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
@@ -70,11 +72,23 @@ def rank(group=None) -> int:
|
||||
return 0
|
||||
|
||||
|
||||
def local_rank() -> int:
|
||||
"""Returns the local rank of the current process."""
|
||||
if "LOCAL_RANK" in os.environ:
|
||||
return int(os.environ["LOCAL_RANK"])
|
||||
raise RuntimeError("LOCAL_RANK environment variable not found.")
|
||||
|
||||
|
||||
def is_master(group=None) -> bool:
|
||||
"""Returns whether the current process is the master process."""
|
||||
return rank(group=group) == 0
|
||||
|
||||
|
||||
def is_last_process(group=None) -> bool:
|
||||
"""Returns whether the current process is the last process."""
|
||||
return rank(group=group) == size(group=group) - 1
|
||||
|
||||
|
||||
def _serialize(obj: Any) -> torch.Tensor:
|
||||
buffer = io.BytesIO()
|
||||
torch.save(obj, buffer)
|
||||
@@ -184,6 +198,21 @@ def master_only(func):
|
||||
return wrapper
|
||||
|
||||
|
||||
def setup(timeout: timedelta | None = None):
|
||||
"""Sets up the distributed environment."""
|
||||
torch.cuda.set_device(local_rank())
|
||||
if not is_initialized():
|
||||
torch.distributed.init_process_group("cpu:gloo,cuda:nccl", timeout=timeout)
|
||||
|
||||
|
||||
def cleanup():
|
||||
"""Cleans up the distributed environment."""
|
||||
if is_initialized():
|
||||
with suppress(Exception):
|
||||
barrier()
|
||||
torch.distributed.destroy_process_group()
|
||||
|
||||
|
||||
class DistributedProcessGroup:
|
||||
"""A convenient wrapper around torch.distributed.ProcessGroup objects."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user