mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 09:38:23 +02:00
fix stable-video-diffusion dtype mismatch
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
+4
-6
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user