mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
cleanup
This commit is contained in:
@@ -871,7 +871,7 @@ def get_diffusers_task(pipe: diffusers.DiffusionPipeline) -> DiffusersTaskType:
|
||||
return DiffusersTaskType.TEXT_2_IMAGE
|
||||
|
||||
|
||||
def switch_pipe(cls: diffusers.DiffusionPipeline, pipeline: diffusers.DiffusionPipeline = None, force = False, args = {}):
|
||||
def switch_pipe(cls: diffusers.DiffusionPipeline, pipeline: diffusers.DiffusionPipeline = None, force = False, args: dict = None):
|
||||
"""
|
||||
args:
|
||||
- cls: can be pipeline class or a string from custom pipelines
|
||||
@@ -881,6 +881,8 @@ def switch_pipe(cls: diffusers.DiffusionPipeline, pipeline: diffusers.DiffusionP
|
||||
for example: { 'vae': None }
|
||||
"""
|
||||
try:
|
||||
if args is None:
|
||||
args = {}
|
||||
if isinstance(cls, str):
|
||||
shared.log.debug(f'Pipeline switch: custom={cls}')
|
||||
cls = diffusers.utils.get_class_from_dynamic_module(cls, module_file='pipeline.py')
|
||||
|
||||
@@ -288,7 +288,9 @@ def get_pipe_variants(pipe=None):
|
||||
return variants
|
||||
|
||||
|
||||
def get_module_names(pipe=None, exclude=[]):
|
||||
def get_module_names(pipe=None, exclude=None):
|
||||
if exclude is None:
|
||||
exclude = []
|
||||
if pipe is None:
|
||||
if shared.sd_loaded:
|
||||
pipe = shared.sd_model
|
||||
@@ -304,7 +306,9 @@ def get_module_names(pipe=None, exclude=[]):
|
||||
return modules_names
|
||||
|
||||
|
||||
def get_module_sizes(pipe=None, exclude=[]):
|
||||
def get_module_sizes(pipe=None, exclude=None):
|
||||
if exclude is None:
|
||||
exclude = []
|
||||
modules = {}
|
||||
for module_name in get_module_names(pipe, exclude):
|
||||
module_size = offload_hook_instance.offload_map.get(module_name, None)
|
||||
@@ -409,7 +413,7 @@ def report_model_stats(module_name, module):
|
||||
shared.log.error(f'Module stats: name={module_name} {e}')
|
||||
|
||||
|
||||
def apply_balanced_offload(sd_model=None, exclude=[], force=False):
|
||||
def apply_balanced_offload(sd_model=None, exclude=None, force=False):
|
||||
global offload_hook_instance # pylint: disable=global-statement
|
||||
if shared.opts.diffusers_offload_mode != "balanced":
|
||||
return sd_model
|
||||
@@ -419,6 +423,8 @@ def apply_balanced_offload(sd_model=None, exclude=[], force=False):
|
||||
sd_model = shared.sd_model
|
||||
if sd_model is None:
|
||||
return sd_model
|
||||
if exclude is None:
|
||||
exclude = []
|
||||
t0 = time.time()
|
||||
if sd_model.__class__.__name__ in balanced_offload_exclude:
|
||||
return sd_model
|
||||
|
||||
Reference in New Issue
Block a user