Support official I2V
This commit is contained in:
@@ -14,6 +14,12 @@ __all__ = [
|
||||
"TEXT_PROJECTION",
|
||||
"DATA_TYPE",
|
||||
"NEGATIVE_PROMPT",
|
||||
"NEGATIVE_PROMPT_I2V",
|
||||
"FLOW_PATH_TYPE",
|
||||
"FLOW_PREDICT_TYPE",
|
||||
"FLOW_LOSS_WEIGHT",
|
||||
"FLOW_SNR_TYPE",
|
||||
"FLOW_SOLVER",
|
||||
]
|
||||
|
||||
PRECISION_TO_TYPE = {
|
||||
@@ -46,7 +52,26 @@ PROMPT_TEMPLATE_ENCODE_VIDEO = (
|
||||
"<|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_I2V = "deformation, a poor composition and deformed video, bad teeth, bad eyes, bad limbs"
|
||||
|
||||
PROMPT_TEMPLATE = {
|
||||
"dit-llm-encode": {
|
||||
@@ -57,6 +82,22 @@ PROMPT_TEMPLATE = {
|
||||
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
|
||||
"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 ======================
|
||||
@@ -77,15 +118,48 @@ VAE_PATH = {"884-16c-hy": f"{MODEL_BASE}/hunyuan-video-t2v-720p/vae"}
|
||||
TEXT_ENCODER_PATH = {
|
||||
"clipL": f"{MODEL_BASE}/text_encoder_2",
|
||||
"llm": f"{MODEL_BASE}/text_encoder",
|
||||
"llm-i2v": f"{MODEL_BASE}/text_encoder_i2v",
|
||||
}
|
||||
|
||||
# Tokenizer
|
||||
TOKENIZER_PATH = {
|
||||
"clipL": f"{MODEL_BASE}/text_encoder_2",
|
||||
"llm": f"{MODEL_BASE}/text_encoder",
|
||||
"llm-i2v": f"{MODEL_BASE}/text_encoder_i2v",
|
||||
}
|
||||
|
||||
TEXT_PROJECTION = {
|
||||
"linear", # Default, an nn.Linear() layer
|
||||
"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 ....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
|
||||
ndim = 5 - 2
|
||||
rope_theta = 225
|
||||
@@ -85,6 +85,8 @@ def get_rotary_pos_embed(transformer, latent_video_length, height, width):
|
||||
theta=rope_theta,
|
||||
use_real=True,
|
||||
theta_rescale_factor=1,
|
||||
num_frames=latent_video_length,
|
||||
k=k,
|
||||
)
|
||||
return freqs_cos, freqs_sin
|
||||
def retrieve_timesteps(
|
||||
@@ -233,8 +235,12 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
freenoise=False,
|
||||
context_size=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 = (
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
@@ -283,11 +289,16 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
#print("place_idx:", place_idx, "delta:", delta, "list_idx:", list_idx)
|
||||
noise[:, :, place_idx:place_idx + delta, :, :] = noise[:, :, list_idx, :, :]
|
||||
|
||||
if latents is None:
|
||||
latents = noise
|
||||
elif leapfusion_img2vid:
|
||||
noise[:, :, [0,], :, :] = latents[:, :, [0,], :, :].to(noise)
|
||||
latents = noise.to(device)
|
||||
if i2v_mask is not None:
|
||||
print("i2v_mask shape:", i2v_mask.shape)
|
||||
if image_cond_latents.shape[2] == 1:
|
||||
image_cond_latents = image_cond_latents.repeat(1, 1, video_length, 1, 1)
|
||||
t = torch.tensor([0.999]).to(device=device)
|
||||
latents = noise * t + image_cond_latents * (1 - t)
|
||||
latents = 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:
|
||||
latents = latents.to(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,
|
||||
leapfusion_img2vid: Optional[bool] = False,
|
||||
image_cond_latents: Optional[torch.Tensor] = None,
|
||||
riflex_freq_index: Optional[int] = None,
|
||||
**kwargs,
|
||||
):
|
||||
r"""
|
||||
@@ -574,16 +586,25 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
else:
|
||||
disable_enhance()
|
||||
|
||||
i2v_mask = None
|
||||
image_latents = None
|
||||
if image_cond_latents is not None:
|
||||
padding_shape = (
|
||||
batch_size,
|
||||
16,
|
||||
latent_video_length - 1,
|
||||
int(height) // 8,
|
||||
int(width) // 8,
|
||||
# Expand to video length and zero-pad remaining frames
|
||||
image_latents = torch.zeros(
|
||||
(batch_size, 16, latent_video_length, height//8, width//8),
|
||||
device=device,
|
||||
dtype=self.base_dtype
|
||||
)
|
||||
latent_padding = torch.zeros(padding_shape, device=device, dtype=self.base_dtype)
|
||||
image_latents = torch.cat([image_cond_latents, latent_padding], dim=2)
|
||||
image_latents[:, :, 0:1, ...] = image_cond_latents
|
||||
|
||||
# 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_latents shape:", image_latents.shape)
|
||||
|
||||
@@ -611,7 +632,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
else:
|
||||
# rotary embeddings
|
||||
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:
|
||||
freqs_cos = freqs_cos.to(self.base_dtype).to(device)
|
||||
@@ -641,7 +662,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
freenoise=freenoise,
|
||||
context_size=context_frames,
|
||||
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
|
||||
@@ -721,18 +744,24 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
latent_image_input = (
|
||||
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)
|
||||
|
||||
if cfg_enabled:
|
||||
guidance_expand = (
|
||||
torch.tensor([embedded_guidance_scale] * latents.shape[0] * 2, dtype=self.base_dtype, device=device)
|
||||
* 1000.0
|
||||
)
|
||||
if self.transformer.guidance_embed:
|
||||
if cfg_enabled:
|
||||
guidance_expand = (
|
||||
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:
|
||||
guidance_expand = (
|
||||
torch.tensor([embedded_guidance_scale] * latents.shape[0], dtype=self.base_dtype, device=device)
|
||||
* 1000.0
|
||||
)
|
||||
guidance_expand = None
|
||||
|
||||
if use_context_schedule:
|
||||
counter = torch.zeros_like(latent_model_input)
|
||||
@@ -745,7 +774,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
#print("partial_latent_model_input", partial_latent_model_input.shape)
|
||||
with torch.autocast(
|
||||
device_type="cuda", dtype=self.base_dtype, enabled=True):
|
||||
noise_pred[:, :, c, :, :] += self.transformer(
|
||||
noise_pred_context = self.transformer(
|
||||
partial_latent_model_input,
|
||||
t_expand,
|
||||
text_states=input_prompt_embeds,
|
||||
@@ -758,8 +787,19 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
stg_mode=stg_mode,
|
||||
return_dict=True,
|
||||
)["x"]
|
||||
|
||||
counter[:, :, c, :, :] += 1
|
||||
window_mask = torch.ones_like(noise_pred_context)
|
||||
# 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 /= counter
|
||||
else:
|
||||
@@ -873,4 +913,6 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
|
||||
if leapfusion_img2vid:
|
||||
latents = latents[:, :, 1:, :, :]
|
||||
if i2v_mask is not None:
|
||||
latents = latents[:, :, 4:, :, :]
|
||||
return latents
|
||||
@@ -688,6 +688,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
self.enable_teacache = False
|
||||
self.cnt = 0
|
||||
self.num_steps = 0
|
||||
self.teacache_skipped_steps = 0
|
||||
self.rel_l1_thresh = 0.15
|
||||
self.accumulated_rel_l1_distance = 0
|
||||
self.previous_modulated_input = None
|
||||
@@ -1026,6 +1027,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
self.cnt = 0
|
||||
|
||||
if not should_calc and self.previous_residual is not None:
|
||||
self.teacache_skipped_steps += 1
|
||||
# Verify tensor dimensions match before adding
|
||||
if img.shape == self.previous_residual.shape:
|
||||
img = img + self.previous_residual
|
||||
|
||||
@@ -113,6 +113,8 @@ def get_nd_rotary_pos_embed(
|
||||
use_real=False,
|
||||
theta_rescale_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.
|
||||
@@ -163,6 +165,8 @@ def get_nd_rotary_pos_embed(
|
||||
use_real=use_real,
|
||||
theta_rescale_factor=theta_rescale_factor[i],
|
||||
interpolation_factor=interpolation_factor[i],
|
||||
L_test=num_frames,
|
||||
k=k,
|
||||
) # 2 x [WHD, rope_dim_list[i]]
|
||||
embs.append(emb)
|
||||
|
||||
@@ -182,6 +186,8 @@ def get_1d_rotary_pos_embed(
|
||||
use_real: bool = False,
|
||||
theta_rescale_factor: float = 1.0,
|
||||
interpolation_factor: float = 1.0,
|
||||
L_test: int = 100,
|
||||
k: int = 0,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""
|
||||
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)
|
||||
) # [D/2]
|
||||
# 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]
|
||||
if use_real:
|
||||
freqs_cos = freqs.cos().repeat_interleave(2, dim=1) # [S, D]
|
||||
|
||||
@@ -4,7 +4,8 @@ from copy import deepcopy
|
||||
|
||||
import torch
|
||||
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 ..constants import TEXT_ENCODER_PATH, TOKENIZER_PATH
|
||||
@@ -46,7 +47,7 @@ def load_text_encoder(
|
||||
text_encoder = LlavaForConditionalGeneration.from_pretrained(
|
||||
text_encoder_path,
|
||||
low_cpu_mem_usage=True,
|
||||
quantization_config=quantization_config
|
||||
quantization_config=quantization_config,
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported text encoder type: {text_encoder_type}")
|
||||
@@ -80,6 +81,10 @@ def load_tokenizer(
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
tokenizer_path, padding_side=padding_side
|
||||
)
|
||||
elif tokenizer_type == "llm-i2v":
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
tokenizer_path, padding_side=padding_side
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported tokenizer type: {tokenizer_type}")
|
||||
|
||||
@@ -121,6 +126,7 @@ class TextEncoder(nn.Module):
|
||||
tokenizer_path: Optional[str] = None,
|
||||
output_key: Optional[str] = None,
|
||||
use_attention_mask: bool = True,
|
||||
i2v_mode: bool = False,
|
||||
input_max_length: Optional[int] = None,
|
||||
hidden_state_skip_layer: Optional[int] = None,
|
||||
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:
|
||||
self.output_key = output_key or "last_hidden_state"
|
||||
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:
|
||||
raise ValueError(f"Unsupported text encoder type: {text_encoder_type}")
|
||||
|
||||
@@ -244,7 +253,7 @@ class TextEncoder(nn.Module):
|
||||
return_attention_mask=True,
|
||||
**kwargs,
|
||||
)
|
||||
if self.text_encoder_type == "vlm":
|
||||
if self.text_encoder_type == "vlm" and image1 is not None:
|
||||
raw_images = []
|
||||
if image1 is not None:
|
||||
raw_images.append(image1.squeeze(0)*255)
|
||||
@@ -277,6 +286,9 @@ class TextEncoder(nn.Module):
|
||||
return_texts=False,
|
||||
prompt_template=None,
|
||||
image_token_selection_expr="::4",
|
||||
data_type="image",
|
||||
semantic_images=None,
|
||||
image_embed_interleave=2,
|
||||
device=None,
|
||||
):
|
||||
"""
|
||||
@@ -299,78 +311,159 @@ class TextEncoder(nn.Module):
|
||||
hidden_state_skip_layer, self.hidden_state_skip_layer
|
||||
)
|
||||
do_sample = use_default(do_sample, not self.reproduce)
|
||||
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
|
||||
if semantic_images is None:
|
||||
attention_mask = (
|
||||
batch_encoding["attention_mask"].to(device) if use_attention_mask else None
|
||||
)
|
||||
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(
|
||||
self,
|
||||
@@ -390,3 +483,48 @@ class TextEncoder(nn.Module):
|
||||
hidden_state_skip_layer=hidden_state_skip_layer,
|
||||
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 diffusers.video_processor import VideoProcessor
|
||||
from typing import List, Dict, Any, Tuple
|
||||
|
||||
import numpy as np
|
||||
from .hyvideo.constants import PROMPT_TEMPLATE
|
||||
from .hyvideo.text_encoder import TextEncoder
|
||||
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
|
||||
from comfy.utils import load_torch_file, save_torch_file
|
||||
from comfy.clip_vision import clip_preprocess
|
||||
import comfy.model_base
|
||||
import comfy.latent_formats
|
||||
|
||||
@@ -315,6 +316,7 @@ class HyVideoModelLoader:
|
||||
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
|
||||
|
||||
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
|
||||
factor_kwargs = {"device": transformer_load_device, "dtype": base_dtype}
|
||||
@@ -325,7 +327,7 @@ class HyVideoModelLoader:
|
||||
"hidden_size": 3072,
|
||||
"heads_num": 24,
|
||||
"mlp_width_ratio": 4,
|
||||
"guidance_embed": True,
|
||||
"guidance_embed": guidance_embed,
|
||||
}
|
||||
with init_empty_weights():
|
||||
transformer = HYVideoDiffusionTransformer(
|
||||
@@ -559,6 +561,11 @@ class HyVideoVAELoader:
|
||||
vae_config = json.load(f)
|
||||
model_path = folder_paths.get_full_path("vae", model_name)
|
||||
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.load_state_dict(vae_sd)
|
||||
@@ -784,7 +791,7 @@ class HyVideoTextEncode:
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "HunyuanVideoWrapper"
|
||||
|
||||
def process(self, text_encoders, prompt, force_offload=True, prompt_template="video", custom_prompt_template=None, clip_l=None, image_token_selection_expr="::4", hyvid_cfg=None, 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:
|
||||
clip_text_override = None
|
||||
device = mm.text_encoder_device()
|
||||
@@ -810,6 +817,10 @@ class HyVideoTextEncode:
|
||||
prompt_template_dict = PROMPT_TEMPLATE["dit-llm-encode-video"]
|
||||
elif prompt_template == "image":
|
||||
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:
|
||||
raise ValueError(f"Invalid prompt_template: {prompt_template_dict}")
|
||||
assert (
|
||||
@@ -827,16 +838,32 @@ class HyVideoTextEncode:
|
||||
batch_size = 1
|
||||
num_videos_per_prompt = 1
|
||||
|
||||
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,
|
||||
device=device
|
||||
)
|
||||
if image is not None:
|
||||
#pixel_values = clip_preprocess(image.to(device), size=336, crop=True).float() * 255
|
||||
#print(pixel_values.min(), pixel_values.max())
|
||||
|
||||
text_inputs = text_encoder.text2tokens(prompt,
|
||||
prompt_template=prompt_template_dict)
|
||||
prompt_outputs = text_encoder.encode(text_inputs,
|
||||
prompt_template=prompt_template_dict,
|
||||
image_token_selection_expr=image_token_selection_expr,
|
||||
semantic_images = [image.squeeze(0) * 255] if text_encoder.text_encoder_type == "vlm" else None,
|
||||
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
|
||||
|
||||
attention_mask = prompt_outputs.attention_mask
|
||||
@@ -984,6 +1011,27 @@ class HyVideoTextImageEncode(HyVideoTextEncode):
|
||||
FUNCTION = "process"
|
||||
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
|
||||
class HyVideoCFG:
|
||||
@classmethod
|
||||
@@ -1150,6 +1198,7 @@ class HyVideoSampler:
|
||||
{
|
||||
"default": 'FlowMatchDiscreteScheduler'
|
||||
}),
|
||||
"riflex_freq_index": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1, "tooltip": "Frequency index for RIFLEX, disabled when 0, default 4. Allows for new frames to be generated after 129 without looping"}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1159,7 +1208,8 @@ class HyVideoSampler:
|
||||
CATEGORY = "HunyuanVideoWrapper"
|
||||
|
||||
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
|
||||
|
||||
device = mm.get_torch_device()
|
||||
@@ -1301,6 +1351,7 @@ class HyVideoSampler:
|
||||
feta_args=feta_args,
|
||||
leapfusion_img2vid = leapfusion_img2vid,
|
||||
image_cond_latents = image_cond_latents["samples"] * VAE_SCALING_FACTOR if image_cond_latents is not None else None,
|
||||
riflex_freq_index = riflex_freq_index
|
||||
)
|
||||
|
||||
print_memory(device)
|
||||
@@ -1309,6 +1360,9 @@ class HyVideoSampler:
|
||||
except:
|
||||
pass
|
||||
|
||||
if teacache_args is not None:
|
||||
log.info(f"TeaCache skipped {transformer.teacache_skipped_steps} steps")
|
||||
|
||||
if force_offload:
|
||||
if model["manual_offloading"]:
|
||||
transformer.to(offload_device)
|
||||
@@ -1460,6 +1514,47 @@ class HyVideoEncode:
|
||||
|
||||
|
||||
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:
|
||||
@classmethod
|
||||
@@ -1558,6 +1653,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"HyVideoContextOptions": HyVideoContextOptions,
|
||||
"HyVideoEnhanceAVideo": HyVideoEnhanceAVideo,
|
||||
"HyVideoTeaCache": HyVideoTeaCache,
|
||||
"HyVideoGetClosestBucketSize": HyVideoGetClosestBucketSize,
|
||||
"HyVideoI2VEncode": HyVideoI2VEncode
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"HyVideoSampler": "HunyuanVideo Sampler",
|
||||
@@ -1581,4 +1678,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"HyVideoContextOptions": "HunyuanVideo Context Options",
|
||||
"HyVideoEnhanceAVideo": "HunyuanVideo Enhance A Video",
|
||||
"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
|
||||
max_memory = torch.cuda.max_memory_allocated(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"Max allocated memory: {max_memory=:.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)
|
||||
#log.info(f"Memory Summary:\n{memory_summary}")
|
||||
Reference in New Issue
Block a user