Initial support for the fixed model (might break other things for now)

This commit is contained in:
kijai
2025-03-07 21:41:38 +02:00
parent 5291557fd6
commit ab1fcc6b31
5 changed files with 224 additions and 81 deletions
@@ -235,8 +235,10 @@ class HunyuanVideoPipeline(DiffusionPipeline):
freenoise=False,
context_size=None,
context_overlap=None,
official_i2v=False,
i2v_condition_type=None,
image_cond_latents=None,
i2v_stability=True,
):
shape = (
batch_size,
@@ -288,19 +290,18 @@ class HunyuanVideoPipeline(DiffusionPipeline):
noise[:, :, place_idx:place_idx + delta, :, :] = noise[:, :, list_idx, :, :]
i2v_mask = None
if official_i2v:
# Create mask
i2v_mask = torch.zeros(shape[0], 1, shape[2], shape[3], shape[4], device=device)
i2v_mask[:, :, 0, ...] = 1.0
if image_cond_latents is not None:
if image_cond_latents.shape[2] == 1:
padding = torch.zeros(shape, device=device)
padding[:, :, 0:1, :, :] = image_cond_latents
image_cond_latents = padding
if i2v_condition_type == "latent_concat":
# Create mask
i2v_mask = torch.zeros(shape[0], 1, shape[2], shape[3], shape[4], device=device)
i2v_mask[:, :, 0, ...] = 1.0
if image_cond_latents.shape[2] == 1:
padding = torch.zeros(shape, device=device)
padding[:, :, 0:1, :, :] = image_cond_latents
image_cond_latents = padding
if denoise_strength < 1.0:
if official_i2v:
if i2v_condition_type == "latent_concat":
latents = torch.cat((latents[:,:,0].unsqueeze(2), latents), dim=2)
timesteps, num_inference_steps = self.get_timesteps(num_inference_steps, denoise_strength, device)
latent_timestep = timesteps[:1]
@@ -316,9 +317,11 @@ class HunyuanVideoPipeline(DiffusionPipeline):
latents = latents[:, :frames_needed, :, :, :]
latents = latents * (1 - latent_timestep / 1000) + latent_timestep / 1000 * noise
print("latents shape:", latents.shape)
elif official_i2v:
elif i2v_stability:
if image_cond_latents.shape[2] == 1:
img_latents = image_cond_latents.repeat(1, 1, video_length, 1, 1)
t = torch.tensor([0.999]).to(device=device)
latents = noise * t + image_cond_latents * (1 - t)
latents = noise * t + img_latents * (1 - t)
latents = latents.to(dtype=self.base_dtype)
else:
latents = noise
@@ -442,6 +445,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
leapfusion_img2vid: Optional[bool] = False,
image_cond_latents: Optional[torch.Tensor] = None,
riflex_freq_index: Optional[int] = None,
i2v_stability=True,
**kwargs,
):
r"""
@@ -584,11 +588,10 @@ class HunyuanVideoPipeline(DiffusionPipeline):
latent_video_length = (video_length - 1) // 4 + 1
official_i2v = False
if self.transformer.in_channels == 33:
official_i2v = True
latent_video_length += 1
original_image_latents = image_cond_latents
i2v_condition_type = self.transformer.i2v_condition_type
#if i2v_condition_type == "latent_concat":
#latent_video_length += 1
if feta_args is not None:
set_enhance_weight(feta_args["weight"])
feta_start_percent = feta_args["start_percent"]
@@ -650,7 +653,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
freenoise=freenoise,
context_size=context_frames,
context_overlap=context_overlap,
official_i2v=official_i2v,
i2v_condition_type=i2v_condition_type,
i2v_stability=i2v_stability,
image_cond_latents=image_cond_latents,
)
@@ -728,13 +732,16 @@ class HunyuanVideoPipeline(DiffusionPipeline):
latent_model_input[:, :, [0,], :, :] = original_latents[:, :, [0,], :, :].to(latent_model_input)
if image_cond_latents is not None and not use_context_schedule:
latent_image_input = (
torch.cat([image_cond_latents] * 2) if cfg_enabled else image_cond_latents
)
if i2v_mask is not None:
if i2v_condition_type == "latent_concat":
latent_image_input = (torch.cat([image_cond_latents] * 2) if cfg_enabled else image_cond_latents)
i2v_mask = torch.cat([i2v_mask] * 2) if cfg_enabled else i2v_mask
latent_image_input = torch.cat([latent_image_input, i2v_mask], dim=1)
latent_model_input = torch.cat([latent_model_input, latent_image_input], dim=1)
latent_model_input = torch.cat([latent_model_input, latent_image_input], dim=1)
elif i2v_condition_type == "token_replace":
latent_image_input = (torch.cat([original_image_latents] * 2) if cfg_enabled else original_image_latents)
latent_model_input = torch.cat([latent_image_input, latent_model_input[:, :, 1:, :, :]], dim=2)
else:
latent_model_input = torch.cat([latent_model_input, latent_image_input], dim=1)
if self.transformer.guidance_embed:
if cfg_enabled:
@@ -884,9 +891,17 @@ class HunyuanVideoPipeline(DiffusionPipeline):
)
# compute the previous noisy sample x_t -> x_t-1
latents = self.scheduler.step(
noise_pred, t, latents, **extra_step_kwargs, return_dict=False
)[0]
if i2v_condition_type == "token_replace":
latents = self.scheduler.step(
noise_pred[:, :, 1:, :, :], t, latents[:, :, 1:, :, :], **extra_step_kwargs, return_dict=False
)[0]
latents = torch.concat(
[original_image_latents, latents], dim=2
)
else:
latents = self.scheduler.step(
noise_pred, t, latents, **extra_step_kwargs, return_dict=False
)[0]
if callback_on_step_end is not None:
callback_kwargs = {}
@@ -920,6 +935,6 @@ class HunyuanVideoPipeline(DiffusionPipeline):
else:
comfy_pbar.update(1)
if leapfusion_img2vid or official_i2v:
if leapfusion_img2vid or i2v_condition_type == "latent_concat":
latents = latents[:, :, 1:, :, :]
return latents
+107 -28
View File
@@ -197,15 +197,36 @@ class MMDoubleStreamBlock(nn.Module):
freqs_cis: tuple = None,
attn_mask: Optional[torch.Tensor] = None,
upcast_rope: bool = True,
token_replace_vec: torch.Tensor = None,
first_frame_token_num: int = None,
condition_type: str = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
(
img_mod1_shift,
img_mod1_scale,
img_mod1_gate,
img_mod2_shift,
img_mod2_scale,
img_mod2_gate,
) = self.img_mod(vec).chunk(6, dim=-1)
if condition_type == "token_replace":
img_mod1, token_replace_img_mod1 = self.img_mod(vec, condition_type=condition_type, \
token_replace_vec=token_replace_vec)
(img_mod1_shift,
img_mod1_scale,
img_mod1_gate,
img_mod2_shift,
img_mod2_scale,
img_mod2_gate) = img_mod1.chunk(6, dim=-1)
(tr_img_mod1_shift,
tr_img_mod1_scale,
tr_img_mod1_gate,
tr_img_mod2_shift,
tr_img_mod2_scale,
tr_img_mod2_gate) = token_replace_img_mod1.chunk(6, dim=-1)
else:
(
img_mod1_shift,
img_mod1_scale,
img_mod1_gate,
img_mod2_shift,
img_mod2_scale,
img_mod2_gate,
) = self.img_mod(vec).chunk(6, dim=-1)
(
txt_mod1_shift,
txt_mod1_scale,
@@ -217,9 +238,16 @@ class MMDoubleStreamBlock(nn.Module):
# Prepare image for attention.
img_modulated = self.img_norm1(img)
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale
)
if condition_type == "token_replace":
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale, condition_type=condition_type,
tr_shift=tr_img_mod1_shift, tr_scale=tr_img_mod1_scale,
first_frame_token_num=first_frame_token_num
)
else:
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale
)
img_qkv = self.img_attn_qkv(img_modulated)
img_q, img_k, img_v = rearrange(
img_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
@@ -273,15 +301,29 @@ class MMDoubleStreamBlock(nn.Module):
img_attn *= feta_scores
# Calculate the img bloks.
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale
)
),
gate=img_mod2_gate,
)
if condition_type == "token_replace":
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate, condition_type=condition_type,
tr_gate=tr_img_mod1_gate, first_frame_token_num=first_frame_token_num)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale, condition_type=condition_type,
tr_shift=tr_img_mod2_shift, tr_scale=tr_img_mod2_scale, first_frame_token_num=first_frame_token_num
)
),
gate=img_mod2_gate, condition_type=condition_type,
tr_gate=tr_img_mod2_gate, first_frame_token_num=first_frame_token_num
)
else:
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale
)
),
gate=img_mod2_gate,
)
# Calculate the txt bloks.
txt = txt + apply_gate(self.txt_attn_proj(txt_attn), gate=txt_mod1_gate)
@@ -382,10 +424,29 @@ class MMSingleStreamBlock(nn.Module):
freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None,
attn_mask: Optional[torch.Tensor] = None,
upcast_rope: bool = True,
token_replace_vec: torch.Tensor = None,
first_frame_token_num: int = None,
condition_type: str = None,
stg_mode: Optional[str] = None,
) -> torch.Tensor:
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale)
if condition_type == "token_replace":
mod, tr_mod = self.modulation(vec,
condition_type=condition_type,
token_replace_vec=token_replace_vec)
(mod_shift,
mod_scale,
mod_gate) = mod.chunk(3, dim=-1)
(tr_mod_shift,
tr_mod_scale,
tr_mod_gate) = tr_mod.chunk(3, dim=-1)
else:
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
if condition_type == "token_replace":
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale, condition_type=condition_type,
tr_shift=tr_mod_shift, tr_scale=tr_mod_scale, first_frame_token_num=first_frame_token_num)
else:
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale)
qkv, mlp = torch.split(
self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1
)
@@ -473,10 +534,12 @@ class MMSingleStreamBlock(nn.Module):
# Compute activation in mlp stream, cat again and run second linear layer.
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
output = x + apply_gate(output, gate=mod_gate)
return output
if condition_type == "token_replace":
output = x + apply_gate(output, gate=mod_gate, condition_type=condition_type,
tr_gate=tr_mod_gate, first_frame_token_num=first_frame_token_num)
return output
else:
return x + apply_gate(output, gate=mod_gate)
class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
@@ -552,6 +615,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
use_attention_mask: bool = True,
text_states_dim: int = 4096,
text_states_dim_2: int = 768,
i2v_condition_type: str = "latent_concat",
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
main_device: Optional[torch.device] = None,
@@ -571,6 +635,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
self.main_device = main_device
self.offload_device = offload_device
self.attention_mode = attention_mode
self.i2v_condition_type = i2v_condition_type
# Text projection. Default to linear projection.
# Alternative: TokenRefiner. See more details (LI-DiT): http://arxiv.org/abs/2406.11831
@@ -935,9 +1000,23 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
# Prepare modulation vectors.
vec = self.time_in(t)
if self.i2v_condition_type == "token_replace":
token_replace_t = torch.zeros_like(t)
token_replace_vec = self.time_in(token_replace_t)
first_frame_token_num = th * tw
else:
token_replace_vec = None
first_frame_token_num = None
# token_replace_mask_img = None
# token_replace_mask_txt = None
# text modulation
if text_states_2 is not None:
vec = vec + self.vector_in(text_states_2)
vec_2 = self.vector_in(text_states_2)
vec = vec + vec_2
if self.i2v_condition_type == "token_replace":
token_replace_vec = token_replace_vec + vec_2
# guidance modulation
if guidance is not None:
@@ -987,7 +1066,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
block_args = [cu_seqlens_q, cu_seqlens_kv, max_seqlen_q, max_seqlen_kv, freqs_cis, attn_mask, self.upcast_rope]
block_args = [cu_seqlens_q, cu_seqlens_kv, max_seqlen_q, max_seqlen_kv, freqs_cis, attn_mask, self.upcast_rope, token_replace_vec, first_frame_token_num, self.i2v_condition_type]
#tea_cache
if self.enable_teacache:
+44 -17
View File
@@ -3,7 +3,6 @@ from typing import Callable
import torch
import torch.nn as nn
class ModulateDiT(nn.Module):
"""Modulation layer for DiT."""
def __init__(
@@ -24,11 +23,19 @@ class ModulateDiT(nn.Module):
nn.init.zeros_(self.linear.weight)
nn.init.zeros_(self.linear.bias)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.linear(self.act(x))
def forward(self, x: torch.Tensor, condition_type=None, token_replace_vec=None) -> torch.Tensor:
x_out = self.linear(self.act(x))
def modulate(x, shift=None, scale=None):
if condition_type == "token_replace":
x_token_replace_out = self.linear(self.act(token_replace_vec))
return x_out, x_token_replace_out
else:
return x_out
def modulate(x, shift=None, scale=None, condition_type=None,
tr_shift=None, tr_scale=None,
first_frame_token_num=None):
"""modulate by shift and scale
Args:
@@ -39,17 +46,23 @@ def modulate(x, shift=None, scale=None):
Returns:
torch.Tensor: the output tensor after modulate.
"""
if scale is None and shift is None:
if condition_type == "token_replace":
x_zero = x[:, :first_frame_token_num] * (1 + tr_scale.unsqueeze(1)) + tr_shift.unsqueeze(1)
x_orig = x[:, first_frame_token_num:] * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
x = torch.concat((x_zero, x_orig), dim=1)
return x
elif shift is None:
return x * (1 + scale.unsqueeze(1))
elif scale is None:
return x + shift.unsqueeze(1)
else:
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
if scale is None and shift is None:
return x
elif shift is None:
return x * (1 + scale.unsqueeze(1))
elif scale is None:
return x + shift.unsqueeze(1)
else:
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
def apply_gate(x, gate=None, tanh=False):
def apply_gate(x, gate=None, tanh=False, condition_type=None, tr_gate=None, first_frame_token_num=None):
"""AI is creating summary for apply_gate
Args:
@@ -60,12 +73,26 @@ def apply_gate(x, gate=None, tanh=False):
Returns:
torch.Tensor: the output tensor after apply gate.
"""
if gate is None:
return x
if tanh:
return x * gate.unsqueeze(1).tanh()
if condition_type == "token_replace":
if gate is None:
return x
if tanh:
x_zero = x[:, :first_frame_token_num] * tr_gate.unsqueeze(1).tanh()
x_orig = x[:, first_frame_token_num:] * gate.unsqueeze(1).tanh()
x = torch.concat((x_zero, x_orig), dim=1)
return x
else:
x_zero = x[:, :first_frame_token_num] * tr_gate.unsqueeze(1)
x_orig = x[:, first_frame_token_num:] * gate.unsqueeze(1)
x = torch.concat((x_zero, x_orig), dim=1)
return x
else:
return x * gate.unsqueeze(1)
if gate is None:
return x
if tanh:
return x * gate.unsqueeze(1).tanh()
else:
return x * gate.unsqueeze(1)
def ckpt_wrapper(module):
@@ -73,4 +100,4 @@ def ckpt_wrapper(module):
outputs = module(*inputs)
return outputs
return ckpt_forward
return ckpt_forward
+4
View File
@@ -455,6 +455,8 @@ class TextEncoder(nn.Module):
image_last_hidden_state = torch.stack(image_last_hidden_state)
image_attention_mask = torch.stack(image_attention_mask)
print("image_embed_interleave", image_embed_interleave)
if semantic_images is not None and 0 < image_embed_interleave < 6:
image_last_hidden_state = image_last_hidden_state[:, ::image_embed_interleave, :]
image_attention_mask = image_attention_mask[:, ::image_embed_interleave]
@@ -477,6 +479,7 @@ class TextEncoder(nn.Module):
do_sample=False,
hidden_state_skip_layer=None,
return_texts=False,
image_embed_interleave=2,
):
batch_encoding = self.text2tokens(text)
return self.encode(
@@ -486,6 +489,7 @@ class TextEncoder(nn.Module):
do_sample=do_sample,
hidden_state_skip_layer=hidden_state_skip_layer,
return_texts=return_texts,
image_embed_interleave=image_embed_interleave
)
xtuner_config={
+26 -8
View File
@@ -316,6 +316,12 @@ class HyVideoModelLoader:
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
in_channels = sd["img_in.proj.weight"].shape[1]
print("In channels: ", in_channels)
if in_channels == 16:
i2v_condition_type = "token_replace"
elif in_channels == 33:
i2v_condition_type = "latent_concat"
guidance_embed = sd.get("guidance_in.mlp.0.weight", False) is not False
out_channels = 16
@@ -328,6 +334,7 @@ class HyVideoModelLoader:
"heads_num": 24,
"mlp_width_ratio": 4,
"guidance_embed": guidance_embed,
"i2v_condition_type": i2v_condition_type,
}
with init_empty_weights():
transformer = HYVideoDiffusionTransformer(
@@ -350,7 +357,7 @@ class HyVideoModelLoader:
)
scheduler_config = {
"flow_shift": 9.0,
"flow_shift": 7.0,
"reverse": True,
"solver": "euler",
"use_flow_sigmas": True,
@@ -791,7 +798,8 @@ class HyVideoTextEncode:
FUNCTION = "process"
CATEGORY = "HunyuanVideoWrapper"
def process(self, text_encoders, prompt, force_offload=True, prompt_template="video", custom_prompt_template=None, clip_l=None, image_token_selection_expr="::4", hyvid_cfg=None, image=None, image1=None, image2=None, clip_text_override=None):
def process(self, text_encoders, prompt, force_offload=True, prompt_template="video", custom_prompt_template=None, clip_l=None, image_token_selection_expr="::4",
hyvid_cfg=None, image=None, image1=None, image2=None, clip_text_override=None, image_embed_interleave=2):
if clip_text_override is not None and len(clip_text_override) == 0:
clip_text_override = None
device = mm.text_encoder_device()
@@ -834,7 +842,7 @@ class HyVideoTextEncode:
else:
prompt_template_dict = None
def encode_prompt(self, prompt, negative_prompt, text_encoder, image_token_selection_expr="::4", image1=None, image2=None, clip_text_override=None):
def encode_prompt(self, prompt, negative_prompt, text_encoder, image_token_selection_expr="::4", semantic_images=None, image1=None, image2=None, clip_text_override=None, image_embed_interleave=2):
batch_size = 1
num_videos_per_prompt = 1
@@ -847,9 +855,10 @@ class HyVideoTextEncode:
prompt_outputs = text_encoder.encode(text_inputs,
prompt_template=prompt_template_dict,
image_token_selection_expr=image_token_selection_expr,
semantic_images = [image.squeeze(0) * 255] if text_encoder.text_encoder_type == "vlm" else None,
semantic_images = [semantic_images.squeeze(0) * 255] if text_encoder.text_encoder_type == "vlm" else None,
image_embed_interleave=image_embed_interleave,
device=device,
data_type=prompt_template
data_type=prompt_template,
)
else:
text_inputs = text_encoder.text2tokens(prompt,
@@ -935,7 +944,9 @@ class HyVideoTextEncode:
text_encoder_1,
image_token_selection_expr=image_token_selection_expr,
image1=image1,
image2=image2)
image2=image2,
semantic_images=image,
image_embed_interleave=image_embed_interleave,)
if force_offload:
text_encoder_1.to(offload_device)
mm.soft_empty_cache()
@@ -1024,6 +1035,7 @@ class HyVideoI2VEncode(HyVideoTextEncode):
"clip_l": ("CLIP", {"tooltip": "Use comfy clip model instead, in this case the text encoder loader's clip_l should be disabled"}),
"image": ("IMAGE", {"default": None}),
"hyvid_cfg": ("HYVID_CFG", ),
"image_embed_interleave": ("INT", {"default": 2}),
}
}
@@ -1199,6 +1211,7 @@ class HyVideoSampler:
"default": 'FlowMatchDiscreteScheduler'
}),
"riflex_freq_index": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1, "tooltip": "Frequency index for RIFLEX, disabled when 0, default 4. Allows for new frames to be generated after 129 without looping"}),
"i2v_mode": (["stability", "dynamic"], {"default": "disabled", "tooltip": "I2V mode for image2video process"}),
}
}
@@ -1209,7 +1222,7 @@ class HyVideoSampler:
def process(self, model, hyvid_embeds, flow_shift, steps, embedded_guidance_scale, seed, width, height, num_frames,
samples=None, denoise_strength=1.0, force_offload=True, stg_args=None, context_options=None, feta_args=None,
teacache_args=None, scheduler=None, image_cond_latents=None, riflex_freq_index=0):
teacache_args=None, scheduler=None, image_cond_latents=None, riflex_freq_index=0, i2v_mode="stability"):
model = model.model
device = mm.get_torch_device()
@@ -1238,6 +1251,10 @@ class HyVideoSampler:
if embedded_guidance_scale == 0.0:
embedded_guidance_scale = None
i2v_stability = False
if i2v_mode == "stability":
i2v_stability = True
generator = torch.Generator(device=torch.device("cpu")).manual_seed(seed)
if width <= 0 or height <= 0 or num_frames <= 0:
@@ -1351,7 +1368,8 @@ class HyVideoSampler:
feta_args=feta_args,
leapfusion_img2vid = leapfusion_img2vid,
image_cond_latents = image_cond_latents["samples"] * VAE_SCALING_FACTOR if image_cond_latents is not None else None,
riflex_freq_index = riflex_freq_index
riflex_freq_index = riflex_freq_index,
i2v_stability = i2v_stability,
)
print_memory(device)