diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index 980f81c..7509c17 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -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}" diff --git a/nodes.py b/nodes.py index 088f0d2..5990b90 100644 --- a/nodes.py +++ b/nodes.py @@ -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", }