video tab alpha releasee

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2025-03-20 14:31:30 -04:00
parent 9a85e45cb4
commit 5c46904fd8
20 changed files with 874 additions and 172 deletions
+39 -11
View File
@@ -2,21 +2,49 @@
## Update for 2025-03-19
### TODO
- Gemma3 requires `transformers==git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3`
- Remote VAE encode for SD15 and Flux.1: <https://github.com/huggingface/diffusers/issues/11069>
- HunyuanVideo-I2V: <https://github.com/huggingface/diffusers/issues/11118>
- HunyuanVideo: Remote VAE
- HunyuanVideo: Tiny VAE
- LTXVideo-095: Condition input
- LTXVideo-095: Broken offloading
### ToDo/Limitations
### Highlights for 2025-03-17
- VLM Gemma3: requires `transformers==git+https://github.com/huggingface/transformers@v4.49.0-Gemma-3`
- VAE Remote encode: SD15 and Flux.1 issues: <https://github.com/huggingface/diffusers/issues/11069>
- Video: ModernUI support is TBD
- Video: API support is TBD
- Video: Wiki page is TBD
- Video: HunyuanVideo-I2V incompatible with latest transformers <https://github.com/huggingface/diffusers/issues/11118>
- Video: LTXVideo-095 support for conditioned input
- Video: LTXVideo-095 support for offloading
Support for [CogView 4](https://huggingface.co/THUDM/CogView4-6B), new CLiP models, improvements to remote VAE, additional docs/guides.
### Highlights for 2025-03-20
### Details for 2025-03-17
Brand new Video processing module with support for all latest models: **WAN21, Hunyuan, LTX, Cog, Allegro, Mochi1** and more!
Plus support for CogView-4, new CLiP models, improvements to remote VAE, additional docs/guides.
### Details for 2025-03-20
- **Video tab**
- initial release so consider this as alpha version
- new top-level tab, replaces previous *video* script in text/image tabs
old scripts are still present, but will be removed in the future
- support for all latest models:
- [Hunyuan](https://huggingface.co/Tencent/HunyuanVideo): *HunyuanVideo, FastHunyuan, SkyReels* | *T2V, I2V*
- [WAN21](https://huggingface.co/Wan-AI/Wan2.1-T2V-1.3B-Diffusers): *1.3B, 13B* | *T2V, I2V*
- [LTXVideo](https://huggingface.co/Lightricks/LTX-Video): *0.9.0, 0.9.1, 0.9.5* | *T2V, I2V*
- [CogVideoX](https://huggingface.co/THUDM/CogVideoX-5b): *2B, 5B* | *T2V, I2V*
- [Allegro](https://huggingface.co/rhymes-ai/Allegro): *T2V*
- [Mochi1](https://huggingface.co/genmo/mochi-1-preview): *T2V*
- decoding:
- **Default**: use vae from model
- **Tiny VAE**: support for *Hunyuan, WAN, Mochi*
- **Remote VAE**: support for *Hunyuan*
- **LoRA**: support for *Hunyuan, LTX, WAN, Mochi, Cog*
- additional key points:
- all models are auto-downloaded upon first use
- optional video interpolation while creating video files
- optional video preview in ui
- support for balanced offloading and model offloading
- on-the-fly quantization: *BnB, Quanto, TorchAO*
- different video models support different video resolutions, frame counts, etc.
and may require specific settings - see model links for details
- see *ToDo/Limitations* section for additional notes
- **Models**
- [THUDM CogView 4 6B](https://huggingface.co/THUDM/CogView4-6B)
new foundation model for image generation based o GLM-4 text encoder and a flow-based diffusion transformer
+7 -1
View File
@@ -30,8 +30,14 @@ force_models = [ # forced always
'sc',
'kandinsky',
'hunyuandit',
'hunyuanvideo',
'auraflow',
# video models
'hunyuanvideo',
'cogvideo',
'wanvideo',
'ltxvideo',
'mochivideo',
'allegrovideo',
]
force_classes = [ # forced always
+11 -4
View File
@@ -29,8 +29,6 @@ def get_model_type(pipe):
model_type = 'auraflow'
elif "Flux" in name:
model_type = 'f1'
elif "Mochi" in name:
model_type = 'mochi'
elif "Lumina2" in name:
model_type = 'lumina2'
elif "Lumina" in name:
@@ -41,12 +39,21 @@ def get_model_type(pipe):
model_type = 'cogview3'
elif "CogView4" in name:
model_type = 'cogview4'
elif "CogVideo" in name:
model_type = 'cogvideox'
elif "Sana" in name:
model_type = 'sana'
# video models
elif "CogVideo" in name:
model_type = 'cogvideo'
elif 'HunyuanVideoPipeline' in name or 'HunyuanSkyreels' in name:
model_type = 'hunyuanvideo'
elif 'Wan' in name:
model_type = 'wanvideo'
elif 'LTX' in name:
model_type = 'ltxvideo'
elif "Mochi" in name:
model_type = 'mochivideo'
elif "Allegro" in name:
model_type = 'allegrovideo'
else:
model_type = name
return model_type
+2
View File
@@ -239,6 +239,8 @@ def vae_postprocess(tensor, model, output_type='np'):
if len(tensor.shape) == 3 and tensor.shape[0] == 3:
tensor = tensor.unsqueeze(0)
if hasattr(model, 'video_processor'):
if len(tensor.shape) == 6 and tensor.shape[1] == 1:
tensor = tensor.squeeze(0)
images = model.video_processor.postprocess_video(tensor, output_type='pil')
elif hasattr(model, 'image_processor'):
images = model.image_processor.postprocess(tensor, output_type=output_type)
+3 -3
View File
@@ -10,7 +10,7 @@ from modules.timer import process as process_timer
debug_move = shared.log.trace if os.environ.get('SD_MOVE_DEBUG', None) is not None else lambda *args, **kwargs: None
should_offload = ['sc', 'sd3', 'f1', 'hunyuandit', 'auraflow', 'omnigen', 'hunyuanvideo', 'cogvideox', 'mochi', 'cogview4']
should_offload = ['sc', 'sd3', 'f1', 'hunyuandit', 'auraflow', 'omnigen', 'cogview4']
offload_hook_instance = None
balanced_offload_exclude = ['OmniGenPipeline', 'CogView4Pipeline']
@@ -66,7 +66,7 @@ def set_diffuser_offload(sd_model, op:str='model', quiet:bool=False):
if not (hasattr(sd_model, "has_accelerate") and sd_model.has_accelerate):
sd_model.has_accelerate = False
if shared.opts.diffusers_offload_mode == "none":
if shared.sd_model_type in should_offload:
if shared.sd_model_type in should_offload or 'video' in shared.sd_model_type:
shared.log.warning(f'Setting {op}: offload={shared.opts.diffusers_offload_mode} type={shared.sd_model.__class__.__name__} large model')
else:
shared.log.quiet(quiet, f'Setting {op}: offload={shared.opts.diffusers_offload_mode} limit={shared.opts.cuda_mem_fraction}')
@@ -183,7 +183,7 @@ class OffloadHook(accelerate.hooks.ModelHook):
return module
def apply_balanced_offload(sd_model, exclude=[]):
def apply_balanced_offload(sd_model=None, exclude=[]):
global offload_hook_instance # pylint: disable=global-statement
if shared.opts.diffusers_offload_mode != "balanced":
return sd_model
+13 -5
View File
@@ -38,11 +38,14 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_
return tensors
t0 = time.time()
modelloader.hf_login()
latents = latents.unsqueeze(0) if len(latents.shape) == 3 else latents
latent_copy = latents.detach().clone().to(device=devices.cpu, dtype=devices.dtype)
latent_copy = latents.unsqueeze(0) if len(latents.shape) == 3 else latents
if model_type == 'hunyuanvideo':
latent_copy = latent_copy.unsqueeze(0)
for i in range(latents.shape[0]):
for i in range(latent_copy.shape[0]):
try:
latent = latents[i].detach().clone().to(device=devices.cpu, dtype=devices.dtype)
latent = latent_copy[i]
if model_type != 'f1':
latent = latent.unsqueeze(0)
params = {
@@ -51,7 +54,12 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_
"dtype": str(latent.dtype).split(".", maxsplit=1)[-1],
}
headers = { "Content-Type": "tensor/binary" }
if shared.opts.remote_vae_type == 'png':
if 'video' in model_type:
params["partial_postprocess"] = False
params["output_type"] = "pt"
params["output_tensor_type"] = "binary"
headers["Accept"] = "tensor/binary"
elif shared.opts.remote_vae_type == 'png':
params["image_format"] = "png"
params["output_type"] = "pil"
headers["Accept"] = "image/png"
@@ -81,7 +89,7 @@ def remote_decode(latents: torch.Tensor, width: int = 0, height: int = 0, model_
shared.log.error(f'Decode: type="remote" model={model_type} code={response.status_code} shape={latent.shape} url="{url}" args={params} headers={response.headers} response={response.json()}')
else:
content += len(response.content)
if shared.opts.remote_vae_type == 'raw':
if shared.opts.remote_vae_type == 'raw' or 'video' in model_type:
shape = json.loads(response.headers["shape"])
dtype = response.headers["dtype"]
tensor = torch.frombuffer(bytearray(response.content), dtype=dtypes[dtype]).reshape(shape)
+36 -19
View File
@@ -16,6 +16,9 @@ TAESD_MODELS = {
'TAESD 1.2 Chocolate-Dipped Shortbread': { 'fn': 'taesd_12_', 'uri': 'https://github.com/madebyollin/taesd/raw/8909b44e3befaa0efa79c5791e4fe1c4d4f7884e', 'model': None },
'TAESD 1.1 Fruit Loops': { 'fn': 'taesd_11_', 'uri': 'https://github.com/madebyollin/taesd/raw/3e8a8a2ab4ad4079db60c1c7dc1379b4cc0c6b31', 'model': None },
'TAESD 1.0': { 'fn': 'taesd_10_', 'uri': 'https://github.com/madebyollin/taesd/raw/88012e67cf0454e6d90f98911fe9d4aef62add86', 'model': None },
'TAE HunyuanVideo': { 'fn': 'taehv.pth', 'uri': 'https://github.com/madebyollin/taehv/raw/refs/heads/main/taehv.pth', 'model': None },
'TAE WanVideo': { 'fn': 'taew1.pth', 'uri': 'https://github.com/madebyollin/taehv/raw/refs/heads/main/taew2_1.pth', 'model': None },
'TAE MochiVideo': { 'fn': 'taem1.pth', 'uri': 'https://github.com/madebyollin/taem1/raw/refs/heads/main/taem1.pth', 'model': None },
}
CQYAN_MODELS = {
'Hybrid-Tiny SD': {
@@ -35,49 +38,63 @@ prev_model = ''
lock = threading.Lock()
def warn_once(msg):
def warn_once(msg, variant=None):
from modules import shared
variant = variant or shared.opts.taesd_variant
global prev_warnings # pylint: disable=global-statement
if not prev_warnings:
prev_warnings = True
shared.log.error(f'Decode: type="taesd" variant="{shared.opts.taesd_variant}": {msg}')
shared.log.error(f'Decode: type="taesd" variant="{variant}": {msg}')
return Image.new('RGB', (8, 8), color = (0, 0, 0))
def get_model(model_type = 'decoder'):
def get_model(model_type = 'decoder', variant = None):
global prev_cls, prev_type, prev_model # pylint: disable=global-statement
from modules import shared
cls = shared.sd_model_type
if cls == 'ldm':
cls = 'sd'
variant = variant or shared.opts.taesd_variant
folder = os.path.join(paths.models_path, "TAESD")
os.makedirs(folder, exist_ok=True)
if 'sd' not in cls and 'f1' not in cls:
if 'video' in cls:
return None
if ('sd' not in cls) and ('f1' not in cls):
warn_once(f'cls={shared.sd_model.__class__.__name__} type={cls} unsuppported')
return None
if shared.opts.taesd_variant.startswith('TAESD'):
cfg = TAESD_MODELS[shared.opts.taesd_variant]
if (cls == prev_cls) and (model_type == prev_type) and (shared.opts.taesd_variant == prev_model) and (cfg['model'] is not None):
if variant.startswith('TAESD'):
cfg = TAESD_MODELS[variant]
if (cls == prev_cls) and (model_type == prev_type) and (variant == prev_model) and (cfg['model'] is not None):
return cfg['model']
fn = os.path.join(folder, cfg['fn'] + cls + '_' + model_type + '.pth')
if not os.path.exists(fn):
uri = cfg['uri'] + '/tae' + cls + '_' + model_type + '.pth'
try:
shared.log.info(f'Decode: type="taesd" variant="{shared.opts.taesd_variant}": uri="{uri}" fn="{fn}" download')
shared.log.info(f'Decode: type="taesd" variant="{variant}": uri="{uri}" fn="{fn}" download')
torch.hub.download_url_to_file(uri, fn)
except Exception as e:
warn_once(f'download uri={uri} {e}')
if os.path.exists(fn):
prev_cls = cls
prev_type = model_type
prev_model = shared.opts.taesd_variant
shared.log.debug(f'Decode: type="taesd" variant="{shared.opts.taesd_variant}" fn="{fn}" load')
from modules.taesd.taesd import TAESD
TAESD_MODELS[shared.opts.taesd_variant]['model'] = TAESD(decoder_path=fn if model_type=='decoder' else None, encoder_path=fn if model_type=='encoder' else None)
return TAESD_MODELS[shared.opts.taesd_variant]['model']
elif shared.opts.taesd_variant.startswith('Hybrid'):
cfg = CQYAN_MODELS[shared.opts.taesd_variant].get(cls, None)
if (cls == prev_cls) and (model_type == prev_type) and (shared.opts.taesd_variant == prev_model) and (cfg['model'] is not None):
prev_model = variant
shared.log.debug(f'Decode: type="taesd" variant="{variant}" fn="{fn}" load')
if 'TAEHV' in variant:
from modules.taesd.taehv import TAEHV
TAESD_MODELS[variant]['model'] = TAEHV(checkpoint_path=fn)
if 'TAEW2' in variant:
from modules.taesd.taehv import TAEHV
TAESD_MODELS[variant]['model'] = TAEHV(checkpoint_path=fn)
elif 'TAEM1' in variant:
from modules.taesd.taem1 import TAEM1
TAESD_MODELS[variant]['model'] = TAEM1(checkpoint_path=fn)
else:
from modules.taesd.taesd import TAESD
TAESD_MODELS[variant]['model'] = TAESD(decoder_path=fn if model_type=='decoder' else None, encoder_path=fn if model_type=='encoder' else None)
return TAESD_MODELS[variant]['model']
elif variant.startswith('Hybrid'):
cfg = CQYAN_MODELS[variant].get(cls, None)
if (cls == prev_cls) and (model_type == prev_type) and (variant == prev_model) and (cfg['model'] is not None):
return cfg['model']
if cfg is None:
warn_once(f'cls={shared.sd_model.__class__.__name__} type={cls} unsuppported')
@@ -85,8 +102,8 @@ def get_model(model_type = 'decoder'):
repo = cfg['repo']
prev_cls = cls
prev_type = model_type
prev_model = shared.opts.taesd_variant
shared.log.debug(f'Decode: type="taesd" variant="{shared.opts.taesd_variant}" id="{repo}" load')
prev_model = variant
shared.log.debug(f'Decode: type="taesd" variant="{variant}" id="{repo}" load')
dtype = devices.dtype_vae if devices.dtype_vae != torch.bfloat16 else torch.float16 # taesd does not support bf16
if 'tiny' in repo:
from diffusers.models import AutoencoderTiny
@@ -95,7 +112,7 @@ def get_model(model_type = 'decoder'):
from modules.taesd.hybrid_small import AutoencoderSmall
vae = AutoencoderSmall.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir, torch_dtype=dtype)
vae = vae.to(devices.device, dtype=dtype)
CQYAN_MODELS[shared.opts.taesd_variant][cls]['model'] = vae
CQYAN_MODELS[variant][cls]['model'] = vae
return vae
else:
warn_once(f'cls={shared.sd_model.__class__.__name__} type={cls} unsuppported')
+284
View File
@@ -0,0 +1,284 @@
#!/usr/bin/env python3
"""
Tiny AutoEncoder for Hunyuan Video
(DNN for encoding / decoding videos to Hunyuan Video's latent space)
"""
from collections import namedtuple
import torch
import torch.nn as nn
import torch.nn.functional as F
from tqdm.auto import tqdm
DecoderResult = namedtuple("DecoderResult", ("frame", "memory"))
TWorkItem = namedtuple("TWorkItem", ("input_tensor", "block_index"))
def conv(n_in, n_out, **kwargs):
return nn.Conv2d(n_in, n_out, 3, padding=1, **kwargs)
class Clamp(nn.Module):
def forward(self, x):
return torch.tanh(x / 3) * 3
class MemBlock(nn.Module):
def __init__(self, n_in, n_out):
super().__init__()
self.conv = nn.Sequential(conv(n_in * 2, n_out), nn.ReLU(inplace=True), conv(n_out, n_out), nn.ReLU(inplace=True), conv(n_out, n_out))
self.skip = nn.Conv2d(n_in, n_out, 1, bias=False) if n_in != n_out else nn.Identity()
self.act = nn.ReLU(inplace=True)
def forward(self, x, past):
return self.act(self.conv(torch.cat([x, past], 1)) + self.skip(x))
class TPool(nn.Module):
def __init__(self, n_f, stride):
super().__init__()
self.stride = stride
self.conv = nn.Conv2d(n_f*stride,n_f, 1, bias=False)
def forward(self, x):
_NT, C, H, W = x.shape
return self.conv(x.reshape(-1, self.stride * C, H, W))
class TGrow(nn.Module):
def __init__(self, n_f, stride):
super().__init__()
self.stride = stride
self.conv = nn.Conv2d(n_f, n_f*stride, 1, bias=False)
def forward(self, x):
_NT, C, H, W = x.shape
x = self.conv(x)
return x.reshape(-1, C, H, W)
def apply_model_with_memblocks(model, x, parallel, show_progress_bar):
"""
Apply a sequential model with memblocks to the given input.
Args:
- model: nn.Sequential of blocks to apply
- x: input data, of dimensions NTCHW
- parallel: if True, parallelize over timesteps (fast but uses O(T) memory)
if False, each timestep will be processed sequentially (slow but uses O(1) memory)
- show_progress_bar: if True, enables tqdm progressbar display
Returns NTCHW tensor of output data.
"""
assert x.ndim == 5, f"TAEHV operates on NTCHW tensors, but got {x.ndim}-dim tensor"
N, T, C, H, W = x.shape
if parallel:
x = x.reshape(N*T, C, H, W)
# parallel over input timesteps, iterate over blocks
for b in tqdm(model, disable=not show_progress_bar):
if isinstance(b, MemBlock):
NT, C, H, W = x.shape
T = NT // N
_x = x.reshape(N, T, C, H, W)
mem = F.pad(_x, (0,0,0,0,0,0,1,0), value=0)[:,:T].reshape(x.shape)
x = b(x, mem)
else:
x = b(x)
NT, C, H, W = x.shape
T = NT // N
x = x.view(N, T, C, H, W)
else:
# TODO(oboerbohan): at least on macos this still gradually uses more memory during decode...
# need to fix :(
out = []
# iterate over input timesteps and also iterate over blocks.
# because of the cursed TPool/TGrow blocks, this is not a nested loop,
# it's actually a ***graph traversal*** problem! so let's make a queue
work_queue = [TWorkItem(xt, 0) for t, xt in enumerate(x.reshape(N, T * C, H, W).chunk(T, dim=1))]
# in addition to manually managing our queue, we also need to manually manage our progressbar.
# we'll update it for every source node that we consume.
progress_bar = tqdm(range(T), disable=not show_progress_bar)
# we'll also need a separate addressable memory per node as well
mem = [None] * len(model)
while work_queue:
xt, i = work_queue.pop(0)
if i == 0:
# new source node consumed
progress_bar.update(1)
if i == len(model):
# reached end of the graph, append result to output list
out.append(xt)
else:
# fetch the block to process
b = model[i]
if isinstance(b, MemBlock):
# mem blocks are simple since we're visiting the graph in causal order
if mem[i] is None:
xt_new = b(xt, xt * 0)
mem[i] = xt
else:
xt_new = b(xt, mem[i])
mem[i].copy_(xt) # inplace might reduce mysterious pytorch memory allocations? doesn't help though
# add successor to work queue
work_queue.insert(0, TWorkItem(xt_new, i+1))
elif isinstance(b, TPool):
# pool blocks are miserable
if mem[i] is None:
mem[i] = [] # pool memory is itself a queue of inputs to pool
mem[i].append(xt)
if len(mem[i]) > b.stride:
# pool mem is in invalid state, we should have pooled before this
raise ValueError("???")
elif len(mem[i]) < b.stride:
# pool mem is not yet full, go back to processing the work queue
pass
else:
# pool mem is ready, run the pool block
N, C, H, W = xt.shape
xt = b(torch.cat(mem[i], 1).view(N*b.stride, C, H, W))
# reset the pool mem
mem[i] = []
# add successor to work queue
work_queue.insert(0, TWorkItem(xt, i+1))
elif isinstance(b, TGrow):
xt = b(xt)
NT, C, H, W = xt.shape
# each tgrow has multiple successor nodes
for xt_next in reversed(xt.view(N, b.stride*C, H, W).chunk(b.stride, 1)):
# add successor to work queue
work_queue.insert(0, TWorkItem(xt_next, i+1))
else:
# normal block with no funny business
xt = b(xt)
# add successor to work queue
work_queue.insert(0, TWorkItem(xt, i+1))
progress_bar.close()
x = torch.stack(out, 1)
return x
class TAEHV(nn.Module):
latent_channels = 16
image_channels = 3
def __init__(self, checkpoint_path="taehv.pth", decoder_time_upscale=(True, True), decoder_space_upscale=(True, True, True)):
"""Initialize pretrained TAEHV from the given checkpoint.
Arg:
checkpoint_path: path to weight file to load. taehv.pth for Hunyuan, taew2_1.pth for Wan 2.1.
decoder_time_upscale: whether temporal upsampling is enabled for each block. upsampling can be disabled for a cheaper preview.
decoder_space_upscale: whether spatial upsampling is enabled for each block. upsampling can be disabled for a cheaper preview.
"""
super().__init__()
self.encoder = nn.Sequential(
conv(TAEHV.image_channels, 64), nn.ReLU(inplace=True),
TPool(64, 2), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64),
TPool(64, 2), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64),
TPool(64, 1), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64),
conv(64, TAEHV.latent_channels),
)
n_f = [256, 128, 64, 64]
self.frames_to_trim = 2**sum(decoder_time_upscale) - 1
self.decoder = nn.Sequential(
Clamp(), conv(TAEHV.latent_channels, n_f[0]), nn.ReLU(inplace=True),
MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), nn.Upsample(scale_factor=2 if decoder_space_upscale[0] else 1), TGrow(n_f[0], 1), conv(n_f[0], n_f[1], bias=False),
MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]), nn.Upsample(scale_factor=2 if decoder_space_upscale[1] else 1), TGrow(n_f[1], 2 if decoder_time_upscale[0] else 1), conv(n_f[1], n_f[2], bias=False),
MemBlock(n_f[2], n_f[2]), MemBlock(n_f[2], n_f[2]), MemBlock(n_f[2], n_f[2]), nn.Upsample(scale_factor=2 if decoder_space_upscale[2] else 1), TGrow(n_f[2], 2 if decoder_time_upscale[1] else 1), conv(n_f[2], n_f[3], bias=False),
nn.ReLU(inplace=True), conv(n_f[3], TAEHV.image_channels),
)
if checkpoint_path is not None:
self.load_state_dict(self.patch_tgrow_layers(torch.load(checkpoint_path, map_location="cpu", weights_only=True)))
def patch_tgrow_layers(self, sd):
"""Patch TGrow layers to use a smaller kernel if needed.
Args:
sd: state dict to patch
"""
new_sd = self.state_dict()
for i, layer in enumerate(self.decoder):
if isinstance(layer, TGrow):
key = f"decoder.{i}.conv.weight"
if sd[key].shape[0] > new_sd[key].shape[0]:
# take the last-timestep output channels
sd[key] = sd[key][-new_sd[key].shape[0]:]
return sd
def encode_video(self, x, parallel=True, show_progress_bar=True):
"""Encode a sequence of frames.
Args:
x: input NTCHW RGB (C=3) tensor with values in [0, 1].
parallel: if True, all frames will be processed at once.
(this is faster but may require more memory).
if False, frames will be processed sequentially.
Returns NTCHW latent tensor with ~Gaussian values.
"""
return apply_model_with_memblocks(self.encoder, x, parallel, show_progress_bar)
def decode_video(self, x, parallel=True, show_progress_bar=True):
"""Decode a sequence of frames.
Args:
x: input NTCHW latent (C=12) tensor with ~Gaussian values.
parallel: if True, all frames will be processed at once.
(this is faster but may require more memory).
if False, frames will be processed sequentially.
Returns NTCHW RGB tensor with ~[0, 1] values.
"""
x = apply_model_with_memblocks(self.decoder, x, parallel, show_progress_bar)
return x[:, self.frames_to_trim:]
def forward(self, x):
return self.c(x)
@torch.no_grad()
def main():
"""Run TAEHV roundtrip reconstruction on the given video paths."""
import sys
import cv2 # no highly esteemed deed is commemorated here
class VideoTensorReader:
def __init__(self, video_file_path):
self.cap = cv2.VideoCapture(video_file_path)
assert self.cap.isOpened(), f"Could not load {video_file_path}"
self.fps = self.cap.get(cv2.CAP_PROP_FPS)
def __iter__(self):
return self
def __next__(self):
ret, frame = self.cap.read()
if not ret:
self.cap.release()
raise StopIteration # End of video or error
return torch.from_numpy(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)).permute(2, 0, 1) # BGR HWC -> RGB CHW
class VideoTensorWriter:
def __init__(self, video_file_path, width_height, fps=30):
self.writer = cv2.VideoWriter(video_file_path, cv2.VideoWriter_fourcc(*'mp4v'), fps, width_height)
assert self.writer.isOpened(), f"Could not create writer for {video_file_path}"
def write(self, frame_tensor):
assert frame_tensor.ndim == 3 and frame_tensor.shape[0] == 3, f"{frame_tensor.shape}??"
self.writer.write(cv2.cvtColor(frame_tensor.permute(1, 2, 0).numpy(), cv2.COLOR_RGB2BGR)) # RGB CHW -> BGR HWC
def __del__(self):
if hasattr(self, 'writer'):
self.writer.release()
dev = torch.device("cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu")
dtype = torch.float16
print("Using device", dev, "and dtype", dtype)
taehv = TAEHV().to(dev, dtype)
for video_path in sys.argv[1:]:
print(f"Processing {video_path}...")
video_in = VideoTensorReader(video_path)
video = torch.stack(list(video_in), 0)[None]
vid_dev = video.to(dev, dtype).div_(255.0)
# convert to device tensor
if video.numel() < 100_000_000:
print(f" {video_path} seems small enough, will process all frames in parallel")
# convert to device tensor
vid_enc = taehv.encode_video(vid_dev)
print(f" Encoded {video_path} -> {vid_enc.shape}. Decoding...")
vid_dec = taehv.decode_video(vid_enc)
print(f" Decoded {video_path} -> {vid_dec.shape}")
else:
print(f" {video_path} seems large, will process each frame sequentially")
# convert to device tensor
vid_enc = taehv.encode_video(vid_dev, parallel=False)
print(f" Encoded {video_path} -> {vid_enc.shape}. Decoding...")
vid_dec = taehv.decode_video(vid_enc, parallel=False)
print(f" Decoded {video_path} -> {vid_dec.shape}")
video_out_path = video_path + ".reconstructed_by_taehv.mp4"
video_out = VideoTensorWriter(video_out_path, (vid_dec.shape[-1], vid_dec.shape[-2]), fps=int(round(video_in.fps)))
for frame in vid_dec.clamp_(0, 1).mul_(255).round_().byte().cpu()[0]:
video_out.write(frame)
print(f" Saved to {video_out_path}")
if __name__ == "__main__":
main()
+272
View File
@@ -0,0 +1,272 @@
#!/usr/bin/env python3
"""
Tiny AutoEncoder for Mochi 1
(DNN for encoding / decoding videos to Mochi 1's latent space)
"""
from collections import namedtuple
import torch
import torch.nn as nn
import torch.nn.functional as F
from tqdm.auto import tqdm
DecoderResult = namedtuple("DecoderResult", ("frame", "memory"))
TWorkItem = namedtuple("TWorkItem", ("input_tensor", "block_index"))
def conv(n_in, n_out, **kwargs):
return nn.Conv2d(n_in, n_out, 3, padding=1, **kwargs)
class Clamp(nn.Module):
def forward(self, x):
return torch.tanh(x / 3) * 3
class MemBlock(nn.Module):
def __init__(self, n_in, n_out):
super().__init__()
self.conv = nn.Sequential(conv(n_in * 2, n_out), nn.ReLU(inplace=True), conv(n_out, n_out), nn.ReLU(inplace=True), conv(n_out, n_out))
self.skip = nn.Conv2d(n_in, n_out, 1, bias=False) if n_in != n_out else nn.Identity()
self.act = nn.ReLU(inplace=True)
def forward(self, x, past):
return self.act(self.conv(torch.cat([x, past], 1)) + self.skip(x))
class TPool(nn.Module):
def __init__(self, n_f, stride):
super().__init__()
self.stride = stride
self.conv = nn.Conv2d(n_f*stride,n_f, 1, bias=False)
def forward(self, x):
_NT, C, H, W = x.shape
return self.conv(x.reshape(-1, self.stride * C, H, W))
class TGrow(nn.Module):
def __init__(self, n_f, stride):
super().__init__()
self.stride = stride
self.conv = nn.Conv2d(n_f, n_f*stride, 1, bias=False)
def forward(self, x):
_NT, C, H, W = x.shape
x = self.conv(x)
return x.reshape(-1, C, H, W)
def apply_model_with_memblocks(model, x, parallel, show_progress_bar):
"""
Apply a sequential model with memblocks to the given input.
Args:
- model: nn.Sequential of blocks to apply
- x: input data, of dimensions NTCHW
- parallel: if True, parallelize over timesteps (fast but uses O(T) memory)
if False, each timestep will be processed sequentially (slow but uses O(1) memory)
- show_progress_bar: if True, enables tqdm progressbar display
Returns NTCHW tensor of output data.
"""
assert x.ndim == 5, f"TAEM1 operates on NTCHW tensors, but got {x.ndim}-dim tensor"
N, T, C, H, W = x.shape
if parallel:
x = x.reshape(N*T, C, H, W)
# parallel over input timesteps, iterate over blocks
for b in tqdm(model, disable=not show_progress_bar):
if isinstance(b, MemBlock):
NT, C, H, W = x.shape
T = NT // N
_x = x.reshape(N, T, C, H, W)
mem = F.pad(_x, (0,0,0,0,0,0,1,0), value=0)[:,:T].reshape(x.shape)
x = b(x, mem)
else:
x = b(x)
NT, C, H, W = x.shape
T = NT // N
x = x.view(N, T, C, H, W)
else:
# TODO(oboerbohan): at least on macos this still gradually uses more memory during decode...
# need to fix :(
out = []
# iterate over input timesteps and also iterate over blocks.
# because of the cursed TPool/TGrow blocks, this is not a nested loop,
# it's actually a ***graph traversal*** problem! so let's make a queue
work_queue = [TWorkItem(xt, 0) for t, xt in enumerate(x.reshape(N, T * C, H, W).chunk(T, dim=1))]
# in addition to manually managing our queue, we also need to manually manage our progressbar.
# we'll update it for every source node that we consume.
progress_bar = tqdm(range(T), disable=not show_progress_bar)
# we'll also need a separate addressable memory per node as well
mem = [None] * len(model)
while work_queue:
xt, i = work_queue.pop(0)
if i == 0:
# new source node consumed
progress_bar.update(1)
if i == len(model):
# reached end of the graph, append result to output list
out.append(xt)
else:
# fetch the block to process
b = model[i]
if isinstance(b, MemBlock):
# mem blocks are simple since we're visiting the graph in causal order
if mem[i] is None:
xt_new = b(xt, xt * 0)
mem[i] = xt
else:
xt_new = b(xt, mem[i])
mem[i].copy_(xt) # inplace might reduce mysterious pytorch memory allocations? doesn't help though
# add successor to work queue
work_queue.insert(0, TWorkItem(xt_new, i+1))
elif isinstance(b, TPool):
# pool blocks are miserable
if mem[i] is None:
mem[i] = [] # pool memory is itself a queue of inputs to pool
mem[i].append(xt)
if len(mem[i]) > b.stride:
# pool mem is in invalid state, we should have pooled before this
raise ValueError("???")
elif len(mem[i]) < b.stride:
# pool mem is not yet full, go back to processing the work queue
pass
else:
# pool mem is ready, run the pool block
N, C, H, W = xt.shape
xt = b(torch.cat(mem[i], 1).view(N*b.stride, C, H, W))
# reset the pool mem
mem[i] = []
# add successor to work queue
work_queue.insert(0, TWorkItem(xt, i+1))
elif isinstance(b, TGrow):
xt = b(xt)
NT, C, H, W = xt.shape
# each tgrow has multiple successor nodes
for xt_next in reversed(xt.view(N, b.stride*C, H, W).chunk(b.stride, 1)):
# add successor to work queue
work_queue.insert(0, TWorkItem(xt_next, i+1))
else:
# normal block with no funny business
xt = b(xt)
# add successor to work queue
work_queue.insert(0, TWorkItem(xt, i+1))
progress_bar.close()
x = torch.stack(out, 1)
return x
class TAEM1(nn.Module):
latent_channels = 12
image_channels = 3
def __init__(self, checkpoint_path="taem1.pth"):
"""Initialize pretrained TAEM1 from the given checkpoints."""
super().__init__()
self.encoder = nn.Sequential(
conv(TAEM1.image_channels, 64), nn.ReLU(inplace=True),
TPool(64, 3), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64),
TPool(64, 2), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64),
TPool(64, 1), conv(64, 64, stride=2, bias=False), MemBlock(64, 64), MemBlock(64, 64), MemBlock(64, 64),
conv(64, TAEM1.latent_channels),
)
n_f = [256, 128, 64, 64]
self.decoder = nn.Sequential(
Clamp(), conv(TAEM1.latent_channels, n_f[0]), nn.ReLU(inplace=True),
MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), nn.Upsample(scale_factor=2), TGrow(n_f[0], 1), conv(n_f[0], n_f[1], bias=False),
MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]), nn.Upsample(scale_factor=2), TGrow(n_f[1], 2), conv(n_f[1], n_f[2], bias=False),
MemBlock(n_f[2], n_f[2]), MemBlock(n_f[2], n_f[2]), MemBlock(n_f[2], n_f[2]), nn.Upsample(scale_factor=2), TGrow(n_f[2], 3), conv(n_f[2], n_f[3], bias=False),
nn.ReLU(inplace=True), conv(n_f[3], TAEM1.image_channels),
)
if checkpoint_path is not None:
self.load_state_dict(torch.load(checkpoint_path, map_location="cpu", weights_only=True))
def encode_video(self, x, parallel=True, show_progress_bar=True):
"""Encode a sequence of frames.
Args:
x: input NTCHW RGB (C=3) tensor with values in [0, 1].
parallel: if True, all frames will be processed at once.
(this is faster but may require more memory).
if False, frames will be processed sequentially.
Returns NTCHW latent tensor with ~Gaussian values.
"""
return apply_model_with_memblocks(self.encoder, x, parallel, show_progress_bar)
def decode_video(self, x, parallel=True, show_progress_bar=True):
"""Decode a sequence of frames.
Args:
x: input NTCHW latent (C=12) tensor with ~Gaussian values.
parallel: if True, all frames will be processed at once.
(this is faster but may require more memory).
if False, frames will be processed sequentially.
Returns NTCHW RGB tensor with ~[0, 1] values.
"""
x = apply_model_with_memblocks(self.decoder, x, parallel, show_progress_bar)
# NOTE:
# the Mochi VAE does not preserve shape along the time axis;
# videos are encoded to floor((n_in - 1)/6)+1 latent frames
# (which makes sense, it's stride 6, so 12 -> 2 and 13->3)
# but then they're decoded to only the *minimal* number
# of input frames (3 latents get decoded to 13 frames, not 18)
# in order to achieve the intended causal structure...
# anyway, that's why we have to remove some frames here.
# mochi-VAE does the slicing at each TGrow (save compute/mem?)
# but I think it's basically the same
return x[:, 5:]
def forward(self, x):
return self.c(x)
@torch.no_grad()
def main():
"""Run TAEM1 roundtrip reconstruction on the given video paths."""
import sys
import cv2 # no highly esteemed deed is commemorated here
class VideoTensorReader:
def __init__(self, video_file_path):
self.cap = cv2.VideoCapture(video_file_path)
assert self.cap.isOpened(), f"Could not load {video_file_path}"
self.fps = self.cap.get(cv2.CAP_PROP_FPS)
def __iter__(self):
return self
def __next__(self):
ret, frame = self.cap.read()
if not ret:
self.cap.release()
raise StopIteration # End of video or error
return torch.from_numpy(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)).permute(2, 0, 1) # BGR HWC -> RGB CHW
class VideoTensorWriter:
def __init__(self, video_file_path, width_height, fps=30):
self.writer = cv2.VideoWriter(video_file_path, cv2.VideoWriter_fourcc(*'mp4v'), fps, width_height)
assert self.writer.isOpened(), f"Could not create writer for {video_file_path}"
def write(self, frame_tensor):
assert frame_tensor.ndim == 3 and frame_tensor.shape[0] == 3, f"{frame_tensor.shape}??"
self.writer.write(cv2.cvtColor(frame_tensor.permute(1, 2, 0).numpy(), cv2.COLOR_RGB2BGR)) # RGB CHW -> BGR HWC
def __del__(self):
if hasattr(self, 'writer'):
self.writer.release()
dev = torch.device("cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu")
dtype = torch.float16
print("Using device", dev, "and dtype", dtype)
taem1 = TAEM1().to(dev, dtype)
for video_path in sys.argv[1:]:
print(f"Processing {video_path}...")
video_in = VideoTensorReader(video_path)
video = torch.stack(list(video_in), 0)[None]
vid_dev = video.to(dev, dtype).div_(255.0)
# convert to device tensor
if video.numel() < 100_000_000:
print(f" {video_path} seems small enough, will process all frames in parallel")
# convert to device tensor
vid_enc = taem1.encode_video(vid_dev)
print(f" Encoded {video_path} -> {vid_enc.shape}. Decoding...")
vid_dec = taem1.decode_video(vid_enc)
print(f" Decoded {video_path} -> {vid_dec.shape}")
else:
print(f" {video_path} seems large, will process each frame sequentially")
# convert to device tensor
vid_enc = taem1.encode_video(vid_dev, parallel=False)
print(f" Encoded {video_path} -> {vid_enc.shape}. Decoding...")
vid_dec = taem1.decode_video(vid_enc, parallel=False)
print(f" Decoded {video_path} -> {vid_dec.shape}")
video_out_path = video_path + ".reconstructed_by_taem1.mp4"
video_out = VideoTensorWriter(video_out_path, (vid_dec.shape[-1], vid_dec.shape[-2]), fps=int(round(video_in.fps)))
for frame in vid_dec.clamp_(0, 1).mul_(255).round_().byte().cpu()[0]:
video_out.write(frame)
print(f" Saved to {video_out_path}")
if __name__ == "__main__":
main()
+3 -3
View File
@@ -1,7 +1,7 @@
import gradio as gr
from modules import shared, sd_models, timer, images, ui_common, ui_sections, ui_symbols, call_queue, generation_parameters_copypaste
from modules.ui_components import ToolButton
from modules.video_models import models_def, video_utils
from modules.video_models import models_def, video_utils, video_load
def engine_change(engine):
@@ -21,7 +21,7 @@ def model_change(engine, model):
sd_models.unload_model_weights()
msg = 'Video model unloaded'
else:
msg = video_utils.load_model(selected)
msg = video_load.load_model(selected)
else:
sd_models.unload_model_weights()
msg = 'Video model unloaded'
@@ -84,7 +84,7 @@ def create_ui():
steps, sampler_index = ui_sections.create_sampler_and_steps_selection(None, "video")
with gr.Row():
sampler_shift = gr.Slider(label='Sampler shift', minimum=0.0, maximum=20.0, step=0.1, value=7.0, elem_id="video_scheduler_shift")
dynamic_shift = gr.Checkbox(label='Dynamic shift', value=False, elem_id="video_dynamic_shift", interactive=False)
dynamic_shift = gr.Checkbox(label='Dynamic shift', value=False, elem_id="video_dynamic_shift", interactive=False) # TODO video: dynamic shift
with gr.Row():
guidance_scale = gr.Slider(label='Guidance scale', minimum=0.0, maximum=14.0, step=0.1, value=6.0, elem_id="video_guidance_scale")
guidance_true = gr.Slider(label='True guidance', minimum=0.0, maximum=14.0, step=0.1, value=1.0, elem_id="video_guidance_true")
+33
View File
@@ -17,6 +17,7 @@ class Model():
te_folder: str = 'text_encoder'
te_hijack: bool = True
vae_hijack: bool = True
vae_remote: bool = False
models = {
@@ -25,18 +26,21 @@ models = {
Model(name='None'),
Model(name='Hunyuan Video T2V',
url='https://huggingface.co/tencent/HunyuanVideo',
vae_remote=True,
repo='hunyuanvideo-community/HunyuanVideo',
repo_cls=diffusers.HunyuanVideoPipeline,
te_cls=transformers.LlamaModel,
dit_cls=diffusers.HunyuanVideoTransformer3DModel),
Model(name='Hunyuan Video I2V', # https://github.com/huggingface/diffusers/pull/10983
url='https://huggingface.co/tencent/HunyuanVideo-I2V',
vae_remote=True,
repo='hunyuanvideo-community/HunyuanVideo-I2V',
repo_cls=diffusers.HunyuanVideoImageToVideoPipeline,
te_cls=transformers.LlavaForConditionalGeneration,
dit_cls=diffusers.HunyuanVideoTransformer3DModel),
Model(name='SkyReels Hunyuan T2V', # https://github.com/huggingface/diffusers/pull/10837
url='https://huggingface.co/Skywork/SkyReels-V1-Hunyuan-T2V',
vae_remote=True,
repo='hunyuanvideo-community/HunyuanVideo',
repo_cls=diffusers.HunyuanVideoPipeline,
te_cls=transformers.LlamaModel,
@@ -45,6 +49,7 @@ models = {
dit_cls=diffusers.HunyuanVideoTransformer3DModel),
Model(name='SkyReels Hunyuan I2V', # https://github.com/huggingface/diffusers/pull/10837
url='https://huggingface.co/Skywork/SkyReels-V1-Hunyuan-I2V',
vae_remote=True,
repo='hunyuanvideo-community/HunyuanVideo',
te_cls=transformers.LlamaModel,
dit='Skywork/SkyReels-V1-Hunyuan-I2V',
@@ -52,6 +57,7 @@ models = {
dit_cls=diffusers.HunyuanVideoTransformer3DModel),
Model(name='Fast Hunyuan T2V', # https://github.com/hao-ai-lab/FastVideo/blob/8a77cf22c9b9e7f931f42bc4b35d21fd91d24e45/fastvideo/models/hunyuan/inference.py#L213
url='https://huggingface.co/FastVideo/FastHunyuan',
vae_remote=True,
repo='hunyuanvideo-community/HunyuanVideo',
repo_cls=diffusers.HunyuanVideoPipeline,
te_cls=transformers.LlamaModel,
@@ -97,6 +103,33 @@ models = {
te_cls=transformers.T5EncoderModel,
dit_cls=diffusers.LTXVideoTransformer3DModel),
],
'WAN Video': [
Model(name='None'),
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',
repo_cls=diffusers.WanPipeline,
te_cls=transformers.T5EncoderModel,
dit_cls=diffusers.WanTransformer3DModel),
Model(name='WAN 2.1 14B T2V',
url='https://huggingface.co/Wan-AI/Wan2.1-T2V-14B-Diffusers',
repo='Wan-AI/Wan2.1-T2V-14B-Diffusers',
repo_cls=diffusers.WanPipeline,
te_cls=transformers.T5EncoderModel,
dit_cls=diffusers.WanTransformer3DModel),
Model(name='WAN 2.1 14B I2V 480p',
url='https://huggingface.co/Wan-AI/Wan2.1-I2V-14B-480P-Diffusers',
repo='Wan-AI/Wan2.1-I2V-14B-480P-Diffusers',
repo_cls=diffusers.WanImageToVideoPipeline,
te_cls=transformers.T5EncoderModel,
dit_cls=diffusers.WanTransformer3DModel),
Model(name='WAN 2.1 14B I2V 720p',
url='https://huggingface.co/Wan-AI/Wan2.1-I2V-14B-720P-Diffusers',
repo='Wan-AI/Wan2.1-I2V-14B-720P-Diffusers',
repo_cls=diffusers.WanImageToVideoPipeline,
te_cls=transformers.T5EncoderModel,
dit_cls=diffusers.WanTransformer3DModel),
],
'Mochi Video': [
Model(name='None'),
Model(name='Mochi 1 T2V',
+4 -4
View File
@@ -1,7 +1,7 @@
import os
import time
from modules import shared, errors, sd_models, processing, devices, images, ui_common
from modules.video_models import models_def, video_utils
from modules.video_models import models_def, video_utils, video_load, video_vae
debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None
@@ -14,7 +14,7 @@ def generate(*args, **kwargs):
found = [model.name for model in models_def.models.get(engine, [])]
selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None
if not shared.sd_loaded or 'Allegro' not in shared.sd_model.__class__.__name__:
video_utils.load_model(selected)
video_load.load_model(selected)
if not shared.sd_loaded or 'Allegro' not in shared.sd_model.__class__.__name__:
return video_utils.queue_err('model not loaded')
debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}')
@@ -54,13 +54,13 @@ def generate(*args, **kwargs):
# set args
processing.fix_seed(p)
video_utils.set_vae_params(p.frames, vae_tile_frames)
video_vae.set_vae_params(p)
video_utils.set_prompt(p)
p.task_args['output_type'] = 'pil'
p.ops.append('video')
orig_dynamic_shift = shared.opts.schedulers_dynamic_shift
orig_sampler_shift = shared.opts.schedulers_shift
shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift
shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift
shared.opts.data['schedulers_shift'] = sampler_shift
debug(f'Video: task_args={p.task_args}')
+4 -4
View File
@@ -1,7 +1,7 @@
import os
import time
from modules import shared, errors, sd_models, processing, devices, images, ui_common
from modules.video_models import models_def, video_utils
from modules.video_models import models_def, video_utils, video_load, video_vae
debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None
@@ -14,7 +14,7 @@ def generate(*args, **kwargs):
found = [model.name for model in models_def.models.get(engine, [])]
selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None
if not shared.sd_loaded or 'Cog' not in shared.sd_model.__class__.__name__:
video_utils.load_model(selected)
video_load.load_model(selected)
if not shared.sd_loaded or 'Cog' not in shared.sd_model.__class__.__name__:
return video_utils.queue_err('model not loaded')
debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}')
@@ -54,13 +54,13 @@ def generate(*args, **kwargs):
# set args
processing.fix_seed(p)
video_utils.set_vae_params(p.frames, vae_tile_frames)
video_vae.set_vae_params(p)
video_utils.set_prompt(p)
p.task_args['output_type'] = 'pil'
p.ops.append('video')
orig_dynamic_shift = shared.opts.schedulers_dynamic_shift
orig_sampler_shift = shared.opts.schedulers_shift
shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift
shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift
shared.opts.data['schedulers_shift'] = sampler_shift
debug(f'Video: task_args={p.task_args}')
+8 -5
View File
@@ -1,7 +1,7 @@
import os
import time
from modules import shared, errors, sd_models, processing, devices, images, ui_common
from modules.video_models import models_def, video_utils
from modules.video_models import models_def, video_utils, video_load, video_vae
debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None
@@ -14,7 +14,7 @@ def generate(*args, **kwargs):
found = [model.name for model in models_def.models.get(engine, [])]
selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None
if not shared.sd_loaded or 'Hunyuan' not in shared.sd_model.__class__.__name__:
video_utils.load_model(selected)
video_load.load_model(selected)
if not shared.sd_loaded or 'Hunyuan' not in shared.sd_model.__class__.__name__:
return video_utils.queue_err('model not loaded')
debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}')
@@ -38,6 +38,9 @@ def generate(*args, **kwargs):
vae_tile_frames=int(vae_tile_frames),
override_settings=override_settings,
)
if p.vae_type == 'Remote' and not selected.vae_remote:
shared.log.warning(f'Video: model={selected.name} remote vae not supported')
p.vae_type = 'Default'
p.scripts = None
p.script_args = None
p.state = ui_state
@@ -54,13 +57,13 @@ def generate(*args, **kwargs):
# set args
processing.fix_seed(p)
video_utils.set_vae_params(p.frames, vae_tile_frames)
video_vae.set_vae_params(p)
video_utils.set_prompt(p)
p.task_args['output_type'] = 'pil'
p.task_args['output_type'] = 'latent' if (p.vae_type == 'Remote') else 'pil'
p.ops.append('video')
orig_dynamic_shift = shared.opts.schedulers_dynamic_shift
orig_sampler_shift = shared.opts.schedulers_shift
shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift
shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift
shared.opts.data['schedulers_shift'] = sampler_shift
debug(f'Video: task_args={p.task_args}')
+4 -4
View File
@@ -1,7 +1,7 @@
import os
import time
from modules import shared, errors, sd_models, processing, devices, images, ui_common
from modules.video_models import models_def, video_utils
from modules.video_models import models_def, video_utils, video_load, video_vae
debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None
@@ -14,7 +14,7 @@ def generate(*args, **kwargs):
found = [model.name for model in models_def.models.get(engine, [])]
selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None
if not shared.sd_loaded or 'LTX' not in shared.sd_model.__class__.__name__:
video_utils.load_model(selected)
video_load.load_model(selected)
if not shared.sd_loaded or 'LTX' not in shared.sd_model.__class__.__name__:
return video_utils.queue_err('model not loaded')
debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}')
@@ -54,13 +54,13 @@ def generate(*args, **kwargs):
# set args
processing.fix_seed(p)
video_utils.set_vae_params(p.frames, vae_tile_frames)
video_vae.set_vae_params(p)
video_utils.set_prompt(p)
p.task_args['output_type'] = 'pil'
p.ops.append('video')
orig_dynamic_shift = shared.opts.schedulers_dynamic_shift
orig_sampler_shift = shared.opts.schedulers_shift
shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift
shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift
shared.opts.data['schedulers_shift'] = sampler_shift
debug(f'Video: task_args={p.task_args}')
+4 -4
View File
@@ -1,7 +1,7 @@
import os
import time
from modules import shared, errors, sd_models, processing, devices, images, ui_common
from modules.video_models import models_def, video_utils
from modules.video_models import models_def, video_utils, video_load, video_vae
debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None
@@ -14,7 +14,7 @@ def generate(*args, **kwargs):
found = [model.name for model in models_def.models.get(engine, [])]
selected: models_def.Model = [m for m in models_def.models[engine] if m.name == model][0] if len(found) > 0 else None
if not shared.sd_loaded or 'Mochi' not in shared.sd_model.__class__.__name__:
video_utils.load_model(selected)
video_load.load_model(selected)
if not shared.sd_loaded or 'Mochi' not in shared.sd_model.__class__.__name__:
return video_utils.queue_err('model not loaded')
debug(f'Video generate: task={task_id} args={args} kwargs={kwargs}')
@@ -54,13 +54,13 @@ def generate(*args, **kwargs):
# set args
processing.fix_seed(p)
video_utils.set_vae_params(p.frames, vae_tile_frames)
video_vae.set_vae_params(p)
video_utils.set_prompt(p)
p.task_args['output_type'] = 'pil'
p.ops.append('video')
orig_dynamic_shift = shared.opts.schedulers_dynamic_shift
orig_sampler_shift = shared.opts.schedulers_shift
shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift # todo video sampler dynamic shift
shared.opts.data['schedulers_dynamic_shift'] = dynamic_shift
shared.opts.data['schedulers_shift'] = sampler_shift
debug(f'Video: task_args={p.task_args}')
+80
View File
@@ -0,0 +1,80 @@
import os
import time
from modules import shared, errors, sd_models, sd_checkpoint, model_quant, devices
from modules.video_models import models_def, video_utils, video_vae
loaded_model = None
debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None
def load_model(selected: models_def.Model):
if selected is None:
return
global loaded_model # pylint: disable=global-statement
if loaded_model == selected.name:
return
sd_models.unload_model_weights()
t0 = time.time()
# text encoder
try:
quant_args = model_quant.create_config(module='Text Encoder')
debug(f'Video load: module=te repo="{selected.te or selected.repo}" folder="{selected.te_folder}" cls={selected.te_cls.__name__} quant={video_utils.get_quant(quant_args)}')
text_encoder = selected.te_cls.from_pretrained(
pretrained_model_name_or_path=selected.te or selected.repo,
subfolder=selected.te_folder,
cache_dir=shared.opts.hfcache_dir,
torch_dtype=devices.dtype,
**quant_args
)
except Exception as e:
shared.log.error(f'video load: module=te cls={selected.te_cls.__name__} {e}')
errors.display(e, 'video')
text_encoder = None
# transformer
try:
quant_args = model_quant.create_config(module='Model')
debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" folder="{selected.dit_folder}" cls={selected.dit_cls.__name__} quant={video_utils.get_quant(quant_args)}')
transformer = selected.dit_cls.from_pretrained(
pretrained_model_name_or_path=selected.dit or selected.repo,
subfolder=selected.dit_folder,
torch_dtype=devices.dtype,
cache_dir=shared.opts.hfcache_dir,
**quant_args
)
except Exception as e:
shared.log.error(f'video load: module=transformer cls={selected.dit_cls.__name__} {e}')
errors.display(e, 'video')
transformer = None
# model
try:
debug(f'Video load: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__}')
shared.sd_model = selected.repo_cls.from_pretrained(
pretrained_model_name_or_path=selected.repo,
transformer=transformer,
text_encoder=text_encoder,
cache_dir=shared.opts.hfcache_dir,
torch_dtype=devices.dtype,
)
except Exception as e:
shared.log.error(f'video load: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__} {e}')
errors.display(e, 'video')
t1 = time.time()
shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(selected.repo)
shared.sd_model.sd_model_hash = None
sd_models.set_diffuser_options(shared.sd_model)
if selected.vae_hijack:
shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode
shared.sd_model.vae.decode = video_vae.hijack_vae_decode
if selected.te_hijack:
shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt
shared.sd_model.encode_prompt = video_utils.hijack_encode_prompt
shared.sd_model.vae.enable_slicing()
loaded_model = selected.name
msg = f'Video load: cls={shared.sd_model.__class__.__name__} model="{selected.name}" time={t1-t0:.2f}'
shared.log.info(msg)
return msg
+1 -101
View File
@@ -1,7 +1,6 @@
import os
import time
from modules import shared, errors, timer, sd_models, sd_checkpoint, model_quant, devices
from modules.video_models import models_def
from modules import shared, sd_models, timer
debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None
@@ -29,31 +28,6 @@ def set_prompt(p):
p.task_args['negative_prompt'] = p.negative_prompt
def set_vae_params(frames, tile_frames):
if tile_frames > frames:
if hasattr(shared.sd_model.vae, 'tile_sample_min_num_frames'):
shared.sd_model.vae.tile_sample_min_num_frames = tile_frames
if hasattr(shared.sd_model.vae, 'use_framewise_decoding'):
shared.sd_model.vae.use_framewise_decoding = True
if hasattr(shared.sd_model.vae, 'enable_tiling'):
shared.sd_model.vae.enable_tiling()
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()
def hijack_vae_decode(*args, **kwargs):
t0 = time.time()
shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae'])
res = shared.sd_model.vae.orig_decode(*args, **kwargs)
t1 = time.time()
timer.process.add('vae', t1-t0)
debug(f'Video decode: vae={shared.sd_model.vae.__class__.__name__} time={t1-t0:.2f}')
return res
def hijack_encode_prompt(*args, **kwargs):
t0 = time.time()
res = shared.sd_model.orig_encode_prompt(*args, **kwargs)
@@ -62,77 +36,3 @@ def hijack_encode_prompt(*args, **kwargs):
debug(f'Video encode: te={shared.sd_model.text_encoder.__class__.__name__} time={t1-t0:.2f}')
shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model)
return res
loaded_model = None
def load_model(selected: models_def.Model):
if selected is None:
return
global loaded_model # pylint: disable=global-statement
if loaded_model == selected.name:
return
sd_models.unload_model_weights()
t0 = time.time()
# text encoder
try:
quant_args = model_quant.create_config(module='Text Encoder')
debug(f'Video load: module=te repo="{selected.te or selected.repo}" folder="{selected.te_folder}" cls={selected.te_cls.__name__} quant={get_quant(quant_args)}')
text_encoder = selected.te_cls.from_pretrained(
pretrained_model_name_or_path=selected.te or selected.repo,
subfolder=selected.te_folder,
cache_dir=shared.opts.hfcache_dir,
torch_dtype=devices.dtype,
**quant_args
)
except Exception as e:
shared.log.error(f'video load: module=te cls={selected.te_cls.__name__} {e}')
errors.display(e, 'video')
text_encoder = None
# transformer
try:
quant_args = model_quant.create_config(module='Model')
debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" folder="{selected.dit_folder}" cls={selected.dit_cls.__name__} quant={get_quant(quant_args)}')
transformer = selected.dit_cls.from_pretrained(
pretrained_model_name_or_path=selected.dit or selected.repo,
subfolder=selected.dit_folder,
torch_dtype=devices.dtype,
cache_dir=shared.opts.hfcache_dir,
**quant_args
)
except Exception as e:
shared.log.error(f'video load: module=transformer cls={selected.dit_cls.__name__} {e}')
errors.display(e, 'video')
transformer = None
# model
try:
debug(f'Video load: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__}')
shared.sd_model = selected.repo_cls.from_pretrained(
pretrained_model_name_or_path=selected.repo,
transformer=transformer,
text_encoder=text_encoder,
cache_dir=shared.opts.hfcache_dir,
torch_dtype=devices.dtype,
)
except Exception as e:
shared.log.error(f'video load: module=pipe repo="{selected.repo}" cls={selected.repo_cls.__name__} {e}')
errors.display(e, 'video')
t1 = time.time()
sd_models.set_diffuser_options(shared.sd_model)
shared.sd_model.sd_checkpoint_info = sd_checkpoint.CheckpointInfo(selected.repo)
shared.sd_model.sd_model_hash = None
if selected.vae_hijack:
shared.sd_model.vae.orig_decode = shared.sd_model.vae.decode
shared.sd_model.vae.decode = hijack_vae_decode
if selected.te_hijack:
shared.sd_model.orig_encode_prompt = shared.sd_model.encode_prompt
shared.sd_model.encode_prompt = hijack_encode_prompt
shared.sd_model.vae.enable_slicing()
loaded_model = selected.name
msg = f'Video load: cls={shared.sd_model.__class__.__name__} model="{selected.name}" time={t1-t0:.2f}'
shared.log.info(msg)
return msg
+62
View File
@@ -0,0 +1,62 @@
import os
import time
from modules import shared, sd_models, devices, timer
debug = shared.log.trace if os.environ.get('SD_VIDEO_DEBUG', None) is not None else lambda *args, **kwargs: None
vae_type = None
def set_vae_params(p):
global vae_type # pylint: disable=global-statement
vae_type = p.vae_type
if p.vae_tile_frames > p.frames:
if hasattr(shared.sd_model.vae, 'tile_sample_min_num_frames'):
shared.sd_model.vae.tile_sample_min_num_frames = p.vae_tile_frames
if hasattr(shared.sd_model.vae, 'use_framewise_decoding'):
shared.sd_model.vae.use_framewise_decoding = True
if hasattr(shared.sd_model.vae, 'enable_tiling'):
shared.sd_model.vae.enable_tiling()
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()
def vae_decode_tiny(latents):
if 'Hunyuan' in shared.sd_model.__class__.__name__:
variant = 'TAE HunyuanVideo'
elif 'Mochi' in shared.sd_model.__class__.__name__:
variant = 'TAE MochiVideo'
elif 'WAN' in shared.sd_model.__class__.__name__:
variant = 'TAE WanVideo'
else:
shared.log.warning(f'Video VAE: type=Tiny cls={shared.sd_model.__class__.__name__} not supported')
return None
from modules import sd_vae_taesd
vae = sd_vae_taesd.get_model(variant)
if vae is None:
return None
debug(f'Video VAE: 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 = images.transpose(1, 2).mul_(2).sub_(1)
return (images, None)
def hijack_vae_decode(*args, **kwargs):
t0 = time.time()
res = None
if vae_type == 'Tiny':
res = vae_decode_tiny(args[0])
if vae_type == 'Remote':
pass
if res is None:
shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model, exclude=['vae'])
res = shared.sd_model.vae.orig_decode(*args, **kwargs)
t1 = time.time()
timer.process.add('vae', t1-t0)
debug(f'Video decode: type={vae_type} vae={shared.sd_model.vae.__class__.__name__} latents={args[0].shape} time={t1-t0:.2f}')
return res
+4 -4
View File
@@ -50,7 +50,7 @@ class Script(scripts.Script):
return [model, sampler, frames, guidance, offload, override, video_type, duration, loop, pad, interpolate, image, video]
def load(self, model):
if (shared.sd_model_type != 'cogvideox' or shared.sd_model.sd_model_checkpoint != model) and model != 'None':
if (shared.sd_model_type != 'cogvideo' or shared.sd_model.sd_model_checkpoint != model) and model != 'None':
sd_models.unload_model_weights('model')
shared.log.info(f'CogVideoX load: model="{model}"')
try:
@@ -64,7 +64,7 @@ class Script(scripts.Script):
shared.log.error(f'Load CogVideoX: {e}')
if debug:
errors.display(e, 'CogVideoX')
if shared.sd_model_type == 'cogvideox' and model != 'None':
if shared.sd_model_type == 'cogvideo' and model != 'None':
shared.sd_model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m', ncols=80, colour='#327fba')
shared.log.debug(f'CogVideoX load: class="{shared.sd_model.__class__.__name__}"')
if shared.sd_model is not None and model == 'None':
@@ -74,7 +74,7 @@ class Script(scripts.Script):
devices.torch_gc()
def offload(self, offload):
if shared.sd_model_type != 'cogvideox':
if shared.sd_model_type != 'cogvideo':
return
if offload == 'none':
sd_models.move_model(shared.sd_model, devices.device)
@@ -131,7 +131,7 @@ class Script(scripts.Script):
return img
def generate(self, p: processing.StableDiffusionProcessing, model: str):
if shared.sd_model_type != 'cogvideox':
if shared.sd_model_type != 'cogvideo':
return []
shared.log.info(f'CogVideoX: sampler={p.sampler} steps={p.steps} frames={p.frames} width={p.width} height={p.height} seed={p.seed} guidance={p.guidance}')
if p.sampler == 'DDIM':