Merge pull request #5038 from vladmandic/fix/video-tiny-vae

Fix/video tiny vae
This commit is contained in:
Vladimir Mandic
2026-08-18 10:47:40 +02:00
committed by GitHub
5 changed files with 33 additions and 16 deletions
+4 -1
View File
@@ -27,7 +27,10 @@ def hijack_vae_decode(*args, **kwargs):
latents = args[0].to(device=devices.device, dtype=shared.sd_model.vae.dtype) # upcast to vae dtype
if hasattr(shared.sd_model.vae, '_asymmetric_upscale_vae'):
res = hijack_vae_upscale(latents, *args[1:], **kwargs)
else:
elif getattr(shared.sd_model, 'sdnext_vae_type', None) == 'Tiny':
from modules.video_models import video_vae
res = video_vae.vae_decode_tiny(latents) # None when the model has no tiny counterpart, and it says so
if res is None:
res = shared.sd_model.vae.orig_decode(latents, *args[1:], **kwargs)
t1 = time.time()
try:
+2 -2
View File
@@ -58,8 +58,8 @@ def warn_once(msg, variant=None):
def get_model(model_cls, variant=None):
if variant is not None:
pass
if variant is not None: # the caller named the variant it wants; the ladder below only derives one
return model_cls, variant
if model_cls in {'sd'}:
model_cls = 'sd'
variant = shared.opts.taesd_variant
+1 -1
View File
@@ -205,7 +205,7 @@ def atomic_save_video(
if metadata is None:
metadata = {}
av = check_av()
if av is None or av is False:
if av is None:
log.error('Video: ffmpeg/av not available')
return
savejob = shared.state.begin('Save video')
+2 -1
View File
@@ -54,13 +54,14 @@ def supports_last_frame(model):
def check_av():
"""The av module, or None when it is unavailable; callers guard on the None."""
install('av')
try:
import av
av.logging.set_level(av.logging.ERROR) # pylint: disable=c-extension-no-member
except Exception as e:
log.error(f'av package: {e}')
return False
return None
return av
+24 -11
View File
@@ -4,14 +4,13 @@ from modules.logger import log
debug = log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None
vae_type = None
UNIT_RANGE_DECODERS = {'TAEM1'} # taehv shifts its output to [-1,1] on the way out and taem1 leaves it at [0,1]
def set_vae_params(p, slicing:bool=True, tiling:bool=True, framewise:bool=True) -> None:
global vae_type # pylint: disable=global-statement
vae_type = p.vae_type
if not hasattr(shared.sd_model, 'vae'):
return
shared.sd_model.sdnext_vae_type = p.vae_type # the decode hijack reads the choice off the pipe
if slicing and hasattr(shared.sd_model.vae, 'enable_slicing'):
shared.sd_model.vae.enable_slicing()
if (p.frames > p.vae_tile_frames) and (p.vae_tile_frames > 0):
@@ -21,33 +20,47 @@ def set_vae_params(p, slicing:bool=True, tiling:bool=True, framewise:bool=True)
shared.sd_model.vae.use_framewise_decoding = True
if tiling and hasattr(shared.sd_model.vae, 'enable_tiling'):
shared.sd_model.vae.enable_tiling()
debug(f'VAE params: type={vae_type} tiling=True frames={p.frames} tile_frames={p.vae_tile_frames} framewise={getattr(shared.sd_model.vae, "use_framewise_decoding", None)}')
debug(f'VAE params: type={p.vae_type} tiling=True frames={p.frames} tile_frames={p.vae_tile_frames} framewise={getattr(shared.sd_model.vae, "use_framewise_decoding", None)}')
else:
if hasattr(shared.sd_model.vae, 'use_framewise_decoding'):
shared.sd_model.vae.use_framewise_decoding = False
if hasattr(shared.sd_model.vae, 'disable_tiling'):
shared.sd_model.vae.disable_tiling()
debug(f'VAE params: type={vae_type} tiling=False frames={p.frames} tile_frames={p.vae_tile_frames} framewise={getattr(shared.sd_model.vae, "use_framewise_decoding", None)}')
debug(f'VAE params: type={p.vae_type} tiling=False frames={p.frames} tile_frames={p.vae_tile_frames} framewise={getattr(shared.sd_model.vae, "use_framewise_decoding", None)}')
def vae_decode_tiny(latents):
if 'Hunyuan' in shared.sd_model.__class__.__name__:
"""Decode through the tiny counterpart of the model's vae, or None when there is none to use.
Returning None leaves the caller on the full vae, so every rejection here is a fallback
rather than a failure.
"""
cls = shared.sd_model.__class__.__name__
if 'Hunyuan' in cls:
variant = 'TAE HunyuanVideo'
elif 'Mochi' in shared.sd_model.__class__.__name__:
elif 'Mochi' in cls:
variant = 'TAE MochiVideo'
elif 'WAN' in shared.sd_model.__class__.__name__:
elif 'Wan' in cls:
variant = 'TAE WanVideo'
elif 'Kandinsky' in shared.sd_model.__class__.__name__:
elif 'Kandinsky' in cls:
variant = 'TAE HunyuanVideo'
else:
log.warning(f'Decode: type=Tiny cls={shared.sd_model.__class__.__name__} not supported')
log.warning(f'Decode: type=Tiny cls={cls} not supported')
return None
from modules.vae import sd_vae_taesd
vae, variant = sd_vae_taesd.load_model(variant=variant)
if vae is None:
return None
expected = getattr(vae, 'latent_channels', None) # 16 on taehv and 12 on taem1, so ask the decoder rather than assume
channels = latents.shape[1] if latents.ndim == 5 else None # the pipes hand the decoder NCTHW
if expected is not None and channels is not None and channels != expected:
log.warning(f'Decode: type=Tiny cls={cls} latents={channels}ch expected={expected}ch not supported')
return None
log.debug(f'Decode: type=Tiny cls={vae.__class__.__name__} variant="{variant}" latents={latents.shape}')
vae = vae.to(device=devices.device, dtype=devices.dtype)
latents = latents.transpose(1, 2).to(device=devices.device, dtype=devices.dtype)
images = vae.decode_video(latents, parallel=False).transpose(1, 2).mul_(2).sub_(1)
images = vae.decode_video(latents, parallel=False).transpose(1, 2)
if type(vae).__name__ in UNIT_RANGE_DECODERS: # the pipelines expect a decode in [-1,1]
images = images.mul_(2).sub_(1)
log.debug(f'Decode: type=Tiny decoded={list(images.shape)} range={images.min().item():.3f}..{images.max().item():.3f}')
return (images, None)