From 41bb446697fe2be2c3ba9424f13b0ec808986fbd Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 22 Sep 2025 13:18:29 -0400 Subject: [PATCH] support configurable multi-stage models in video tab Signed-off-by: Vladimir Mandic --- CHANGELOG.md | 4 +- modules/video_models/models_def.py | 5 +- modules/video_models/video_load.py | 41 ++++---- modules/video_models/video_overrides.py | 9 ++ pipelines/model_wanai.py | 2 + pipelines/wan/wan_image.py | 119 ++++++++++++++++++++++++ webui.py | 6 +- 7 files changed, 156 insertions(+), 30 deletions(-) create mode 100644 pipelines/wan/wan_image.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 2f490ed68..22b3ac43b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,6 @@ # Change Log for SD.Next -## Update for 2025-09-21 +## Update for 2025-09-22 - **Models** - [WAN 2.2 14B VACE](https://huggingface.co/alibaba-pai/Wan2.2-VACE-Fun-A14B) @@ -33,6 +33,8 @@ - styles and wildcards now use same seed as main generate for reproducible results - **api** new endpoint POST `/sdapi/v1/civitai` to trigger civitai models metadata update accepts optional `page` parameter to search specific networks page + - **reference models** additional example images, thanks @liutyi + - **video** support for configurable multi-stage models such as WAN-2.2-14B - **Fixes** - framepack: add explicit hf-login before framepack load - benchmark: remove forced sampler from system info benchmark diff --git a/modules/video_models/models_def.py b/modules/video_models/models_def.py index b9ff5c260..e2605ae50 100644 --- a/modules/video_models/models_def.py +++ b/modules/video_models/models_def.py @@ -170,11 +170,12 @@ models = { dit_cls=diffusers.WanTransformer3DModel, dit_folder=("transformer", "transformer_2")), Model(name='WAN 2.2 14B VACE', - url='https://huggingface.co/Wan-AI/Wan2.2-14B-VACE-T2V-Diffusers', + url='https://huggingface.co/linoyts/Wan2.2-VACE-Fun-14B-diffusers', repo='linoyts/Wan2.2-VACE-Fun-14B-diffusers', repo_cls=diffusers.WanVACEPipeline, te_cls=transformers.T5EncoderModel, - dit_cls=diffusers.WanVACETransformer3DModel), + dit_cls=diffusers.WanVACETransformer3DModel, + dit_folder=("transformer", "transformer_2")), Model(name='WAN 2.1 1.3B T2V', url='https://huggingface.co/Wan-AI/Wan2.1-T2V-1.3B-Diffusers', repo='Wan-AI/Wan2.1-T2V-1.3B-Diffusers', diff --git a/modules/video_models/video_load.py b/modules/video_models/video_load.py index 72c411166..874df0231 100644 --- a/modules/video_models/video_load.py +++ b/modules/video_models/video_load.py @@ -40,31 +40,24 @@ def load_model(selected: models_def.Model): # transformer try: + if selected.dit_folder is None: + selected.dit_folder = ['transformer'] if isinstance(selected.dit_folder, list) or isinstance(selected.dit_folder, tuple): - # wan a14b has transformer and transformer_2 - for dit_folder in selected.dit_folder: - # 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) - shared.log.debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" folder="{dit_folder}" cls={selected.dit_cls.__name__} quant={model_quant.get_quant_type(quant_args)}') - kwargs[dit_folder] = 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 - ) - else: - load_args, quant_args = model_quant.get_dit_args({}, module='Model', device_map=True) - shared.log.debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" folder="{selected.dit_folder}" cls={selected.dit_cls.__name__} quant={model_quant.get_quant_type(quant_args)}') - kwargs["transformer"] = selected.dit_cls.from_pretrained( - pretrained_model_name_or_path=selected.dit or selected.repo, - subfolder=selected.dit_folder, - revision=selected.dit_revision or selected.repo_revision, - cache_dir=shared.opts.hfcache_dir, - **load_args, - **quant_args - ) + for dit_folder in selected.dit_folder: # wan a14b has transformer and transformer_2 + if dit_folder is not None and dit_folder 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) + shared.log.debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" module="{dit_folder}" folder="{dit_folder}" cls={selected.dit_cls.__name__} quant={model_quant.get_quant_type(quant_args)}') + kwargs[dit_folder] = 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 + ) + else: + shared.log.debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" module="{dit_folder}" folder="{dit_folder}" cls={selected.dit_cls.__name__} skip') except Exception as e: shared.log.error(f'video load: module=transformer cls={selected.dit_cls.__name__} {e}') errors.display(e, 'video') diff --git a/modules/video_models/video_overrides.py b/modules/video_models/video_overrides.py index 057ae3f8a..856da62e6 100644 --- a/modules/video_models/video_overrides.py +++ b/modules/video_models/video_overrides.py @@ -19,6 +19,15 @@ def load_override(selected: Model): # WAN if 'WAN 2.1 14B' in selected.name: kwargs['vae'] = diffusers.AutoencoderKLWan.from_pretrained(selected.repo, subfolder="vae", torch_dtype=torch.float32, cache_dir=shared.opts.hfcache_dir) + if 'A14B' in selected.name or '14B VACE' in selected.name: + if shared.opts.model_wan_stage == 'combined': + kwargs['boundary_ratio'] = shared.opts.model_wan_boundary + elif shared.opts.model_wan_stage == 'high noise': + kwargs['transformer_2'] = None + kwargs['boundary_ratio'] = 0.0 + elif shared.opts.model_wan_stage == 'low noise': + kwargs['boundary_ratio'] = 1.0 + kwargs['transformer'] = None debug(f'Video overrides: model="{selected.name}" kwargs={list(kwargs)}') return kwargs diff --git a/pipelines/model_wanai.py b/pipelines/model_wanai.py index e44d1eb78..718ec63c6 100644 --- a/pipelines/model_wanai.py +++ b/pipelines/model_wanai.py @@ -96,8 +96,10 @@ def load_wan(checkpoint_info, diffusers_load_config={}): diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["wanai"] = diffusers.WanVACEPipeline diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["wanai"] = diffusers.WanVACEPipeline else: + from pipelines.wan.wan_image import WanImagePipeline pipe_cls = diffusers.WanPipeline diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["wanai"] = diffusers.WanPipeline + diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["wanai"] = WanImagePipeline shared.log.debug(f'Load model: type=WanAI model="{checkpoint_info.name}" repo="{repo_id}" cls={pipe_cls.__name__} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args} stage="{shared.opts.model_wan_stage}" boundary={boundary_ratio}') pipe = pipe_cls.from_pretrained( repo_id, diff --git a/pipelines/wan/wan_image.py b/pipelines/wan/wan_image.py new file mode 100644 index 000000000..bd9923e5b --- /dev/null +++ b/pipelines/wan/wan_image.py @@ -0,0 +1,119 @@ +from typing import Any, Callable, Dict, List, Optional, Union +import torch +import diffusers +from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback +from diffusers.image_processor import PipelineImageInput + +from modules import devices + + +class WanImagePipeline(diffusers.WanPipeline): + def __call__( + self, + prompt: Union[str, List[str]] = None, + negative_prompt: Union[str, List[str]] = None, + height: int = 480, + width: int = 832, + num_frames: int = 81, + num_inference_steps: int = 50, + guidance_scale: float = 5.0, + guidance_scale_2: Optional[float] = None, + num_videos_per_prompt: Optional[int] = 1, + generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, + latents: Optional[torch.Tensor] = None, + prompt_embeds: Optional[torch.Tensor] = None, + negative_prompt_embeds: Optional[torch.Tensor] = None, + output_type: Optional[str] = "np", + return_dict: bool = True, + attention_kwargs: Optional[Dict[str, Any]] = None, + callback_on_step_end: Optional[Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]] = None, + callback_on_step_end_tensor_inputs: List[str] = ["latents"], + max_sequence_length: int = 512, + strength: float = 0.3, # new + image: PipelineImageInput = None, # new + ): + # get img2img timesteps + self.scheduler.set_timesteps(num_inference_steps, device=devices.device) + timesteps, num_inference_steps = self.get_timesteps(num_inference_steps, strength) + # monkey patch original pipeline + self.scheduler.timesteps = timesteps + # self.scheduler._step_index = 0 + self.scheduler.orig_set_timesteps = self.scheduler.set_timesteps + self.scheduler.set_timesteps = lambda *args, **kwargs: None + + # prepare latents + latents = self.img2img_prepare_latents( + image=image, + timesteps=timesteps, + dtype=devices.dtype, + device=devices.device, + generator=generator, + ) + + # call original pipeline + result = super().__call__( # pylint: disable=no-member + prompt=prompt, + negative_prompt=negative_prompt, + height=height, + width=width, + num_frames=num_frames, + num_inference_steps=num_inference_steps, + guidance_scale=guidance_scale, + guidance_scale_2=guidance_scale_2, + num_videos_per_prompt=num_videos_per_prompt, + generator=generator, + latents=latents, + prompt_embeds=prompt_embeds, + negative_prompt_embeds=negative_prompt_embeds, + output_type=output_type, + return_dict=return_dict, + attention_kwargs=attention_kwargs, + callback_on_step_end=callback_on_step_end, + callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs, + max_sequence_length=max_sequence_length, + ) + + # un-monkey patch original pipeline + self.scheduler.set_timesteps = self.scheduler.orig_set_timesteps + return result + + def get_timesteps(self, num_inference_steps, strength): + init_timestep = min(int(num_inference_steps * strength), num_inference_steps) + t_start = max(num_inference_steps - init_timestep, 0) + timesteps = self.scheduler.timesteps[t_start * self.scheduler.order :] + if hasattr(self.scheduler, "set_begin_index"): + # self.scheduler.set_begin_index(t_start * self.scheduler.order) + self.scheduler.set_begin_index(0) + return timesteps, num_inference_steps - t_start + + def img2img_prepare_latents( + self, + image: torch.Tensor = None, + timesteps: torch.Tensor = None, + dtype: Optional[torch.dtype] = None, + device: Optional[torch.device] = None, + generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, + ) -> torch.Tensor: + from diffusers.utils.torch_utils import randn_tensor + from diffusers.video_processor import VideoProcessor + + if isinstance(image, list): + image = image[0] # ignore batch for now + + video_processor = VideoProcessor(vae_scale_factor=self.vae_scale_factor_spatial) + image_tensor = video_processor.preprocess(image, None, None) # convert PIL to [B, C, H, W] # channels may need rearrange + image_tensor = image_tensor.squeeze(0).to(device=device, dtype=dtype) + image_tensor = image_tensor[None, :, None, :, :] # expand before encode to [B, C, N, H, W] + encoder_output = self.vae.encode(image_tensor) + # init_latents = encoder_output.latent_dist.mode() # argmax or sample? + init_latents = encoder_output.latent_dist.sample(generator) + + latents_mean = torch.tensor(self.vae.config.latents_mean, device=device, dtype=torch.float32).view(1, self.vae.config.z_dim, 1, 1, 1) + latents_std = 1.0 / torch.tensor(self.vae.config.latents_std, device=device, dtype=torch.float32).view(1, self.vae.config.z_dim, 1, 1, 1) + init_latents = ((init_latents.float() - latents_mean) * latents_std).to(dtype) # normalized to standard distribution range + + init_noise = randn_tensor(init_latents.shape, generator=generator, device=device, dtype=dtype) + init_timestep = timesteps[:1] + noised_latents = self.scheduler.add_noise(init_latents, init_noise, init_timestep) + + return noised_latents diff --git a/webui.py b/webui.py index ec33c908e..8dc29d5bc 100644 --- a/webui.py +++ b/webui.py @@ -152,9 +152,9 @@ def initialize(): def load_model(): modeldata.model_data.locked = False - if not shared.opts.sd_checkpoint_autoload and shared.cmd_opts.ckpt is None: - log.info('Model: autoload=False') - else: + autoload = shared.opts.sd_checkpoint_autoload or shared.cmd_opts.ckpt is not None + log.info(f'Model: autoload={autoload} selected="{shared.opts.sd_model_checkpoint}"') + if autoload: jobid = shared.state.begin('Load model') thread_model = Thread(target=lambda: shared.sd_model) thread_model.start()