From a01b0796be154ee224caa15d0709d1d5f6a51fe8 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 9 Dec 2024 19:20:45 +0200 Subject: [PATCH] VAE fixes and more options --- configs/hy_vae_config.json | 5 +- examples/hyvideo_lowvram_blockswap_test.json | 108 ++++++++++-------- examples/hyvideo_t2v_example_01.json | 112 +++++++++++-------- examples/hyvideo_v2v_example_01.json | 108 ++++++++++-------- hyvideo/vae/autoencoder_kl_causal_3d.py | 44 ++++---- nodes.py | 23 +++- 6 files changed, 229 insertions(+), 171 deletions(-) diff --git a/configs/hy_vae_config.json b/configs/hy_vae_config.json index 73d4cc0..b9b595d 100644 --- a/configs/hy_vae_config.json +++ b/configs/hy_vae_config.json @@ -19,7 +19,7 @@ "layers_per_block": 2, "norm_num_groups": 32, "out_channels": 3, - "sample_size": 256, + "tile_sample_min_size": 256, "sample_tsize": 64, "up_block_types": [ "UpDecoderBlockCausal3D", @@ -29,6 +29,5 @@ ], "scaling_factor": 0.476986, "time_compression_ratio": 4, - "mid_block_add_attention": true, - "mid_block_causal_attn": true + "mid_block_add_attention": true } diff --git a/examples/hyvideo_lowvram_blockswap_test.json b/examples/hyvideo_lowvram_blockswap_test.json index 3bafb06..47e5fa9 100644 --- a/examples/hyvideo_lowvram_blockswap_test.json +++ b/examples/hyvideo_lowvram_blockswap_test.json @@ -195,6 +195,18 @@ "name": "text_encoders", "type": "HYVIDTEXTENCODER", "link": 35 + }, + { + "name": "custom_prompt_template", + "type": "PROMPT_TEMPLATE", + "link": null, + "shape": 7 + }, + { + "name": "clip_l", + "type": "CLIP", + "link": null, + "shape": 7 } ], "outputs": [ @@ -394,50 +406,6 @@ "color": "#432", "bgcolor": "#653" }, - { - "id": 5, - "type": "HyVideoDecode", - "pos": [ - 920, - -279 - ], - "size": [ - 345.4285888671875, - 102 - ], - "flags": {}, - "order": 11, - "mode": 0, - "inputs": [ - { - "name": "vae", - "type": "VAE", - "link": 6 - }, - { - "name": "samples", - "type": "LATENT", - "link": 4 - } - ], - "outputs": [ - { - "name": "images", - "type": "IMAGE", - "links": [ - 42 - ], - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "HyVideoDecode" - }, - "widgets_values": [ - true, - 4 - ] - }, { "id": 34, "type": "VHS_VideoCombine", @@ -510,6 +478,52 @@ "muted": false } } + }, + { + "id": 5, + "type": "HyVideoDecode", + "pos": [ + 920, + -279 + ], + "size": [ + 345.4285888671875, + 150 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "vae", + "type": "VAE", + "link": 6 + }, + { + "name": "samples", + "type": "LATENT", + "link": 4 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 42 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "HyVideoDecode" + }, + "widgets_values": [ + true, + 64, + 128, + false + ] } ], "links": [ @@ -574,10 +588,10 @@ "config": {}, "extra": { "ds": { - "scale": 0.7513148009015777, + "scale": 1, "offset": [ - 856.2226738131614, - 635.642485451872 + 499.8829252621208, + 382.71199061362694 ] } }, diff --git a/examples/hyvideo_t2v_example_01.json b/examples/hyvideo_t2v_example_01.json index 456812c..b30ac56 100644 --- a/examples/hyvideo_t2v_example_01.json +++ b/examples/hyvideo_t2v_example_01.json @@ -191,50 +191,6 @@ "disabled" ] }, - { - "id": 5, - "type": "HyVideoDecode", - "pos": [ - 651, - -285 - ], - "size": [ - 345.4285888671875, - 102 - ], - "flags": {}, - "order": 5, - "mode": 0, - "inputs": [ - { - "name": "vae", - "type": "VAE", - "link": 6 - }, - { - "name": "samples", - "type": "LATENT", - "link": 4 - } - ], - "outputs": [ - { - "name": "images", - "type": "IMAGE", - "links": [ - 42 - ], - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "HyVideoDecode" - }, - "widgets_values": [ - true, - 8 - ] - }, { "id": 30, "type": "HyVideoTextEncode", @@ -254,6 +210,18 @@ "name": "text_encoders", "type": "HYVIDTEXTENCODER", "link": 35 + }, + { + "name": "custom_prompt_template", + "type": "PROMPT_TEMPLATE", + "link": null, + "shape": 7 + }, + { + "name": "clip_l", + "type": "CLIP", + "link": null, + "shape": 7 } ], "outputs": [ @@ -278,8 +246,8 @@ "id": 34, "type": "VHS_VideoCombine", "pos": [ - 657, - -112 + 673.133544921875, + -37.19999694824219 ], "size": [ 580.7774658203125, @@ -346,6 +314,52 @@ "muted": false } } + }, + { + "id": 5, + "type": "HyVideoDecode", + "pos": [ + 651, + -285 + ], + "size": [ + 345.4285888671875, + 150 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "vae", + "type": "VAE", + "link": 6 + }, + { + "name": "samples", + "type": "LATENT", + "link": 4 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 42 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "HyVideoDecode" + }, + "widgets_values": [ + true, + 64, + 256, + true + ] } ], "links": [ @@ -402,10 +416,10 @@ "config": {}, "extra": { "ds": { - "scale": 0.8264462809917354, + "scale": 0.9090909090909091, "offset": [ - 1003.830865103555, - 591.1597820691316 + 740.8495512386833, + 610.811990613627 ] } }, diff --git a/examples/hyvideo_v2v_example_01.json b/examples/hyvideo_v2v_example_01.json index 6d9effc..ffa1746 100644 --- a/examples/hyvideo_v2v_example_01.json +++ b/examples/hyvideo_v2v_example_01.json @@ -268,6 +268,18 @@ "name": "text_encoders", "type": "HYVIDTEXTENCODER", "link": 35 + }, + { + "name": "custom_prompt_template", + "type": "PROMPT_TEMPLATE", + "link": null, + "shape": 7 + }, + { + "name": "clip_l", + "type": "CLIP", + "link": null, + "shape": 7 } ], "outputs": [ @@ -402,50 +414,6 @@ "bf16" ] }, - { - "id": 5, - "type": "HyVideoDecode", - "pos": [ - 712, - -282 - ], - "size": [ - 345.4285888671875, - 102 - ], - "flags": {}, - "order": 9, - "mode": 0, - "inputs": [ - { - "name": "vae", - "type": "VAE", - "link": 6 - }, - { - "name": "samples", - "type": "LATENT", - "link": 4 - } - ], - "outputs": [ - { - "name": "images", - "type": "IMAGE", - "links": [ - 59 - ], - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "HyVideoDecode" - }, - "widgets_values": [ - true, - 8 - ] - }, { "id": 44, "type": "ImageConcatMulti", @@ -681,6 +649,52 @@ "muted": false } } + }, + { + "id": 5, + "type": "HyVideoDecode", + "pos": [ + 683.6054077148438, + -282.8873291015625 + ], + "size": [ + 345.4285888671875, + 150 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "vae", + "type": "VAE", + "link": 6 + }, + { + "name": "samples", + "type": "LATENT", + "link": 4 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 59 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "HyVideoDecode" + }, + "widgets_values": [ + true, + 64, + 256, + true + ] } ], "links": [ @@ -817,10 +831,10 @@ "config": {}, "extra": { "ds": { - "scale": 0.620921323059155, + "scale": 0.7513148009015777, "offset": [ - 1327.1114051035552, - 721.9111720691317 + 1018.584830413488, + 698.393990613627 ] } }, diff --git a/hyvideo/vae/autoencoder_kl_causal_3d.py b/hyvideo/vae/autoencoder_kl_causal_3d.py index 0261568..1793eb0 100644 --- a/hyvideo/vae/autoencoder_kl_causal_3d.py +++ b/hyvideo/vae/autoencoder_kl_causal_3d.py @@ -41,7 +41,8 @@ from diffusers.models.attention_processor import ( from diffusers.models.modeling_outputs import AutoencoderKLOutput from diffusers.models.modeling_utils import ModelMixin from .vae import DecoderCausal3D, BaseOutput, DecoderOutput, DiagonalGaussianDistribution, EncoderCausal3D - +from tqdm import tqdm +from comfy.utils import ProgressBar @dataclass class DecoderOutput2(BaseOutput): @@ -71,8 +72,9 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin): act_fn: str = "silu", latent_channels: int = 4, norm_num_groups: int = 32, - sample_size: int = 32, + tile_sample_min_size: int = 256, sample_tsize: int = 64, + overlap_factor: float = 0.25, scaling_factor: float = 0.18215, force_upcast: float = True, spatial_compression_ratio: int = 8, @@ -125,15 +127,13 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin): self.tile_sample_min_tsize = self.sample_tsize self.tile_latent_min_tsize = self.sample_tsize // time_compression_ratio - self.tile_sample_min_size = self.config.sample_size - sample_size = ( - self.config.sample_size[0] - if isinstance(self.config.sample_size, (list, tuple)) - else self.config.sample_size - ) + self.tile_sample_min_size = tile_sample_min_size + self.tile_latent_min_size = int( - sample_size / (2 ** (len(self.config.block_out_channels) - 1))) - self.tile_overlap_factor = 0.25 + self.tile_sample_min_size / (2 ** (len(self.config.block_out_channels) - 1))) + + self.tile_overlap_factor = overlap_factor + self.t_tile_overlap_factor = overlap_factor def _set_gradient_checkpointing(self, module, value=False): if isinstance(module, (EncoderCausal3D, DecoderCausal3D)): @@ -336,7 +336,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin): If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is returned. - """ + """ if self.use_slicing and z.shape[0] > 1: decoded_slices = [self._decode( z_slice).sample for z_slice in z.split(1)] @@ -449,24 +449,26 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin): self.tile_overlap_factor) row_limit = self.tile_sample_min_size - blend_extent - # Split z into overlapping tiles and decode them separately. - # The tiles have an overlap to avoid seams between tiles. + total_rows = (z.shape[-2] + overlap_size - 1) // overlap_size + comfy_pbar = ProgressBar(total_rows) + + # Split z into overlapping tiles with progress bar rows = [] - for i in range(0, z.shape[-2], overlap_size): + for i in tqdm(range(0, z.shape[-2], overlap_size), desc="Decoding rows", total=total_rows): row = [] for j in range(0, z.shape[-1], overlap_size): - tile = z[:, :, :, i: i + self.tile_latent_min_size, - j: j + self.tile_latent_min_size] + tile = z[:, :, :, i:i + self.tile_latent_min_size, j:j + self.tile_latent_min_size] tile = self.post_quant_conv(tile) decoded = self.decoder(tile) row.append(decoded) rows.append(row) + comfy_pbar.update(1) + + # Process results with progress bar result_rows = [] - for i, row in enumerate(rows): + for i, row in tqdm(enumerate(rows), desc="Blending tiles", total=len(rows)): result_row = [] for j, tile in enumerate(row): - # blend the above tile and the left tile - # to the current tile and add the current tile to the result row if i > 0: tile = self.blend_v(rows[i - 1][j], tile, blend_extent) if j > 0: @@ -522,9 +524,9 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin): B, C, T, H, W = z.shape overlap_size = int(self.tile_latent_min_tsize * - (1 - self.tile_overlap_factor)) + (1 - self.t_tile_overlap_factor)) blend_extent = int(self.tile_sample_min_tsize * - self.tile_overlap_factor) + self.t_tile_overlap_factor) t_limit = self.tile_sample_min_tsize - blend_extent row = [] diff --git a/nodes.py b/nodes.py index a633f1e..698c95a 100644 --- a/nodes.py +++ b/nodes.py @@ -290,7 +290,8 @@ class HyVideoModelLoader: pipe = HunyuanVideoPipeline( transformer=transformer, scheduler=scheduler, - progress_bar_config=None + progress_bar_config=None, + base_dtype=base_dtype ) pipeline = { @@ -864,7 +865,9 @@ class HyVideoDecode: "vae": ("VAE",), "samples": ("LATENT",), "enable_vae_tiling": ("BOOLEAN", {"default": True, "tooltip": "Drastically reduces memory use but may introduce seams"}), - "temporal_tiling_sample_size": ("INT", {"default": 16, "min": 4, "max": 256, "tooltip": "Smaller values use less VRAM, model default is 64 which doesn't fit on most GPUs"}), + "temporal_tiling_sample_size": ("INT", {"default": 16, "min": 4, "max": 256, "tooltip": "Smaller values use less VRAM, model default is 64, any other value will cause stutter"}), + "spatial_tile_sample_min_size": ("INT", {"default": 256, "min": 32, "max": 2048, "step": 32, "tooltip": "Spatial tile minimum size in pixels, smaller values use less VRAM, may introduce more seams"}), + "auto_tile_size": ("BOOLEAN", {"default": True, "tooltip": "Automatically set tile size based on defaults, above settings are ignored"}), }, } @@ -873,14 +876,26 @@ class HyVideoDecode: FUNCTION = "decode" CATEGORY = "HunyuanVideoWrapper" - def decode(self, vae, samples, enable_vae_tiling, temporal_tiling_sample_size): + def decode(self, vae, samples, enable_vae_tiling, temporal_tiling_sample_size, spatial_tile_sample_min_size, auto_tile_size): device = mm.get_torch_device() offload_device = mm.unet_offload_device() mm.soft_empty_cache() latents = samples["samples"] generator = torch.Generator(device=torch.device("cpu"))#.manual_seed(seed) vae.to(device) - vae.sample_tsize = temporal_tiling_sample_size + if not auto_tile_size: + vae.tile_latent_min_tsize = temporal_tiling_sample_size // 4 + vae.tile_sample_min_size = spatial_tile_sample_min_size + vae.tile_latent_min_size = spatial_tile_sample_min_size // 8 + if temporal_tiling_sample_size != 64: + vae.t_tile_overlap_factor = 0.0 + else: + vae.t_tile_overlap_factor = 0.25 + else: + #defaults + vae.tile_latent_min_tsize = 64 + vae.tile_sample_min_size = 256 + vae.tile_latent_min_size = 32 expand_temporal_dim = False