diff --git a/CHANGELOG.md b/CHANGELOG.md index 265255431..45c20f9b2 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -15,6 +15,7 @@ in *settings -> networks -> network scan* comma-separate list of regex patterns to skip - fix debug logging +- add explicit offload after encode prompt ## Update for 2025-04-12 diff --git a/modules/cfgzero/__init__.py b/modules/cfgzero/__init__.py index ee2882c9c..f7d350ee3 100644 --- a/modules/cfgzero/__init__.py +++ b/modules/cfgzero/__init__.py @@ -44,7 +44,9 @@ def apply(p: processing.StableDiffusionProcessing): shared.sd_model = sd_models.switch_pipe(WanCFGZeroPipeline, shared.sd_model) if cls == 'HunyuanVideoPipeline': from modules.cfgzero.hunyuan_t2v_pipeline import HunyuanVideoCFGZeroPipeline + from modules.model_hidream import init_hijack shared.sd_model = sd_models.switch_pipe(HunyuanVideoCFGZeroPipeline, shared.sd_model) + init_hijack(shared.sd_model) shared.log.debug(f'Apply CFGZero: cls={cls} init={shared.opts.cfgzero_enabled} star={shared.opts.cfgzero_star} steps={shared.opts.cfgzero_steps}') p.task_args['use_zero_init'] = shared.opts.cfgzero_enabled diff --git a/modules/model_flux.py b/modules/model_flux.py index 7f1ae035d..5c0395cf9 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -5,7 +5,7 @@ import diffusers import transformers from safetensors.torch import load_file from huggingface_hub import hf_hub_download -from modules import shared, errors, devices, modelloader, sd_models, sd_unet, model_te, model_quant +from modules import shared, errors, devices, modelloader, sd_models, sd_unet, model_te, model_quant, sd_hijack_te debug = shared.log.trace if os.environ.get('SD_LOAD_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -123,34 +123,6 @@ def load_quants(kwargs, repo_id, cache_dir, allow_quant): return kwargs -""" -def load_flux_gguf(file_path): - transformer = None - ggml.install_gguf() - from accelerate import init_empty_weights - from diffusers.loaders.single_file_utils import convert_flux_transformer_checkpoint_to_diffusers - from modules import ggml, sd_hijack_accelerate - with init_empty_weights(): - config = diffusers.FluxTransformer2DModel.load_config(os.path.join('configs', 'flux'), subfolder="transformer") - transformer = diffusers.FluxTransformer2DModel.from_config(config).to(devices.dtype) - expected_state_dict_keys = list(transformer.state_dict().keys()) - state_dict, stats = ggml.load_gguf_state_dict(file_path, devices.dtype) - state_dict = convert_flux_transformer_checkpoint_to_diffusers(state_dict) - applied, skipped = 0, 0 - for param_name, param in state_dict.items(): - if param_name not in expected_state_dict_keys: - # shared.log.warning(f'Load model: type=Unet/Transformer param={param_name} unexpected') - skipped += 1 - continue - applied += 1 - sd_hijack_accelerate.hijack_set_module_tensor_simple(transformer, tensor_name=param_name, value=param, device=0) - transformer.gguf = 'gguf' - state_dict[param_name] = None - shared.log.debug(f'Load model: type=Unet/Transformer applied={applied} skipped={skipped} stats={stats}') - return transformer, None -""" - - def load_transformer(file_path): # triggered by opts.sd_unet change if file_path is None or not os.path.exists(file_path): return None @@ -345,5 +317,6 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch vae = None for k in kwargs.keys(): kwargs[k] = None + sd_hijack_te.init_hijack(pipe) devices.torch_gc(force=True) return pipe diff --git a/modules/model_hidream.py b/modules/model_hidream.py index 286d61408..c9702e7cf 100644 --- a/modules/model_hidream.py +++ b/modules/model_hidream.py @@ -1,20 +1,6 @@ -import os -import time import transformers import diffusers -from modules import shared, devices, sd_models, timer, model_quant, modelloader - - -def hijack_encode_prompt(*args, **kwargs): - t0 = time.time() - if 'max_sequence_length' in kwargs: - kwargs['max_sequence_length'] = os.environ.get('HIDREAM_MAX_SEQUENCE_LENGTH', 256) - res = shared.sd_model.orig_encode_prompt(*args, **kwargs) - t1 = time.time() - timer.process.add('te', t1-t0) - # shared.log.debug(f'Hijack: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - return res +from modules import shared, devices, sd_models, model_quant, modelloader, sd_hijack_te def load_hidream(checkpoint_info, diffusers_load_config={}): @@ -82,9 +68,7 @@ def load_hidream(checkpoint_info, diffusers_load_config={}): cache_dir=shared.opts.diffusers_dir, **load_args, ) - - pipe.orig_encode_prompt = pipe.encode_prompt - pipe.encode_prompt = hijack_encode_prompt + sd_hijack_te.init_hijack(pipe) devices.torch_gc() return pipe diff --git a/modules/sd_hijack_te.py b/modules/sd_hijack_te.py new file mode 100644 index 000000000..f80ac29ee --- /dev/null +++ b/modules/sd_hijack_te.py @@ -0,0 +1,26 @@ +import os +import time +from modules import shared, errors, timer, sd_models + + +def hijack_encode_prompt(*args, **kwargs): + t0 = time.time() + if 'max_sequence_length' in kwargs: + kwargs['max_sequence_length'] = max(kwargs['max_sequence_length'], os.environ.get('HIDREAM_MAX_SEQUENCE_LENGTH', 256)) + try: + res = shared.sd_model.orig_encode_prompt(*args, **kwargs) + except Exception as e: + shared.log.error(f'Eencode prompt: {e}') + errors.display(e, 'Video encode prompt') + res = None + t1 = time.time() + timer.process.add('te', t1-t0) + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) + return res + + +def init_hijack(pipe): + if pipe is not None and not hasattr(pipe, 'orig_encode_prompt') and hasattr(pipe, 'encode_prompt'): + shared.log.debug(f'Model: cls={pipe.__class__.__name__} hijack encode') + pipe.orig_encode_prompt = pipe.encode_prompt + pipe.encode_prompt = hijack_encode_prompt diff --git a/modules/sd_models.py b/modules/sd_models.py index 1ac396106..e06e56272 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -8,9 +8,8 @@ from enum import Enum import diffusers import diffusers.loaders.single_file_utils import torch - from installer import log -from modules import paths, shared, shared_state, shared_items, modelloader, devices, script_callbacks, sd_vae, sd_unet, errors, sd_models_config, sd_models_compile, sd_hijack_accelerate, sd_detect, model_quant +from modules import paths, shared, shared_state, shared_items, modelloader, devices, script_callbacks, sd_vae, sd_unet, errors, sd_models_config, sd_models_compile, sd_hijack_accelerate, sd_detect, model_quant, sd_hijack_te from modules.timer import Timer, process as process_timer from modules.memstats import memory_stats from modules.modeldata import model_data @@ -754,6 +753,7 @@ def switch_pipe(cls: diffusers.DiffusionPipeline, pipeline: diffusers.DiffusionP components_skipped.append(k) if new_pipe is not None: copy_diffuser_options(new_pipe, pipeline) + sd_hijack_te.init_hijack(new_pipe) if hasattr(new_pipe, "watermark"): new_pipe.watermark = NoWatermark() if switch_mode == 'auto': diff --git a/modules/video_models/video_load.py b/modules/video_models/video_load.py index 12f02caf1..33813df82 100644 --- a/modules/video_models/video_load.py +++ b/modules/video_models/video_load.py @@ -1,6 +1,6 @@ import os import time -from modules import shared, errors, sd_models, sd_checkpoint, model_quant, devices +from modules import shared, errors, sd_models, sd_checkpoint, model_quant, devices, sd_hijack_te from modules.video_models import models_def, video_utils, video_vae, video_overrides, video_cache @@ -80,7 +80,7 @@ def load_model(selected: models_def.Model): shared.sd_model.vae.encode = video_vae.hijack_vae_encode if selected.te_hijack and hasattr(shared.sd_model, 'encode_prompt'): shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt - shared.sd_model.encode_prompt = video_utils.hijack_encode_prompt + sd_hijack_te.init_hijack(shared.sd_model) if selected.image_hijack and hasattr(shared.sd_model, 'encode_image'): shared.sd_model.orig_encode_image = shared.sd_model.encode_image shared.sd_model.encode_image = video_utils.hijack_encode_image diff --git a/modules/video_models/video_utils.py b/modules/video_models/video_utils.py index 16edb070d..b9510dd36 100644 --- a/modules/video_models/video_utils.py +++ b/modules/video_models/video_utils.py @@ -24,22 +24,6 @@ def set_prompt(p): p.task_args['negative_prompt'] = p.negative_prompt -def hijack_encode_prompt(*args, **kwargs): - t0 = time.time() - try: - sd_models.move_model(shared.sd_model.text_encoder, devices.device) - res = shared.sd_model.orig_encode_prompt(*args, **kwargs) - except Exception as e: - shared.log.error(f'Video encode prompt: {e}') - errors.display(e, 'Video encode prompt') - res = None - t1 = time.time() - timer.process.add('te', t1-t0) - debug(f'Video encode prompt: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - return res - - def hijack_encode_image(*args, **kwargs): t0 = time.time() try: diff --git a/scripts/allegrovideo.py b/scripts/allegrovideo.py index f1d1ab45f..673627ddb 100644 --- a/scripts/allegrovideo.py +++ b/scripts/allegrovideo.py @@ -2,7 +2,7 @@ import time import gradio as gr import transformers import diffusers -from modules import scripts, processing, shared, images, devices, sd_models, sd_checkpoint, model_quant, timer +from modules import scripts, processing, shared, images, devices, sd_models, sd_checkpoint, model_quant, timer, sd_hijack_te repo_id = 'rhymes-ai/Allegro' @@ -19,16 +19,6 @@ def hijack_decode(*args, **kwargs): return res -def hijack_encode_prompt(*args, **kwargs): - t0 = time.time() - res = shared.sd_model.orig_encode_prompt(*args, **kwargs) - t1 = time.time() - timer.process.add('te', t1-t0) - shared.log.debug(f'Video: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - return res - - class Script(scripts.Script): def title(self): return 'Video: Allegro (Legacy)' @@ -94,9 +84,9 @@ class Script(scripts.Script): shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt shared.sd_model.vae.decode = hijack_decode - shared.sd_model.encode_prompt = hijack_encode_prompt shared.sd_model.vae.enable_tiling() # shared.sd_model.vae.enable_slicing() + sd_hijack_te.init_hijack(shared.sd_model) shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) devices.torch_gc(force=True) diff --git a/scripts/hunyuanvideo.py b/scripts/hunyuanvideo.py index 50cbf567f..0291c4124 100644 --- a/scripts/hunyuanvideo.py +++ b/scripts/hunyuanvideo.py @@ -3,7 +3,7 @@ import torch import gradio as gr import transformers import diffusers -from modules import scripts, processing, shared, images, devices, sd_models, sd_checkpoint, sd_samplers, model_quant, timer +from modules import scripts, processing, shared, images, devices, sd_models, sd_checkpoint, sd_samplers, model_quant, timer, sd_hijack_te default_template = """Describe the video by detailing the following aspects: @@ -48,16 +48,6 @@ def hijack_decode(*args, **kwargs): return res -def hijack_encode_prompt(*args, **kwargs): - t0 = time.time() - res = shared.sd_model.orig_encode_prompt(*args, **kwargs) - t1 = time.time() - timer.process.add('te', t1-t0) - shared.log.debug(f'Video: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - return res - - class Script(scripts.Script): def title(self): return 'Video: Hunyuan Video (Legacy)' @@ -135,10 +125,10 @@ class Script(scripts.Script): shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt shared.sd_model.vae.decode = hijack_decode - shared.sd_model.encode_prompt = hijack_encode_prompt shared.sd_model.vae.enable_slicing() shared.sd_model.vae.enable_tiling() shared.sd_model.vae.use_framewise_decoding = True + sd_hijack_te.init_hijack(shared.sd_model) loaded_model = model def run(self, p: processing.StableDiffusionProcessing, model, num_frames, tile_frames, override_scheduler, scheduler_shift, template, video_type, duration, gif_loop, mp4_pad, mp4_interpolate): # pylint: disable=arguments-differ, unused-argument diff --git a/scripts/legacy_allegrovideo.py b/scripts/legacy_allegrovideo.py index f1d1ab45f..1403e0765 100644 --- a/scripts/legacy_allegrovideo.py +++ b/scripts/legacy_allegrovideo.py @@ -2,7 +2,7 @@ import time import gradio as gr import transformers import diffusers -from modules import scripts, processing, shared, images, devices, sd_models, sd_checkpoint, model_quant, timer +from modules import scripts, processing, shared, images, devices, sd_models, sd_checkpoint, model_quant, timer, sd_hijack_te repo_id = 'rhymes-ai/Allegro' @@ -19,16 +19,6 @@ def hijack_decode(*args, **kwargs): return res -def hijack_encode_prompt(*args, **kwargs): - t0 = time.time() - res = shared.sd_model.orig_encode_prompt(*args, **kwargs) - t1 = time.time() - timer.process.add('te', t1-t0) - shared.log.debug(f'Video: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - return res - - class Script(scripts.Script): def title(self): return 'Video: Allegro (Legacy)' @@ -94,8 +84,8 @@ class Script(scripts.Script): shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt shared.sd_model.vae.decode = hijack_decode - shared.sd_model.encode_prompt = hijack_encode_prompt shared.sd_model.vae.enable_tiling() + sd_hijack_te.init_hijack(shared.sd_model) # shared.sd_model.vae.enable_slicing() shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) diff --git a/scripts/ltxvideo.py b/scripts/ltxvideo.py index 84b5953c4..43ea75595 100644 --- a/scripts/ltxvideo.py +++ b/scripts/ltxvideo.py @@ -4,7 +4,7 @@ import torch import gradio as gr import diffusers import transformers -from modules import scripts, processing, shared, images, devices, sd_models, sd_checkpoint, model_quant, timer +from modules import scripts, processing, shared, images, devices, sd_models, sd_checkpoint, model_quant, timer, sd_hijack_te repos = { @@ -39,16 +39,6 @@ def hijack_decode(*args, **kwargs): return res -def hijack_encode_prompt(*args, **kwargs): - t0 = time.time() - res = shared.sd_model.orig_encode_prompt(*args, **kwargs) - t1 = time.time() - timer.process.add('te', t1-t0) - shared.log.debug(f'Video: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}') - shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) - return res - - class Script(scripts.Script): def title(self): return 'Video: LTX Video (Legacy)' @@ -131,9 +121,10 @@ class Script(scripts.Script): shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt shared.sd_model.vae.decode = hijack_decode - shared.sd_model.encode_prompt = hijack_encode_prompt shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(repo_id) shared.sd_model.sd_model_hash = None + sd_hijack_te.init_hijack(shared.sd_model) + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) shared.sd_model.vae.enable_slicing() shared.sd_model.vae.enable_tiling()