From 314a84dfdde7d4f23693ad0eb7d4e19ebded7392 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=88=98=E9=9B=AA=E5=B3=B0?= Date: Mon, 10 Mar 2025 14:31:20 +0800 Subject: [PATCH] sync hunyuanvideo and ltxvideo code of comfyui --- nodes/node_utils.py | 16 +- nodes/patch_lib/HunYuanVideoPatch.py | 55 ++++- nodes/patch_lib/LTXVideoPatch.py | 57 ++--- nodes/patch_lib/old/HunYuanVideoPatch.py | 292 ++++++++++++++++++++++ nodes/patch_lib/old/LTXVideoPatch.py | 299 +++++++++++++++++++++++ pyproject.toml | 2 +- requirements.txt | 3 +- 7 files changed, 667 insertions(+), 57 deletions(-) create mode 100644 nodes/patch_lib/old/HunYuanVideoPatch.py create mode 100644 nodes/patch_lib/old/LTXVideoPatch.py diff --git a/nodes/node_utils.py b/nodes/node_utils.py index 27b5953..db23690 100644 --- a/nodes/node_utils.py +++ b/nodes/node_utils.py @@ -1,6 +1,18 @@ +from packaging import version as version +import comfyui_version from .patch_lib.FluxPatch import flux_forward_orig -from .patch_lib.HunYuanVideoPatch import hunyuan_forward_orig -from .patch_lib.LTXVideoPatch import ltx_forward_orig +comfyui_ver = version.parse(comfyui_version.__version__) + +if comfyui_ver >= version.parse('0.3.25'): + from .patch_lib.HunYuanVideoPatch import hunyuan_forward_orig +else: + from .patch_lib.old.HunYuanVideoPatch import hunyuan_forward_orig + +if comfyui_ver > version.parse('0.3.19'): + # support LTXV 0.9.5 + from .patch_lib.LTXVideoPatch import ltx_forward_orig +else: + from .patch_lib.old.LTXVideoPatch import ltx_forward_orig from .patch_lib.MochiVideoPatch import mochi_forward from .patch_lib.WanVideoPatch import wan_forward_orig from .patch_util import is_hunyuan_video_model, is_ltxv_video_model, is_flux_model, is_mochi_video_model, \ diff --git a/nodes/patch_lib/HunYuanVideoPatch.py b/nodes/patch_lib/HunYuanVideoPatch.py index e4c335d..f2fd2ed 100644 --- a/nodes/patch_lib/HunYuanVideoPatch.py +++ b/nodes/patch_lib/HunYuanVideoPatch.py @@ -15,6 +15,7 @@ def hunyuan_forward_orig( timesteps: Tensor, y: Tensor, guidance: Tensor = None, + guiding_frame_index=None, control=None, transformer_options={}, ) -> Tensor: @@ -43,7 +44,17 @@ def hunyuan_forward_orig( img = self.img_in(img) vec = self.time_in(timestep_embedding(timesteps, 256, time_factor=1.0).to(img.dtype)) - vec = vec + self.vector_in(y[:, :self.params.vec_in_dim]) + if guiding_frame_index is not None: + token_replace_vec = self.time_in(timestep_embedding(guiding_frame_index, 256, time_factor=1.0)) + vec_ = self.vector_in(y[:, :self.params.vec_in_dim]) + vec = torch.cat([(vec_ + token_replace_vec).unsqueeze(1), (vec_ + vec).unsqueeze(1)], dim=1) + frame_tokens = (initial_shape[-1] // self.patch_size[-1]) * (initial_shape[-2] // self.patch_size[-2]) + modulation_dims = [(0, frame_tokens, 0), (frame_tokens, None, 1)] + modulation_dims_txt = [(0, None, 1)] + else: + vec = vec + self.vector_in(y[:, :self.params.vec_in_dim]) + modulation_dims = None + modulation_dims_txt = None if self.params.guidance_embed: if guidance is not None: @@ -72,7 +83,8 @@ def hunyuan_forward_orig( for blocks_before in patch_blocks_before: img, txt, vec, ids, pe = blocks_before(img, txt, vec, ids, pe, transformer_options) - def double_blocks_wrap(img, txt, vec, pe, control=None, attn_mask=None, transformer_options={}): + def double_blocks_wrap(img, txt, vec, pe, control=None, attn_mask=None, transformer_options={}, + modulation_dims_img=None, modulation_dims_txt=None): running_net_model = transformer_options[PatchKeys.running_net_model] patch_double_blocks_with_control_replace = patches_point.get(PatchKeys.dit_double_block_with_control_replace) for i, block in enumerate(running_net_model.double_blocks): @@ -84,7 +96,9 @@ def hunyuan_forward_orig( 'vec': vec, 'pe': pe, 'control': control, - 'attn_mask': attn_mask + 'attn_mask': attn_mask, + 'modulation_dims_img': modulation_dims_img, + 'modulation_dims_txt': modulation_dims_txt }, { "original_func": double_block_and_control_replace, @@ -99,6 +113,8 @@ def hunyuan_forward_orig( pe=pe, control=control, attn_mask=attn_mask, + modulation_dims_img=modulation_dims_img, + modulation_dims_txt=modulation_dims_txt, transformer_options=transformer_options ) @@ -114,6 +130,8 @@ def hunyuan_forward_orig( "pe": pe, "control": control, "attn_mask": attn_mask, + "modulation_dims_img": modulation_dims, + "modulation_dims_txt": modulation_dims_txt, }, { "original_blocks": double_blocks_wrap, @@ -126,6 +144,8 @@ def hunyuan_forward_orig( pe=pe, control=control, attn_mask=attn_mask, + modulation_dims_img=modulation_dims, + modulation_dims_txt=modulation_dims_txt, transformer_options=transformer_options ) @@ -155,7 +175,7 @@ def hunyuan_forward_orig( for patch_single_blocks_before in patches_single_blocks_before: img, txt = patch_single_blocks_before(img, txt, transformer_options) - def single_blocks_wrap(img, txt, vec, pe, control=None, attn_mask=None, transformer_options={}): + def single_blocks_wrap(img, txt, vec, pe, control=None, attn_mask=None, transformer_options={}, modulation_dims=None): running_net_model = transformer_options[PatchKeys.running_net_model] for i, block in enumerate(running_net_model.single_blocks): if ("single_block", i) in blocks_replace: @@ -164,20 +184,22 @@ def hunyuan_forward_orig( out["img"] = block(args["img"], vec=args["vec"], pe=args["pe"], - attn_mask=args.get("attention_mask")) + attn_mask=args.get("attention_mask"), + modulation_dims=args.get("modulation_dims")) return out out = blocks_replace[("single_block", i)]({"img": img, "vec": vec, "pe": pe, - "attention_mask": attn_mask}, + "attention_mask": attn_mask, + 'modulation_dims': modulation_dims}, { "original_block": block_wrap, "transformer_options": transformer_options }) img = out["img"] else: - img = block(img, vec=vec, pe=pe, attn_mask=attn_mask) + img = block(img, vec=vec, pe=pe, attn_mask=attn_mask, modulation_dims=modulation_dims) if control is not None: # Controlnet control_o = control.get("output") @@ -196,7 +218,8 @@ def hunyuan_forward_orig( "vec": vec, "pe": pe, "control": control, - "attn_mask": attn_mask + "attn_mask": attn_mask, + "modulation_dims": modulation_dims, }, { "original_blocks": single_blocks_wrap, @@ -209,6 +232,7 @@ def hunyuan_forward_orig( pe=pe, control=control, attn_mask=attn_mask, + modulation_dims=modulation_dims, transformer_options=transformer_options ) @@ -237,7 +261,7 @@ def hunyuan_forward_orig( for patch_final_layer_before in patches_final_layer_before: img = patch_final_layer_before(img, txt, transformer_options) - img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels) + img = self.final_layer(img, vec, modulation_dims=modulation_dims) # (N, T, patch_size ** 2 * out_channels) shape = initial_shape[-3:] for i in range(len(shape)): @@ -255,7 +279,8 @@ def hunyuan_forward_orig( return img -def double_block_and_control_replace(i, block, img, txt=None, vec=None, pe=None, control=None, attn_mask=None, transformer_options={}): + +def double_block_and_control_replace(i, block, img, txt=None, vec=None, pe=None, control=None, attn_mask=None, transformer_options={}, modulation_dims_img=None, modulation_dims_txt=None): blocks_replace = transformer_options.get("patches_replace", {}).get("dit", {}) if ("double_block", i) in blocks_replace: def block_wrap(args): @@ -264,14 +289,18 @@ def double_block_and_control_replace(i, block, img, txt=None, vec=None, pe=None, txt=args["txt"], vec=args["vec"], pe=args["pe"], - attn_mask=args.get("attention_mask")) + attn_mask=args.get("attention_mask"), + modulation_dims_img=args["modulation_dims_img"], + modulation_dims_txt=args["modulation_dims_txt"]) return out out = blocks_replace[("double_block", i)]({"img": img, "txt": txt, "vec": vec, "pe": pe, - "attention_mask": attn_mask + "attention_mask": attn_mask, + 'modulation_dims_img': modulation_dims_img, + 'modulation_dims_txt': modulation_dims_txt }, { "original_block": block_wrap, @@ -280,7 +309,7 @@ def double_block_and_control_replace(i, block, img, txt=None, vec=None, pe=None, txt = out["txt"] img = out["img"] else: - img, txt = block(img=img, txt=txt, vec=vec, pe=pe, attn_mask=attn_mask) + img, txt = block(img=img, txt=txt, vec=vec, pe=pe, attn_mask=attn_mask, modulation_dims_img=modulation_dims_img, modulation_dims_txt=modulation_dims_txt) if control is not None: # Controlnet control_i = control.get("input") if i < len(control_i): diff --git a/nodes/patch_lib/LTXVideoPatch.py b/nodes/patch_lib/LTXVideoPatch.py index 7a76fdb..60e8853 100644 --- a/nodes/patch_lib/LTXVideoPatch.py +++ b/nodes/patch_lib/LTXVideoPatch.py @@ -4,6 +4,7 @@ import torch from torch import Tensor from comfy.ldm.lightricks.model import precompute_freqs_cis +from comfy.ldm.lightricks.symmetric_patchifier import latent_to_pixel_coords from ..patch_util import PatchKeys @@ -15,8 +16,8 @@ def ltx_forward_orig( attention_mask, frame_rate=25, guiding_latent=None, - guiding_latent_noise_scale=0, transformer_options={}, + keyframe_idxs=None, **kwargs ) -> Tensor: patches_replace = transformer_options.get("patches_replace", {}) @@ -27,50 +28,31 @@ def ltx_forward_orig( patches_enter = patches_point.get(PatchKeys.dit_enter, []) if patches_enter is not None and len(patches_enter) > 0: for patch_enter in patches_enter: - x, timestep, context, attention_mask, frame_rate, guiding_latent, guiding_latent_noise_scale = patch_enter( + x, timestep, context, attention_mask, frame_rate, guiding_latent, keyframe_idxs = patch_enter( x, timestep, context, attention_mask, frame_rate, guiding_latent, - guiding_latent_noise_scale, + keyframe_idxs, transformer_options ) - indices_grid = self.patchifier.get_grid( - orig_num_frames=x.shape[2], - orig_height=x.shape[3], - orig_width=x.shape[4], - batch_size=x.shape[0], - scale_grid=((1 / frame_rate) * 8, 32, 32), - device=x.device, - ) - - if guiding_latent is not None: - ts = torch.ones([x.shape[0], 1, x.shape[2], x.shape[3], x.shape[4]], device=x.device, dtype=x.dtype) - input_ts = timestep.view([timestep.shape[0]] + [1] * (x.ndim - 1)) - ts *= input_ts - ts[:, :, 0] = guiding_latent_noise_scale * (input_ts[:, :, 0] ** 2) - timestep = self.patchifier.patchify(ts) - input_x = x.clone() - x[:, :, 0] = guiding_latent[:, :, 0] - if guiding_latent_noise_scale > 0: - if self.generator is None: - self.generator = torch.Generator(device=x.device).manual_seed(42) - elif self.generator.device != x.device: - self.generator = torch.Generator(device=x.device).set_state(self.generator.get_state()) - - noise_shape = [guiding_latent.shape[0], guiding_latent.shape[1], 1, guiding_latent.shape[3], guiding_latent.shape[4]] - scale = guiding_latent_noise_scale * (input_ts ** 2) - guiding_noise = scale * torch.randn(size=noise_shape, device=x.device, generator=self.generator) - - x[:, :, 0] = guiding_noise[:, :, 0] + x[:, :, 0] * (1.0 - scale[:, :, 0]) - - orig_shape = list(x.shape) - x = self.patchifier.patchify(x) + x, latent_coords = self.patchifier.patchify(x) + pixel_coords = latent_to_pixel_coords( + latent_coords=latent_coords, + scale_factors=self.vae_scale_factors, + causal_fix=self.causal_temporal_positioning, + ) + + if keyframe_idxs is not None: + pixel_coords[:, :, -keyframe_idxs.shape[2]:] = keyframe_idxs + + fractional_coords = pixel_coords.to(torch.float32) + fractional_coords[:, 0] = fractional_coords[:, 0] * (1.0 / frame_rate) x = self.patchify_proj(x) timestep = timestep * 1000.0 @@ -78,7 +60,7 @@ def ltx_forward_orig( if attention_mask is not None and not torch.is_floating_point(attention_mask): attention_mask = (attention_mask - 1).to(x.dtype).reshape((attention_mask.shape[0], 1, -1, attention_mask.shape[-1])) * torch.finfo(x.dtype).max - pe = precompute_freqs_cis(indices_grid, dim=self.inner_dim, out_dtype=x.dtype) + pe = precompute_freqs_cis(fractional_coords, dim=self.inner_dim, out_dtype=x.dtype) batch_size = x.shape[0] timestep, embedded_timestep = self.adaln_single( @@ -101,8 +83,6 @@ def ltx_forward_orig( batch_size, -1, x.shape[-1] ) - blocks_replace = patches_replace.get("dit", {}) - patch_blocks_before = patches_point.get(PatchKeys.dit_blocks_before, []) if patch_blocks_before is not None and len(patch_blocks_before) > 0: for blocks_before in patch_blocks_before: @@ -261,9 +241,6 @@ def ltx_forward_orig( out_channels=orig_shape[1] // math.prod(self.patchifier.patch_size), ) - if guiding_latent is not None: - x[:, :, 0] = (input_x[:, :, 0] - guiding_latent[:, :, 0]) / input_ts[:, :, 0] - patches_exit = patches_point.get(PatchKeys.dit_exit, []) if patches_exit is not None and len(patches_exit) > 0: for patch_exit in patches_exit: diff --git a/nodes/patch_lib/old/HunYuanVideoPatch.py b/nodes/patch_lib/old/HunYuanVideoPatch.py new file mode 100644 index 0000000..6eeb73c --- /dev/null +++ b/nodes/patch_lib/old/HunYuanVideoPatch.py @@ -0,0 +1,292 @@ +import torch +from torch import Tensor + +from ...patch_util import PatchKeys +from comfy.ldm.flux.layers import timestep_embedding + + +def hunyuan_forward_orig( + self, + img: Tensor, + img_ids: Tensor, + txt: Tensor, + txt_ids: Tensor, + txt_mask: Tensor, + timesteps: Tensor, + y: Tensor, + guidance: Tensor = None, + control=None, + transformer_options={}, + **kwargs +) -> Tensor: + patches_replace = transformer_options.get("patches_replace", {}) + patches_point = transformer_options.get(PatchKeys.options_key, {}) + + transformer_options[PatchKeys.running_net_model] = self + + patches_enter = patches_point.get(PatchKeys.dit_enter, []) + if patches_enter is not None and len(patches_enter) > 0: + for patch_enter in patches_enter: + img, img_ids, txt, txt_ids, timesteps, y, guidance, control, txt_mask = patch_enter(img, + img_ids, + txt, + txt_ids, + timesteps, + y, + guidance, + control, + attn_mask=txt_mask, + transformer_options=transformer_options + ) + + initial_shape = list(img.shape) + # running on sequences img + img = self.img_in(img) + vec = self.time_in(timestep_embedding(timesteps, 256, time_factor=1.0).to(img.dtype)) + + vec = vec + self.vector_in(y[:, :self.params.vec_in_dim]) + + if self.params.guidance_embed: + if guidance is not None: + vec = vec + self.guidance_in(timestep_embedding(guidance, 256).to(img.dtype)) + + if txt_mask is not None and not torch.is_floating_point(txt_mask): + txt_mask = (txt_mask - 1).to(img.dtype) * torch.finfo(img.dtype).max + + txt = self.txt_in(txt, timesteps, txt_mask) + + ids = torch.cat((img_ids, txt_ids), dim=1) + pe = self.pe_embedder(ids) + + img_len = img.shape[1] + if txt_mask is not None: + attn_mask_len = img_len + txt.shape[1] + attn_mask = torch.zeros((1, 1, attn_mask_len), dtype=img.dtype, device=img.device) + attn_mask[:, 0, img_len:] = txt_mask + else: + attn_mask = None + + blocks_replace = patches_replace.get("dit", {}) + + patch_blocks_before = patches_point.get(PatchKeys.dit_blocks_before, []) + if patch_blocks_before is not None and len(patch_blocks_before) > 0: + for blocks_before in patch_blocks_before: + img, txt, vec, ids, pe = blocks_before(img, txt, vec, ids, pe, transformer_options) + + def double_blocks_wrap(img, txt, vec, pe, control=None, attn_mask=None, transformer_options={}): + running_net_model = transformer_options[PatchKeys.running_net_model] + patch_double_blocks_with_control_replace = patches_point.get(PatchKeys.dit_double_block_with_control_replace) + for i, block in enumerate(running_net_model.double_blocks): + if patch_double_blocks_with_control_replace is not None: + img, txt = patch_double_blocks_with_control_replace({'i': i, + 'block': block, + 'img': img, + 'txt': txt, + 'vec': vec, + 'pe': pe, + 'control': control, + 'attn_mask': attn_mask + }, + { + "original_func": double_block_and_control_replace, + "transformer_options": transformer_options + }) + else: + img, txt = double_block_and_control_replace(i=i, + block=block, + img=img, + txt=txt, + vec=vec, + pe=pe, + control=control, + attn_mask=attn_mask, + transformer_options=transformer_options + ) + + del patch_double_blocks_with_control_replace + return img, txt + + patch_double_blocks_replace = patches_point.get(PatchKeys.dit_double_blocks_replace) + + if patch_double_blocks_replace is not None: + img, txt = patch_double_blocks_replace({"img": img, + "txt": txt, + "vec": vec, + "pe": pe, + "control": control, + "attn_mask": attn_mask, + }, + { + "original_blocks": double_blocks_wrap, + "transformer_options": transformer_options + }) + else: + img, txt = double_blocks_wrap(img=img, + txt=txt, + vec=vec, + pe=pe, + control=control, + attn_mask=attn_mask, + transformer_options=transformer_options + ) + + patches_double_blocks_after = patches_point.get(PatchKeys.dit_double_blocks_after, []) + if patches_double_blocks_after is not None and len(patches_double_blocks_after) > 0: + for patch_double_blocks_after in patches_double_blocks_after: + img, txt = patch_double_blocks_after(img, txt, transformer_options) + + patch_blocks_transition = patches_point.get(PatchKeys.dit_blocks_transition_replace) + + def blocks_transition_wrap(**kwargs): + txt = kwargs["txt"] + img = kwargs["img"] + return torch.cat((img, txt), 1) + + if patch_blocks_transition is not None: + img = patch_blocks_transition({"img": img, "txt": txt, "vec": vec, "pe": pe}, + { + "original_func": blocks_transition_wrap, + "transformer_options": transformer_options + }) + else: + img = blocks_transition_wrap(img=img, txt=txt) + + patches_single_blocks_before = patches_point.get(PatchKeys.dit_single_blocks_before, []) + if patches_single_blocks_before is not None and len(patches_single_blocks_before) > 0: + for patch_single_blocks_before in patches_single_blocks_before: + img, txt = patch_single_blocks_before(img, txt, transformer_options) + + def single_blocks_wrap(img, txt, vec, pe, control=None, attn_mask=None, transformer_options={}): + running_net_model = transformer_options[PatchKeys.running_net_model] + for i, block in enumerate(running_net_model.single_blocks): + if ("single_block", i) in blocks_replace: + def block_wrap(args): + out = {} + out["img"] = block(args["img"], + vec=args["vec"], + pe=args["pe"], + attn_mask=args.get("attention_mask")) + return out + + out = blocks_replace[("single_block", i)]({"img": img, + "vec": vec, + "pe": pe, + "attention_mask": attn_mask}, + { + "original_block": block_wrap, + "transformer_options": transformer_options + }) + img = out["img"] + else: + img = block(img, vec=vec, pe=pe, attn_mask=attn_mask) + + if control is not None: # Controlnet + control_o = control.get("output") + if i < len(control_o): + add = control_o[i] + if add is not None: + img[:, : img_len] += add + + return img + + patch_single_blocks_replace = patches_point.get(PatchKeys.dit_single_blocks_replace) + + if patch_single_blocks_replace is not None: + img, txt = patch_single_blocks_replace({"img": img, + "txt": txt, + "vec": vec, + "pe": pe, + "control": control, + "attn_mask": attn_mask + }, + { + "original_blocks": single_blocks_wrap, + "transformer_options": transformer_options + }) + else: + img = single_blocks_wrap(img=img, + txt=txt, + vec=vec, + pe=pe, + control=control, + attn_mask=attn_mask, + transformer_options=transformer_options + ) + + patch_blocks_exit = patches_point.get(PatchKeys.dit_blocks_after, []) + if patch_blocks_exit is not None and len(patch_blocks_exit) > 0: + for blocks_after in patch_blocks_exit: + img, txt = blocks_after(img, txt, transformer_options) + + def final_transition_wrap(**kwargs): + img = kwargs["img"] + img_len = kwargs["img_len"] + return img[:, : img_len] + + patch_blocks_after_transition_replace = patches_point.get(PatchKeys.dit_blocks_after_transition_replace) + if patch_blocks_after_transition_replace is not None: + img = patch_blocks_after_transition_replace({"img": img, "txt": txt, "vec": vec, "pe": pe, "img_len": img_len}, + { + "original_func": final_transition_wrap, + "transformer_options": transformer_options + }) + else: + img = final_transition_wrap(img=img, img_len=img_len) + + patches_final_layer_before = patches_point.get(PatchKeys.dit_final_layer_before, []) + if patches_final_layer_before is not None and len(patches_final_layer_before) > 0: + for patch_final_layer_before in patches_final_layer_before: + img = patch_final_layer_before(img, txt, transformer_options) + + img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels) + + shape = initial_shape[-3:] + for i in range(len(shape)): + shape[i] = shape[i] // self.patch_size[i] + img = img.reshape([img.shape[0]] + shape + [self.out_channels] + self.patch_size) + img = img.permute(0, 4, 1, 5, 2, 6, 3, 7) + img = img.reshape(initial_shape[0], self.out_channels, initial_shape[2], initial_shape[3], initial_shape[4]) + + patches_exit = patches_point.get(PatchKeys.dit_exit, []) + if patches_exit is not None and len(patches_exit) > 0: + for patch_exit in patches_exit: + img = patch_exit(img, transformer_options) + + del transformer_options[PatchKeys.running_net_model] + + return img + +def double_block_and_control_replace(i, block, img, txt=None, vec=None, pe=None, control=None, attn_mask=None, transformer_options={}): + blocks_replace = transformer_options.get("patches_replace", {}).get("dit", {}) + if ("double_block", i) in blocks_replace: + def block_wrap(args): + out = {} + out["img"], out["txt"] = block(img=args["img"], + txt=args["txt"], + vec=args["vec"], + pe=args["pe"], + attn_mask=args.get("attention_mask")) + return out + + out = blocks_replace[("double_block", i)]({"img": img, + "txt": txt, + "vec": vec, + "pe": pe, + "attention_mask": attn_mask + }, + { + "original_block": block_wrap, + "transformer_options": transformer_options + }) + txt = out["txt"] + img = out["img"] + else: + img, txt = block(img=img, txt=txt, vec=vec, pe=pe, attn_mask=attn_mask) + if control is not None: # Controlnet + control_i = control.get("input") + if i < len(control_i): + add = control_i[i] + if add is not None: + img += add + + return img, txt diff --git a/nodes/patch_lib/old/LTXVideoPatch.py b/nodes/patch_lib/old/LTXVideoPatch.py new file mode 100644 index 0000000..f47b4b0 --- /dev/null +++ b/nodes/patch_lib/old/LTXVideoPatch.py @@ -0,0 +1,299 @@ +import math + +import torch +from torch import Tensor + +from comfy.ldm.lightricks.model import precompute_freqs_cis +from ...patch_util import PatchKeys + +# changed in comfyui hash commit 93fedd92fe0eb67a09e29069b05adebb40678639 (between comfyui version 0.3.19 and 1.3.20) +def ltx_forward_orig( + self, + x, + timestep, + context, + attention_mask, + frame_rate=25, + guiding_latent=None, + guiding_latent_noise_scale=0, + transformer_options={}, + **kwargs +) -> Tensor: + patches_point = transformer_options.get(PatchKeys.options_key, {}) + + transformer_options[PatchKeys.running_net_model] = self + + patches_enter = patches_point.get(PatchKeys.dit_enter, []) + if patches_enter is not None and len(patches_enter) > 0: + for patch_enter in patches_enter: + x, timestep, context, attention_mask, frame_rate, guiding_latent, guiding_latent_noise_scale = patch_enter( + x, + timestep, + context, + attention_mask, + frame_rate, + guiding_latent, + guiding_latent_noise_scale, + transformer_options + ) + + indices_grid = self.patchifier.get_grid( + orig_num_frames=x.shape[2], + orig_height=x.shape[3], + orig_width=x.shape[4], + batch_size=x.shape[0], + scale_grid=((1 / frame_rate) * 8, 32, 32), + device=x.device, + ) + + if guiding_latent is not None: + ts = torch.ones([x.shape[0], 1, x.shape[2], x.shape[3], x.shape[4]], device=x.device, dtype=x.dtype) + input_ts = timestep.view([timestep.shape[0]] + [1] * (x.ndim - 1)) + ts *= input_ts + ts[:, :, 0] = guiding_latent_noise_scale * (input_ts[:, :, 0] ** 2) + timestep = self.patchifier.patchify(ts) + input_x = x.clone() + x[:, :, 0] = guiding_latent[:, :, 0] + if guiding_latent_noise_scale > 0: + if self.generator is None: + self.generator = torch.Generator(device=x.device).manual_seed(42) + elif self.generator.device != x.device: + self.generator = torch.Generator(device=x.device).set_state(self.generator.get_state()) + + noise_shape = [guiding_latent.shape[0], guiding_latent.shape[1], 1, guiding_latent.shape[3], guiding_latent.shape[4]] + scale = guiding_latent_noise_scale * (input_ts ** 2) + guiding_noise = scale * torch.randn(size=noise_shape, device=x.device, generator=self.generator) + + x[:, :, 0] = guiding_noise[:, :, 0] + x[:, :, 0] * (1.0 - scale[:, :, 0]) + + + orig_shape = list(x.shape) + + x = self.patchifier.patchify(x) + + x = self.patchify_proj(x) + timestep = timestep * 1000.0 + + if attention_mask is not None and not torch.is_floating_point(attention_mask): + attention_mask = (attention_mask - 1).to(x.dtype).reshape((attention_mask.shape[0], 1, -1, attention_mask.shape[-1])) * torch.finfo(x.dtype).max + + pe = precompute_freqs_cis(indices_grid, dim=self.inner_dim, out_dtype=x.dtype) + + batch_size = x.shape[0] + timestep, embedded_timestep = self.adaln_single( + timestep.flatten(), + {"resolution": None, "aspect_ratio": None}, + batch_size=batch_size, + hidden_dtype=x.dtype, + ) + # Second dimension is 1 or number of tokens (if timestep_per_token) + timestep = timestep.view(batch_size, -1, timestep.shape[-1]) + embedded_timestep = embedded_timestep.view( + batch_size, -1, embedded_timestep.shape[-1] + ) + + # 2. Blocks + if self.caption_projection is not None: + batch_size = x.shape[0] + context = self.caption_projection(context) + context = context.view( + batch_size, -1, x.shape[-1] + ) + + patch_blocks_before = patches_point.get(PatchKeys.dit_blocks_before, []) + if patch_blocks_before is not None and len(patch_blocks_before) > 0: + for blocks_before in patch_blocks_before: + x, context, timestep, ids, pe = blocks_before(img=x, txt=context, vec=timestep, ids=None, pe=pe, transformer_options=transformer_options) + + def double_blocks_wrap(img, txt, vec, pe, control=None, attn_mask=None, transformer_options={}): + running_net_model = transformer_options[PatchKeys.running_net_model] + patch_double_blocks_with_control_replace = patches_point.get(PatchKeys.dit_double_block_with_control_replace) + for i, block in enumerate(running_net_model.transformer_blocks): + if patch_double_blocks_with_control_replace is not None: + img, txt = patch_double_blocks_with_control_replace({'i': i, + 'block': block, + 'img': img, + 'txt': txt, + 'vec': vec, + 'pe': pe, + 'control': control, + 'attn_mask': attn_mask + }, + { + "original_func": double_block_and_control_replace, + "transformer_options": transformer_options + }) + else: + img, txt = double_block_and_control_replace(i=i, + block=block, + img=img, + txt=txt, + vec=vec, + pe=pe, + control=control, + attn_mask=attn_mask, + transformer_options=transformer_options + ) + + del patch_double_blocks_with_control_replace + return img, txt + + patch_double_blocks_replace = patches_point.get(PatchKeys.dit_double_blocks_replace) + + if patch_double_blocks_replace is not None: + x, context = patch_double_blocks_replace({"img": x, + "txt": context, + "vec": timestep, + "pe": pe, + "control": None, + "attn_mask": attention_mask, + }, + { + "original_blocks": double_blocks_wrap, + "transformer_options": transformer_options + }) + else: + x, context = double_blocks_wrap(img=x, + txt=context, + vec=timestep, + pe=pe, + control=None, + attn_mask=attention_mask, + transformer_options=transformer_options + ) + + patches_double_blocks_after = patches_point.get(PatchKeys.dit_double_blocks_after, []) + if patches_double_blocks_after is not None and len(patches_double_blocks_after) > 0: + for patch_double_blocks_after in patches_double_blocks_after: + x, context = patch_double_blocks_after(x, context, transformer_options) + + patch_blocks_transition = patches_point.get(PatchKeys.dit_blocks_transition_replace) + + def blocks_transition_wrap(**kwargs): + x = kwargs["img"] + return x + + if patch_blocks_transition is not None: + x = patch_blocks_transition({"img": x, "txt": context, "vec": timestep, "pe": pe}, + { + "original_func": blocks_transition_wrap, + "transformer_options": transformer_options + }) + else: + x = blocks_transition_wrap(img=x, txt=context) + + patches_single_blocks_before = patches_point.get(PatchKeys.dit_single_blocks_before, []) + if patches_single_blocks_before is not None and len(patches_single_blocks_before) > 0: + for patch_single_blocks_before in patches_single_blocks_before: + x, context = patch_single_blocks_before(x, context, transformer_options) + + def single_blocks_wrap(img, **kwargs): + return img + + patch_single_blocks_replace = patches_point.get(PatchKeys.dit_single_blocks_replace) + + if patch_single_blocks_replace is not None: + x, context = patch_single_blocks_replace({"img": x, + "txt": context, + "vec": timestep, + "pe": pe, + "control": None, + "attn_mask": attention_mask + }, + { + "original_blocks": single_blocks_wrap, + "transformer_options": transformer_options + }) + else: + x = single_blocks_wrap(img=x, + txt=context, + vec=timestep, + pe=pe, + control=None, + attn_mask=attention_mask, + transformer_options=transformer_options + ) + + patch_blocks_exit = patches_point.get(PatchKeys.dit_blocks_after, []) + if patch_blocks_exit is not None and len(patch_blocks_exit) > 0: + for blocks_after in patch_blocks_exit: + x, context = blocks_after(x, context, transformer_options) + + # 3. Output + def final_transition_wrap(**kwargs): + running_net_model = transformer_options[PatchKeys.running_net_model] + x = kwargs["img"] + embedded_timestep = kwargs["embedded_timestep"] + scale_shift_values = ( + running_net_model.scale_shift_table[None, None].to(device=x.device, dtype=x.dtype) + embedded_timestep[:, :, None] + ) + shift, scale = scale_shift_values[:, :, 0], scale_shift_values[:, :, 1] + x = running_net_model.norm_out(x) + # Modulation + x = x * (1 + scale) + shift + return x + + patch_blocks_after_transition_replace = patches_point.get(PatchKeys.dit_blocks_after_transition_replace) + if patch_blocks_after_transition_replace is not None: + x = patch_blocks_after_transition_replace({"img": x, "txt": context, "vec": timestep, "pe": pe, "embedded_timestep": embedded_timestep}, + { + "original_func": final_transition_wrap, + "transformer_options": transformer_options + }) + else: + x = final_transition_wrap(img=x, embedded_timestep=embedded_timestep) + + patches_final_layer_before = patches_point.get(PatchKeys.dit_final_layer_before, []) + if patches_final_layer_before is not None and len(patches_final_layer_before) > 0: + for patch_final_layer_before in patches_final_layer_before: + x = patch_final_layer_before(img=x, txt=context, transformer_options=transformer_options) + + x = self.proj_out(x) + + x = self.patchifier.unpatchify( + latents=x, + output_height=orig_shape[3], + output_width=orig_shape[4], + output_num_frames=orig_shape[2], + out_channels=orig_shape[1] // math.prod(self.patchifier.patch_size), + ) + + if guiding_latent is not None: + x[:, :, 0] = (input_x[:, :, 0] - guiding_latent[:, :, 0]) / input_ts[:, :, 0] + + patches_exit = patches_point.get(PatchKeys.dit_exit, []) + if patches_exit is not None and len(patches_exit) > 0: + for patch_exit in patches_exit: + x = patch_exit(x, transformer_options) + + del transformer_options[PatchKeys.running_net_model] + + return x + +def double_block_and_control_replace(i, block, img, txt=None, vec=None, pe=None, control=None, attn_mask=None, transformer_options={}): + blocks_replace = transformer_options.get("patches_replace", {}).get("dit", {}) + if ("double_block", i) in blocks_replace: + def block_wrap(args): + out = {} + out["img"] = block(x=args["img"], + context=args["txt"], + timestep=args["vec"], + pe=args["pe"], + attention_mask=args.get("attention_mask")) + return out + + out = blocks_replace[("double_block", i)]({"img": img, + "txt": txt, + "vec": vec, + "pe": pe, + "attention_mask": attn_mask, + }, + { + "original_block": block_wrap, + "transformer_options": transformer_options + }) + img = out["img"] + else: + img = block(x=img, context=txt, timestep=vec, pe=pe, attention_mask=attn_mask) + + return img, txt diff --git a/pyproject.toml b/pyproject.toml index 47f8fd9..2843069 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui_patches_ll" description = "Some patches for Flux|HunYuanVideo|LTXVideo|MochiVideo|WanVideo etc, support TeaCache, PuLID, First Block Cache." -version = "1.1.0" +version = "1.1.1" license = {file = "LICENSE"} dependencies = [] diff --git a/requirements.txt b/requirements.txt index 296d654..0a1814e 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1 +1,2 @@ -numpy \ No newline at end of file +numpy +packaging \ No newline at end of file