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:
@@ -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}"
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user