support configurable multi-stage models in video tab

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2025-09-22 13:18:29 -04:00
parent 8b47d72610
commit 41bb446697
7 changed files with 156 additions and 30 deletions
+3 -1
View File
@@ -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
+3 -2
View File
@@ -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',
+17 -24
View File
@@ -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')
+9
View File
@@ -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
+2
View File
@@ -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,
+119
View File
@@ -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
+3 -3
View File
@@ -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()