mirror of
https://github.com/vladmandic/automatic
synced 2026-09-06 04:50:44 +02:00
0d1882eca4
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.
265 lines
14 KiB
Python
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
|