mirror of
https://github.com/vladmandic/automatic
synced 2026-09-13 18:18:44 +02:00
balanced offload improvements
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
+12
-8
@@ -73,16 +73,20 @@ def wrap_gradio_call(func, extra_outputs=None, add_stats=False, name=None):
|
||||
elapsed_m = int(elapsed // 60)
|
||||
elapsed_s = elapsed % 60
|
||||
elapsed_text = f"{elapsed_m}m {elapsed_s:.2f}s" if elapsed_m > 0 else f"{elapsed_s:.2f}s"
|
||||
summary = timer.process.summary(min_time=0.1, total=False).replace('=', ' ')
|
||||
vram_html = ''
|
||||
summary = timer.process.summary(min_time=0.25, total=False).replace('=', ' ')
|
||||
gpu = ''
|
||||
cpu = ''
|
||||
if not shared.mem_mon.disabled:
|
||||
vram = {k: -(v//-(1024*1024)) for k, v in shared.mem_mon.read().items()}
|
||||
used = round(100 * vram['used'] / (vram['total'] + 0.001))
|
||||
if vram.get('active_peak', 0) > 0:
|
||||
vram_html = " | "
|
||||
vram_html += f"GPU {max(vram['active_peak'], vram['reserved_peak'])} MB {used}%"
|
||||
vram_html += f" | retries {vram['retries']} oom {vram['oom']}" if vram.get('retries', 0) > 0 or vram.get('oom', 0) > 0 else ''
|
||||
peak = max(vram['active_peak'], vram['reserved_peak'], vram['used'])
|
||||
used = round(100.0 * peak / vram['total']) if vram['total'] > 0 else 0
|
||||
if used > 0:
|
||||
gpu += f"| GPU {peak} MB {used}%"
|
||||
gpu += f" | retries {vram['retries']} oom {vram['oom']}" if vram.get('retries', 0) > 0 or vram.get('oom', 0) > 0 else ''
|
||||
ram = shared.ram_stats()
|
||||
if ram['used'] > 0:
|
||||
cpu += f"| RAM {ram['used']} GB {round(100.0 * ram['used'] / ram['total'])}%"
|
||||
if isinstance(res, list):
|
||||
res[-1] += f"<div class='performance'><p>Time: {elapsed_text} | {summary}{vram_html}</p></div>"
|
||||
res[-1] += f"<div class='performance'><p>Time: {elapsed_text} | {summary} {gpu} {cpu}</p></div>"
|
||||
return tuple(res)
|
||||
return f
|
||||
|
||||
+2
-2
@@ -224,7 +224,7 @@ def torch_gc(force=False, fast=False):
|
||||
timer.process.records['gc'] = 0
|
||||
timer.process.records['gc'] += t1 - t0
|
||||
if not force or collected == 0:
|
||||
return used_gpu
|
||||
return used_gpu, used_ram
|
||||
mem = memstats.memory_stats()
|
||||
saved = round(gpu.get('used', 0) - mem.get('gpu', {}).get('used', 0), 2)
|
||||
before = { 'gpu': gpu.get('used', 0), 'ram': ram.get('used', 0) }
|
||||
@@ -233,7 +233,7 @@ def torch_gc(force=False, fast=False):
|
||||
results = { 'collected': collected, 'saved': saved }
|
||||
fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access
|
||||
log.debug(f'GC: utilization={utilization} gc={results} before={before} after={after} device={torch.device(get_optimal_device_name())} fn={fn} time={round(t1 - t0, 2)}') # pylint: disable=protected-access
|
||||
return used_gpu
|
||||
return used_gpu, used_ram
|
||||
|
||||
|
||||
def set_cuda_sync_mode(mode):
|
||||
|
||||
@@ -447,8 +447,8 @@ def network_process():
|
||||
if component is not None and hasattr(component, 'named_modules'):
|
||||
modules += list(component.named_modules())
|
||||
if len(loaded_networks) > 0:
|
||||
pbar = rp.Progress(rp.TextColumn('[cyan]{task.description}'), rp.BarColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), console=shared.console)
|
||||
task = pbar.add_task(description='Apply network: type=LoRA' , total=len(modules))
|
||||
pbar = rp.Progress(rp.TextColumn('[cyan]Apply network: type=LoRA'), rp.BarColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), rp.TextColumn('[cyan]{task.description}'), console=shared.console)
|
||||
task = pbar.add_task(description='' , total=len(modules))
|
||||
else:
|
||||
task = None
|
||||
pbar = nullcontext()
|
||||
@@ -463,7 +463,8 @@ def network_process():
|
||||
current_names = getattr(module, "network_current_names", ())
|
||||
if shared.state.interrupted or network_layer_name is None or current_names == wanted_names:
|
||||
continue
|
||||
weight = module.weight.to(devices.device, non_blocking=True) if hasattr(module, 'weight') else None
|
||||
weight = getattr(module, 'weight', None)
|
||||
weight = weight.to(devices.device, non_blocking=True) if weight is not None else None
|
||||
backup_size += network_backup_weights(module, weight, network_layer_name, wanted_names)
|
||||
batch_updown, batch_ex_bias = network_calc_weights(module, weight, network_layer_name)
|
||||
del weight
|
||||
@@ -472,13 +473,13 @@ def network_process():
|
||||
weights_dtypes.append(weights_dtype)
|
||||
module.network_current_names = wanted_names
|
||||
if task is not None:
|
||||
pbar.update(task, advance=1) # progress bar becomes visible if operation takes more than 1sec
|
||||
pbar.update(task, advance=1, description=f'networks={len(loaded_networks)} modules={len(modules)} apply={applied} backup={backup_size}')
|
||||
if batch_updown is not None or batch_ex_bias is not None:
|
||||
applied += 1
|
||||
# pbar.remove_task(task)
|
||||
weights_devices, weights_dtypes = list(set([x for x in weights_devices if x is not None])), list(set([x for x in weights_dtypes if x is not None])) # noqa: C403
|
||||
if debug and len(loaded_networks) > 0:
|
||||
shared.log.debug(f'Load network: type=LoRA modules={len(modules)} networks={len(loaded_networks)} apply={applied} device={weights_devices} dtype={weights_dtypes} backup={backup_size} time={get_timers()}')
|
||||
shared.log.debug(f'Load network: type=LoRA networks={len(loaded_networks)} modules={len(modules)} apply={applied} device={weights_devices} dtype={weights_dtypes} backup={backup_size} time={get_timers()}')
|
||||
modules.clear()
|
||||
if shared.opts.diffusers_offload_mode == "sequential":
|
||||
sd_models.set_diffuser_offload(sd_model, op="model")
|
||||
|
||||
+15
-3
@@ -5,11 +5,12 @@ from modules import shared, errors
|
||||
|
||||
fail_once = False
|
||||
|
||||
def gb(val: float):
|
||||
return round(val / 1024 / 1024 / 1024, 2)
|
||||
|
||||
|
||||
def memory_stats():
|
||||
global fail_once # pylint: disable=global-statement
|
||||
def gb(val: float):
|
||||
return round(val / 1024 / 1024 / 1024, 2)
|
||||
|
||||
mem = {}
|
||||
try:
|
||||
process = psutil.Process(os.getpid())
|
||||
@@ -38,3 +39,14 @@ def memory_stats():
|
||||
except Exception:
|
||||
pass
|
||||
return mem
|
||||
|
||||
|
||||
def ram_stats():
|
||||
try:
|
||||
process = psutil.Process(os.getpid())
|
||||
res = process.memory_info()
|
||||
ram_total = 100 * res.rss / process.memory_percent()
|
||||
ram = { 'used': gb(res.rss), 'total': gb(ram_total) }
|
||||
return ram
|
||||
except Exception:
|
||||
return { 'used': 0, 'total': 0 }
|
||||
|
||||
@@ -83,6 +83,8 @@ def process_base(p: processing.StableDiffusionProcessing):
|
||||
try:
|
||||
t0 = time.time()
|
||||
sd_models_compile.check_deepcache(enable=True)
|
||||
if shared.opts.diffusers_offload_mode == "balanced":
|
||||
shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model)
|
||||
sd_models.move_model(shared.sd_model, devices.device)
|
||||
if hasattr(shared.sd_model, 'unet'):
|
||||
sd_models.move_model(shared.sd_model.unet, devices.device)
|
||||
|
||||
@@ -39,8 +39,8 @@ def prepare_model(pipe = None):
|
||||
pipe = pipe.pipe
|
||||
if not hasattr(pipe, "text_encoder"):
|
||||
return None
|
||||
if shared.opts.diffusers_offload_mode == "balanced":
|
||||
pipe = sd_models.apply_balanced_offload(pipe)
|
||||
# if shared.opts.diffusers_offload_mode == "balanced":
|
||||
# pipe = sd_models.apply_balanced_offload(pipe)
|
||||
elif hasattr(pipe, "maybe_free_model_hooks"):
|
||||
pipe.maybe_free_model_hooks()
|
||||
devices.torch_gc()
|
||||
@@ -79,8 +79,8 @@ class PromptEmbedder:
|
||||
self.scheduled_encode(pipe, batchidx)
|
||||
else:
|
||||
self.encode(pipe, prompt, negative_prompt, batchidx)
|
||||
if shared.opts.diffusers_offload_mode == "balanced":
|
||||
pipe = sd_models.apply_balanced_offload(pipe)
|
||||
# if shared.opts.diffusers_offload_mode == "balanced":
|
||||
# pipe = sd_models.apply_balanced_offload(pipe)
|
||||
self.checkcache(p)
|
||||
debug(f"Prompt encode: time={(time.time() - t0):.3f}")
|
||||
|
||||
@@ -199,8 +199,6 @@ class PromptEmbedder:
|
||||
|
||||
|
||||
def compel_hijack(self, token_ids: torch.Tensor, attention_mask: typing.Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||
if not devices.same_device(self.text_encoder.device, devices.device):
|
||||
sd_models.move_model(self.text_encoder, devices.device)
|
||||
needs_hidden_states = self.returned_embeddings_type != 1
|
||||
text_encoder_output = self.text_encoder(token_ids, attention_mask, output_hidden_states=needs_hidden_states, return_dict=True)
|
||||
|
||||
@@ -377,25 +375,31 @@ def prepare_embedding_providers(pipe, clip_skip) -> list[EmbeddingsProvider]:
|
||||
embedding_type = -(clip_skip + 1)
|
||||
else:
|
||||
embedding_type = clip_skip
|
||||
embedding_args = {
|
||||
'truncate': False,
|
||||
'returned_embeddings_type': embedding_type,
|
||||
'device': device,
|
||||
'dtype_for_device_getter': lambda device: devices.dtype,
|
||||
}
|
||||
if getattr(pipe, "prior_pipe", None) is not None and getattr(pipe.prior_pipe, "tokenizer", None) is not None and getattr(pipe.prior_pipe, "text_encoder", None) is not None:
|
||||
provider = EmbeddingsProvider(padding_attention_mask_value=0, tokenizer=pipe.prior_pipe.tokenizer, text_encoder=pipe.prior_pipe.text_encoder, truncate=False, returned_embeddings_type=embedding_type, device=device)
|
||||
provider = EmbeddingsProvider(padding_attention_mask_value=0, tokenizer=pipe.prior_pipe.tokenizer, text_encoder=pipe.prior_pipe.text_encoder, **embedding_args)
|
||||
embeddings_providers.append(provider)
|
||||
no_mask_provider = EmbeddingsProvider(padding_attention_mask_value=1 if "sote" in pipe.sd_checkpoint_info.name.lower() else 0, tokenizer=pipe.prior_pipe.tokenizer, text_encoder=pipe.prior_pipe.text_encoder, truncate=False, returned_embeddings_type=embedding_type, device=device)
|
||||
no_mask_provider = EmbeddingsProvider(padding_attention_mask_value=1 if "sote" in pipe.sd_checkpoint_info.name.lower() else 0, tokenizer=pipe.prior_pipe.tokenizer, text_encoder=pipe.prior_pipe.text_encoder, **embedding_args)
|
||||
embeddings_providers.append(no_mask_provider)
|
||||
elif getattr(pipe, "tokenizer", None) is not None and getattr(pipe, "text_encoder", None) is not None:
|
||||
if not devices.same_device(pipe.text_encoder.device, devices.device):
|
||||
sd_models.move_model(pipe.text_encoder, devices.device)
|
||||
provider = EmbeddingsProvider(tokenizer=pipe.tokenizer, text_encoder=pipe.text_encoder, truncate=False, returned_embeddings_type=embedding_type, device=device)
|
||||
if pipe.text_encoder.__class__.__name__.startswith('CLIP'):
|
||||
sd_models.move_model(pipe.text_encoder, devices.device, force=True)
|
||||
provider = EmbeddingsProvider(tokenizer=pipe.tokenizer, text_encoder=pipe.text_encoder, **embedding_args)
|
||||
embeddings_providers.append(provider)
|
||||
if getattr(pipe, "tokenizer_2", None) is not None and getattr(pipe, "text_encoder_2", None) is not None:
|
||||
if not devices.same_device(pipe.text_encoder_2.device, devices.device):
|
||||
sd_models.move_model(pipe.text_encoder_2, devices.device)
|
||||
provider = EmbeddingsProvider(tokenizer=pipe.tokenizer_2, text_encoder=pipe.text_encoder_2, truncate=False, returned_embeddings_type=embedding_type, device=device)
|
||||
if pipe.text_encoder_2.__class__.__name__.startswith('CLIP'):
|
||||
sd_models.move_model(pipe.text_encoder_2, devices.device, force=True)
|
||||
provider = EmbeddingsProvider(tokenizer=pipe.tokenizer_2, text_encoder=pipe.text_encoder_2, **embedding_args)
|
||||
embeddings_providers.append(provider)
|
||||
if getattr(pipe, "tokenizer_3", None) is not None and getattr(pipe, "text_encoder_3", None) is not None:
|
||||
if not devices.same_device(pipe.text_encoder_3.device, devices.device):
|
||||
sd_models.move_model(pipe.text_encoder_3, devices.device)
|
||||
provider = EmbeddingsProvider(tokenizer=pipe.tokenizer_3, text_encoder=pipe.text_encoder_3, truncate=False, returned_embeddings_type=embedding_type, device=device)
|
||||
if pipe.text_encoder_3.__class__.__name__.startswith('CLIP'):
|
||||
sd_models.move_model(pipe.text_encoder_3, devices.device, force=True)
|
||||
provider = EmbeddingsProvider(tokenizer=pipe.tokenizer_3, text_encoder=pipe.text_encoder_3, **embedding_args)
|
||||
embeddings_providers.append(provider)
|
||||
return embeddings_providers
|
||||
|
||||
|
||||
+9
-11
@@ -401,7 +401,6 @@ class OffloadHook(accelerate.hooks.ModelHook):
|
||||
module._hf_hook.execution_device = torch.device(devices.device) # pylint: disable=protected-access
|
||||
module.balanced_offload_device_map = device_map
|
||||
module.balanced_offload_max_memory = max_memory
|
||||
module.balanced_offload_active = True
|
||||
return args, kwargs
|
||||
|
||||
def post_forward(self, module, output):
|
||||
@@ -429,7 +428,7 @@ def apply_balanced_offload(sd_model):
|
||||
checkpoint_name = sd_model.__class__.__name__
|
||||
|
||||
def apply_balanced_offload_to_module(pipe):
|
||||
used_gpu = devices.torch_gc(fast=True)
|
||||
used_gpu, used_ram = devices.torch_gc(fast=True)
|
||||
if hasattr(pipe, "pipe"):
|
||||
apply_balanced_offload_to_module(pipe.pipe)
|
||||
if hasattr(pipe, "_internal_dict"):
|
||||
@@ -438,20 +437,21 @@ def apply_balanced_offload(sd_model):
|
||||
keys = get_signature(pipe).keys()
|
||||
for module_name in keys: # pylint: disable=protected-access
|
||||
module = getattr(pipe, module_name, None)
|
||||
balanced_offload_active = getattr(module, "balanced_offload_active", None)
|
||||
if isinstance(module, torch.nn.Module) and (balanced_offload_active is None or balanced_offload_active):
|
||||
if isinstance(module, torch.nn.Module):
|
||||
network_layer_name = getattr(module, "network_layer_name", None)
|
||||
device_map = getattr(module, "balanced_offload_device_map", None)
|
||||
max_memory = getattr(module, "balanced_offload_max_memory", None)
|
||||
module = accelerate.hooks.remove_hook_from_module(module, recurse=True)
|
||||
try:
|
||||
if used_gpu > 100 * shared.opts.diffusers_offload_min_gpu_memory:
|
||||
debug_move(f'Balanced offload: gpu={used_gpu} ram={used_ram} current={module.device} target={devices.cpu} component={module.__class__.__name__}')
|
||||
module = module.to(devices.cpu, non_blocking=True)
|
||||
used_gpu = devices.torch_gc(fast=True)
|
||||
used_gpu, used_ram = devices.torch_gc(fast=True)
|
||||
else:
|
||||
debug_move(f'Balanced offload: gpu={used_gpu} ram={used_ram} current={module.device} target={devices.cpu} component={module.__class__.__name__}')
|
||||
module.offload_dir = os.path.join(shared.opts.accelerate_offload_path, checkpoint_name, module_name)
|
||||
module = accelerate.hooks.add_hook_to_module(module, offload_hook_instance, append=True)
|
||||
module._hf_hook.execution_device = torch.device(devices.device) # pylint: disable=protected-access
|
||||
module.balanced_offload_active = False
|
||||
if network_layer_name:
|
||||
module.network_layer_name = network_layer_name
|
||||
if device_map and max_memory:
|
||||
@@ -515,13 +515,13 @@ def move_model(model, device=None, force=False):
|
||||
shared.log.error(f'Model move execution device: device={device} {e}')
|
||||
if getattr(model, 'has_accelerate', False) and not force:
|
||||
return
|
||||
if hasattr(model, "device") and devices.normalize_device(model.device) == devices.normalize_device(device):
|
||||
if hasattr(model, "device") and devices.normalize_device(model.device) == devices.normalize_device(device) and not force:
|
||||
return
|
||||
try:
|
||||
t0 = time.time()
|
||||
try:
|
||||
if hasattr(model, 'to'):
|
||||
model.to(device)
|
||||
model.to(device, non_blocking=True)
|
||||
if hasattr(model, "prior_pipe"):
|
||||
model.prior_pipe.to(device)
|
||||
except Exception as e0:
|
||||
@@ -551,7 +551,7 @@ def move_model(model, device=None, force=False):
|
||||
if 'move' not in process_timer.records:
|
||||
process_timer.records['move'] = 0
|
||||
process_timer.records['move'] += t1 - t0
|
||||
if os.environ.get('SD_MOVE_DEBUG', None) or (t1-t0) > 1:
|
||||
if os.environ.get('SD_MOVE_DEBUG', None) or (t1-t0) > 2:
|
||||
shared.log.debug(f'Model move: device={device} class={model.__class__.__name__} accelerate={getattr(model, "has_accelerate", False)} fn={fn} time={t1-t0:.2f}') # pylint: disable=protected-access
|
||||
devices.torch_gc()
|
||||
|
||||
@@ -1492,8 +1492,6 @@ def disable_offload(sd_model):
|
||||
module = getattr(sd_model, module_name, None)
|
||||
if isinstance(module, torch.nn.Module):
|
||||
network_layer_name = getattr(module, "network_layer_name", None)
|
||||
if getattr(module, "balanced_offload_active", None) is not None:
|
||||
module.balanced_offload_active = None
|
||||
module = remove_hook_from_module(module, recurse=True)
|
||||
if network_layer_name:
|
||||
module.network_layer_name = network_layer_name
|
||||
|
||||
+1
-1
@@ -20,7 +20,7 @@ from modules import errors, devices, shared_items, shared_state, cmd_args, theme
|
||||
from modules.paths import models_path, script_path, data_path, sd_configs_path, sd_default_config, sd_model_file, default_sd_model_file, extensions_dir, extensions_builtin_dir # pylint: disable=W0611
|
||||
from modules.dml import memory_providers, default_memory_provider, directml_do_hijack
|
||||
from modules.onnx_impl import initialize_onnx, execution_providers
|
||||
from modules.memstats import memory_stats
|
||||
from modules.memstats import memory_stats, ram_stats
|
||||
from modules.ui_components import DropdownEditable
|
||||
import modules.interrogate
|
||||
import modules.memmon
|
||||
|
||||
+12
-7
@@ -29,15 +29,20 @@ def return_stats(t: float = None):
|
||||
elapsed_m = int(elapsed // 60)
|
||||
elapsed_s = elapsed % 60
|
||||
elapsed_text = f"Time: {elapsed_m}m {elapsed_s:.2f}s |" if elapsed_m > 0 else f"Time: {elapsed_s:.2f}s |"
|
||||
summary = timer.process.summary(min_time=0.1, total=False).replace('=', ' ')
|
||||
vram_html = ''
|
||||
summary = timer.process.summary(min_time=0.25, total=False).replace('=', ' ')
|
||||
gpu = ''
|
||||
cpu = ''
|
||||
if not shared.mem_mon.disabled:
|
||||
vram = {k: -(v//-(1024*1024)) for k, v in shared.mem_mon.read().items()}
|
||||
used = round(100 * vram['used'] / (vram['total'] + 0.001))
|
||||
if vram.get('active_peak', 0) > 0:
|
||||
vram_html += f"| GPU {max(vram['active_peak'], vram['reserved_peak'])} MB {used}%"
|
||||
vram_html += f" | retries {vram['retries']} oom {vram['oom']}" if vram.get('retries', 0) > 0 or vram.get('oom', 0) > 0 else ''
|
||||
return f"<div class='performance'><p>{elapsed_text} {summary} {vram_html}</p></div>"
|
||||
peak = max(vram['active_peak'], vram['reserved_peak'], vram['used'])
|
||||
used = round(100.0 * peak / vram['total']) if vram['total'] > 0 else 0
|
||||
if used > 0:
|
||||
gpu += f"| GPU {peak} MB {used}%"
|
||||
gpu += f" | retries {vram['retries']} oom {vram['oom']}" if vram.get('retries', 0) > 0 or vram.get('oom', 0) > 0 else ''
|
||||
ram = shared.ram_stats()
|
||||
if ram['used'] > 0:
|
||||
cpu += f"| RAM {ram['used']} GB {round(100.0 * ram['used'] / ram['total'])}%"
|
||||
return f"<div class='performance'><p>Time: {elapsed_text} | {summary} {gpu} {cpu}</p></div>"
|
||||
|
||||
|
||||
def return_controls(res, t: float = None):
|
||||
|
||||
Reference in New Issue
Block a user