Initial support for the fixed model (might break other things for now)
This commit is contained in:
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
|
||||
@@ -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={
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user