general te hijack

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2025-04-14 09:46:03 -04:00
parent cf1606ea11
commit a38f7cbca0
12 changed files with 46 additions and 115 deletions
+1
View File
@@ -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
+2
View File
@@ -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
+2 -29
View File
@@ -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
+2 -18
View File
@@ -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
+26
View File
@@ -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
+2 -2
View File
@@ -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':
+2 -2
View File
@@ -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
-16
View File
@@ -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:
+2 -12
View File
@@ -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)
+2 -12
View File
@@ -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
+2 -12
View File
@@ -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)
+3 -12
View File
@@ -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()