Add nodes for text embed saving/loading

This way you can encode prompts in separate workflow to avoid ever having to load the text encoders for sampling if you're low on RAM
This commit is contained in:
kijai
2024-12-13 13:07:34 +02:00
parent 764dad6975
commit 706c3a07ef
2 changed files with 82 additions and 4 deletions
+3 -3
View File
@@ -665,7 +665,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
out = {}
img = x
txt = text_states
txt = text_states.to(x.device)
_, _, ot, oh, ow = x.shape
tt, th, tw = (
ot // self.patch_size[0],
@@ -678,7 +678,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
# text modulation
if text_states_2 is not None:
vec = vec + self.vector_in(text_states_2)
vec = vec + self.vector_in(text_states_2.to(x.device))
# guidance modulation
if self.guidance_embed:
@@ -700,7 +700,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
if self.text_projection == "linear":
txt = self.txt_in(txt)
elif self.text_projection == "single_refiner":
txt = self.txt_in(txt, t, text_mask if self.use_attention_mask else None)
txt = self.txt_in(txt, t, text_mask.to(x.device) if self.use_attention_mask else None)
else:
raise NotImplementedError(
f"Unsupported text_projection: {self.text_projection}"
+79 -1
View File
@@ -17,8 +17,10 @@ from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
import folder_paths
folder_paths.add_model_folder_path("hyvid_embeds", os.path.join(folder_paths.get_output_directory(), "hyvid_embeds"))
import comfy.model_management as mm
from comfy.utils import load_torch_file
from comfy.utils import load_torch_file, save_torch_file
import comfy.model_base
import comfy.latent_formats
@@ -873,6 +875,78 @@ class HyVideoTextEncode:
}
return (prompt_embeds_dict,)
class HyVideoTextEmbedsSave:
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
@classmethod
def INPUT_TYPES(s):
return {"required": {
"hyvid_embeds": ("HYVIDEMBEDS",),
"filename_prefix": ("STRING", {"default": "hyvid_embeds/hyvid_embed"}),
},
"hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"},
}
RETURN_TYPES = ("STRING", )
RETURN_NAMES = ("output_path",)
FUNCTION = "save"
CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "Save the text embeds"
def save(self, hyvid_embeds, prompt, filename_prefix, extra_pnginfo=None):
from comfy.cli_args import args
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir)
file = f"{filename}_{counter:05}_.safetensors"
file = os.path.join(full_output_folder, file)
tensors_to_save = {}
for key, value in hyvid_embeds.items():
if value is not None:
tensors_to_save[key] = value
prompt_info = ""
if prompt is not None:
prompt_info = json.dumps(prompt)
metadata = None
if not args.disable_metadata:
metadata = {"prompt": prompt_info}
if extra_pnginfo is not None:
for x in extra_pnginfo:
metadata[x] = json.dumps(extra_pnginfo[x])
save_torch_file(tensors_to_save, file, metadata=metadata)
return (file,)
class HyVideoTextEmbedsLoad:
@classmethod
def INPUT_TYPES(s):
return {"required": {"embeds": (folder_paths.get_filename_list("hyvid_embeds"), {"tooltip": "The saved embeds to load from output/hyvid_embeds."})}}
RETURN_TYPES = ("HYVIDEMBEDS", )
RETURN_NAMES = ("hyvid_embeds",)
FUNCTION = "load"
CATEGORY = "HunyuanVideoWrapper"
DESCTIPTION = "Load the saved text embeds"
def load(self, embeds):
embed_path = folder_paths.get_full_path_or_raise("hyvid_embeds", embeds)
loaded_tensors = load_torch_file(embed_path)
# Reconstruct original dictionary with None for missing keys
prompt_embeds_dict = {
"prompt_embeds": loaded_tensors.get("prompt_embeds", None),
"negative_prompt_embeds": loaded_tensors.get("negative_prompt_embeds", None),
"attention_mask": loaded_tensors.get("attention_mask", None),
"negative_attention_mask": loaded_tensors.get("negative_attention_mask", None),
"prompt_embeds_2": loaded_tensors.get("prompt_embeds_2", None),
"negative_prompt_embeds_2": loaded_tensors.get("negative_prompt_embeds_2", None),
"attention_mask_2": loaded_tensors.get("attention_mask_2", None),
"negative_attention_mask_2": loaded_tensors.get("negative_attention_mask_2", None)
}
return (prompt_embeds_dict,)
#region Sampler
class HyVideoSampler:
@@ -1236,6 +1310,8 @@ NODE_CLASS_MAPPINGS = {
"HyVideoLatentPreview": HyVideoLatentPreview,
"HyVideoLoraSelect": HyVideoLoraSelect,
"HyVideoLoraBlockEdit": HyVideoLoraBlockEdit,
"HyVideoTextEmbedsSave": HyVideoTextEmbedsSave,
"HyVideoTextEmbedsLoad": HyVideoTextEmbedsLoad,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"HyVideoSampler": "HunyuanVideo Sampler",
@@ -1252,4 +1328,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"HyVideoLatentPreview": "HunyuanVideo Latent Preview",
"HyVideoLoraSelect": "HunyuanVideo Lora Select",
"HyVideoLoraBlockEdit": "HunyuanVideo Lora Block Edit",
"HyVideoTextEmbedsSave": "HunyuanVideo TextEmbeds Save",
"HyVideoTextEmbedsLoad": "HunyuanVideo TextEmbeds Load",
}