fix stable-video-diffusion dtype mismatch

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2025-11-13 08:56:14 -05:00
parent 007e265740
commit 2612ad95a9
5 changed files with 68 additions and 33 deletions
+4 -6
View File
@@ -1,16 +1,13 @@
# Change Log for SD.Next
## Update for 2025-11-12
## Update for 2025-11-13
### Highlights for 2025-11-12
### Highlights for 2025-11-13
TBD
### Details for 2025-11-12
### Details for 2025-11-13
- **Models**
- [WAN 2.2 Animate 14B](https://huggingface.co/Wan-AI/Wan2.2-Animate-14B)
available for *text-to-video* and *image-to-video* workflows
- **Features**
- **kanvas**: new module for native canvas-based image manipulation
kanvas is a full replacement for *img2img, inpaint and outpaint* controls
@@ -41,6 +38,7 @@ TBD
- detailer: using lora in detailer prompt
- detailer: fail on unsupported models instead of corrputing results
- ui: fix collapsible panels
- svd: fix stable-video-diffusion dtype mismatch
- process: improve send-to functionality
- control: safe load non-sparse controlnet
- control: fix marigold preprocessor with bfloat16
+2
View File
@@ -14,6 +14,8 @@ def get_model_type(pipe):
model_type = 'sdxl'
elif "StableDiffusion" in name:
model_type = 'sd'
elif "StableVideoDiffusion" in name:
model_type = 'svd'
elif "LatentConsistencyModel" in name:
model_type = 'sd' # lcm is compatible with sd
elif "InstaFlowPipeline" in name:
+2
View File
@@ -127,6 +127,8 @@ def guess_by_name(fn, current_guess):
new_guess = 'X-Omni'
elif 'sdxl-turbo' in fn.lower() or 'stable-diffusion-xl' in fn.lower():
new_guess = 'Stable Diffusion XL'
elif 'stable-video-diffusion' in fn.lower():
new_guess = 'StableVideoDiffusion'
if debug_load:
shared.log.trace(f'Autodetect: method=name file="{fn}" previous="{current_guess}" current="{new_guess}"')
return new_guess or current_guess
+31 -25
View File
@@ -84,33 +84,39 @@ def set_vae_options(sd_model, vae=None, op:str='model', quiet:bool=False):
ops['no-half'] = True
if hasattr(sd_model, 'vae') and hasattr(sd_model.vae, 'enable_slicing') and hasattr(sd_model.vae, 'disable_slicing'):
ops['slicing'] = shared.opts.diffusers_vae_slicing
if shared.opts.diffusers_vae_slicing:
sd_model.vae.enable_slicing()
else:
sd_model.vae.disable_slicing()
try:
if shared.opts.diffusers_vae_slicing:
sd_model.vae.enable_slicing()
else:
sd_model.vae.disable_slicing()
except Exception:
pass
if hasattr(sd_model, 'vae') and hasattr(sd_model.vae, 'enable_tiling') and hasattr(sd_model.vae, 'disable_tiling'):
ops['tiling'] = shared.opts.diffusers_vae_tiling
if shared.opts.diffusers_vae_tiling:
if hasattr(sd_model, 'vae') and hasattr(sd_model.vae, 'config') and hasattr(sd_model.vae.config, 'sample_size') and isinstance(sd_model.vae.config.sample_size, int):
if getattr(sd_model.vae, "tile_sample_min_size_backup", None) is None:
sd_model.vae.tile_sample_min_size_backup = sd_model.vae.tile_sample_min_size
sd_model.vae.tile_latent_min_size_backup = sd_model.vae.tile_latent_min_size
sd_model.vae.tile_overlap_factor_backup = sd_model.vae.tile_overlap_factor
if shared.opts.diffusers_vae_tile_size > 0:
sd_model.vae.tile_sample_min_size = int(shared.opts.diffusers_vae_tile_size)
sd_model.vae.tile_latent_min_size = int(shared.opts.diffusers_vae_tile_size / (2 ** (len(sd_model.vae.config.block_out_channels) - 1)))
else:
sd_model.vae.tile_sample_min_size = getattr(sd_model.vae, "tile_sample_min_size_backup", sd_model.vae.tile_sample_min_size)
sd_model.vae.tile_latent_min_size = getattr(sd_model.vae, "tile_latent_min_size_backup", sd_model.vae.tile_latent_min_size)
if shared.opts.diffusers_vae_tile_overlap != 0.25:
sd_model.vae.tile_overlap_factor = float(shared.opts.diffusers_vae_tile_overlap)
else:
sd_model.vae.tile_overlap_factor = getattr(sd_model.vae, "tile_overlap_factor_backup", sd_model.vae.tile_overlap_factor)
ops['tile'] = sd_model.vae.tile_sample_min_size
ops['overlap'] = sd_model.vae.tile_overlap_factor
sd_model.vae.enable_tiling()
else:
sd_model.vae.disable_tiling()
try:
if shared.opts.diffusers_vae_tiling:
if hasattr(sd_model, 'vae') and hasattr(sd_model.vae, 'config') and hasattr(sd_model.vae.config, 'sample_size') and isinstance(sd_model.vae.config.sample_size, int):
if getattr(sd_model.vae, "tile_sample_min_size_backup", None) is None:
sd_model.vae.tile_sample_min_size_backup = sd_model.vae.tile_sample_min_size
sd_model.vae.tile_latent_min_size_backup = sd_model.vae.tile_latent_min_size
sd_model.vae.tile_overlap_factor_backup = sd_model.vae.tile_overlap_factor
if shared.opts.diffusers_vae_tile_size > 0:
sd_model.vae.tile_sample_min_size = int(shared.opts.diffusers_vae_tile_size)
sd_model.vae.tile_latent_min_size = int(shared.opts.diffusers_vae_tile_size / (2 ** (len(sd_model.vae.config.block_out_channels) - 1)))
else:
sd_model.vae.tile_sample_min_size = getattr(sd_model.vae, "tile_sample_min_size_backup", sd_model.vae.tile_sample_min_size)
sd_model.vae.tile_latent_min_size = getattr(sd_model.vae, "tile_latent_min_size_backup", sd_model.vae.tile_latent_min_size)
if shared.opts.diffusers_vae_tile_overlap != 0.25:
sd_model.vae.tile_overlap_factor = float(shared.opts.diffusers_vae_tile_overlap)
else:
sd_model.vae.tile_overlap_factor = getattr(sd_model.vae, "tile_overlap_factor_backup", sd_model.vae.tile_overlap_factor)
ops['tile'] = sd_model.vae.tile_sample_min_size
ops['overlap'] = sd_model.vae.tile_overlap_factor
sd_model.vae.enable_tiling()
else:
sd_model.vae.disable_tiling()
except Exception:
pass
if hasattr(sd_model, "vqvae"):
ops['upcast'] = True
sd_model.vqvae.to(torch.float32) # vqvae is producing nans in fp16
+29 -2
View File
@@ -32,7 +32,7 @@ class Script(scripts_manager.Script):
min_guidance_scale = gr.Slider(label='Min guidance', minimum=0.0, maximum=10.0, step=0.1, value=1.0)
max_guidance_scale = gr.Slider(label='Max guidance', minimum=0.0, maximum=10.0, step=0.1, value=3.0)
with gr.Row():
decode_chunk_size = gr.Slider(label='Decode chunks', minimum=1, maximum=25, step=1, value=6)
decode_chunk_size = gr.Slider(label='Decode chunks', minimum=1, maximum=25, step=1, value=1)
motion_bucket_id = gr.Slider(label='Motion level', minimum=0, maximum=1, step=0.05, value=0.5)
noise_aug_strength = gr.Slider(label='Noise strength', minimum=0.0, maximum=1.0, step=0.01, value=0.1)
with gr.Row():
@@ -42,6 +42,31 @@ class Script(scripts_manager.Script):
video_type, duration, gif_loop, mp4_pad, mp4_interpolate = create_video_inputs(tab='img2img' if is_img2img else 'txt2img')
return [model, num_frames, override_resolution, min_guidance_scale, max_guidance_scale, decode_chunk_size, motion_bucket_id, noise_aug_strength, video_type, duration, gif_loop, mp4_pad, mp4_interpolate]
def _encode_image(self, image: torch.Tensor, device, num_videos_per_prompt, do_classifier_free_guidance):
image = image.to(device=device, dtype=shared.sd_model.vae.dtype)
shared.log.debug(f'Video encode: type=svd input={image.shape} dtype={image.dtype} device={image.device}')
image_latents = shared.sd_model.vae.encode(image).latent_dist.mode()
image_latents = image_latents.repeat(num_videos_per_prompt, 1, 1, 1)
if do_classifier_free_guidance:
negative_image_latents = torch.zeros_like(image_latents)
image_latents = torch.cat([negative_image_latents, image_latents])
return image_latents
def _decode_latents(self, latents: torch.Tensor, num_frames: int, decode_chunk_size: int = 14):
shared.log.debug(f'Video decode: type=svd input={latents.shape} dtype={latents.dtype} device={latents.device} chunk={decode_chunk_size} frames={num_frames}')
latents = latents.flatten(0, 1)
latents = 1 / shared.sd_model.vae.config.scaling_factor * latents
frames = []
for i in range(0, latents.shape[0], decode_chunk_size):
num_frames_in = latents[i : i + decode_chunk_size].shape[0]
decode_kwargs = { "num_frames": num_frames_in }
frame = shared.sd_model.vae.decode(latents[i : i + decode_chunk_size], **decode_kwargs).sample
frames.append(frame)
frames = torch.cat(frames, dim=0)
frames = frames.reshape(-1, num_frames, *frames.shape[1:]).permute(0, 2, 1, 3, 4)
frames = frames.float()
return frames
def run(self, p: processing.StableDiffusionProcessing, model, num_frames, override_resolution, min_guidance_scale, max_guidance_scale, decode_chunk_size, motion_bucket_id, noise_aug_strength, video_type, duration, gif_loop, mp4_pad, mp4_interpolate): # pylint: disable=arguments-differ, unused-argument
image = getattr(p, 'init_images', None)
if image is None or len(image) == 0:
@@ -60,9 +85,11 @@ class Script(scripts_manager.Script):
c = shared.sd_model.__class__.__name__
model_loaded = shared.sd_model.sd_checkpoint_info.model_name if shared.sd_loaded else None
if model_name != model_loaded or c != 'StableVideoDiffusionPipeline':
from diffusers import StableVideoDiffusionPipeline # pylint: disable=unused-import
shared.opts.sd_model_checkpoint = model_path
sd_models.reload_model_weights()
shared.sd_model = shared.sd_model.to(torch.float32) # must run in fp32 due to dtype mismatch
shared.sd_model._encode_vae_image = self._encode_image # pylint: disable=protected-access
shared.sd_model.decode_latents = self._decode_latents # pylint: disable=protected-access
# set params
if override_resolution: