Files
Vladimir Mandic cd25a5296a video loader use generic methods and auth
Signed-off-by: Vladimir Mandic <mandic00@live.com>
2026-08-16 09:02:40 +02:00

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