diff --git a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py index 4c68f8c..b334991 100644 --- a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py +++ b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py @@ -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 \ No newline at end of file diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index f7daffe..1f16ca3 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -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: diff --git a/hyvideo/modules/modulate_layers.py b/hyvideo/modules/modulate_layers.py index 93a57c6..a8d70db 100644 --- a/hyvideo/modules/modulate_layers.py +++ b/hyvideo/modules/modulate_layers.py @@ -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 \ No newline at end of file diff --git a/hyvideo/text_encoder/__init__.py b/hyvideo/text_encoder/__init__.py index e72dc87..000585e 100644 --- a/hyvideo/text_encoder/__init__.py +++ b/hyvideo/text_encoder/__init__.py @@ -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={ diff --git a/nodes.py b/nodes.py index d39637e..3972264 100644 --- a/nodes.py +++ b/nodes.py @@ -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)