mirror of
https://github.com/vladmandic/automatic
synced 2026-08-27 15:41:00 +02:00
cd25a5296a
Signed-off-by: Vladimir Mandic <mandic00@live.com>
230 lines
11 KiB
Python
230 lines
11 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
|
|
from pipelines import generic
|
|
|
|
|
|
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)
|
|
sd_models.hf_auth_check(selected.repo)
|
|
|
|
# text encoder
|
|
if selected.te_cls is not None:
|
|
te_repo, te_folder, te_revision = selected.te, selected.te_folder, selected.te_revision
|
|
kwargs["text_encoder"] = generic.load_text_encoder(
|
|
te_repo or selected.repo,
|
|
cls_name=selected.te_cls,
|
|
subfolder=te_folder,
|
|
revision=te_revision or selected.repo_revision,
|
|
)
|
|
|
|
# transformer
|
|
if selected.dit_cls is not None:
|
|
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:
|
|
kwargs[dit_kwarg] = generic.load_transformer(
|
|
selected.dit or selected.repo,
|
|
cls_name=selected.dit_cls,
|
|
subfolder=dit_folder,
|
|
revision=selected.dit_revision or selected.repo_revision,
|
|
)
|
|
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)
|
|
|
|
# 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
|