Support official I2V
This commit is contained in:
@@ -14,6 +14,12 @@ __all__ = [
|
|||||||
"TEXT_PROJECTION",
|
"TEXT_PROJECTION",
|
||||||
"DATA_TYPE",
|
"DATA_TYPE",
|
||||||
"NEGATIVE_PROMPT",
|
"NEGATIVE_PROMPT",
|
||||||
|
"NEGATIVE_PROMPT_I2V",
|
||||||
|
"FLOW_PATH_TYPE",
|
||||||
|
"FLOW_PREDICT_TYPE",
|
||||||
|
"FLOW_LOSS_WEIGHT",
|
||||||
|
"FLOW_SNR_TYPE",
|
||||||
|
"FLOW_SOLVER",
|
||||||
]
|
]
|
||||||
|
|
||||||
PRECISION_TO_TYPE = {
|
PRECISION_TO_TYPE = {
|
||||||
@@ -46,7 +52,26 @@ PROMPT_TEMPLATE_ENCODE_VIDEO = (
|
|||||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
|
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
PROMPT_TEMPLATE_ENCODE_I2V = (
|
||||||
|
"<|start_header_id|>system<|end_header_id|>\n\n<image>\nDescribe the image by detailing the color, shape, size, texture, "
|
||||||
|
"quantity, text, spatial relationships of the objects and background:<|eot_id|>"
|
||||||
|
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
|
||||||
|
"<|start_header_id|>assistant<|end_header_id|>\n\n"
|
||||||
|
)
|
||||||
|
|
||||||
|
PROMPT_TEMPLATE_ENCODE_VIDEO_I2V = (
|
||||||
|
"<|start_header_id|>system<|end_header_id|>\n\n<image>\nDescribe the video by detailing the following aspects according to the reference image: "
|
||||||
|
"1. The main content and theme of the video."
|
||||||
|
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
|
||||||
|
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
|
||||||
|
"4. background environment, light, style and atmosphere."
|
||||||
|
"5. camera angles, movements, and transitions used in the video:<|eot_id|>\n\n"
|
||||||
|
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
|
||||||
|
"<|start_header_id|>assistant<|end_header_id|>\n\n"
|
||||||
|
)
|
||||||
|
|
||||||
NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
|
NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
|
||||||
|
NEGATIVE_PROMPT_I2V = "deformation, a poor composition and deformed video, bad teeth, bad eyes, bad limbs"
|
||||||
|
|
||||||
PROMPT_TEMPLATE = {
|
PROMPT_TEMPLATE = {
|
||||||
"dit-llm-encode": {
|
"dit-llm-encode": {
|
||||||
@@ -57,6 +82,22 @@ PROMPT_TEMPLATE = {
|
|||||||
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
|
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
|
||||||
"crop_start": 95,
|
"crop_start": 95,
|
||||||
},
|
},
|
||||||
|
"dit-llm-encode-i2v": {
|
||||||
|
"template": PROMPT_TEMPLATE_ENCODE_I2V,
|
||||||
|
"crop_start": 36,
|
||||||
|
"image_emb_start": 5,
|
||||||
|
"image_emb_end": 581,
|
||||||
|
"image_emb_len": 576,
|
||||||
|
"double_return_token_id": 271
|
||||||
|
},
|
||||||
|
"dit-llm-encode-video-i2v": {
|
||||||
|
"template": PROMPT_TEMPLATE_ENCODE_VIDEO_I2V,
|
||||||
|
"crop_start": 103,
|
||||||
|
"image_emb_start": 5,
|
||||||
|
"image_emb_end": 581,
|
||||||
|
"image_emb_len": 576,
|
||||||
|
"double_return_token_id": 271
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
# ======================= Model ======================
|
# ======================= Model ======================
|
||||||
@@ -77,15 +118,48 @@ VAE_PATH = {"884-16c-hy": f"{MODEL_BASE}/hunyuan-video-t2v-720p/vae"}
|
|||||||
TEXT_ENCODER_PATH = {
|
TEXT_ENCODER_PATH = {
|
||||||
"clipL": f"{MODEL_BASE}/text_encoder_2",
|
"clipL": f"{MODEL_BASE}/text_encoder_2",
|
||||||
"llm": f"{MODEL_BASE}/text_encoder",
|
"llm": f"{MODEL_BASE}/text_encoder",
|
||||||
|
"llm-i2v": f"{MODEL_BASE}/text_encoder_i2v",
|
||||||
}
|
}
|
||||||
|
|
||||||
# Tokenizer
|
# Tokenizer
|
||||||
TOKENIZER_PATH = {
|
TOKENIZER_PATH = {
|
||||||
"clipL": f"{MODEL_BASE}/text_encoder_2",
|
"clipL": f"{MODEL_BASE}/text_encoder_2",
|
||||||
"llm": f"{MODEL_BASE}/text_encoder",
|
"llm": f"{MODEL_BASE}/text_encoder",
|
||||||
|
"llm-i2v": f"{MODEL_BASE}/text_encoder_i2v",
|
||||||
}
|
}
|
||||||
|
|
||||||
TEXT_PROJECTION = {
|
TEXT_PROJECTION = {
|
||||||
"linear", # Default, an nn.Linear() layer
|
"linear", # Default, an nn.Linear() layer
|
||||||
"single_refiner", # Single TokenRefiner. Refer to LI-DiT
|
"single_refiner", # Single TokenRefiner. Refer to LI-DiT
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# Flow Matching path type
|
||||||
|
FLOW_PATH_TYPE = {
|
||||||
|
"linear", # Linear trajectory between noise and data
|
||||||
|
"gvp", # Generalized variance-preserving SDE
|
||||||
|
"vp", # Variance-preserving SDE
|
||||||
|
}
|
||||||
|
|
||||||
|
# Flow Matching predict type
|
||||||
|
FLOW_PREDICT_TYPE = {
|
||||||
|
"velocity", # Predict velocity
|
||||||
|
"score", # Predict score
|
||||||
|
"noise", # Predict noise
|
||||||
|
}
|
||||||
|
|
||||||
|
# Flow Matching loss weight
|
||||||
|
FLOW_LOSS_WEIGHT = {
|
||||||
|
"velocity", # Weight loss by velocity
|
||||||
|
"likelihood", # Weight loss by likelihood
|
||||||
|
}
|
||||||
|
|
||||||
|
# Flow Matching SNR type
|
||||||
|
FLOW_SNR_TYPE = {
|
||||||
|
"lognorm", # Log-normal SNR
|
||||||
|
"uniform", # Uniform SNR
|
||||||
|
}
|
||||||
|
|
||||||
|
# Flow Matching solvers
|
||||||
|
FLOW_SOLVER = {
|
||||||
|
"euler", # Euler solver
|
||||||
|
}
|
||||||
@@ -40,7 +40,7 @@ EXAMPLE_DOC_STRING = """"""
|
|||||||
from ...modules.posemb_layers import get_nd_rotary_pos_embed
|
from ...modules.posemb_layers import get_nd_rotary_pos_embed
|
||||||
from ....enhance_a_video.globals import enable_enhance, disable_enhance, set_enhance_weight
|
from ....enhance_a_video.globals import enable_enhance, disable_enhance, set_enhance_weight
|
||||||
|
|
||||||
def get_rotary_pos_embed(transformer, latent_video_length, height, width):
|
def get_rotary_pos_embed(transformer, latent_video_length, height, width, k=0):
|
||||||
target_ndim = 3
|
target_ndim = 3
|
||||||
ndim = 5 - 2
|
ndim = 5 - 2
|
||||||
rope_theta = 225
|
rope_theta = 225
|
||||||
@@ -85,6 +85,8 @@ def get_rotary_pos_embed(transformer, latent_video_length, height, width):
|
|||||||
theta=rope_theta,
|
theta=rope_theta,
|
||||||
use_real=True,
|
use_real=True,
|
||||||
theta_rescale_factor=1,
|
theta_rescale_factor=1,
|
||||||
|
num_frames=latent_video_length,
|
||||||
|
k=k,
|
||||||
)
|
)
|
||||||
return freqs_cos, freqs_sin
|
return freqs_cos, freqs_sin
|
||||||
def retrieve_timesteps(
|
def retrieve_timesteps(
|
||||||
@@ -233,8 +235,12 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
freenoise=False,
|
freenoise=False,
|
||||||
context_size=None,
|
context_size=None,
|
||||||
context_overlap=None,
|
context_overlap=None,
|
||||||
leapfusion_img2vid=False
|
leapfusion_img2vid=False,
|
||||||
|
i2v_mask=None,
|
||||||
|
image_cond_latents=None,
|
||||||
):
|
):
|
||||||
|
#if i2v_mask is not None:
|
||||||
|
# num_channels_latents = (num_channels_latents - 1) // 2
|
||||||
shape = (
|
shape = (
|
||||||
batch_size,
|
batch_size,
|
||||||
num_channels_latents,
|
num_channels_latents,
|
||||||
@@ -283,11 +289,16 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
#print("place_idx:", place_idx, "delta:", delta, "list_idx:", list_idx)
|
#print("place_idx:", place_idx, "delta:", delta, "list_idx:", list_idx)
|
||||||
noise[:, :, place_idx:place_idx + delta, :, :] = noise[:, :, list_idx, :, :]
|
noise[:, :, place_idx:place_idx + delta, :, :] = noise[:, :, list_idx, :, :]
|
||||||
|
|
||||||
if latents is None:
|
if i2v_mask is not None:
|
||||||
latents = noise
|
print("i2v_mask shape:", i2v_mask.shape)
|
||||||
elif leapfusion_img2vid:
|
if image_cond_latents.shape[2] == 1:
|
||||||
noise[:, :, [0,], :, :] = latents[:, :, [0,], :, :].to(noise)
|
image_cond_latents = image_cond_latents.repeat(1, 1, video_length, 1, 1)
|
||||||
latents = noise.to(device)
|
t = torch.tensor([0.999]).to(device=device)
|
||||||
|
latents = noise * t + image_cond_latents * (1 - t)
|
||||||
|
latents = latents.to(dtype=self.base_dtype)
|
||||||
|
elif latents is None:
|
||||||
|
print("No latents provided, generating noise and using it as latents")
|
||||||
|
latents = noise
|
||||||
elif denoise_strength < 1.0:
|
elif denoise_strength < 1.0:
|
||||||
latents = latents.to(device)
|
latents = latents.to(device)
|
||||||
timesteps, num_inference_steps = self.get_timesteps(num_inference_steps, denoise_strength, device)
|
timesteps, num_inference_steps = self.get_timesteps(num_inference_steps, denoise_strength, device)
|
||||||
@@ -424,6 +435,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
feta_args: Optional[Dict] = None,
|
feta_args: Optional[Dict] = None,
|
||||||
leapfusion_img2vid: Optional[bool] = False,
|
leapfusion_img2vid: Optional[bool] = False,
|
||||||
image_cond_latents: Optional[torch.Tensor] = None,
|
image_cond_latents: Optional[torch.Tensor] = None,
|
||||||
|
riflex_freq_index: Optional[int] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
r"""
|
r"""
|
||||||
@@ -574,16 +586,25 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
else:
|
else:
|
||||||
disable_enhance()
|
disable_enhance()
|
||||||
|
|
||||||
|
i2v_mask = None
|
||||||
|
image_latents = None
|
||||||
if image_cond_latents is not None:
|
if image_cond_latents is not None:
|
||||||
padding_shape = (
|
# Expand to video length and zero-pad remaining frames
|
||||||
batch_size,
|
image_latents = torch.zeros(
|
||||||
16,
|
(batch_size, 16, latent_video_length, height//8, width//8),
|
||||||
latent_video_length - 1,
|
device=device,
|
||||||
int(height) // 8,
|
dtype=self.base_dtype
|
||||||
int(width) // 8,
|
|
||||||
)
|
)
|
||||||
latent_padding = torch.zeros(padding_shape, device=device, dtype=self.base_dtype)
|
image_latents[:, :, 0:1, ...] = image_cond_latents
|
||||||
image_latents = torch.cat([image_cond_latents, latent_padding], dim=2)
|
|
||||||
|
# Create mask
|
||||||
|
i2v_mask = torch.zeros(
|
||||||
|
batch_size, 1, latent_video_length, height//8, width//8,
|
||||||
|
device=device
|
||||||
|
)
|
||||||
|
i2v_mask[:, :, 0, ...] = 1.0
|
||||||
|
print("i2v_mask shape:", i2v_mask.shape)
|
||||||
|
|
||||||
print("image_cond_latents shape:", image_cond_latents.shape)
|
print("image_cond_latents shape:", image_cond_latents.shape)
|
||||||
print("image_latents shape:", image_latents.shape)
|
print("image_latents shape:", image_latents.shape)
|
||||||
|
|
||||||
@@ -611,7 +632,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
else:
|
else:
|
||||||
# rotary embeddings
|
# rotary embeddings
|
||||||
freqs_cos, freqs_sin = get_rotary_pos_embed(
|
freqs_cos, freqs_sin = get_rotary_pos_embed(
|
||||||
self.transformer, latent_video_length, height, width
|
self.transformer, latent_video_length, height, width, k=riflex_freq_index
|
||||||
)
|
)
|
||||||
if not self.transformer.upcast_rope:
|
if not self.transformer.upcast_rope:
|
||||||
freqs_cos = freqs_cos.to(self.base_dtype).to(device)
|
freqs_cos = freqs_cos.to(self.base_dtype).to(device)
|
||||||
@@ -641,7 +662,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
freenoise=freenoise,
|
freenoise=freenoise,
|
||||||
context_size=context_frames,
|
context_size=context_frames,
|
||||||
context_overlap=context_overlap,
|
context_overlap=context_overlap,
|
||||||
leapfusion_img2vid=leapfusion_img2vid
|
leapfusion_img2vid=leapfusion_img2vid,
|
||||||
|
i2v_mask=i2v_mask,
|
||||||
|
image_cond_latents=image_latents,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
|
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
|
||||||
@@ -721,18 +744,24 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
latent_image_input = (
|
latent_image_input = (
|
||||||
torch.cat([image_latents] * 2) if cfg_enabled else image_latents
|
torch.cat([image_latents] * 2) if cfg_enabled else image_latents
|
||||||
)
|
)
|
||||||
|
if i2v_mask is not None:
|
||||||
|
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)
|
||||||
|
|
||||||
if cfg_enabled:
|
if self.transformer.guidance_embed:
|
||||||
guidance_expand = (
|
if cfg_enabled:
|
||||||
torch.tensor([embedded_guidance_scale] * latents.shape[0] * 2, dtype=self.base_dtype, device=device)
|
guidance_expand = (
|
||||||
* 1000.0
|
torch.tensor([embedded_guidance_scale] * latents.shape[0] * 2, dtype=self.base_dtype, device=device)
|
||||||
)
|
* 1000.0
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
guidance_expand = (
|
||||||
|
torch.tensor([embedded_guidance_scale] * latents.shape[0], dtype=self.base_dtype, device=device)
|
||||||
|
* 1000.0
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
guidance_expand = (
|
guidance_expand = None
|
||||||
torch.tensor([embedded_guidance_scale] * latents.shape[0], dtype=self.base_dtype, device=device)
|
|
||||||
* 1000.0
|
|
||||||
)
|
|
||||||
|
|
||||||
if use_context_schedule:
|
if use_context_schedule:
|
||||||
counter = torch.zeros_like(latent_model_input)
|
counter = torch.zeros_like(latent_model_input)
|
||||||
@@ -745,7 +774,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
#print("partial_latent_model_input", partial_latent_model_input.shape)
|
#print("partial_latent_model_input", partial_latent_model_input.shape)
|
||||||
with torch.autocast(
|
with torch.autocast(
|
||||||
device_type="cuda", dtype=self.base_dtype, enabled=True):
|
device_type="cuda", dtype=self.base_dtype, enabled=True):
|
||||||
noise_pred[:, :, c, :, :] += self.transformer(
|
noise_pred_context = self.transformer(
|
||||||
partial_latent_model_input,
|
partial_latent_model_input,
|
||||||
t_expand,
|
t_expand,
|
||||||
text_states=input_prompt_embeds,
|
text_states=input_prompt_embeds,
|
||||||
@@ -758,8 +787,19 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
stg_mode=stg_mode,
|
stg_mode=stg_mode,
|
||||||
return_dict=True,
|
return_dict=True,
|
||||||
)["x"]
|
)["x"]
|
||||||
|
window_mask = torch.ones_like(noise_pred_context)
|
||||||
counter[:, :, c, :, :] += 1
|
# Apply left-side blending for all except first chunk
|
||||||
|
if min(c) > 0:
|
||||||
|
ramp_up = torch.linspace(0, 1, context_overlap, device=noise_pred_context.device)
|
||||||
|
ramp_up = ramp_up.view(1, 1, -1, 1, 1)
|
||||||
|
window_mask[:, :, :context_overlap] = ramp_up
|
||||||
|
# Apply right-side blending for all except last chunk
|
||||||
|
if max(c) < latent_video_length - 1:
|
||||||
|
ramp_down = torch.linspace(1, 0, context_overlap, device=noise_pred_context.device)
|
||||||
|
ramp_down = ramp_down.view(1, 1, -1, 1, 1)
|
||||||
|
window_mask[:, :, -context_overlap:] = ramp_down
|
||||||
|
noise_pred[:, :, c, :, :] += noise_pred_context * window_mask
|
||||||
|
counter[:, :, c, :, :] += window_mask
|
||||||
noise_pred = noise_pred.float()
|
noise_pred = noise_pred.float()
|
||||||
noise_pred /= counter
|
noise_pred /= counter
|
||||||
else:
|
else:
|
||||||
@@ -873,4 +913,6 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
|
|
||||||
if leapfusion_img2vid:
|
if leapfusion_img2vid:
|
||||||
latents = latents[:, :, 1:, :, :]
|
latents = latents[:, :, 1:, :, :]
|
||||||
|
if i2v_mask is not None:
|
||||||
|
latents = latents[:, :, 4:, :, :]
|
||||||
return latents
|
return latents
|
||||||
@@ -688,6 +688,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
|||||||
self.enable_teacache = False
|
self.enable_teacache = False
|
||||||
self.cnt = 0
|
self.cnt = 0
|
||||||
self.num_steps = 0
|
self.num_steps = 0
|
||||||
|
self.teacache_skipped_steps = 0
|
||||||
self.rel_l1_thresh = 0.15
|
self.rel_l1_thresh = 0.15
|
||||||
self.accumulated_rel_l1_distance = 0
|
self.accumulated_rel_l1_distance = 0
|
||||||
self.previous_modulated_input = None
|
self.previous_modulated_input = None
|
||||||
@@ -1026,6 +1027,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
|||||||
self.cnt = 0
|
self.cnt = 0
|
||||||
|
|
||||||
if not should_calc and self.previous_residual is not None:
|
if not should_calc and self.previous_residual is not None:
|
||||||
|
self.teacache_skipped_steps += 1
|
||||||
# Verify tensor dimensions match before adding
|
# Verify tensor dimensions match before adding
|
||||||
if img.shape == self.previous_residual.shape:
|
if img.shape == self.previous_residual.shape:
|
||||||
img = img + self.previous_residual
|
img = img + self.previous_residual
|
||||||
|
|||||||
@@ -113,6 +113,8 @@ def get_nd_rotary_pos_embed(
|
|||||||
use_real=False,
|
use_real=False,
|
||||||
theta_rescale_factor: Union[float, List[float]] = 1.0,
|
theta_rescale_factor: Union[float, List[float]] = 1.0,
|
||||||
interpolation_factor: Union[float, List[float]] = 1.0,
|
interpolation_factor: Union[float, List[float]] = 1.0,
|
||||||
|
num_frames: int = 129,
|
||||||
|
k: int = 0,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
|
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
|
||||||
@@ -163,6 +165,8 @@ def get_nd_rotary_pos_embed(
|
|||||||
use_real=use_real,
|
use_real=use_real,
|
||||||
theta_rescale_factor=theta_rescale_factor[i],
|
theta_rescale_factor=theta_rescale_factor[i],
|
||||||
interpolation_factor=interpolation_factor[i],
|
interpolation_factor=interpolation_factor[i],
|
||||||
|
L_test=num_frames,
|
||||||
|
k=k,
|
||||||
) # 2 x [WHD, rope_dim_list[i]]
|
) # 2 x [WHD, rope_dim_list[i]]
|
||||||
embs.append(emb)
|
embs.append(emb)
|
||||||
|
|
||||||
@@ -182,6 +186,8 @@ def get_1d_rotary_pos_embed(
|
|||||||
use_real: bool = False,
|
use_real: bool = False,
|
||||||
theta_rescale_factor: float = 1.0,
|
theta_rescale_factor: float = 1.0,
|
||||||
interpolation_factor: float = 1.0,
|
interpolation_factor: float = 1.0,
|
||||||
|
L_test: int = 100,
|
||||||
|
k: int = 0,
|
||||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
"""
|
"""
|
||||||
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
|
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
|
||||||
@@ -215,6 +221,12 @@ def get_1d_rotary_pos_embed(
|
|||||||
theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)
|
theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)
|
||||||
) # [D/2]
|
) # [D/2]
|
||||||
# assert interpolation_factor == 1.0, f"interpolation_factor: {interpolation_factor}"
|
# assert interpolation_factor == 1.0, f"interpolation_factor: {interpolation_factor}"
|
||||||
|
|
||||||
|
#RIFLEx https://github.com/thu-ml/RIFLEx
|
||||||
|
if k > 0:
|
||||||
|
freqs[k-1] = 0.9 * 2 * torch.pi / L_test
|
||||||
|
|
||||||
|
|
||||||
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
|
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
|
||||||
if use_real:
|
if use_real:
|
||||||
freqs_cos = freqs.cos().repeat_interleave(2, dim=1) # [S, D]
|
freqs_cos = freqs.cos().repeat_interleave(2, dim=1) # [S, D]
|
||||||
|
|||||||
@@ -4,7 +4,8 @@ from copy import deepcopy
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from transformers import CLIPTextModel, CLIPTokenizer, AutoTokenizer, AutoModel, LlavaForConditionalGeneration, AutoProcessor
|
from transformers import CLIPTextModel, CLIPTokenizer, AutoTokenizer, AutoModel, AutoProcessor, CLIPImageProcessor #LlavaForConditionalGeneration
|
||||||
|
from .modeling_llava import LlavaForConditionalGeneration
|
||||||
from transformers.utils import ModelOutput
|
from transformers.utils import ModelOutput
|
||||||
|
|
||||||
from ..constants import TEXT_ENCODER_PATH, TOKENIZER_PATH
|
from ..constants import TEXT_ENCODER_PATH, TOKENIZER_PATH
|
||||||
@@ -46,7 +47,7 @@ def load_text_encoder(
|
|||||||
text_encoder = LlavaForConditionalGeneration.from_pretrained(
|
text_encoder = LlavaForConditionalGeneration.from_pretrained(
|
||||||
text_encoder_path,
|
text_encoder_path,
|
||||||
low_cpu_mem_usage=True,
|
low_cpu_mem_usage=True,
|
||||||
quantization_config=quantization_config
|
quantization_config=quantization_config,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported text encoder type: {text_encoder_type}")
|
raise ValueError(f"Unsupported text encoder type: {text_encoder_type}")
|
||||||
@@ -80,6 +81,10 @@ def load_tokenizer(
|
|||||||
tokenizer = AutoTokenizer.from_pretrained(
|
tokenizer = AutoTokenizer.from_pretrained(
|
||||||
tokenizer_path, padding_side=padding_side
|
tokenizer_path, padding_side=padding_side
|
||||||
)
|
)
|
||||||
|
elif tokenizer_type == "llm-i2v":
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(
|
||||||
|
tokenizer_path, padding_side=padding_side
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported tokenizer type: {tokenizer_type}")
|
raise ValueError(f"Unsupported tokenizer type: {tokenizer_type}")
|
||||||
|
|
||||||
@@ -121,6 +126,7 @@ class TextEncoder(nn.Module):
|
|||||||
tokenizer_path: Optional[str] = None,
|
tokenizer_path: Optional[str] = None,
|
||||||
output_key: Optional[str] = None,
|
output_key: Optional[str] = None,
|
||||||
use_attention_mask: bool = True,
|
use_attention_mask: bool = True,
|
||||||
|
i2v_mode: bool = False,
|
||||||
input_max_length: Optional[int] = None,
|
input_max_length: Optional[int] = None,
|
||||||
hidden_state_skip_layer: Optional[int] = None,
|
hidden_state_skip_layer: Optional[int] = None,
|
||||||
apply_final_norm: bool = False,
|
apply_final_norm: bool = False,
|
||||||
@@ -160,7 +166,10 @@ class TextEncoder(nn.Module):
|
|||||||
elif "llm" in text_encoder_type or "glm" in text_encoder_type or "vlm" in text_encoder_type:
|
elif "llm" in text_encoder_type or "glm" in text_encoder_type or "vlm" in text_encoder_type:
|
||||||
self.output_key = output_key or "last_hidden_state"
|
self.output_key = output_key or "last_hidden_state"
|
||||||
if "glm" in text_encoder_type or "vlm" in text_encoder_type:
|
if "glm" in text_encoder_type or "vlm" in text_encoder_type:
|
||||||
self.processor = AutoProcessor.from_pretrained(text_encoder_path, device=device)
|
#self.processor = AutoProcessor.from_pretrained(text_encoder_path, device=device)
|
||||||
|
self.processor = CLIPImageProcessor.from_pretrained(text_encoder_path, use_fast=False)
|
||||||
|
self.processor.patch_size = None
|
||||||
|
self.processor.vision_feature_select_strategy = None
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported text encoder type: {text_encoder_type}")
|
raise ValueError(f"Unsupported text encoder type: {text_encoder_type}")
|
||||||
|
|
||||||
@@ -244,7 +253,7 @@ class TextEncoder(nn.Module):
|
|||||||
return_attention_mask=True,
|
return_attention_mask=True,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
if self.text_encoder_type == "vlm":
|
if self.text_encoder_type == "vlm" and image1 is not None:
|
||||||
raw_images = []
|
raw_images = []
|
||||||
if image1 is not None:
|
if image1 is not None:
|
||||||
raw_images.append(image1.squeeze(0)*255)
|
raw_images.append(image1.squeeze(0)*255)
|
||||||
@@ -277,6 +286,9 @@ class TextEncoder(nn.Module):
|
|||||||
return_texts=False,
|
return_texts=False,
|
||||||
prompt_template=None,
|
prompt_template=None,
|
||||||
image_token_selection_expr="::4",
|
image_token_selection_expr="::4",
|
||||||
|
data_type="image",
|
||||||
|
semantic_images=None,
|
||||||
|
image_embed_interleave=2,
|
||||||
device=None,
|
device=None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -299,78 +311,159 @@ class TextEncoder(nn.Module):
|
|||||||
hidden_state_skip_layer, self.hidden_state_skip_layer
|
hidden_state_skip_layer, self.hidden_state_skip_layer
|
||||||
)
|
)
|
||||||
do_sample = use_default(do_sample, not self.reproduce)
|
do_sample = use_default(do_sample, not self.reproduce)
|
||||||
attention_mask = (
|
if semantic_images is None:
|
||||||
batch_encoding["attention_mask"].to(device) if use_attention_mask else None
|
attention_mask = (
|
||||||
)
|
batch_encoding["attention_mask"].to(device) if use_attention_mask else None
|
||||||
for k,v in batch_encoding.items():
|
|
||||||
batch_encoding[k] = v.to(device) if isinstance(v, torch.Tensor) else v
|
|
||||||
outputs = self.model(
|
|
||||||
**batch_encoding,
|
|
||||||
output_hidden_states=output_hidden_states
|
|
||||||
or hidden_state_skip_layer is not None,
|
|
||||||
)
|
|
||||||
|
|
||||||
if hidden_state_skip_layer is not None:
|
|
||||||
last_hidden_state = outputs.hidden_states[-(hidden_state_skip_layer + 1)]
|
|
||||||
# Real last hidden state already has layer norm applied. So here we only apply it
|
|
||||||
# for intermediate layers.
|
|
||||||
if hidden_state_skip_layer > 0 and self.apply_final_norm:
|
|
||||||
last_hidden_state = self.model.final_layer_norm(last_hidden_state)
|
|
||||||
else:
|
|
||||||
last_hidden_state = outputs[self.output_key]
|
|
||||||
|
|
||||||
# Remove hidden states of instruction tokens, only keep prompt tokens.
|
|
||||||
if prompt_template is not None and self.text_encoder_type == "llm":
|
|
||||||
crop_start = prompt_template.get("crop_start", -1)
|
|
||||||
if crop_start > 0:
|
|
||||||
last_hidden_state = last_hidden_state[:, crop_start:]
|
|
||||||
attention_mask = (
|
|
||||||
attention_mask[:, crop_start:] if use_attention_mask else None
|
|
||||||
)
|
|
||||||
elif prompt_template is not None and self.text_encoder_type == "vlm":
|
|
||||||
# Temporory implementation for one round chat template to get rid of system prompts aand chat header
|
|
||||||
user_start_tokens = self.tokenizer(
|
|
||||||
text="<|start_header_id|>user<|end_header_id|>",
|
|
||||||
add_special_tokens=False,
|
|
||||||
return_tensors="pt"
|
|
||||||
)
|
|
||||||
image_token = self.tokenizer(
|
|
||||||
text="<image>",
|
|
||||||
add_special_tokens=False,
|
|
||||||
return_tensors="pt"
|
|
||||||
)
|
|
||||||
image_token = image_token["input_ids"].to(device)
|
|
||||||
user_start_tokens["input_ids"] = user_start_tokens["input_ids"].to(device)
|
|
||||||
tk_idx, tk_n, tk_len = find_subsequence(batch_encoding["input_ids"], user_start_tokens["input_ids"])
|
|
||||||
if tk_n != 1:
|
|
||||||
raise ValueError("Template seems not in the required format, do you have <|start_header_id|>user<|end_header_id|> in place, and only one round of user input?")
|
|
||||||
user_tokens = batch_encoding["input_ids"][:,tk_idx[0]+tk_len:]
|
|
||||||
img_idx, img_n, _ = find_subsequence(user_tokens, image_token)
|
|
||||||
img_seq_len=outputs["image_hidden_states"].shape[1]
|
|
||||||
last_hidden_state = last_hidden_state[:, tk_idx[0]+tk_len:]
|
|
||||||
# create image_mask to subset non-image hidden state
|
|
||||||
seq_mask = torch.ones_like(last_hidden_state, device=device, dtype=torch.bool)
|
|
||||||
img_mask=torch.zeros_like(outputs["image_hidden_states"][0:1], device=device, dtype=torch.bool)
|
|
||||||
img_mask[:, multi_slice_to_mask(image_token_selection_expr, img_mask.shape[1])]=True
|
|
||||||
|
|
||||||
drift=0
|
|
||||||
for i in img_idx:
|
|
||||||
i = i+drift
|
|
||||||
seq_mask[:,i:i+img_seq_len,:] = img_mask
|
|
||||||
drift+=img_seq_len
|
|
||||||
|
|
||||||
last_hidden_state = last_hidden_state[seq_mask].view(1,-1,outputs["image_hidden_states"].shape[-1])
|
|
||||||
|
|
||||||
attention_mask = torch.ones(last_hidden_state.shape[0], last_hidden_state.shape[1], device=device, dtype=torch.int64)
|
|
||||||
|
|
||||||
elif prompt_template is None and self.text_encoder_type == "vlm":
|
|
||||||
raise ValueError("Vlm encoders must use compatiable chat template.")
|
|
||||||
|
|
||||||
if output_hidden_states:
|
|
||||||
return TextEncoderModelOutput(
|
|
||||||
last_hidden_state, attention_mask, outputs.hidden_states
|
|
||||||
)
|
)
|
||||||
return TextEncoderModelOutput(last_hidden_state, attention_mask)
|
for k,v in batch_encoding.items():
|
||||||
|
batch_encoding[k] = v.to(device) if isinstance(v, torch.Tensor) else v
|
||||||
|
outputs = self.model(
|
||||||
|
**batch_encoding,
|
||||||
|
output_hidden_states=output_hidden_states
|
||||||
|
or hidden_state_skip_layer is not None,
|
||||||
|
)
|
||||||
|
|
||||||
|
if hidden_state_skip_layer is not None:
|
||||||
|
last_hidden_state = outputs.hidden_states[-(hidden_state_skip_layer + 1)]
|
||||||
|
# Real last hidden state already has layer norm applied. So here we only apply it
|
||||||
|
# for intermediate layers.
|
||||||
|
if hidden_state_skip_layer > 0 and self.apply_final_norm:
|
||||||
|
last_hidden_state = self.model.final_layer_norm(last_hidden_state)
|
||||||
|
else:
|
||||||
|
last_hidden_state = outputs[self.output_key]
|
||||||
|
|
||||||
|
# Remove hidden states of instruction tokens, only keep prompt tokens.
|
||||||
|
if prompt_template is not None and self.text_encoder_type == "llm":
|
||||||
|
crop_start = prompt_template.get("crop_start", -1)
|
||||||
|
if crop_start > 0:
|
||||||
|
last_hidden_state = last_hidden_state[:, crop_start:]
|
||||||
|
attention_mask = (
|
||||||
|
attention_mask[:, crop_start:] if use_attention_mask else None
|
||||||
|
)
|
||||||
|
elif prompt_template is not None and self.text_encoder_type == "vlm":
|
||||||
|
# Temporory implementation for one round chat template to get rid of system prompts aand chat header
|
||||||
|
user_start_tokens = self.tokenizer(
|
||||||
|
text="<|start_header_id|>user<|end_header_id|>",
|
||||||
|
add_special_tokens=False,
|
||||||
|
return_tensors="pt"
|
||||||
|
)
|
||||||
|
image_token = self.tokenizer(
|
||||||
|
text="<image>",
|
||||||
|
add_special_tokens=False,
|
||||||
|
return_tensors="pt"
|
||||||
|
)
|
||||||
|
image_token = image_token["input_ids"].to(device)
|
||||||
|
user_start_tokens["input_ids"] = user_start_tokens["input_ids"].to(device)
|
||||||
|
tk_idx, tk_n, tk_len = find_subsequence(batch_encoding["input_ids"], user_start_tokens["input_ids"])
|
||||||
|
if tk_n != 1:
|
||||||
|
raise ValueError("Template seems not in the required format, do you have <|start_header_id|>user<|end_header_id|> in place, and only one round of user input?")
|
||||||
|
user_tokens = batch_encoding["input_ids"][:,tk_idx[0]+tk_len:]
|
||||||
|
img_idx, img_n, _ = find_subsequence(user_tokens, image_token)
|
||||||
|
img_seq_len=outputs["image_hidden_states"].shape[1]
|
||||||
|
last_hidden_state = last_hidden_state[:, tk_idx[0]+tk_len:]
|
||||||
|
# create image_mask to subset non-image hidden state
|
||||||
|
seq_mask = torch.ones_like(last_hidden_state, device=device, dtype=torch.bool)
|
||||||
|
img_mask=torch.zeros_like(outputs["image_hidden_states"][0:1], device=device, dtype=torch.bool)
|
||||||
|
img_mask[:, multi_slice_to_mask(image_token_selection_expr, img_mask.shape[1])]=True
|
||||||
|
|
||||||
|
drift=0
|
||||||
|
for i in img_idx:
|
||||||
|
i = i+drift
|
||||||
|
seq_mask[:,i:i+img_seq_len,:] = img_mask
|
||||||
|
drift+=img_seq_len
|
||||||
|
|
||||||
|
last_hidden_state = last_hidden_state[seq_mask].view(1,-1,outputs["image_hidden_states"].shape[-1])
|
||||||
|
|
||||||
|
attention_mask = torch.ones(last_hidden_state.shape[0], last_hidden_state.shape[1], device=device, dtype=torch.int64)
|
||||||
|
|
||||||
|
elif prompt_template is None and self.text_encoder_type == "vlm":
|
||||||
|
raise ValueError("Vlm encoders must use compatiable chat template.")
|
||||||
|
|
||||||
|
if output_hidden_states:
|
||||||
|
return TextEncoderModelOutput(
|
||||||
|
last_hidden_state, attention_mask, outputs.hidden_states
|
||||||
|
)
|
||||||
|
return TextEncoderModelOutput(last_hidden_state, attention_mask)
|
||||||
|
else:
|
||||||
|
image_outputs = self.processor(semantic_images, return_tensors='pt')["pixel_values"].to(device)
|
||||||
|
|
||||||
|
attention_mask = (
|
||||||
|
batch_encoding["attention_mask"].to(device) if use_attention_mask else None
|
||||||
|
)
|
||||||
|
#print(prompt_template)
|
||||||
|
outputs = self.model(
|
||||||
|
input_ids=batch_encoding["input_ids"].to(device),
|
||||||
|
attention_mask=attention_mask,
|
||||||
|
output_hidden_states=output_hidden_states or hidden_state_skip_layer is not None,
|
||||||
|
pixel_values=image_outputs,
|
||||||
|
)
|
||||||
|
if hidden_state_skip_layer is not None:
|
||||||
|
last_hidden_state = outputs.hidden_states[-(hidden_state_skip_layer + 1)]
|
||||||
|
# Real last hidden state already has layer norm applied. So here we only apply it
|
||||||
|
# for intermediate layers.
|
||||||
|
if hidden_state_skip_layer > 0 and self.apply_final_norm:
|
||||||
|
last_hidden_state = self.model.final_layer_norm(last_hidden_state)
|
||||||
|
else:
|
||||||
|
last_hidden_state = outputs[self.output_key]
|
||||||
|
if prompt_template is not None:
|
||||||
|
if data_type == 'I2V_image':
|
||||||
|
crop_start = prompt_template.get("crop_start", -1)
|
||||||
|
crop_end = prompt_template.get('assistant_emb_start', -1)
|
||||||
|
elif data_type == 'I2V_video':
|
||||||
|
crop_start = prompt_template.get("crop_start", -1)
|
||||||
|
text_crop_start = crop_start - 1 + prompt_template.get("image_emb_len", 576)
|
||||||
|
image_crop_start = prompt_template.get("image_emb_start", 5)
|
||||||
|
image_crop_end = prompt_template.get('image_emb_end', 581)
|
||||||
|
batch_indices, last_double_return_token_indices = torch.where(
|
||||||
|
batch_encoding["input_ids"] == prompt_template.get('double_return_token_id', 271))
|
||||||
|
last_double_return_token_indices = last_double_return_token_indices.reshape(
|
||||||
|
batch_encoding["input_ids"].shape[0], -1)[:, -1]
|
||||||
|
batch_indices = batch_indices.reshape(batch_encoding["input_ids"].shape[0], -1)[:, -1]
|
||||||
|
assistant_crop_start = last_double_return_token_indices - 1 + prompt_template.get(
|
||||||
|
"image_emb_len", 576) - 4
|
||||||
|
assistant_crop_end = last_double_return_token_indices - 1 + prompt_template.get(
|
||||||
|
"image_emb_len", 576)
|
||||||
|
|
||||||
|
attention_mask_assistant_crop_start = last_double_return_token_indices - 4
|
||||||
|
attention_mask_assistant_crop_end = last_double_return_token_indices
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unsupported data type: {data_type}")
|
||||||
|
|
||||||
|
text_last_hidden_state = []
|
||||||
|
text_attention_mask = []
|
||||||
|
image_last_hidden_state = []
|
||||||
|
image_attention_mask = []
|
||||||
|
for i in range(batch_encoding["input_ids"].shape[0]):
|
||||||
|
text_last_hidden_state.append(torch.cat(
|
||||||
|
[last_hidden_state[i, text_crop_start:assistant_crop_start[i].item()],
|
||||||
|
last_hidden_state[i, assistant_crop_end[i].item():]]))
|
||||||
|
text_attention_mask.append(torch.cat(
|
||||||
|
[attention_mask[i, crop_start:attention_mask_assistant_crop_start[i].item()], attention_mask[i,
|
||||||
|
attention_mask_assistant_crop_end[
|
||||||
|
i].item():]]) if use_attention_mask else None)
|
||||||
|
image_last_hidden_state.append(last_hidden_state[i, image_crop_start:image_crop_end])
|
||||||
|
image_attention_mask.append(
|
||||||
|
torch.ones(image_last_hidden_state[-1].shape[0]).to(last_hidden_state.device).to(
|
||||||
|
attention_mask.dtype) if use_attention_mask else None)
|
||||||
|
|
||||||
|
text_last_hidden_state = torch.stack(text_last_hidden_state)
|
||||||
|
text_attention_mask = torch.stack(text_attention_mask)
|
||||||
|
image_last_hidden_state = torch.stack(image_last_hidden_state)
|
||||||
|
image_attention_mask = torch.stack(image_attention_mask)
|
||||||
|
|
||||||
|
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]
|
||||||
|
|
||||||
|
assert text_last_hidden_state.shape[0] == text_attention_mask.shape[0] and \
|
||||||
|
image_last_hidden_state.shape[0] == image_attention_mask.shape[0]
|
||||||
|
|
||||||
|
last_hidden_state = torch.cat([image_last_hidden_state, text_last_hidden_state], dim=1)
|
||||||
|
attention_mask = torch.cat([image_attention_mask, text_attention_mask], dim=1)
|
||||||
|
if output_hidden_states:
|
||||||
|
return TextEncoderModelOutput(last_hidden_state, attention_mask,
|
||||||
|
hidden_states_list=outputs.hidden_states)
|
||||||
|
return TextEncoderModelOutput(last_hidden_state, attention_mask)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -390,3 +483,48 @@ class TextEncoder(nn.Module):
|
|||||||
hidden_state_skip_layer=hidden_state_skip_layer,
|
hidden_state_skip_layer=hidden_state_skip_layer,
|
||||||
return_texts=return_texts,
|
return_texts=return_texts,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
xtuner_config={
|
||||||
|
"architectures": [
|
||||||
|
"LlavaForConditionalGeneration"
|
||||||
|
],
|
||||||
|
"ignore_index": -100,
|
||||||
|
"image_token_index": 128257,
|
||||||
|
"model_type": "llava",
|
||||||
|
"pad_token_id": 128258,
|
||||||
|
"projector_hidden_act": "gelu",
|
||||||
|
"text_config": {
|
||||||
|
"architectures": [
|
||||||
|
"LlamaForCausalLM"
|
||||||
|
],
|
||||||
|
"bos_token_id": 128000,
|
||||||
|
"eos_token_id": 128001,
|
||||||
|
"intermediate_size": 14336,
|
||||||
|
"max_position_embeddings": 8192,
|
||||||
|
"model_type": "llama",
|
||||||
|
"num_key_value_heads": 8,
|
||||||
|
"rms_norm_eps": 1e-05,
|
||||||
|
"rope_theta": 500000.0,
|
||||||
|
"torch_dtype": "float16",
|
||||||
|
"vocab_size": 128320
|
||||||
|
},
|
||||||
|
"torch_dtype": "float16",
|
||||||
|
"transformers_version": "4.40.1",
|
||||||
|
"vision_config": {
|
||||||
|
"architectures": [
|
||||||
|
"CLIPVisionModel"
|
||||||
|
],
|
||||||
|
"dropout": 0.0,
|
||||||
|
"hidden_size": 1024,
|
||||||
|
"image_size": 336,
|
||||||
|
"intermediate_size": 4096,
|
||||||
|
"model_type": "clip_vision_model",
|
||||||
|
"num_attention_heads": 16,
|
||||||
|
"num_hidden_layers": 24,
|
||||||
|
"patch_size": 14,
|
||||||
|
"projection_dim": 768,
|
||||||
|
"torch_dtype": "float32"
|
||||||
|
},
|
||||||
|
"vision_feature_layer": -2,
|
||||||
|
"vision_feature_select_strategy": "default"
|
||||||
|
}
|
||||||
@@ -0,0 +1,131 @@
|
|||||||
|
# coding=utf-8
|
||||||
|
# Copyright 2023 Microsoft Research & University of Wisconsin-Madison and the HuggingFace Inc. team. All rights reserved.
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
"""Llava model configuration"""
|
||||||
|
|
||||||
|
from transformers.configuration_utils import PretrainedConfig
|
||||||
|
from transformers.utils import logging
|
||||||
|
from transformers.models.auto import CONFIG_MAPPING, AutoConfig
|
||||||
|
|
||||||
|
|
||||||
|
logger = logging.get_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class LlavaConfig(PretrainedConfig):
|
||||||
|
r"""
|
||||||
|
This is the configuration class to store the configuration of a [`LlavaForConditionalGeneration`]. It is used to instantiate an
|
||||||
|
Llava model according to the specified arguments, defining the model architecture. Instantiating a configuration
|
||||||
|
with the defaults will yield a similar configuration to that of the Llava-9B.
|
||||||
|
|
||||||
|
e.g. [llava-hf/llava-9b](https://huggingface.co/llava-hf/llava-9b)
|
||||||
|
|
||||||
|
Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the
|
||||||
|
documentation from [`PretrainedConfig`] for more information.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
vision_config (`Union[AutoConfig, dict]`, *optional*, defaults to `CLIPVisionConfig`):
|
||||||
|
The config object or dictionary of the vision backbone.
|
||||||
|
text_config (`Union[AutoConfig, dict]`, *optional*, defaults to `LlamaConfig`):
|
||||||
|
The config object or dictionary of the text backbone.
|
||||||
|
ignore_index (`int`, *optional*, defaults to -100):
|
||||||
|
The ignore index for the loss function.
|
||||||
|
image_token_index (`int`, *optional*, defaults to 32000):
|
||||||
|
The image token index to encode the image prompt.
|
||||||
|
projector_hidden_act (`str`, *optional*, defaults to `"gelu"`):
|
||||||
|
The activation function used by the multimodal projector.
|
||||||
|
vision_feature_select_strategy (`str`, *optional*, defaults to `"default"`):
|
||||||
|
The feature selection strategy used to select the vision feature from the vision backbone.
|
||||||
|
Can be one of `"default"` or `"full"`.
|
||||||
|
vision_feature_layer (`int`, *optional*, defaults to -2):
|
||||||
|
The index of the layer to select the vision feature.
|
||||||
|
image_seq_length (`int`, *optional*, defaults to 576):
|
||||||
|
Sequence length of one image embedding.
|
||||||
|
|
||||||
|
Example:
|
||||||
|
|
||||||
|
```python
|
||||||
|
>>> from transformers import LlavaForConditionalGeneration, LlavaConfig, CLIPVisionConfig, LlamaConfig
|
||||||
|
|
||||||
|
>>> # Initializing a CLIP-vision config
|
||||||
|
>>> vision_config = CLIPVisionConfig()
|
||||||
|
|
||||||
|
>>> # Initializing a Llama config
|
||||||
|
>>> text_config = LlamaConfig()
|
||||||
|
|
||||||
|
>>> # Initializing a Llava llava-1.5-7b style configuration
|
||||||
|
>>> configuration = LlavaConfig(vision_config, text_config)
|
||||||
|
|
||||||
|
>>> # Initializing a model from the llava-1.5-7b style configuration
|
||||||
|
>>> model = LlavaForConditionalGeneration(configuration)
|
||||||
|
|
||||||
|
>>> # Accessing the model configuration
|
||||||
|
>>> configuration = model.config
|
||||||
|
```"""
|
||||||
|
|
||||||
|
model_type = "llava"
|
||||||
|
sub_configs = {"text_config": AutoConfig, "vision_config": AutoConfig}
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
vision_config=None,
|
||||||
|
text_config=None,
|
||||||
|
ignore_index=-100,
|
||||||
|
image_token_index=32000,
|
||||||
|
projector_hidden_act="gelu",
|
||||||
|
vision_feature_select_strategy="default",
|
||||||
|
vision_feature_layer=-2,
|
||||||
|
image_seq_length=576,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
self.ignore_index = ignore_index
|
||||||
|
self.image_token_index = image_token_index
|
||||||
|
self.projector_hidden_act = projector_hidden_act
|
||||||
|
self.image_seq_length = image_seq_length
|
||||||
|
|
||||||
|
if vision_feature_select_strategy not in ["default", "full"]:
|
||||||
|
raise ValueError(
|
||||||
|
"vision_feature_select_strategy should be one of 'default', 'full'."
|
||||||
|
f"Got: {vision_feature_select_strategy}"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.vision_feature_select_strategy = vision_feature_select_strategy
|
||||||
|
self.vision_feature_layer = vision_feature_layer
|
||||||
|
|
||||||
|
if isinstance(vision_config, dict):
|
||||||
|
vision_config["model_type"] = (
|
||||||
|
vision_config["model_type"] if "model_type" in vision_config else "clip_vision_model"
|
||||||
|
)
|
||||||
|
vision_config = CONFIG_MAPPING[vision_config["model_type"]](**vision_config)
|
||||||
|
elif vision_config is None:
|
||||||
|
vision_config = CONFIG_MAPPING["clip_vision_model"](
|
||||||
|
intermediate_size=4096,
|
||||||
|
hidden_size=1024,
|
||||||
|
patch_size=14,
|
||||||
|
image_size=336,
|
||||||
|
num_hidden_layers=24,
|
||||||
|
num_attention_heads=16,
|
||||||
|
vocab_size=32000,
|
||||||
|
projection_dim=768,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.vision_config = vision_config
|
||||||
|
|
||||||
|
if isinstance(text_config, dict):
|
||||||
|
text_config["model_type"] = text_config["model_type"] if "model_type" in text_config else "llama"
|
||||||
|
text_config = CONFIG_MAPPING[text_config["model_type"]](**text_config)
|
||||||
|
elif text_config is None:
|
||||||
|
text_config = CONFIG_MAPPING["llama"]()
|
||||||
|
|
||||||
|
self.text_config = text_config
|
||||||
|
|
||||||
|
super().__init__(**kwargs)
|
||||||
@@ -0,0 +1,620 @@
|
|||||||
|
# coding=utf-8
|
||||||
|
# Copyright 2023 the HuggingFace Inc. team. All rights reserved.
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
"""PyTorch Llava model."""
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.utils.checkpoint
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
from transformers.activations import ACT2FN
|
||||||
|
from transformers.generation import GenerationMixin
|
||||||
|
from transformers.modeling_outputs import ModelOutput
|
||||||
|
from transformers.modeling_utils import PreTrainedModel
|
||||||
|
from transformers.utils import (
|
||||||
|
add_start_docstrings,
|
||||||
|
add_start_docstrings_to_model_forward,
|
||||||
|
logging,
|
||||||
|
replace_return_docstrings,
|
||||||
|
)
|
||||||
|
from transformers.models.auto import AutoModel, AutoModelForCausalLM
|
||||||
|
from .configuration_llava import LlavaConfig
|
||||||
|
|
||||||
|
|
||||||
|
logger = logging.get_logger(__name__)
|
||||||
|
|
||||||
|
_CONFIG_FOR_DOC = "LlavaConfig"
|
||||||
|
|
||||||
|
# Base docstring
|
||||||
|
_CHECKPOINT_FOR_DOC = "llava-hf/llava-1.5-7b-hf"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class LlavaCausalLMOutputWithPast(ModelOutput):
|
||||||
|
"""
|
||||||
|
Base class for Llava causal language model (or autoregressive) outputs.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
|
||||||
|
Language modeling loss (for next-token prediction).
|
||||||
|
logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
|
||||||
|
Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).
|
||||||
|
past_key_values (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
|
||||||
|
Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of shape
|
||||||
|
`(batch_size, num_heads, sequence_length, embed_size_per_head)`)
|
||||||
|
|
||||||
|
Contains pre-computed hidden-states (key and values in the self-attention blocks) that can be used (see
|
||||||
|
`past_key_values` input) to speed up sequential decoding.
|
||||||
|
hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
|
||||||
|
Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
|
||||||
|
one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.
|
||||||
|
|
||||||
|
Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
|
||||||
|
attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
|
||||||
|
Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
|
||||||
|
sequence_length)`.
|
||||||
|
|
||||||
|
Attentions weights after the attention softmax, used to compute the weighted average in the self-attention
|
||||||
|
heads.
|
||||||
|
image_hidden_states (`torch.FloatTensor`, *optional*):
|
||||||
|
A `torch.FloatTensor` of size (batch_size, num_images, sequence_length, hidden_size)`.
|
||||||
|
image_hidden_states of the model produced by the vision encoder and after projecting the last hidden state.
|
||||||
|
"""
|
||||||
|
|
||||||
|
loss: Optional[torch.FloatTensor] = None
|
||||||
|
logits: torch.FloatTensor = None
|
||||||
|
past_key_values: Optional[List[torch.FloatTensor]] = None
|
||||||
|
hidden_states: Optional[Tuple[torch.FloatTensor]] = None
|
||||||
|
attentions: Optional[Tuple[torch.FloatTensor]] = None
|
||||||
|
image_hidden_states: Optional[torch.FloatTensor] = None
|
||||||
|
|
||||||
|
|
||||||
|
class LlavaMultiModalProjector(nn.Module):
|
||||||
|
def __init__(self, config: LlavaConfig):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.linear_1 = nn.Linear(config.vision_config.hidden_size, config.text_config.hidden_size, bias=True)
|
||||||
|
self.act = ACT2FN[config.projector_hidden_act]
|
||||||
|
self.linear_2 = nn.Linear(config.text_config.hidden_size, config.text_config.hidden_size, bias=True)
|
||||||
|
|
||||||
|
def forward(self, image_features):
|
||||||
|
hidden_states = self.linear_1(image_features)
|
||||||
|
hidden_states = self.act(hidden_states)
|
||||||
|
hidden_states = self.linear_2(hidden_states)
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
|
LLAVA_START_DOCSTRING = r"""
|
||||||
|
This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the
|
||||||
|
library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads
|
||||||
|
etc.)
|
||||||
|
|
||||||
|
This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.
|
||||||
|
Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage
|
||||||
|
and behavior.
|
||||||
|
|
||||||
|
Parameters:
|
||||||
|
config ([`LlavaConfig`] or [`LlavaVisionConfig`]):
|
||||||
|
Model configuration class with all the parameters of the model. Initializing with a config file does not
|
||||||
|
load the weights associated with the model, only the configuration. Check out the
|
||||||
|
[`~PreTrainedModel.from_pretrained`] method to load the model weights.
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
@add_start_docstrings(
|
||||||
|
"The bare LLaMA Model outputting raw hidden-states without any specific head on top.",
|
||||||
|
LLAVA_START_DOCSTRING,
|
||||||
|
)
|
||||||
|
class LlavaPreTrainedModel(PreTrainedModel):
|
||||||
|
config_class = LlavaConfig
|
||||||
|
base_model_prefix = "model"
|
||||||
|
supports_gradient_checkpointing = True
|
||||||
|
_no_split_modules = ["LlavaVisionAttention"]
|
||||||
|
_skip_keys_device_placement = "past_key_values"
|
||||||
|
_supports_cache_class = True
|
||||||
|
_supports_flash_attn_2 = True
|
||||||
|
_supports_sdpa = True
|
||||||
|
|
||||||
|
def _init_weights(self, module):
|
||||||
|
# important: this ported version of Llava isn't meant for training from scratch - only
|
||||||
|
# inference and fine-tuning - so the proper init weights code has been removed - the original codebase
|
||||||
|
# https://github.com/haotian-liu/LLaVA/tree/main/llava should serve for that purpose
|
||||||
|
std = (
|
||||||
|
self.config.initializer_range
|
||||||
|
if hasattr(self.config, "initializer_range")
|
||||||
|
else self.config.text_config.initializer_range
|
||||||
|
)
|
||||||
|
|
||||||
|
if hasattr(module, "class_embedding"):
|
||||||
|
module.class_embedding.data.normal_(mean=0.0, std=std)
|
||||||
|
|
||||||
|
if isinstance(module, (nn.Linear, nn.Conv2d)):
|
||||||
|
module.weight.data.normal_(mean=0.0, std=std)
|
||||||
|
if module.bias is not None:
|
||||||
|
module.bias.data.zero_()
|
||||||
|
elif isinstance(module, nn.Embedding):
|
||||||
|
module.weight.data.normal_(mean=0.0, std=std)
|
||||||
|
if module.padding_idx is not None:
|
||||||
|
module.weight.data[module.padding_idx].zero_()
|
||||||
|
|
||||||
|
|
||||||
|
LLAVA_INPUTS_DOCSTRING = r"""
|
||||||
|
Args:
|
||||||
|
input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
|
||||||
|
Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide
|
||||||
|
it.
|
||||||
|
|
||||||
|
Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
|
||||||
|
[`PreTrainedTokenizer.__call__`] for details.
|
||||||
|
|
||||||
|
[What are input IDs?](../glossary#input-ids)
|
||||||
|
pixel_values (`torch.FloatTensor` of shape `(batch_size, num_channels, image_size, image_size)):
|
||||||
|
The tensors corresponding to the input images. Pixel values can be obtained using
|
||||||
|
[`AutoImageProcessor`]. See [`CLIPImageProcessor.__call__`] for details ([]`LlavaProcessor`] uses
|
||||||
|
[`CLIPImageProcessor`] for processing images).
|
||||||
|
attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
|
||||||
|
Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:
|
||||||
|
|
||||||
|
- 1 for tokens that are **not masked**,
|
||||||
|
- 0 for tokens that are **masked**.
|
||||||
|
|
||||||
|
[What are attention masks?](../glossary#attention-mask)
|
||||||
|
|
||||||
|
Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
|
||||||
|
[`PreTrainedTokenizer.__call__`] for details.
|
||||||
|
|
||||||
|
If `past_key_values` is used, optionally only the last `decoder_input_ids` have to be input (see
|
||||||
|
`past_key_values`).
|
||||||
|
|
||||||
|
If you want to change padding behavior, you should read [`modeling_opt._prepare_decoder_attention_mask`]
|
||||||
|
and modify to your needs. See diagram 1 in [the paper](https://arxiv.org/abs/1910.13461) for more
|
||||||
|
information on the default strategy.
|
||||||
|
|
||||||
|
- 1 indicates the head is **not masked**,
|
||||||
|
- 0 indicates the head is **masked**.
|
||||||
|
position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
|
||||||
|
Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
|
||||||
|
config.n_positions - 1]`. [What are position IDs?](../glossary#position-ids)
|
||||||
|
past_key_values (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
|
||||||
|
Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of shape
|
||||||
|
`(batch_size, num_heads, sequence_length, embed_size_per_head)`) and 2 additional tensors of shape
|
||||||
|
`(batch_size, num_heads, encoder_sequence_length, embed_size_per_head)`.
|
||||||
|
|
||||||
|
Contains pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention
|
||||||
|
blocks) that can be used (see `past_key_values` input) to speed up sequential decoding.
|
||||||
|
|
||||||
|
If `past_key_values` are used, the user can optionally input only the last `decoder_input_ids` (those that
|
||||||
|
don't have their past key value states given to this model) of shape `(batch_size, 1)` instead of all
|
||||||
|
`decoder_input_ids` of shape `(batch_size, sequence_length)`.
|
||||||
|
inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
|
||||||
|
Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This
|
||||||
|
is useful if you want more control over how to convert `input_ids` indices into associated vectors than the
|
||||||
|
model's internal embedding lookup matrix.
|
||||||
|
vision_feature_layer (`int`, *optional*, defaults to -2):
|
||||||
|
The index of the layer to select the vision feature.
|
||||||
|
vision_feature_select_strategy (`str`, *optional*, defaults to `"default"`):
|
||||||
|
The feature selection strategy used to select the vision feature from the vision backbone.
|
||||||
|
Can be one of `"default"` or `"full"`.
|
||||||
|
use_cache (`bool`, *optional*):
|
||||||
|
If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see
|
||||||
|
`past_key_values`).
|
||||||
|
output_attentions (`bool`, *optional*):
|
||||||
|
Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
|
||||||
|
tensors for more detail.
|
||||||
|
output_hidden_states (`bool`, *optional*):
|
||||||
|
Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
|
||||||
|
more detail.
|
||||||
|
return_dict (`bool`, *optional*):
|
||||||
|
Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
|
||||||
|
cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*):
|
||||||
|
Indices depicting the position of the input sequence tokens in the sequence. Contrarily to `position_ids`,
|
||||||
|
this tensor is not affected by padding. It is used to update the cache in the correct position and to infer
|
||||||
|
the complete sequence length.
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
@add_start_docstrings(
|
||||||
|
"""The LLAVA model which consists of a vision backbone and a language model.""",
|
||||||
|
LLAVA_START_DOCSTRING,
|
||||||
|
)
|
||||||
|
class LlavaForConditionalGeneration(LlavaPreTrainedModel, GenerationMixin):
|
||||||
|
def __init__(self, config: LlavaConfig):
|
||||||
|
super().__init__(config)
|
||||||
|
self.vision_tower = AutoModel.from_config(config.vision_config)
|
||||||
|
|
||||||
|
self.multi_modal_projector = LlavaMultiModalProjector(config)
|
||||||
|
self.vocab_size = config.text_config.vocab_size
|
||||||
|
self.language_model = AutoModelForCausalLM.from_config(config.text_config)
|
||||||
|
self.pad_token_id = self.config.pad_token_id if self.config.pad_token_id is not None else -1
|
||||||
|
self.post_init()
|
||||||
|
|
||||||
|
def get_input_embeddings(self):
|
||||||
|
return self.language_model.get_input_embeddings()
|
||||||
|
|
||||||
|
def set_input_embeddings(self, value):
|
||||||
|
self.language_model.set_input_embeddings(value)
|
||||||
|
|
||||||
|
def get_output_embeddings(self):
|
||||||
|
return self.language_model.get_output_embeddings()
|
||||||
|
|
||||||
|
def set_output_embeddings(self, new_embeddings):
|
||||||
|
self.language_model.set_output_embeddings(new_embeddings)
|
||||||
|
|
||||||
|
def set_decoder(self, decoder):
|
||||||
|
self.language_model.set_decoder(decoder)
|
||||||
|
|
||||||
|
def get_decoder(self):
|
||||||
|
return self.language_model.get_decoder()
|
||||||
|
|
||||||
|
def tie_weights(self):
|
||||||
|
return self.language_model.tie_weights()
|
||||||
|
|
||||||
|
def resize_token_embeddings(self, new_num_tokens: Optional[int] = None, pad_to_multiple_of=None) -> nn.Embedding:
|
||||||
|
model_embeds = self.language_model.resize_token_embeddings(new_num_tokens, pad_to_multiple_of)
|
||||||
|
# update vocab size
|
||||||
|
self.config.text_config.vocab_size = model_embeds.num_embeddings
|
||||||
|
self.vocab_size = model_embeds.num_embeddings
|
||||||
|
return model_embeds
|
||||||
|
|
||||||
|
def get_image_features(
|
||||||
|
self, pixel_values: torch.FloatTensor, vision_feature_layer: int, vision_feature_select_strategy: str
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Obtains image last hidden states from the vision tower and apply multimodal projection.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
pixel_values (`torch.FloatTensor]` of shape `(batch_size, channels, height, width)`)
|
||||||
|
The tensors corresponding to the input images.
|
||||||
|
vision_feature_layer (`int`):
|
||||||
|
The index of the layer to select the vision feature.
|
||||||
|
vision_feature_select_strategy (`str`):
|
||||||
|
The feature selection strategy used to select the vision feature from the vision backbone.
|
||||||
|
Can be one of `"default"` or `"full"`
|
||||||
|
Returns:
|
||||||
|
image_features (`torch.Tensor`): Image feature tensor of shape `(num_images, image_length, embed_dim)`).
|
||||||
|
"""
|
||||||
|
image_outputs = self.vision_tower(pixel_values, output_hidden_states=True)
|
||||||
|
# this is not memory efficient at all (output_hidden_states=True) will save all the hidden stated.
|
||||||
|
selected_image_feature = image_outputs.hidden_states[vision_feature_layer]
|
||||||
|
if vision_feature_select_strategy == "default":
|
||||||
|
selected_image_feature = selected_image_feature[:, 1:]
|
||||||
|
elif vision_feature_select_strategy == "full":
|
||||||
|
selected_image_feature = selected_image_feature
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unexpected select feature strategy: {self.config.vision_feature_select_strategy}")
|
||||||
|
image_features = self.multi_modal_projector(selected_image_feature)
|
||||||
|
return image_features
|
||||||
|
|
||||||
|
def _merge_input_ids_with_image_features(self, image_features, inputs_embeds, input_ids, attention_mask, labels):
|
||||||
|
num_images, num_image_patches, embed_dim = image_features.shape
|
||||||
|
batch_size, sequence_length = input_ids.shape
|
||||||
|
left_padding = not torch.sum(input_ids[:, -1] == torch.tensor(self.pad_token_id))
|
||||||
|
# 1. Create a mask to know where special image tokens are
|
||||||
|
special_image_token_mask = input_ids == self.config.image_token_index
|
||||||
|
num_special_image_tokens = torch.sum(special_image_token_mask, dim=-1)
|
||||||
|
# Compute the maximum embed dimension
|
||||||
|
max_embed_dim = (num_special_image_tokens.max() * (num_image_patches - 1)) + sequence_length
|
||||||
|
batch_indices, non_image_indices = torch.where(input_ids != self.config.image_token_index)
|
||||||
|
|
||||||
|
# 2. Compute the positions where text should be written
|
||||||
|
# Calculate new positions for text tokens in merged image-text sequence.
|
||||||
|
# `special_image_token_mask` identifies image tokens. Each image token will be replaced by `nb_text_tokens_per_images - 1` text tokens.
|
||||||
|
# `torch.cumsum` computes how each image token shifts subsequent text token positions.
|
||||||
|
# - 1 to adjust for zero-based indexing, as `cumsum` inherently increases indices by one.
|
||||||
|
new_token_positions = torch.cumsum((special_image_token_mask * (num_image_patches - 1) + 1), -1) - 1
|
||||||
|
nb_image_pad = max_embed_dim - 1 - new_token_positions[:, -1]
|
||||||
|
if left_padding:
|
||||||
|
new_token_positions += nb_image_pad[:, None] # offset for left padding
|
||||||
|
text_to_overwrite = new_token_positions[batch_indices, non_image_indices]
|
||||||
|
|
||||||
|
# 3. Create the full embedding, already padded to the maximum position
|
||||||
|
final_embedding = torch.zeros(
|
||||||
|
batch_size, max_embed_dim, embed_dim, dtype=inputs_embeds.dtype, device=inputs_embeds.device
|
||||||
|
)
|
||||||
|
final_attention_mask = torch.zeros(
|
||||||
|
batch_size, max_embed_dim, dtype=attention_mask.dtype, device=inputs_embeds.device
|
||||||
|
)
|
||||||
|
if labels is not None:
|
||||||
|
final_labels = torch.full(
|
||||||
|
(batch_size, max_embed_dim), self.config.ignore_index, dtype=input_ids.dtype, device=input_ids.device
|
||||||
|
)
|
||||||
|
# In case the Vision model or the Language model has been offloaded to CPU, we need to manually
|
||||||
|
# set the corresponding tensors into their correct target device.
|
||||||
|
target_device = inputs_embeds.device
|
||||||
|
batch_indices, non_image_indices, text_to_overwrite = (
|
||||||
|
batch_indices.to(target_device),
|
||||||
|
non_image_indices.to(target_device),
|
||||||
|
text_to_overwrite.to(target_device),
|
||||||
|
)
|
||||||
|
attention_mask = attention_mask.to(target_device)
|
||||||
|
|
||||||
|
# 4. Fill the embeddings based on the mask. If we have ["hey" "<image>", "how", "are"]
|
||||||
|
# we need to index copy on [0, 577, 578, 579] for the text and [1:576] for the image features
|
||||||
|
final_embedding[batch_indices, text_to_overwrite] = inputs_embeds[batch_indices, non_image_indices]
|
||||||
|
final_attention_mask[batch_indices, text_to_overwrite] = attention_mask[batch_indices, non_image_indices]
|
||||||
|
if labels is not None:
|
||||||
|
final_labels[batch_indices, text_to_overwrite] = labels[batch_indices, non_image_indices]
|
||||||
|
|
||||||
|
# 5. Fill the embeddings corresponding to the images. Anything that is not `text_positions` needs filling (#29835)
|
||||||
|
image_to_overwrite = torch.full(
|
||||||
|
(batch_size, max_embed_dim), True, dtype=torch.bool, device=inputs_embeds.device
|
||||||
|
)
|
||||||
|
image_to_overwrite[batch_indices, text_to_overwrite] = False
|
||||||
|
if left_padding:
|
||||||
|
image_to_overwrite &= image_to_overwrite.cumsum(-1) - 1 >= nb_image_pad[:, None].to(target_device)
|
||||||
|
else:
|
||||||
|
mask = torch.ones_like(image_to_overwrite, dtype=torch.bool).cumsum(-1) - 1
|
||||||
|
padding_mask = mask <= new_token_positions[:, -1:].to(target_device)
|
||||||
|
image_to_overwrite &= padding_mask
|
||||||
|
|
||||||
|
if image_to_overwrite.sum() != image_features.shape[:-1].numel():
|
||||||
|
raise ValueError(
|
||||||
|
f"The input provided to the model are wrong. The number of image tokens is {torch.sum(special_image_token_mask)} while"
|
||||||
|
f" the number of image given to the model is {num_images}. This prevents correct indexing and breaks batch generation."
|
||||||
|
)
|
||||||
|
|
||||||
|
final_embedding[image_to_overwrite] = image_features.contiguous().reshape(-1, embed_dim).to(target_device)
|
||||||
|
final_attention_mask |= image_to_overwrite
|
||||||
|
position_ids = (final_attention_mask.cumsum(-1) - 1).masked_fill_((final_attention_mask == 0), 1)
|
||||||
|
|
||||||
|
# 6. Mask out the embedding at padding positions, as we later use the past_key_value value to determine the non-attended tokens.
|
||||||
|
batch_indices, pad_indices = torch.where(input_ids == self.pad_token_id)
|
||||||
|
indices_to_mask = new_token_positions[batch_indices, pad_indices]
|
||||||
|
|
||||||
|
final_embedding[batch_indices, indices_to_mask] = 0
|
||||||
|
|
||||||
|
if labels is None:
|
||||||
|
final_labels = None
|
||||||
|
|
||||||
|
return final_embedding, final_attention_mask, final_labels, position_ids
|
||||||
|
|
||||||
|
@add_start_docstrings_to_model_forward(LLAVA_INPUTS_DOCSTRING)
|
||||||
|
@replace_return_docstrings(output_type=LlavaCausalLMOutputWithPast, config_class=_CONFIG_FOR_DOC)
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
input_ids: torch.LongTensor = None,
|
||||||
|
pixel_values: torch.FloatTensor = None,
|
||||||
|
attention_mask: Optional[torch.Tensor] = None,
|
||||||
|
position_ids: Optional[torch.LongTensor] = None,
|
||||||
|
past_key_values: Optional[List[torch.FloatTensor]] = None,
|
||||||
|
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
vision_feature_layer: Optional[int] = None,
|
||||||
|
vision_feature_select_strategy: Optional[str] = None,
|
||||||
|
labels: Optional[torch.LongTensor] = None,
|
||||||
|
use_cache: Optional[bool] = None,
|
||||||
|
output_attentions: Optional[bool] = None,
|
||||||
|
output_hidden_states: Optional[bool] = None,
|
||||||
|
return_dict: Optional[bool] = None,
|
||||||
|
cache_position: Optional[torch.LongTensor] = None,
|
||||||
|
num_logits_to_keep: int = 0,
|
||||||
|
) -> Union[Tuple, LlavaCausalLMOutputWithPast]:
|
||||||
|
r"""
|
||||||
|
Args:
|
||||||
|
labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
|
||||||
|
Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
|
||||||
|
config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
|
||||||
|
(masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
|
||||||
|
|
||||||
|
num_logits_to_keep (`int`, *optional*):
|
||||||
|
Calculate logits for the last `num_logits_to_keep` tokens. If `0`, calculate logits for all
|
||||||
|
`input_ids` (special case). Only last token logits are needed for generation, and calculating them only for that
|
||||||
|
token can save memory, which becomes pretty significant for long sequences or large vocabulary size.
|
||||||
|
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
|
||||||
|
Example:
|
||||||
|
|
||||||
|
```python
|
||||||
|
>>> from PIL import Image
|
||||||
|
>>> import requests
|
||||||
|
>>> from transformers import AutoProcessor, LlavaForConditionalGeneration
|
||||||
|
|
||||||
|
>>> model = LlavaForConditionalGeneration.from_pretrained("llava-hf/llava-1.5-7b-hf")
|
||||||
|
>>> processor = AutoProcessor.from_pretrained("llava-hf/llava-1.5-7b-hf")
|
||||||
|
|
||||||
|
>>> prompt = "USER: <image>\nWhat's the content of the image? ASSISTANT:"
|
||||||
|
>>> url = "https://www.ilankelman.org/stopsigns/australia.jpg"
|
||||||
|
>>> image = Image.open(requests.get(url, stream=True).raw)
|
||||||
|
|
||||||
|
>>> inputs = processor(images=image, text=prompt, return_tensors="pt")
|
||||||
|
|
||||||
|
>>> # Generate
|
||||||
|
>>> generate_ids = model.generate(**inputs, max_new_tokens=15)
|
||||||
|
>>> processor.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
|
||||||
|
"USER: \nWhat's the content of the image? ASSISTANT: The image features a busy city street with a stop sign prominently displayed"
|
||||||
|
```"""
|
||||||
|
|
||||||
|
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
|
||||||
|
output_hidden_states = (
|
||||||
|
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
|
||||||
|
)
|
||||||
|
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||||
|
vision_feature_layer = (
|
||||||
|
vision_feature_layer if vision_feature_layer is not None else self.config.vision_feature_layer
|
||||||
|
)
|
||||||
|
vision_feature_select_strategy = (
|
||||||
|
vision_feature_select_strategy
|
||||||
|
if vision_feature_select_strategy is not None
|
||||||
|
else self.config.vision_feature_select_strategy
|
||||||
|
)
|
||||||
|
|
||||||
|
if (input_ids is None) ^ (inputs_embeds is not None):
|
||||||
|
raise ValueError("You must specify exactly one of input_ids or inputs_embeds")
|
||||||
|
|
||||||
|
if pixel_values is not None and inputs_embeds is not None:
|
||||||
|
raise ValueError(
|
||||||
|
"You cannot specify both pixel_values and inputs_embeds at the same time, and must specify either one"
|
||||||
|
)
|
||||||
|
|
||||||
|
legacy_processing = False
|
||||||
|
if inputs_embeds is None:
|
||||||
|
inputs_embeds = self.get_input_embeddings()(input_ids)
|
||||||
|
|
||||||
|
# if the number of image tokens is more than image embeddings seq length, then prob we expanded it in processing
|
||||||
|
# not very reliable, but we don't expect one to actually pass 500+ images for one prompt
|
||||||
|
# In case we're in decoding stage, legacy behavior is checked by presence of pixel values even if use_cache=True
|
||||||
|
legacy_processing = (
|
||||||
|
(input_ids == self.config.image_token_index).sum(1).max() < self.config.image_seq_length
|
||||||
|
) or (input_ids.shape[-1] == 1 and pixel_values is not None)
|
||||||
|
|
||||||
|
image_features = None
|
||||||
|
if pixel_values is not None:
|
||||||
|
image_features = self.get_image_features(
|
||||||
|
pixel_values=pixel_values,
|
||||||
|
vision_feature_layer=vision_feature_layer,
|
||||||
|
vision_feature_select_strategy=vision_feature_select_strategy,
|
||||||
|
)
|
||||||
|
|
||||||
|
if legacy_processing:
|
||||||
|
logger.warning_once(
|
||||||
|
"Expanding inputs for image tokens in LLaVa should be done in processing. "
|
||||||
|
"Please add `patch_size` and `vision_feature_select_strategy` to the model's processing config or set directly "
|
||||||
|
"with `processor.patch_size = {{patch_size}}` and processor.vision_feature_select_strategy = {{vision_feature_select_strategy}}`. "
|
||||||
|
"Using processors without these attributes in the config is deprecated and will throw an error in v4.50."
|
||||||
|
)
|
||||||
|
# prefill stage vs decoding stage (legacy behavior copied)
|
||||||
|
if input_ids.shape[1] != 1:
|
||||||
|
inputs_embeds, attention_mask, labels, position_ids = self._merge_input_ids_with_image_features(
|
||||||
|
image_features, inputs_embeds, input_ids, attention_mask, labels
|
||||||
|
)
|
||||||
|
cache_position = torch.arange(attention_mask.shape[1], device=attention_mask.device)
|
||||||
|
else:
|
||||||
|
# Retrieve the first layer to inspect the logits and mask out the hidden states
|
||||||
|
# that are set to 0
|
||||||
|
first_layer_past_key_value = past_key_values[0][0][:, :, :, 0]
|
||||||
|
|
||||||
|
# Sum all dimensions of head_dim (-2) to avoid random errors such as: https://github.com/huggingface/transformers/pull/28032#issuecomment-1863691941
|
||||||
|
batch_index, non_attended_tokens = torch.where(first_layer_past_key_value.float().sum(-2) == 0)
|
||||||
|
|
||||||
|
# Get the target length
|
||||||
|
target_length = input_ids.shape[1]
|
||||||
|
past_length = first_layer_past_key_value.shape[-1]
|
||||||
|
|
||||||
|
extended_attention_mask = torch.ones(
|
||||||
|
(attention_mask.shape[0], past_length),
|
||||||
|
dtype=attention_mask.dtype,
|
||||||
|
device=attention_mask.device,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Filter out only the tokens that can be un-attended, this can happen
|
||||||
|
# if one uses Llava + Fused modules where the cache on the
|
||||||
|
# first iteration is already big enough, or if one passes custom cache
|
||||||
|
valid_indices = non_attended_tokens < extended_attention_mask.size(-1)
|
||||||
|
new_batch_index = batch_index[valid_indices]
|
||||||
|
new_non_attended_tokens = non_attended_tokens[valid_indices]
|
||||||
|
|
||||||
|
# Zero-out the places where we don't need to attend
|
||||||
|
extended_attention_mask[new_batch_index, new_non_attended_tokens] = 0
|
||||||
|
|
||||||
|
attention_mask = torch.cat((extended_attention_mask, attention_mask[:, -target_length:]), dim=1)
|
||||||
|
position_ids = torch.sum(attention_mask, dim=1).unsqueeze(-1) - 1
|
||||||
|
cache_position = torch.arange(attention_mask.shape[1], device=attention_mask.device)[-target_length:]
|
||||||
|
|
||||||
|
# TODO: @raushan retain only the new behavior after v4.47
|
||||||
|
elif image_features is not None:
|
||||||
|
n_image_tokens = (input_ids == self.config.image_token_index).sum().item()
|
||||||
|
n_image_features = image_features.shape[0] * image_features.shape[1]
|
||||||
|
|
||||||
|
if n_image_tokens != n_image_features:
|
||||||
|
raise ValueError(
|
||||||
|
f"Image features and image tokens do not match: tokens: {n_image_tokens}, features {n_image_features}"
|
||||||
|
)
|
||||||
|
special_image_mask = (
|
||||||
|
(input_ids == self.config.image_token_index)
|
||||||
|
.unsqueeze(-1)
|
||||||
|
.expand_as(inputs_embeds)
|
||||||
|
.to(inputs_embeds.device)
|
||||||
|
)
|
||||||
|
image_features = image_features.to(inputs_embeds.device, inputs_embeds.dtype)
|
||||||
|
inputs_embeds = inputs_embeds.masked_scatter(special_image_mask, image_features)
|
||||||
|
|
||||||
|
outputs = self.language_model(
|
||||||
|
attention_mask=attention_mask,
|
||||||
|
position_ids=position_ids,
|
||||||
|
past_key_values=past_key_values,
|
||||||
|
inputs_embeds=inputs_embeds,
|
||||||
|
use_cache=use_cache,
|
||||||
|
output_attentions=output_attentions,
|
||||||
|
output_hidden_states=output_hidden_states,
|
||||||
|
return_dict=return_dict,
|
||||||
|
cache_position=cache_position,
|
||||||
|
num_logits_to_keep=num_logits_to_keep,
|
||||||
|
)
|
||||||
|
|
||||||
|
logits = outputs[0]
|
||||||
|
|
||||||
|
loss = None
|
||||||
|
if labels is not None:
|
||||||
|
# Shift so that tokens < n predict n
|
||||||
|
if attention_mask is not None:
|
||||||
|
# we use the input attention mask to shift the logits and labels, because it is 2D.
|
||||||
|
# we also crop attn mask in case it is longer, which happens in PrefixTuning with peft
|
||||||
|
shift_attention_mask = attention_mask[:, -(logits.shape[1] - 1) :].to(logits.device)
|
||||||
|
shift_logits = logits[..., :-1, :][shift_attention_mask.to(logits.device) != 0].contiguous()
|
||||||
|
shift_labels = labels[..., 1:][shift_attention_mask.to(labels.device) != 0].contiguous()
|
||||||
|
else:
|
||||||
|
shift_logits = logits[..., :-1, :].contiguous()
|
||||||
|
shift_labels = labels[..., 1:].contiguous()
|
||||||
|
# Flatten the tokens
|
||||||
|
loss_fct = nn.CrossEntropyLoss()
|
||||||
|
loss = loss_fct(
|
||||||
|
shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1).to(shift_logits.device)
|
||||||
|
)
|
||||||
|
|
||||||
|
if not return_dict:
|
||||||
|
output = (logits,) + outputs[1:]
|
||||||
|
return (loss,) + output if loss is not None else output
|
||||||
|
|
||||||
|
return LlavaCausalLMOutputWithPast(
|
||||||
|
loss=loss,
|
||||||
|
logits=logits,
|
||||||
|
past_key_values=outputs.past_key_values,
|
||||||
|
hidden_states=outputs.hidden_states,
|
||||||
|
attentions=outputs.attentions,
|
||||||
|
image_hidden_states=image_features if pixel_values is not None else None,
|
||||||
|
)
|
||||||
|
|
||||||
|
def prepare_inputs_for_generation(
|
||||||
|
self,
|
||||||
|
input_ids,
|
||||||
|
past_key_values=None,
|
||||||
|
inputs_embeds=None,
|
||||||
|
pixel_values=None,
|
||||||
|
attention_mask=None,
|
||||||
|
cache_position=None,
|
||||||
|
num_logits_to_keep=None,
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
# Overwritten -- in specific circumstances we don't want to forward image inputs to the model
|
||||||
|
|
||||||
|
model_inputs = self.language_model.prepare_inputs_for_generation(
|
||||||
|
input_ids,
|
||||||
|
past_key_values=past_key_values,
|
||||||
|
inputs_embeds=inputs_embeds,
|
||||||
|
attention_mask=attention_mask,
|
||||||
|
cache_position=cache_position,
|
||||||
|
num_logits_to_keep=num_logits_to_keep,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
if cache_position[0] == 0:
|
||||||
|
# If we're in cached decoding stage, pixel values should be None because input ids do not contain special image token anymore
|
||||||
|
# Otherwise we need pixel values to be passed to model
|
||||||
|
model_inputs["pixel_values"] = pixel_values
|
||||||
|
|
||||||
|
return model_inputs
|
||||||
@@ -5,7 +5,7 @@ import gc
|
|||||||
from .utils import log, print_memory
|
from .utils import log, print_memory
|
||||||
from diffusers.video_processor import VideoProcessor
|
from diffusers.video_processor import VideoProcessor
|
||||||
from typing import List, Dict, Any, Tuple
|
from typing import List, Dict, Any, Tuple
|
||||||
|
import numpy as np
|
||||||
from .hyvideo.constants import PROMPT_TEMPLATE
|
from .hyvideo.constants import PROMPT_TEMPLATE
|
||||||
from .hyvideo.text_encoder import TextEncoder
|
from .hyvideo.text_encoder import TextEncoder
|
||||||
from .hyvideo.utils.data_utils import align_to
|
from .hyvideo.utils.data_utils import align_to
|
||||||
@@ -35,6 +35,7 @@ folder_paths.add_model_folder_path("hyvid_embeds", os.path.join(folder_paths.get
|
|||||||
|
|
||||||
import comfy.model_management as mm
|
import comfy.model_management as mm
|
||||||
from comfy.utils import load_torch_file, save_torch_file
|
from comfy.utils import load_torch_file, save_torch_file
|
||||||
|
from comfy.clip_vision import clip_preprocess
|
||||||
import comfy.model_base
|
import comfy.model_base
|
||||||
import comfy.latent_formats
|
import comfy.latent_formats
|
||||||
|
|
||||||
@@ -315,6 +316,7 @@ class HyVideoModelLoader:
|
|||||||
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
|
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
|
||||||
|
|
||||||
in_channels = sd["img_in.proj.weight"].shape[1]
|
in_channels = sd["img_in.proj.weight"].shape[1]
|
||||||
|
guidance_embed = sd.get("guidance_in.mlp.0.weight", False) is not False
|
||||||
|
|
||||||
out_channels = 16
|
out_channels = 16
|
||||||
factor_kwargs = {"device": transformer_load_device, "dtype": base_dtype}
|
factor_kwargs = {"device": transformer_load_device, "dtype": base_dtype}
|
||||||
@@ -325,7 +327,7 @@ class HyVideoModelLoader:
|
|||||||
"hidden_size": 3072,
|
"hidden_size": 3072,
|
||||||
"heads_num": 24,
|
"heads_num": 24,
|
||||||
"mlp_width_ratio": 4,
|
"mlp_width_ratio": 4,
|
||||||
"guidance_embed": True,
|
"guidance_embed": guidance_embed,
|
||||||
}
|
}
|
||||||
with init_empty_weights():
|
with init_empty_weights():
|
||||||
transformer = HYVideoDiffusionTransformer(
|
transformer = HYVideoDiffusionTransformer(
|
||||||
@@ -559,6 +561,11 @@ class HyVideoVAELoader:
|
|||||||
vae_config = json.load(f)
|
vae_config = json.load(f)
|
||||||
model_path = folder_paths.get_full_path("vae", model_name)
|
model_path = folder_paths.get_full_path("vae", model_name)
|
||||||
vae_sd = load_torch_file(model_path, safe_load=True)
|
vae_sd = load_torch_file(model_path, safe_load=True)
|
||||||
|
|
||||||
|
if not "decoder.conv_norm_out.weight" in vae_sd:
|
||||||
|
raise ValueError("""
|
||||||
|
Incompatible VAE model selected, the HunyuanVideoWrapper's VAE nodes require using the original VAE model: 'https://huggingface.co/Kijai/HunyuanVideo_comfy/blob/main/hunyuan_video_vae_bf16.safetensors'
|
||||||
|
Alternatively you can also use the ComfyUI native VAELoader and the usual VAE nodes with the wrapper.""")
|
||||||
|
|
||||||
vae = AutoencoderKLCausal3D.from_config(vae_config)
|
vae = AutoencoderKLCausal3D.from_config(vae_config)
|
||||||
vae.load_state_dict(vae_sd)
|
vae.load_state_dict(vae_sd)
|
||||||
@@ -784,7 +791,7 @@ class HyVideoTextEncode:
|
|||||||
FUNCTION = "process"
|
FUNCTION = "process"
|
||||||
CATEGORY = "HunyuanVideoWrapper"
|
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, 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):
|
||||||
if clip_text_override is not None and len(clip_text_override) == 0:
|
if clip_text_override is not None and len(clip_text_override) == 0:
|
||||||
clip_text_override = None
|
clip_text_override = None
|
||||||
device = mm.text_encoder_device()
|
device = mm.text_encoder_device()
|
||||||
@@ -810,6 +817,10 @@ class HyVideoTextEncode:
|
|||||||
prompt_template_dict = PROMPT_TEMPLATE["dit-llm-encode-video"]
|
prompt_template_dict = PROMPT_TEMPLATE["dit-llm-encode-video"]
|
||||||
elif prompt_template == "image":
|
elif prompt_template == "image":
|
||||||
prompt_template_dict = PROMPT_TEMPLATE["dit-llm-encode"]
|
prompt_template_dict = PROMPT_TEMPLATE["dit-llm-encode"]
|
||||||
|
elif prompt_template == "I2V_video":
|
||||||
|
prompt_template_dict = PROMPT_TEMPLATE["dit-llm-encode-video-i2v"]
|
||||||
|
elif prompt_template == "I2V_image":
|
||||||
|
prompt_template_dict = PROMPT_TEMPLATE["dit-llm-encode-i2v"]
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Invalid prompt_template: {prompt_template_dict}")
|
raise ValueError(f"Invalid prompt_template: {prompt_template_dict}")
|
||||||
assert (
|
assert (
|
||||||
@@ -827,16 +838,32 @@ class HyVideoTextEncode:
|
|||||||
batch_size = 1
|
batch_size = 1
|
||||||
num_videos_per_prompt = 1
|
num_videos_per_prompt = 1
|
||||||
|
|
||||||
text_inputs = text_encoder.text2tokens(prompt,
|
if image is not None:
|
||||||
prompt_template=prompt_template_dict,
|
#pixel_values = clip_preprocess(image.to(device), size=336, crop=True).float() * 255
|
||||||
image1=image1,
|
#print(pixel_values.min(), pixel_values.max())
|
||||||
image2=image2,
|
|
||||||
clip_text_override=clip_text_override)
|
text_inputs = text_encoder.text2tokens(prompt,
|
||||||
prompt_outputs = text_encoder.encode(text_inputs,
|
prompt_template=prompt_template_dict)
|
||||||
prompt_template=prompt_template_dict,
|
prompt_outputs = text_encoder.encode(text_inputs,
|
||||||
image_token_selection_expr=image_token_selection_expr,
|
prompt_template=prompt_template_dict,
|
||||||
device=device
|
image_token_selection_expr=image_token_selection_expr,
|
||||||
)
|
semantic_images = [image.squeeze(0) * 255] if text_encoder.text_encoder_type == "vlm" else None,
|
||||||
|
device=device,
|
||||||
|
data_type=prompt_template
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
text_inputs = text_encoder.text2tokens(prompt,
|
||||||
|
prompt_template=prompt_template_dict,
|
||||||
|
image1=image1,
|
||||||
|
image2=image2,
|
||||||
|
clip_text_override=clip_text_override)
|
||||||
|
prompt_outputs = text_encoder.encode(text_inputs,
|
||||||
|
prompt_template=prompt_template_dict,
|
||||||
|
image_token_selection_expr=image_token_selection_expr,
|
||||||
|
semantic_images = None,
|
||||||
|
device=device
|
||||||
|
)
|
||||||
|
|
||||||
prompt_embeds = prompt_outputs.hidden_state
|
prompt_embeds = prompt_outputs.hidden_state
|
||||||
|
|
||||||
attention_mask = prompt_outputs.attention_mask
|
attention_mask = prompt_outputs.attention_mask
|
||||||
@@ -984,6 +1011,27 @@ class HyVideoTextImageEncode(HyVideoTextEncode):
|
|||||||
FUNCTION = "process"
|
FUNCTION = "process"
|
||||||
CATEGORY = "HunyuanVideoWrapper"
|
CATEGORY = "HunyuanVideoWrapper"
|
||||||
|
|
||||||
|
class HyVideoI2VEncode(HyVideoTextEncode):
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"text_encoders": ("HYVIDTEXTENCODER",),
|
||||||
|
"prompt": ("STRING", {"default": "", "multiline": True} ),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"force_offload": ("BOOLEAN", {"default": True}),
|
||||||
|
"prompt_template": (["I2V_video", "I2V_image", "disabled"], {"default": "I2V_video", "tooltip": "Use the default prompt templates for the llm text encoder"}),
|
||||||
|
"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", ),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("HYVIDEMBEDS", )
|
||||||
|
RETURN_NAMES = ("hyvid_embeds",)
|
||||||
|
FUNCTION = "process"
|
||||||
|
CATEGORY = "HunyuanVideoWrapper"
|
||||||
|
|
||||||
# region CFG
|
# region CFG
|
||||||
class HyVideoCFG:
|
class HyVideoCFG:
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -1150,6 +1198,7 @@ class HyVideoSampler:
|
|||||||
{
|
{
|
||||||
"default": 'FlowMatchDiscreteScheduler'
|
"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"}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1159,7 +1208,8 @@ class HyVideoSampler:
|
|||||||
CATEGORY = "HunyuanVideoWrapper"
|
CATEGORY = "HunyuanVideoWrapper"
|
||||||
|
|
||||||
def process(self, model, hyvid_embeds, flow_shift, steps, embedded_guidance_scale, seed, width, height, num_frames,
|
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):
|
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):
|
||||||
model = model.model
|
model = model.model
|
||||||
|
|
||||||
device = mm.get_torch_device()
|
device = mm.get_torch_device()
|
||||||
@@ -1301,6 +1351,7 @@ class HyVideoSampler:
|
|||||||
feta_args=feta_args,
|
feta_args=feta_args,
|
||||||
leapfusion_img2vid = leapfusion_img2vid,
|
leapfusion_img2vid = leapfusion_img2vid,
|
||||||
image_cond_latents = image_cond_latents["samples"] * VAE_SCALING_FACTOR if image_cond_latents is not None else None,
|
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
|
||||||
)
|
)
|
||||||
|
|
||||||
print_memory(device)
|
print_memory(device)
|
||||||
@@ -1309,6 +1360,9 @@ class HyVideoSampler:
|
|||||||
except:
|
except:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
if teacache_args is not None:
|
||||||
|
log.info(f"TeaCache skipped {transformer.teacache_skipped_steps} steps")
|
||||||
|
|
||||||
if force_offload:
|
if force_offload:
|
||||||
if model["manual_offloading"]:
|
if model["manual_offloading"]:
|
||||||
transformer.to(offload_device)
|
transformer.to(offload_device)
|
||||||
@@ -1460,6 +1514,47 @@ class HyVideoEncode:
|
|||||||
|
|
||||||
|
|
||||||
return ({"samples": latents},)
|
return ({"samples": latents},)
|
||||||
|
|
||||||
|
class HyVideoGetClosestBucketSize:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"image": ("IMAGE",),
|
||||||
|
"base_size": (["360", "540", "720", "960"], {"default": "540", "tooltip": "Resizes the input image to closest original training bucket size"}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("INT","INT",)
|
||||||
|
RETURN_NAMES = ("width", "height",)
|
||||||
|
FUNCTION = "encode"
|
||||||
|
CATEGORY = "HunyuanVideoWrapper"
|
||||||
|
|
||||||
|
def encode(self, image, base_size):
|
||||||
|
B, H, W, C = image.shape
|
||||||
|
crop_size_list = self.generate_crop_size_list(int(base_size), 32)
|
||||||
|
aspect_ratios = np.array([round(float(h)/float(w), 5) for h, w in crop_size_list])
|
||||||
|
closest_size, closest_ratio = self.get_closest_ratio(H, W, aspect_ratios, crop_size_list)
|
||||||
|
log.info(f"ImageResizeToBucket: Closest size = {closest_size}, closest ratio = {closest_ratio}")
|
||||||
|
return (closest_size[1], closest_size[0],)
|
||||||
|
|
||||||
|
def generate_crop_size_list(self, base_size=256, patch_size=16, max_ratio=4.0):
|
||||||
|
num_patches = round((base_size / patch_size) ** 2)
|
||||||
|
assert max_ratio >= 1.
|
||||||
|
crop_size_list = []
|
||||||
|
wp, hp = num_patches, 1
|
||||||
|
while wp > 0:
|
||||||
|
if max(wp, hp) / min(wp, hp) <= max_ratio:
|
||||||
|
crop_size_list.append((wp * patch_size, hp * patch_size))
|
||||||
|
if (hp + 1) * wp <= num_patches:
|
||||||
|
hp += 1
|
||||||
|
else:
|
||||||
|
wp -= 1
|
||||||
|
return crop_size_list
|
||||||
|
def get_closest_ratio(self, height: float, width: float, ratios: list, buckets: list):
|
||||||
|
aspect_ratio = float(height)/float(width)
|
||||||
|
closest_ratio_id = np.abs(ratios - aspect_ratio).argmin()
|
||||||
|
closest_ratio = min(ratios, key=lambda ratio: abs(float(ratio) - aspect_ratio))
|
||||||
|
return buckets[closest_ratio_id], float(closest_ratio)
|
||||||
|
|
||||||
class HyVideoLatentPreview:
|
class HyVideoLatentPreview:
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -1558,6 +1653,8 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
"HyVideoContextOptions": HyVideoContextOptions,
|
"HyVideoContextOptions": HyVideoContextOptions,
|
||||||
"HyVideoEnhanceAVideo": HyVideoEnhanceAVideo,
|
"HyVideoEnhanceAVideo": HyVideoEnhanceAVideo,
|
||||||
"HyVideoTeaCache": HyVideoTeaCache,
|
"HyVideoTeaCache": HyVideoTeaCache,
|
||||||
|
"HyVideoGetClosestBucketSize": HyVideoGetClosestBucketSize,
|
||||||
|
"HyVideoI2VEncode": HyVideoI2VEncode
|
||||||
}
|
}
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
"HyVideoSampler": "HunyuanVideo Sampler",
|
"HyVideoSampler": "HunyuanVideo Sampler",
|
||||||
@@ -1581,4 +1678,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
|||||||
"HyVideoContextOptions": "HunyuanVideo Context Options",
|
"HyVideoContextOptions": "HunyuanVideo Context Options",
|
||||||
"HyVideoEnhanceAVideo": "HunyuanVideo Enhance A Video",
|
"HyVideoEnhanceAVideo": "HunyuanVideo Enhance A Video",
|
||||||
"HyVideoTeaCache": "HunyuanVideo TeaCache",
|
"HyVideoTeaCache": "HunyuanVideo TeaCache",
|
||||||
|
"HyVideoGetClosestBucketSize": "HunyuanVideo Get Closest Bucket Size",
|
||||||
|
"HyVideoI2VEncode": "HyVideo I2V Encode"
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -17,8 +17,10 @@ def print_memory(device):
|
|||||||
memory = torch.cuda.memory_allocated(device) / 1024**3
|
memory = torch.cuda.memory_allocated(device) / 1024**3
|
||||||
max_memory = torch.cuda.max_memory_allocated(device) / 1024**3
|
max_memory = torch.cuda.max_memory_allocated(device) / 1024**3
|
||||||
max_reserved = torch.cuda.max_memory_reserved(device) / 1024**3
|
max_reserved = torch.cuda.max_memory_reserved(device) / 1024**3
|
||||||
|
log.info(f"-------------------------------")
|
||||||
log.info(f"Allocated memory: {memory=:.3f} GB")
|
log.info(f"Allocated memory: {memory=:.3f} GB")
|
||||||
log.info(f"Max allocated memory: {max_memory=:.3f} GB")
|
log.info(f"Max allocated memory: {max_memory=:.3f} GB")
|
||||||
log.info(f"Max reserved memory: {max_reserved=:.3f} GB")
|
log.info(f"Max reserved memory: {max_reserved=:.3f} GB")
|
||||||
|
log.info(f"-------------------------------")
|
||||||
#memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False)
|
#memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False)
|
||||||
#log.info(f"Memory Summary:\n{memory_summary}")
|
#log.info(f"Memory Summary:\n{memory_summary}")
|
||||||
Reference in New Issue
Block a user