VAE fixes and more options

This commit is contained in:
kijai
2024-12-09 19:20:45 +02:00
parent fb00591ced
commit a01b0796be
6 changed files with 229 additions and 171 deletions
+2 -3
View File
@@ -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
}
+61 -47
View File
@@ -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
]
}
},
+63 -49
View File
@@ -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
]
}
},
+61 -47
View File
@@ -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
]
}
},
+23 -21
View File
@@ -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 = []
+19 -4
View File
@@ -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