balanced offload improvements

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2024-12-01 15:34:25 -05:00
parent 507636d0a1
commit 023b13b6cb
9 changed files with 80 additions and 54 deletions
+12 -8
View File
@@ -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
View File
@@ -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):
+6 -5
View File
@@ -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
View File
@@ -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 }
+2
View File
@@ -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)
+21 -17
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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):