From ec004433957e6d37501a39fc89f18d9994ddbdce Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 8 Dec 2024 17:08:06 +0200 Subject: [PATCH] Allow using custom prompt_templates --- hyvideo/text_encoder/__init__.py | 58 ++-------- nodes.py | 187 ++++++++++++++++++------------- 2 files changed, 118 insertions(+), 127 deletions(-) diff --git a/hyvideo/text_encoder/__init__.py b/hyvideo/text_encoder/__init__.py index 4564024..537aa28 100644 --- a/hyvideo/text_encoder/__init__.py +++ b/hyvideo/text_encoder/__init__.py @@ -115,8 +115,6 @@ class TextEncoder(nn.Module): output_key: Optional[str] = None, use_attention_mask: bool = True, input_max_length: Optional[int] = None, - prompt_template: Optional[dict] = None, - prompt_template_video: Optional[dict] = None, hidden_state_skip_layer: Optional[int] = None, apply_final_norm: bool = False, reproduce: bool = False, @@ -137,43 +135,15 @@ class TextEncoder(nn.Module): tokenizer_path if tokenizer_path is not None else text_encoder_path ) self.use_attention_mask = use_attention_mask - if prompt_template_video is not None: - assert ( - use_attention_mask is True - ), "Attention mask is True required when training videos." + self.input_max_length = ( input_max_length if input_max_length is not None else max_length ) - self.prompt_template = prompt_template - self.prompt_template_video = prompt_template_video self.hidden_state_skip_layer = hidden_state_skip_layer self.apply_final_norm = apply_final_norm self.reproduce = reproduce self.logger = logger - self.use_template = self.prompt_template is not None - if self.use_template: - assert ( - isinstance(self.prompt_template, dict) - and "template" in self.prompt_template - ), f"`prompt_template` must be a dictionary with a key 'template', got {self.prompt_template}" - assert "{}" in str(self.prompt_template["template"]), ( - "`prompt_template['template']` must contain a placeholder `{}` for the input text, " - f"got {self.prompt_template['template']}" - ) - - self.use_video_template = self.prompt_template_video is not None - if self.use_video_template: - if self.prompt_template_video is not None: - assert ( - isinstance(self.prompt_template_video, dict) - and "template" in self.prompt_template_video - ), f"`prompt_template_video` must be a dictionary with a key 'template', got {self.prompt_template_video}" - assert "{}" in str(self.prompt_template_video["template"]), ( - "`prompt_template_video['template']` must contain a placeholder `{}` for the input text, " - f"got {self.prompt_template_video['template']}" - ) - if "t5" in text_encoder_type: self.output_key = output_key or "last_hidden_state" elif "clip" in text_encoder_type: @@ -222,7 +192,7 @@ class TextEncoder(nn.Module): else: raise TypeError(f"Unsupported template type: {type(template)}") - def text2tokens(self, text, data_type="image"): + def text2tokens(self, text, prompt_template): """ Tokenize the input text. @@ -230,22 +200,16 @@ class TextEncoder(nn.Module): text (str or list): Input text. """ tokenize_input_type = "str" - if self.use_template: - if data_type == "image": - prompt_template = self.prompt_template["template"] - elif data_type == "video": - prompt_template = self.prompt_template_video["template"] - else: - raise ValueError(f"Unsupported data type: {data_type}") + if prompt_template is not None and self.text_encoder_type == "llm": if isinstance(text, (list, tuple)): text = [ - self.apply_text_to_template(one_text, prompt_template) + self.apply_text_to_template(one_text, prompt_template["template"]) for one_text in text ] if isinstance(text[0], list): tokenize_input_type = "list" elif isinstance(text, str): - text = self.apply_text_to_template(text, prompt_template) + text = self.apply_text_to_template(text, prompt_template["template"]) if isinstance(text, list): tokenize_input_type = "list" else: @@ -284,7 +248,7 @@ class TextEncoder(nn.Module): do_sample=None, hidden_state_skip_layer=None, return_texts=False, - data_type="image", + prompt_template=None, device=None, ): """ @@ -326,13 +290,9 @@ class TextEncoder(nn.Module): last_hidden_state = outputs[self.output_key] # Remove hidden states of instruction tokens, only keep prompt tokens. - if self.use_template: - if data_type == "image": - crop_start = self.prompt_template.get("crop_start", -1) - elif data_type == "video": - crop_start = self.prompt_template_video.get("crop_start", -1) - else: - raise ValueError(f"Unsupported data type: {data_type}") + 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 = ( diff --git a/nodes.py b/nodes.py index 003ab5b..4ab6720 100644 --- a/nodes.py +++ b/nodes.py @@ -472,14 +472,6 @@ class DownloadAndLoadHyVideoTextEncoder: local_dir=base_path, local_dir_use_symlinks=False, ) - # prompt_template - prompt_template = ( - PROMPT_TEMPLATE["dit-llm-encode"] - ) - # prompt_template_video - prompt_template_video = ( - PROMPT_TEMPLATE["dit-llm-encode-video"] - ) text_encoder = TextEncoder( text_encoder_path=base_path, @@ -487,8 +479,6 @@ class DownloadAndLoadHyVideoTextEncoder: max_length=256, text_encoder_precision=precision, tokenizer_type="llm", - prompt_template=prompt_template, - prompt_template_video=prompt_template_video, hidden_state_skip_layer=hidden_state_skip_layer, apply_final_norm=apply_final_norm, logger=log, @@ -504,7 +494,28 @@ class DownloadAndLoadHyVideoTextEncoder: } return (hyvid_text_encoders,) - + +class HyVideoCustomPromptTemplate: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "custom_prompt_template": ("STRING", {"default": f"{PROMPT_TEMPLATE['dit-llm-encode-video']["template"]}", "multiline": True}), + "crop_start": ("INT", {"default": PROMPT_TEMPLATE['dit-llm-encode-video']["crop_start"], "tooltip": "To cropt the system prompt"}), + }, + } + + RETURN_TYPES = ("PROMPT_TEMPLATE", ) + RETURN_NAMES = ("hyvid_prompt_template",) + FUNCTION = "process" + CATEGORY = "HunyuanVideoWrapper" + + def process(self, custom_prompt_template, crop_start): + prompt_template_dict = { + "template": custom_prompt_template, + "crop_start": crop_start, + } + return (prompt_template_dict,) + class HyVideoTextEncode: @classmethod def INPUT_TYPES(s): @@ -515,7 +526,8 @@ class HyVideoTextEncode: }, "optional": { "force_offload": ("BOOLEAN", {"default": True}), - "prompt_template": (["video", "image", "disabled"], {"default": "video", "tooltip": "Use the default prompt templates for the llm text encoder"}), + "prompt_template": (["video", "image", "custom", "disabled"], {"default": "video", "tooltip": "Use the default prompt templates for the llm text encoder"}), + "custom_prompt_template": ("PROMPT_TEMPLATE", {"default": PROMPT_TEMPLATE["dit-llm-encode-video"], "multiline": True}), } } @@ -524,7 +536,7 @@ class HyVideoTextEncode: FUNCTION = "process" CATEGORY = "HunyuanVideoWrapper" - def process(self, text_encoders, prompt, force_offload=True, prompt_template="video"): + def process(self, text_encoders, prompt, force_offload=True, prompt_template="video", custom_prompt_template=None): device = mm.text_encoder_device() offload_device = mm.text_encoder_offload_device() @@ -532,22 +544,39 @@ class HyVideoTextEncode: text_encoder_2 = text_encoders["text_encoder_2"] negative_prompt = None - - text_encoder_1.use_template = True if prompt_template != "disabled" else False + + if prompt_template != "disabled": + if prompt_template == "custom": + prompt_template_dict = custom_prompt_template + elif prompt_template == "video": + prompt_template_dict = PROMPT_TEMPLATE["dit-llm-encode-video"] + elif prompt_template == "image": + prompt_template_dict = PROMPT_TEMPLATE["dit-llm-encode"] + else: + raise ValueError(f"Invalid prompt_template: {prompt_template_dict}") + assert ( + isinstance(prompt_template_dict, dict) + and "template" in prompt_template_dict + ), f"`prompt_template` must be a dictionary with a key 'template', got {prompt_template_dict}" + assert "{}" in str(prompt_template_dict["template"]), ( + "`prompt_template['template']` must contain a placeholder `{}` for the input text, " + f"got {prompt_template_dict['template']}" + ) + else: + prompt_template_dict = None def encode_prompt(self, prompt, negative_prompt, text_encoder): batch_size = 1 num_videos_per_prompt = 1 - do_classifier_free_guidance = False - data_type = prompt_template + do_classifier_free_guidance = False # not implemented, for now we only have cfg distilled model - text_inputs = text_encoder.text2tokens(prompt, data_type=data_type) + text_inputs = text_encoder.text2tokens(prompt, prompt_template=prompt_template_dict) - prompt_outputs = text_encoder.encode(text_inputs, data_type=data_type, device=device) + prompt_outputs = text_encoder.encode(text_inputs, prompt_template=prompt_template_dict, device=device) prompt_embeds = prompt_outputs.hidden_state attention_mask = prompt_outputs.attention_mask - print("prompt attention_mask: ", attention_mask.shape) + log.info(f"{text_encoder.text_encoder_type} prompt attention_mask shape: {attention_mask.shape}, masked tokens: {attention_mask[0].sum().item()}") if attention_mask is not None: attention_mask = attention_mask.to(device) bs_embed, seq_len = attention_mask.shape @@ -579,70 +608,70 @@ class HyVideoTextEncode: ) # get unconditional embeddings for classifier free guidance - if do_classifier_free_guidance: - uncond_tokens: List[str] - if negative_prompt is None: - uncond_tokens = [""] * batch_size - elif prompt is not None and type(prompt) is not type(negative_prompt): - raise TypeError( - f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !=" - f" {type(prompt)}." - ) - elif isinstance(negative_prompt, str): - uncond_tokens = [negative_prompt] - elif batch_size != len(negative_prompt): - raise ValueError( - f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:" - f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches" - " the batch size of `prompt`." - ) - else: - uncond_tokens = negative_prompt + # if do_classifier_free_guidance: + # uncond_tokens: List[str] + # if negative_prompt is None: + # uncond_tokens = [""] * batch_size + # elif prompt is not None and type(prompt) is not type(negative_prompt): + # raise TypeError( + # f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !=" + # f" {type(prompt)}." + # ) + # elif isinstance(negative_prompt, str): + # uncond_tokens = [negative_prompt] + # elif batch_size != len(negative_prompt): + # raise ValueError( + # f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:" + # f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches" + # " the batch size of `prompt`." + # ) + # else: + # uncond_tokens = negative_prompt - # max_length = prompt_embeds.shape[1] - uncond_input = text_encoder.text2tokens(uncond_tokens, data_type=data_type) + # # max_length = prompt_embeds.shape[1] + # uncond_input = text_encoder.text2tokens(uncond_tokens, data_type=data_type) - negative_prompt_outputs = text_encoder.encode( - uncond_input, data_type=data_type, device=device - ) - negative_prompt_embeds = negative_prompt_outputs.hidden_state + # negative_prompt_outputs = text_encoder.encode( + # uncond_input, data_type=data_type, device=device + # ) + # negative_prompt_embeds = negative_prompt_outputs.hidden_state - negative_attention_mask = negative_prompt_outputs.attention_mask - if negative_attention_mask is not None: - negative_attention_mask = negative_attention_mask.to(device) - _, seq_len = negative_attention_mask.shape - negative_attention_mask = negative_attention_mask.repeat( - 1, num_videos_per_prompt - ) - negative_attention_mask = negative_attention_mask.view( - batch_size * num_videos_per_prompt, seq_len - ) - else: - negative_prompt_embeds = None - negative_attention_mask = None + # negative_attention_mask = negative_prompt_outputs.attention_mask + # if negative_attention_mask is not None: + # negative_attention_mask = negative_attention_mask.to(device) + # _, seq_len = negative_attention_mask.shape + # negative_attention_mask = negative_attention_mask.repeat( + # 1, num_videos_per_prompt + # ) + # negative_attention_mask = negative_attention_mask.view( + # batch_size * num_videos_per_prompt, seq_len + # ) + # else: + negative_prompt_embeds = None + negative_attention_mask = None - if do_classifier_free_guidance: - # duplicate unconditional embeddings for each generation per prompt, using mps friendly method - seq_len = negative_prompt_embeds.shape[1] + # if do_classifier_free_guidance: + # # duplicate unconditional embeddings for each generation per prompt, using mps friendly method + # seq_len = negative_prompt_embeds.shape[1] - negative_prompt_embeds = negative_prompt_embeds.to( - dtype=prompt_embeds_dtype, device=device - ) + # negative_prompt_embeds = negative_prompt_embeds.to( + # dtype=prompt_embeds_dtype, device=device + # ) - if negative_prompt_embeds.ndim == 2: - negative_prompt_embeds = negative_prompt_embeds.repeat( - 1, num_videos_per_prompt - ) - negative_prompt_embeds = negative_prompt_embeds.view( - batch_size * num_videos_per_prompt, -1 - ) - else: - negative_prompt_embeds = negative_prompt_embeds.repeat( - 1, num_videos_per_prompt, 1 - ) - negative_prompt_embeds = negative_prompt_embeds.view( - batch_size * num_videos_per_prompt, seq_len, -1 - ) + # if negative_prompt_embeds.ndim == 2: + # negative_prompt_embeds = negative_prompt_embeds.repeat( + # 1, num_videos_per_prompt + # ) + # negative_prompt_embeds = negative_prompt_embeds.view( + # batch_size * num_videos_per_prompt, -1 + # ) + # else: + # negative_prompt_embeds = negative_prompt_embeds.repeat( + # 1, num_videos_per_prompt, 1 + # ) + # negative_prompt_embeds = negative_prompt_embeds.view( + # batch_size * num_videos_per_prompt, seq_len, -1 + # ) return ( prompt_embeds, @@ -995,6 +1024,7 @@ NODE_CLASS_MAPPINGS = { "HyVideoBlockSwap": HyVideoBlockSwap, "HyVideoTorchCompileSettings": HyVideoTorchCompileSettings, "HyVideoSTG": HyVideoSTG, + "HyVideoCustomPromptTemplate": HyVideoCustomPromptTemplate, } NODE_DISPLAY_NAME_MAPPINGS = { "HyVideoSampler": "HunyuanVideo Sampler", @@ -1007,4 +1037,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "HyVideoBlockSwap": "HunyuanVideo BlockSwap", "HyVideoTorchCompileSettings": "HunyuanVideo Torch Compile Settings", "HyVideoSTG": "HunyuanVideo STG", + "HyVideoCustomPromptTemplate": "HunyuanVideo Custom Prompt Template", }