diff --git a/hyvideo/constants.py b/hyvideo/constants.py index db4db71..6099214 100644 --- a/hyvideo/constants.py +++ b/hyvideo/constants.py @@ -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\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\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 +} \ No newline at end of file diff --git a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py index 389f05d..6cbd18c 100644 --- a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py +++ b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py @@ -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 \ No newline at end of file diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index 884c898..f7daffe 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -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 diff --git a/hyvideo/modules/posemb_layers.py b/hyvideo/modules/posemb_layers.py index 2ac4ce7..0fbbada 100644 --- a/hyvideo/modules/posemb_layers.py +++ b/hyvideo/modules/posemb_layers.py @@ -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] diff --git a/hyvideo/text_encoder/__init__.py b/hyvideo/text_encoder/__init__.py index f877593..873a739 100644 --- a/hyvideo/text_encoder/__init__.py +++ b/hyvideo/text_encoder/__init__.py @@ -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="", - 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="", + 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" +} \ No newline at end of file diff --git a/hyvideo/text_encoder/configuration_llava.py b/hyvideo/text_encoder/configuration_llava.py new file mode 100644 index 0000000..fd86c9e --- /dev/null +++ b/hyvideo/text_encoder/configuration_llava.py @@ -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) diff --git a/hyvideo/text_encoder/modeling_llava.py b/hyvideo/text_encoder/modeling_llava.py new file mode 100644 index 0000000..2d1cb09 --- /dev/null +++ b/hyvideo/text_encoder/modeling_llava.py @@ -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" "", "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: \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 diff --git a/nodes.py b/nodes.py index 049076f..0db983b 100644 --- a/nodes.py +++ b/nodes.py @@ -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" } diff --git a/utils.py b/utils.py index ac263e9..bce6e2b 100644 --- a/utils.py +++ b/utils.py @@ -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}") \ No newline at end of file