sync hunyuanvideo and ltxvideo code of comfyui
This commit is contained in:
+14
-2
@@ -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, \
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
+1
-1
@@ -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 = []
|
||||
|
||||
|
||||
+2
-1
@@ -1 +1,2 @@
|
||||
numpy
|
||||
numpy
|
||||
packaging
|
||||
Reference in New Issue
Block a user