diff --git a/comfy/cli_args.py b/comfy/cli_args.py index 72eeaea9a..59adbe56c 100644 --- a/comfy/cli_args.py +++ b/comfy/cli_args.py @@ -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("--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: args = parser.parse_args() else: diff --git a/comfy/model_management.py b/comfy/model_management.py index d08aee1fe..2dae00f6c 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -424,6 +424,47 @@ def module_size(module): module_mem += t.nelement() * t.element_size() 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: def __init__(self, model): self._set_model(model) @@ -434,7 +475,12 @@ class LoadedModel: self._patcher_finalizer = None 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: self._parent_model = weakref.ref(model.parent) self._patcher_finalizer = weakref.finalize(model, self._switch_parent)