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
+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: