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, modular_load, sd_hijack_te, sd_hijack_vae, sd_hijack_modular 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 if 'gemini-omni' in model_name: from modules.video_models.google_omni import load_omni pipe = load_omni(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): sd_hijack_modular.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