Files
automatic/modules/video_models/video_load.py
T
CalamitousFelicitousness 0d1882eca4 fix(video): keep text encoder dedup out of the registry rows
The shared repo was written back onto the registry row, a module-level
singleton, so turning the setting off left the row pointing at the shared copy
for the rest of the session. It is chosen into locals instead.
2026-08-14 03:39:01 +01:00

265 lines
14 KiB
Python

import os
import sys
import copy
import time
import torch
import transformers
import diffusers
from modules import shared, errors, sd_models, sd_checkpoint, model_quant, devices, sd_hijack_te, sd_hijack_vae, modular_load
from modules.logger import log
from modules.video_models import models_def, video_utils, video_overrides, video_cache
def _loader(component):
"""Return loader type for log messages."""
if sys.platform != 'linux':
return 'default'
if component == 'diffusers':
return 'runai' if shared.opts.runai_streamer_diffusers else 'default'
return 'runai' if shared.opts.runai_streamer_transformers else 'default'
loaded_model = None
def snapshot_without(repo: str, ignore_patterns, revision: str | None = None, passed=None, **offline_args) -> str:
"""Materialize a repo snapshot minus ignore_patterns and return the local folder.
DiffusionPipeline.download builds its own ignore list from the passed components and never
reads the caller's, so repo files that no component claims are fetched regardless. Doing the
snapshot here and handing from_pretrained a folder skips its download path entirely, which
also means the component folders it would have pruned have to be pruned here instead.
"""
if not ignore_patterns:
return repo
from huggingface_hub import snapshot_download
patterns = list(ignore_patterns) + [f'{name}/*' for name in (passed or [])]
try:
folder = snapshot_download(repo, revision=revision, cache_dir=shared.opts.hfcache_dir, ignore_patterns=patterns, **offline_args)
except Exception as e:
log.warning(f'Load video: module=snapshot repo="{repo}" ignore={patterns} {e}')
return repo
log.debug(f'Load video: module=snapshot repo="{repo}" ignore={patterns}')
return folder
def load_custom(model_name: str):
log.debug(f'Load video: module=pipe repo="{model_name}" cls=Custom')
if 'veo-3.1' in model_name:
from modules.video_models.google_veo import load_veo
pipe = load_veo(model_name)
return pipe
return None
def load_model(selected: models_def.Model):
import sdnq # pylint: disable=unused-import
if selected is None or selected.repo is None:
return ''
if isinstance(selected.repo_cls, str):
selected.repo_cls=models_def.getpipe(diffusers, selected.repo_cls, None)
if isinstance(selected.dit_cls, str):
selected.dit_cls=models_def.getpipe(diffusers, selected.dit_cls, None)
if isinstance(selected.te_cls, str):
selected.te_cls=models_def.getpipe(transformers, selected.te_cls, None)
global loaded_model # pylint: disable=global-statement
if not shared.sd_loaded:
loaded_model = None
elif loaded_model == selected.name and selected.repo_cls is not None and not isinstance(shared.sd_model, selected.repo_cls):
# shared.sd_model auto-reloads the default checkpoint when model_data.sd_model is None,
# which silently swaps the pipe class behind the name-based cache. Pipe-class mismatch
# is the reliable signal that the cached name no longer maps to the cached object.
log.warning(f'Load video: cached model="{selected.name}" cls={type(shared.sd_model).__name__} mismatch forcing reload')
loaded_model = None
if loaded_model == selected.name:
return ''
if shared.sd_loaded:
sd_models.unload_model_weights()
t0 = time.time()
jobid = shared.state.begin('Load model')
video_cache.apply_teacache_patch(selected.dit_cls)
# overrides
offline_args = {}
if shared.opts.offline_mode:
offline_args["local_files_only"] = True
os.environ['HF_HUB_OFFLINE'] = '1'
else:
os.environ.pop('HF_HUB_OFFLINE', None)
os.unsetenv('HF_HUB_OFFLINE')
kwargs = video_overrides.load_override(selected, **offline_args)
# text encoder
if selected.te_cls is not None:
try:
load_args, quant_args = model_quant.get_dit_args({}, module='TE', device_map=True)
# loader deduplication of text-encoder models: picked per load, not written back onto
# the registry row where it would outlive the setting
te_repo, te_folder, te_revision = selected.te, selected.te_folder, selected.te_revision
if shared.opts.te_shared_te:
te_cls_name = selected.te_cls.__name__
if te_cls_name == 'T5EncoderModel':
te_repo, te_folder, te_revision = 'Disty0/t5-xxl', '', None
elif te_cls_name == 'UMT5EncoderModel':
te_repo = 'Disty0/Wan2.2-T2V-A14B-SDNQ-uint4-svd-r32' if 'SDNQ' in selected.name else 'Wan-AI/Wan2.2-TI2V-5B-Diffusers'
te_folder, te_revision = 'text_encoder', None
elif te_cls_name == 'LlamaModel':
te_repo, te_folder, te_revision = 'hunyuanvideo-community/HunyuanVideo', 'text_encoder', None
elif te_cls_name == 'Qwen2_5_VLForConditionalGeneration':
te_repo, te_folder, te_revision = 'ai-forever/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers', 'text_encoder', None
elif te_cls_name == 'Gemma3ForConditionalGeneration':
te_repo = 'OzzyGT/LTX-2.3-sdnq-dynamic-int4' if 'SDNQ' in selected.name else 'OzzyGT/LTX-2.3'
te_folder, te_revision = 'text_encoder', None
log.debug(f'Load video: module=te repo="{te_repo or selected.repo}" folder="{te_folder}" cls={selected.te_cls.__name__} quant={model_quant.get_quant_type(quant_args)} loader={_loader("transformers")}')
kwargs["text_encoder"] = selected.te_cls.from_pretrained(
pretrained_model_name_or_path=te_repo or selected.repo,
subfolder=te_folder,
revision=te_revision or selected.repo_revision,
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
**offline_args,
)
except Exception as e:
log.error(f'video load: module=te cls={selected.te_cls.__name__} {e}')
errors.display(e, 'video')
# transformer
if selected.dit_cls is not None:
try:
def load_dit_folder(dit_folder, dit_kwarg=None):
dit_kwarg = dit_kwarg or dit_folder # ltx-2.5 keeps its dev transformer in transformer_full
if dit_folder is not None and dit_kwarg not in kwargs:
# get a new quant arg on every loop to prevent the quant config classes getting entangled
load_args, quant_args = model_quant.get_dit_args({}, module='Model', device_map=True)
log.debug(f'Load video: module=transformer repo="{selected.dit or selected.repo}" module="{dit_kwarg}" folder="{dit_folder}" cls={selected.dit_cls.__name__} quant={model_quant.get_quant_type(quant_args)} loader={_loader("diffusers")}')
kwargs[dit_kwarg] = selected.dit_cls.from_pretrained(
pretrained_model_name_or_path=selected.dit or selected.repo,
subfolder=dit_folder,
revision=selected.dit_revision or selected.repo_revision,
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args,
**offline_args,
)
else:
log.debug(f'Load video: module=transformer repo="{selected.dit or selected.repo}" module="{dit_kwarg}" folder="{dit_folder}" cls={selected.dit_cls.__name__} loader={_loader("diffusers")} skip')
if selected.dit_folder is None:
selected.dit_folder = ['transformer']
if isinstance(selected.dit_folder, list) or isinstance(selected.dit_folder, tuple):
if selected.dit_kwarg is not None:
log.warning(f'Load video: model="{selected.name}" dit_kwarg unsupported with multiple folders')
for dit_folder in selected.dit_folder: # wan a14b has transformer and transformer_2
load_dit_folder(dit_folder)
else:
load_dit_folder(selected.dit_folder, selected.dit_kwarg)
except Exception as e:
log.error(f'video load: module=transformer cls={selected.dit_cls.__name__} {e}')
errors.display(e, 'video')
# model
try:
if selected.workflow is not None or modular_load.is_modular(selected.repo_cls):
from modules.modular_load import load_modular_pipe
shared.sd_model = load_modular_pipe(selected.repo_cls, selected.repo, workflow=selected.workflow, revision=selected.repo_revision, offline_args=offline_args, base=selected.base)
elif selected.repo_cls is None:
shared.sd_model = load_custom(selected.repo)
else:
log.debug(f'Load video: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__}')
sd_models.hf_prefetch_configs(selected.repo, {}, 'video')
passed = [k for k, v in kwargs.items() if isinstance(v, torch.nn.Module)]
repo_path = snapshot_without(selected.repo, kwargs.pop('ignore_patterns', None), selected.repo_revision, passed=passed, **offline_args)
shared.sd_model = selected.repo_cls.from_pretrained(
pretrained_model_name_or_path=repo_path,
revision=selected.repo_revision,
cache_dir=shared.opts.hfcache_dir,
torch_dtype=devices.dtype,
**kwargs,
**offline_args,
)
except Exception as e:
log.error(f'video load: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__} {e}')
errors.display(e, 'video')
if shared.sd_model is None:
msg = f'Load video: model="{selected.name}" failed'
log.error(msg)
return msg
t1 = time.time()
cls_name = shared.sd_model.__class__.__name__
# LTX 0.9.x is plain linear; pin use_dynamic_shifting=False against upstream config drift.
# LTX-2.x canonical is token-count-based dynamic shift (base_shift=0.95, max_shift=2.05);
# disabling it there would take the model off-distribution.
if cls_name.startswith("LTX") and not cls_name.startswith("LTX2"):
shared.sd_model.scheduler.config.use_dynamic_shifting = False
shared.sd_model.default_scheduler = copy.deepcopy(shared.sd_model.scheduler) if hasattr(shared.sd_model, "scheduler") else None
shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(selected.repo)
shared.sd_model.sd_model_hash = None
sd_models.set_diffuser_options(shared.sd_model, offload=False)
decode, text, image, slicing, tiling, framewise = False, False, False, False, False, False
if selected.vae_hijack and hasattr(shared.sd_model, 'vae') and hasattr(shared.sd_model.vae, 'decode'):
sd_hijack_vae.init_hijack(shared.sd_model)
decode = True
if selected.te_hijack and hasattr(shared.sd_model, 'encode_prompt'):
sd_hijack_te.init_hijack(shared.sd_model)
text = True
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
image = True
if hasattr(shared.sd_model, 'vae') and hasattr(shared.sd_model.vae, 'use_framewise_decoding'):
shared.sd_model.vae.use_framewise_decoding = True
framewise = True
if hasattr(shared.sd_model, 'vae') and hasattr(shared.sd_model.vae, 'enable_slicing'):
shared.sd_model.vae.enable_slicing()
slicing = True
if hasattr(shared.sd_model, 'vae') and hasattr(shared.sd_model.vae, 'enable_tiling'):
shared.sd_model.vae.enable_tiling()
tiling = True
if hasattr(shared.sd_model, "set_progress_bar_config"):
shared.sd_model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar:15} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m', ncols=120, colour='#327fba')
shared.sd_model = model_quant.do_post_load_quant(shared.sd_model, allow=False)
sd_models.set_diffuser_offload(shared.sd_model)
if modular_load.is_modular(shared.sd_model):
modular_load.install_state_hook(shared.sd_model)
loaded_model = selected.name
msg = f'Load video: cls={shared.sd_model.__class__.__name__} model="{selected.name}" time={t1-t0:.2f}'
log.info(msg)
log.debug(f'Video hijacks: decode={decode} text={text} image={image} slicing={slicing} tiling={tiling} framewise={framewise}')
shared.state.end(jobid)
return msg
def load_upscale_vae():
if not hasattr(shared.sd_model, 'vae'):
return
if hasattr(shared.sd_model.vae, '_asymmetric_upscale_vae'):
return # already loaded
if shared.sd_model.vae.__class__.__name__ != 'AutoencoderKLWan':
log.warning('Video decode: upscale VAE unsupported')
return
repo_id = 'spacepxl/Wan2.1-VAE-upscale2x'
subfolder = "diffusers/Wan2.1_VAE_upscale2x_imageonly_real_v1"
vae_decode = diffusers.AutoencoderKLWan.from_pretrained(repo_id, subfolder=subfolder, cache_dir=shared.opts.hfcache_dir)
vae_decode.requires_grad_(False)
vae_decode = vae_decode.to(device=devices.device, dtype=devices.dtype)
vae_decode.eval()
log.debug(f'Decode: load="{repo_id}"')
shared.sd_model.orig_vae = shared.sd_model.vae
shared.sd_model.vae = vae_decode
shared.sd_model.vae._asymmetric_upscale_vae = True # pylint: disable=protected-access
sd_hijack_vae.init_hijack(shared.sd_model)
sd_models.apply_balanced_offload(shared.sd_model, force=True) # reapply offload