mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-16 01:36:41 +08:00
Add fsdp
This commit is contained in:
parent
c9ebe70072
commit
1a8221507b
@ -214,6 +214,9 @@ database_default_path = os.path.abspath(
|
|||||||
)
|
)
|
||||||
parser.add_argument("--database-url", type=str, default=f"sqlite:///{database_default_path}", help="Specify the database URL, e.g. for an in-memory database you can use 'sqlite:///:memory:'.")
|
parser.add_argument("--database-url", type=str, default=f"sqlite:///{database_default_path}", help="Specify the database URL, e.g. for an in-memory database you can use 'sqlite:///:memory:'.")
|
||||||
|
|
||||||
|
parser.add_argument("--fsdp", action="store_true", help="Enable Fully Sharded Data Parallel (FSDP) for distributed inference.")
|
||||||
|
parser.add_argument("--world-size", type=int, default=1, help="Number of processes to use for distributed inference. Default is 1.")
|
||||||
|
|
||||||
if comfy.options.args_parsing:
|
if comfy.options.args_parsing:
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
else:
|
else:
|
||||||
|
|||||||
@ -424,6 +424,47 @@ 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():
|
||||||
|
if args.fsdp:
|
||||||
|
fsdp = True
|
||||||
|
return fsdp
|
||||||
|
|
||||||
|
def get_world_size():
|
||||||
|
if args.world_size:
|
||||||
|
return args.world_size
|
||||||
|
return 1
|
||||||
|
|
||||||
|
def init_distributed():
|
||||||
|
world_size = get_world_size()
|
||||||
|
|
||||||
|
if world_size > 1 and not torch.distributed.is_initialized():
|
||||||
|
torch.distributed.init_process_group(backend="nccl", world_size=world_size)
|
||||||
|
|
||||||
|
def get_distributed_model(
|
||||||
|
model,
|
||||||
|
param_dtype=torch.bfloat16,
|
||||||
|
reduce_dtype=torch.float32,
|
||||||
|
buffer_dtype=torch.float32,
|
||||||
|
sync_module_states=True,
|
||||||
|
):
|
||||||
|
from torch.distributed.fsdp import (
|
||||||
|
FullyShardedDataParallel,
|
||||||
|
ShardingStrategy,
|
||||||
|
MixedPrecision,
|
||||||
|
)
|
||||||
|
|
||||||
|
model = FullyShardedDataParallel(
|
||||||
|
model,
|
||||||
|
sharding_strategy=ShardingStrategy.FULL_SHARD,
|
||||||
|
mixed_precision=MixedPrecision(
|
||||||
|
param_dtype=param_dtype,
|
||||||
|
reduce_dtype=reduce_dtype,
|
||||||
|
buffer_dtype=buffer_dtype,
|
||||||
|
),
|
||||||
|
sync_module_states=sync_module_states,
|
||||||
|
)
|
||||||
|
return model
|
||||||
|
|
||||||
class LoadedModel:
|
class LoadedModel:
|
||||||
def __init__(self, model):
|
def __init__(self, model):
|
||||||
self._set_model(model)
|
self._set_model(model)
|
||||||
@ -434,7 +475,12 @@ class LoadedModel:
|
|||||||
self._patcher_finalizer = None
|
self._patcher_finalizer = None
|
||||||
|
|
||||||
def _set_model(self, model):
|
def _set_model(self, model):
|
||||||
self._model = weakref.ref(model)
|
if is_fsdp():
|
||||||
|
init_distributed()
|
||||||
|
dist_model = get_distributed_model(model)
|
||||||
|
self._model = weakref.ref(dist_model)
|
||||||
|
else:
|
||||||
|
self._model = weakref.ref(model)
|
||||||
if model.parent is not None:
|
if model.parent is not None:
|
||||||
self._parent_model = weakref.ref(model.parent)
|
self._parent_model = weakref.ref(model.parent)
|
||||||
self._patcher_finalizer = weakref.finalize(model, self._switch_parent)
|
self._patcher_finalizer = weakref.finalize(model, self._switch_parent)
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user