mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-16 01:36:41 +08:00
Add rank settings for fsdp
This commit is contained in:
parent
1a8221507b
commit
d006ee2553
@ -16,6 +16,7 @@
|
|||||||
along with this program. If not, see <https://www.gnu.org/licenses/>.
|
along with this program. If not, see <https://www.gnu.org/licenses/>.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
import psutil
|
import psutil
|
||||||
import logging
|
import logging
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
@ -424,7 +425,7 @@ def module_size(module):
|
|||||||
module_mem += t.nelement() * t.element_size()
|
module_mem += t.nelement() * t.element_size()
|
||||||
return module_mem
|
return module_mem
|
||||||
|
|
||||||
def is_fsdp():
|
def is_fsdp_enabled():
|
||||||
if args.fsdp:
|
if args.fsdp:
|
||||||
fsdp = True
|
fsdp = True
|
||||||
return fsdp
|
return fsdp
|
||||||
@ -438,7 +439,11 @@ def init_distributed():
|
|||||||
world_size = get_world_size()
|
world_size = get_world_size()
|
||||||
|
|
||||||
if world_size > 1 and not torch.distributed.is_initialized():
|
if world_size > 1 and not torch.distributed.is_initialized():
|
||||||
torch.distributed.init_process_group(backend="nccl", world_size=world_size)
|
rank = int(os.getenv("RANK"), 0)
|
||||||
|
local_rank = int(os.getenv("LOCAL_RANK"), 0)
|
||||||
|
torch.cuda.set_device(local_rank)
|
||||||
|
|
||||||
|
torch.distributed.init_process_group(backend="nccl", world_size=world_size, rank=rank)
|
||||||
|
|
||||||
def get_distributed_model(
|
def get_distributed_model(
|
||||||
model,
|
model,
|
||||||
@ -466,16 +471,22 @@ def get_distributed_model(
|
|||||||
return model
|
return model
|
||||||
|
|
||||||
class LoadedModel:
|
class LoadedModel:
|
||||||
def __init__(self, model):
|
def __init__(
|
||||||
|
self,
|
||||||
|
model,
|
||||||
|
# use_fsdp=False
|
||||||
|
):
|
||||||
self._set_model(model)
|
self._set_model(model)
|
||||||
self.device = model.load_device
|
self.device = model.load_device
|
||||||
self.real_model = None
|
self.real_model = None
|
||||||
self.currently_used = True
|
self.currently_used = True
|
||||||
self.model_finalizer = None
|
self.model_finalizer = None
|
||||||
self._patcher_finalizer = None
|
self._patcher_finalizer = None
|
||||||
|
# self.use_fsdp = use_fsdp
|
||||||
|
self.use_fsdp = is_fsdp_enabled()
|
||||||
|
|
||||||
def _set_model(self, model):
|
def _set_model(self, model):
|
||||||
if is_fsdp():
|
if self.use_fsdp and is_fsdp_enabled():
|
||||||
init_distributed()
|
init_distributed()
|
||||||
dist_model = get_distributed_model(model)
|
dist_model = get_distributed_model(model)
|
||||||
self._model = weakref.ref(dist_model)
|
self._model = weakref.ref(dist_model)
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user