diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..7b171e7 --- /dev/null +++ b/__init__.py @@ -0,0 +1,8 @@ +from .nodes import NODE_CLASS_MAPPINGS as NODES_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as NODES_DISPLAY_MAPPINGS +from .nodes_F1 import NODE_CLASS_MAPPINGS as F1_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as F1_DISPLAY_MAPPINGS + +# Combine the mappings +NODE_CLASS_MAPPINGS = {**NODES_MAPPINGS, **F1_MAPPINGS} +NODE_DISPLAY_NAME_MAPPINGS = {**NODES_DISPLAY_MAPPINGS, **F1_DISPLAY_MAPPINGS} + +__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file diff --git a/fp8_optimization.py b/fp8_optimization.py new file mode 100644 index 0000000..0688ee6 --- /dev/null +++ b/fp8_optimization.py @@ -0,0 +1,39 @@ +#based on ComfyUI's and MinusZoneAI's fp8_linear optimization + +import torch +import torch.nn as nn + +def fp8_linear_forward(cls, original_dtype, input): + weight_dtype = cls.weight.dtype + if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: + if len(input.shape) == 3: + target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn + inn = input.reshape(-1, input.shape[2]).to(target_dtype) + w = cls.weight.t() + + scale = torch.ones((1), device=input.device, dtype=torch.float32) + bias = cls.bias.to(original_dtype) if cls.bias is not None else None + + if bias is not None: + o = torch._scaled_mm(inn, w, out_dtype=original_dtype, bias=bias, scale_a=scale, scale_b=scale) + else: + o = torch._scaled_mm(inn, w, out_dtype=original_dtype, scale_a=scale, scale_b=scale) + + if isinstance(o, tuple): + o = o[0] + + return o.reshape((-1, input.shape[1], cls.weight.shape[0])) + else: + return cls.original_forward(input.to(original_dtype)) + else: + return cls.original_forward(input) + +def convert_fp8_linear(module, original_dtype, params_to_keep={}): + setattr(module, "fp8_matmul_enabled", True) + + for name, module in module.named_modules(): + if not any(keyword in name for keyword in params_to_keep): + if isinstance(module, nn.Linear): + original_forward = module.forward + setattr(module, "original_forward", original_forward) + setattr(module, "forward", lambda input, m=module: fp8_linear_forward(m, original_dtype, input)) diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..38fa86c --- /dev/null +++ b/nodes.py @@ -0,0 +1,937 @@ +import os +import torch +import math +from tqdm import tqdm +import sys +import logging +from pathlib import Path + +logger = logging.getLogger(__name__) + +from accelerate import init_empty_weights +from accelerate.utils import set_module_tensor_to_device + +import folder_paths +import comfy.model_management as mm +from comfy.utils import load_torch_file, ProgressBar, common_upscale +import comfy.model_base +import comfy.latent_formats +from comfy.cli_args import args, LatentPreviewMethod + +from .utils import log + +script_directory = os.path.dirname(os.path.abspath(__file__)) +vae_scaling_factor = 0.476986 + +from .diffusers_helper.models.hunyuan_video_packed import HunyuanVideoTransformer3DModel +from .diffusers_helper.memory import DynamicSwapInstaller, move_model_to_device_with_memory_preservation +from .diffusers_helper.pipelines.k_diffusion_hunyuan import sample_hunyuan +from .diffusers_helper.utils import crop_or_pad_yield_mask +from .diffusers_helper.bucket_tools import find_nearest_bucket + +# Import original function for fallback +from diffusers.loaders.lora_conversion_utils import _convert_hunyuan_video_lora_to_diffusers + +def patched_convert_hunyuan_video_lora(original_state_dict): + """Patched version that filters out problematic tensors before conversion""" + try: + # Make a copy of the original state dict to avoid modifying it + state_dict_copy = {} + + # Remove scalar (0-dimensional) tensors that cause problems + for key, value in original_state_dict.items(): + if isinstance(value, torch.Tensor): + if value.dim() == 0: + print(f"Skipping 0-dimensional tensor: {key}") + continue + state_dict_copy[key] = value + else: + print(f"Skipping non-tensor value: {key}") + + print(f"After filtering: {len(state_dict_copy)} valid keys") + + # Try the original conversion with the filtered state dict + try: + from diffusers.loaders.lora_conversion_utils import _convert_hunyuan_video_lora_to_diffusers + result = _convert_hunyuan_video_lora_to_diffusers(state_dict_copy) + print("Successfully converted LoRA weights") + return result + except Exception as e: + print(f"Error in standard conversion: {e}") + # Fall back to empty dict if conversion fails + print("Conversion failed, returning empty state dict") + return {} + + except Exception as e: + print(f"LoRA conversion failed: {str(e)}") + # Return empty state dict as fallback + return {} + + def remap_img_attn_qkv_(key, state_dict): + try: + weight = state_dict.pop(key) + + # Add dimension check + if weight.dim() == 0: + logger.warning(f"Invalid tensor dimensions for {key}: scalar tensor. Skipping.") + return + + if "lora_A" in key: + state_dict[key.replace("img_attn_qkv", "attn.to_q")] = weight + state_dict[key.replace("img_attn_qkv", "attn.to_k")] = weight + state_dict[key.replace("img_attn_qkv", "attn.to_v")] = weight + else: + # Ensure tensor is properly sized before chunking + if weight.dim() == 0 or weight.size(0) < 3: + logger.warning(f"Invalid tensor size for {key}: {weight.shape}. Using equal splits.") + # Create minimal placeholders + if weight.dim() > 0 and weight.size(0) > 0: + to_q = weight[:1] + to_k = weight[:1] if weight.size(0) == 1 else weight[1:2] + to_v = weight[:1] if weight.size(0) <= 2 else weight[2:3] + else: + # For zero-dim tensors, create basic ones + to_q = torch.ones(1, dtype=weight.dtype, device=weight.device) + to_k = torch.ones(1, dtype=weight.dtype, device=weight.device) + to_v = torch.ones(1, dtype=weight.dtype, device=weight.device) + else: + to_q, to_k, to_v = weight.chunk(3, dim=0) + + state_dict[key.replace("img_attn_qkv", "attn.to_q")] = to_q + state_dict[key.replace("img_attn_qkv", "attn.to_k")] = to_k + state_dict[key.replace("img_attn_qkv", "attn.to_v")] = to_v + except Exception as e: + logger.warning(f"Error processing key {key}: {str(e)}. Skipping.") + # Just skip the problematic key + + def remap_txt_attn_qkv_(key, state_dict): + try: + weight = state_dict.pop(key) + + # Add dimension check + if weight.dim() == 0: + logger.warning(f"Invalid tensor dimensions for {key}: scalar tensor. Skipping.") + return + + if "lora_A" in key: + state_dict[key.replace("txt_attn_qkv", "attn.add_q_proj")] = weight + state_dict[key.replace("txt_attn_qkv", "attn.add_k_proj")] = weight + state_dict[key.replace("txt_attn_qkv", "attn.add_v_proj")] = weight + else: + # Ensure tensor is properly sized before chunking + if weight.dim() == 0 or weight.size(0) < 3: + logger.warning(f"Invalid tensor size for {key}: {weight.shape}. Using equal splits.") + # Create minimal placeholders + if weight.dim() > 0 and weight.size(0) > 0: + to_q = weight[:1] + to_k = weight[:1] if weight.size(0) == 1 else weight[1:2] + to_v = weight[:1] if weight.size(0) <= 2 else weight[2:3] + else: + # For zero-dim tensors, create basic ones + to_q = torch.ones(1, dtype=weight.dtype, device=weight.device) + to_k = torch.ones(1, dtype=weight.dtype, device=weight.device) + to_v = torch.ones(1, dtype=weight.dtype, device=weight.device) + else: + to_q, to_k, to_v = weight.chunk(3, dim=0) + + state_dict[key.replace("txt_attn_qkv", "attn.add_q_proj")] = to_q + state_dict[key.replace("txt_attn_qkv", "attn.add_k_proj")] = to_k + state_dict[key.replace("txt_attn_qkv", "attn.add_v_proj")] = to_v + except Exception as e: + logger.warning(f"Error processing key {key}: {str(e)}. Skipping.") + # Just skip the problematic key + + def remap_txt_in_(key, state_dict): + def rename_key(key): + new_key = key.replace("individual_token_refiner.blocks", "token_refiner.refiner_blocks") + new_key = new_key.replace("adaLN_modulation.1", "norm_out.linear") + new_key = new_key.replace("txt_in", "context_embedder") + new_key = new_key.replace("t_embedder.mlp.0", "time_text_embed.timestep_embedder.linear_1") + new_key = new_key.replace("t_embedder.mlp.2", "time_text_embed.timestep_embedder.linear_2") + new_key = new_key.replace("c_embedder", "time_text_embed.text_embedder") + new_key = new_key.replace("mlp", "ff") + return new_key + + try: + if "self_attn_qkv" in key: + weight = state_dict.pop(key) + # Ensure tensor is at least 1D before chunking + if weight.dim() == 0 or weight.size(0) < 3: + logger.warning(f"Invalid tensor dimensions for {key}: {weight.shape}. Skipping.") + return + + to_q, to_k, to_v = weight.chunk(3, dim=0) + state_dict[rename_key(key.replace("self_attn_qkv", "attn.to_q"))] = to_q + state_dict[rename_key(key.replace("self_attn_qkv", "attn.to_k"))] = to_k + state_dict[rename_key(key.replace("self_attn_qkv", "attn.to_v"))] = to_v + else: + state_dict[rename_key(key)] = state_dict.pop(key) + except Exception as e: + logger.warning(f"Error processing key {key}: {str(e)}") + # Skip if we can't process this key properly + if key in state_dict: + state_dict.pop(key) + + def remap_single_transformer_blocks_(key, state_dict): + try: + hidden_size = 3072 + + if "linear1.lora_A.weight" in key or "linear1.lora_B.weight" in key: + linear1_weight = state_dict.pop(key) + if "lora_A" in key: + new_key = key.replace("single_blocks", "single_transformer_blocks") + if new_key.endswith(".linear1.lora_A.weight"): + new_key = new_key[:-len(".linear1.lora_A.weight")] + state_dict[f"{new_key}.attn.to_q.lora_A.weight"] = linear1_weight + state_dict[f"{new_key}.attn.to_k.lora_A.weight"] = linear1_weight + state_dict[f"{new_key}.attn.to_v.lora_A.weight"] = linear1_weight + state_dict[f"{new_key}.proj_mlp.lora_A.weight"] = linear1_weight + else: + # Ensure tensor size is sufficient for splitting + if linear1_weight.dim() == 0 or linear1_weight.size(0) < 3 * hidden_size: + logger.warning(f"Invalid tensor size for {key}: {linear1_weight.shape}. Skipping splitting.") + return + + split_size = (hidden_size, hidden_size, hidden_size, linear1_weight.size(0) - 3 * hidden_size) + q, k, v, mlp = torch.split(linear1_weight, split_size, dim=0) + new_key = key.replace("single_blocks", "single_transformer_blocks") + if new_key.endswith(".linear1.lora_B.weight"): + new_key = new_key[:-len(".linear1.lora_B.weight")] + state_dict[f"{new_key}.attn.to_q.lora_B.weight"] = q + state_dict[f"{new_key}.attn.to_k.lora_B.weight"] = k + state_dict[f"{new_key}.attn.to_v.lora_B.weight"] = v + state_dict[f"{new_key}.proj_mlp.lora_B.weight"] = mlp + + elif "linear1.lora_A.bias" in key or "linear1.lora_B.bias" in key: + linear1_bias = state_dict.pop(key) + if "lora_A" in key: + new_key = key.replace("single_blocks", "single_transformer_blocks") + if new_key.endswith(".linear1.lora_A.bias"): + new_key = new_key[:-len(".linear1.lora_A.bias")] + state_dict[f"{new_key}.attn.to_q.lora_A.bias"] = linear1_bias + state_dict[f"{new_key}.attn.to_k.lora_A.bias"] = linear1_bias + state_dict[f"{new_key}.attn.to_v.lora_A.bias"] = linear1_bias + state_dict[f"{new_key}.proj_mlp.lora_A.bias"] = linear1_bias + else: + # Ensure tensor size is sufficient for splitting + if linear1_bias.dim() == 0 or linear1_bias.size(0) < 3 * hidden_size: + logger.warning(f"Invalid tensor size for {key}: {linear1_bias.shape}. Skipping splitting.") + return + + split_size = (hidden_size, hidden_size, hidden_size, linear1_bias.size(0) - 3 * hidden_size) + q_bias, k_bias, v_bias, mlp_bias = torch.split(linear1_bias, split_size, dim=0) + new_key = key.replace("single_blocks", "single_transformer_blocks") + if new_key.endswith(".linear1.lora_B.bias"): + new_key = new_key[:-len(".linear1.lora_B.bias")] + state_dict[f"{new_key}.attn.to_q.lora_B.bias"] = q_bias + state_dict[f"{new_key}.attn.to_k.lora_B.bias"] = k_bias + state_dict[f"{new_key}.attn.to_v.lora_B.bias"] = v_bias + state_dict[f"{new_key}.proj_mlp.lora_B.bias"] = mlp_bias + + else: + new_key = key.replace("single_blocks", "single_transformer_blocks") + new_key = new_key.replace("linear2", "proj_out") + new_key = new_key.replace("q_norm", "attn.norm_q") + new_key = new_key.replace("k_norm", "attn.norm_k") + state_dict[new_key] = state_dict.pop(key) + except Exception as e: + logger.warning(f"Error processing key {key}: {str(e)}") + # Skip if we can't process this key properly + if key in state_dict: + state_dict.pop(key) + + TRANSFORMER_KEYS_RENAME_DICT = { + "img_in": "x_embedder", + "time_in.mlp.0": "time_text_embed.timestep_embedder.linear_1", + "time_in.mlp.2": "time_text_embed.timestep_embedder.linear_2", + "guidance_in.mlp.0": "time_text_embed.guidance_embedder.linear_1", + "guidance_in.mlp.2": "time_text_embed.guidance_embedder.linear_2", + "vector_in.in_layer": "time_text_embed.text_embedder.linear_1", + "vector_in.out_layer": "time_text_embed.text_embedder.linear_2", + "double_blocks": "transformer_blocks", + "img_attn_q_norm": "attn.norm_q", + "img_attn_k_norm": "attn.norm_k", + "img_attn_proj": "attn.to_out.0", + "txt_attn_q_norm": "attn.norm_added_q", + "txt_attn_k_norm": "attn.norm_added_k", + "txt_attn_proj": "attn.to_add_out", + "img_mod.linear": "norm1.linear", + "img_norm1": "norm1.norm", + "img_norm2": "norm2", + "img_mlp": "ff", + "txt_mod.linear": "norm1_context.linear", + "txt_norm1": "norm1.norm", + "txt_norm2": "norm2_context", + "txt_mlp": "ff_context", + "self_attn_proj": "attn.to_out.0", + "modulation.linear": "norm.linear", + "pre_norm": "norm.norm", + "final_layer.norm_final": "norm_out.norm", + "final_layer.linear": "proj_out", + "fc1": "net.0.proj", + "fc2": "net.2", + "input_embedder": "proj_in", + } + + TRANSFORMER_SPECIAL_KEYS_REMAP = { + "txt_in": remap_txt_in_, + "img_attn_qkv": remap_img_attn_qkv_, + "txt_attn_qkv": remap_txt_attn_qkv_, + "single_blocks": remap_single_transformer_blocks_, + "final_layer.adaLN_modulation.1": remap_norm_scale_shift_, + } + + # Some folks attempt to make their state dict compatible with diffusers by adding "transformer." prefix to all keys + # and use their custom code. To make sure both "original" and "attempted diffusers" loras work as expected, we make + # sure that both follow the same initial format by stripping off the "transformer." prefix. + for key in list(converted_state_dict.keys()): + try: + if key.startswith("transformer."): + converted_state_dict[key[len("transformer.") :]] = converted_state_dict.pop(key) + if key.startswith("diffusion_model."): + converted_state_dict[key[len("diffusion_model.") :]] = converted_state_dict.pop(key) + except Exception as e: + logger.warning(f"Error processing key {key}: {str(e)}") + # Skip if we can't process this key properly + + # Rename and remap the state dict keys + for key in list(converted_state_dict.keys()): + try: + new_key = key[:] + for replace_key, rename_key in TRANSFORMER_KEYS_RENAME_DICT.items(): + new_key = new_key.replace(replace_key, rename_key) + converted_state_dict[new_key] = converted_state_dict.pop(key) + except Exception as e: + logger.warning(f"Error processing key {key}: {str(e)}") + # Skip if we can't process this key properly + + for key in list(converted_state_dict.keys()): + try: + for special_key, handler_fn_inplace in TRANSFORMER_SPECIAL_KEYS_REMAP.items(): + if special_key not in key: + continue + handler_fn_inplace(key, converted_state_dict) + except Exception as e: + logger.warning(f"Error processing key {key}: {str(e)}") + # Skip if we can't process this key properly + + # Add back the "transformer." prefix + for key in list(converted_state_dict.keys()): + try: + converted_state_dict[f"transformer.{key}"] = converted_state_dict.pop(key) + except Exception as e: + logger.warning(f"Error processing key {key}: {str(e)}") + # Skip if we can't process this key properly + + return converted_state_dict + except Exception as e: + logger.error(f"LoRA conversion failed: {str(e)}") + # Return empty state dict as fallback + return {} + + +class HyVideoModel(comfy.model_base.BaseModel): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.pipeline = {} + self.load_device = mm.get_torch_device() + + def __getitem__(self, k): + return self.pipeline[k] + + def __setitem__(self, k, v): + self.pipeline[k] = v + + +class HyVideoModelConfig: + def __init__(self, dtype): + self.unet_config = {} + self.unet_extra_config = {} + self.latent_format = comfy.latent_formats.HunyuanVideo + self.latent_format.latent_channels = 16 + self.manual_cast_dtype = dtype + self.sampling_settings = {"multiplier": 1.0} + self.memory_usage_factor = 2.0 + self.unet_config["disable_unet_model_creation"] = True + +class FramePackTorchCompileSettings: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "backend": (["inductor","cudagraphs"], {"default": "inductor"}), + "fullgraph": ("BOOLEAN", {"default": False, "tooltip": "Enable full graph mode"}), + "mode": (["default", "max-autotune", "max-autotune-no-cudagraphs", "reduce-overhead"], {"default": "default"}), + "dynamic": ("BOOLEAN", {"default": False, "tooltip": "Enable dynamic mode"}), + "dynamo_cache_size_limit": ("INT", {"default": 64, "min": 0, "max": 1024, "step": 1, "tooltip": "torch._dynamo.config.cache_size_limit"}), + "compile_single_blocks": ("BOOLEAN", {"default": True, "tooltip": "Enable single block compilation"}), + "compile_double_blocks": ("BOOLEAN", {"default": True, "tooltip": "Enable double block compilation"}), + }, + } + RETURN_TYPES = ("FRAMEPACKCOMPILEARGS",) + RETURN_NAMES = ("torch_compile_args",) + FUNCTION = "loadmodel" + CATEGORY = "HunyuanVideoWrapper" + DESCRIPTION = "torch.compile settings, when connected to the model loader, torch.compile of the selected layers is attempted. Requires Triton and torch 2.5.0 is recommended" + + def loadmodel(self, backend, fullgraph, mode, dynamic, dynamo_cache_size_limit, compile_single_blocks, compile_double_blocks): + + compile_args = { + "backend": backend, + "fullgraph": fullgraph, + "mode": mode, + "dynamic": dynamic, + "dynamo_cache_size_limit": dynamo_cache_size_limit, + "compile_single_blocks": compile_single_blocks, + "compile_double_blocks": compile_double_blocks + } + + return (compile_args, ) + +#region Model loading +class DownloadAndLoadFramePackModel: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": (["lllyasviel/FramePackI2V_HY"],), + + "base_precision": (["fp32", "bf16", "fp16"], {"default": "bf16"}), + "quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2'], {"default": 'disabled', "tooltip": "optional quantization method"}), + }, + "optional": { + "attention_mode": ([ + "sdpa", + "flash_attn", + "sageattn", + ], {"default": "sdpa"}), + "compile_args": ("FRAMEPACKCOMPILEARGS", ), + } + } + + RETURN_TYPES = ("FramePackMODEL",) + RETURN_NAMES = ("model", ) + FUNCTION = "loadmodel" + CATEGORY = "FramePackWrapper" + + def loadmodel(self, model, base_precision, quantization, + compile_args=None, attention_mode="sdpa"): + + base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision] + + device = mm.get_torch_device() + + model_path = os.path.join(folder_paths.models_dir, "diffusers", "lllyasviel", "FramePackI2V_HY") + if not os.path.exists(model_path): + print(f"Downloading clip model to: {model_path}") + from huggingface_hub import snapshot_download + snapshot_download( + repo_id=model, + local_dir=model_path, + local_dir_use_symlinks=False, + ) + + transformer = HunyuanVideoTransformer3DModel.from_pretrained(model_path, torch_dtype=base_dtype, attention_mode=attention_mode).cpu() + params_to_keep = {"norm", "bias", "time_in", "vector_in", "guidance_in", "txt_in", "img_in"} + if quantization == 'fp8_e4m3fn' or quantization == 'fp8_e4m3fn_fast': + transformer = transformer.to(torch.float8_e4m3fn) + if quantization == "fp8_e4m3fn_fast": + from .fp8_optimization import convert_fp8_linear + convert_fp8_linear(transformer, base_dtype, params_to_keep=params_to_keep) + elif quantization == 'fp8_e5m2': + transformer = transformer.to(torch.float8_e5m2) + else: + transformer = transformer.to(base_dtype) + + DynamicSwapInstaller.install_model(transformer, device=device) + + if compile_args is not None: + if compile_args["compile_single_blocks"]: + for i, block in enumerate(transformer.single_transformer_blocks): + transformer.single_transformer_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + if compile_args["compile_double_blocks"]: + for i, block in enumerate(transformer.transformer_blocks): + transformer.transformer_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + + #transformer = torch.compile(transformer, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + + pipe = { + "transformer": transformer.eval(), + "dtype": base_dtype, + } + return (pipe, ) + +class FramePackLoraSelect: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "lora": (folder_paths.get_filename_list("loras"), + {"tooltip": "LORA models are expected to be in ComfyUI/models/loras with .safetensors extension"}), + "strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.0001, "tooltip": "LORA strength, set to 0.0 to unmerge the LORA"}), + "fuse_lora": ("BOOLEAN", {"default": True, "tooltip": "Fuse the LORA model with the base model. This is recommended for better performance."}), + }, + "optional": { + "prev_lora":("FPLORA", {"default": None, "tooltip": "For loading multiple LoRAs"}), + } + } + + RETURN_TYPES = ("FPLORA",) + RETURN_NAMES = ("lora", ) + FUNCTION = "getlorapath" + CATEGORY = "FramePackWrapper" + DESCRIPTION = "Select a LoRA model from ComfyUI/models/loras" + + def getlorapath(self, lora, strength, prev_lora=None, fuse_lora=True): + loras_list = [] + + lora = { + "path": folder_paths.get_full_path("loras", lora), + "strength": strength, + "name": lora.split(".")[0], + "fuse_lora": fuse_lora, + } + if prev_lora is not None: + loras_list.extend(prev_lora) + + loras_list.append(lora) + return (loras_list,) + +class LoadFramePackModel: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}), + + "base_precision": (["fp32", "bf16", "fp16"], {"default": "bf16"}), + "quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2'], {"default": 'disabled', "tooltip": "optional quantization method"}), + "load_device": (["main_device", "offload_device"], {"default": "cuda", "tooltip": "Initialize the model on the main device or offload device"}), + }, + "optional": { + "attention_mode": ([ + "sdpa", + "flash_attn", + "sageattn", + ], {"default": "sdpa"}), + "compile_args": ("FRAMEPACKCOMPILEARGS", ), + "lora": ("FPLORA", {"default": None, "tooltip": "LORA model to load"}), + } + } + + RETURN_TYPES = ("FramePackMODEL",) + RETURN_NAMES = ("model", ) + FUNCTION = "loadmodel" + CATEGORY = "FramePackWrapper" + + def loadmodel(self, model, base_precision, quantization, + compile_args=None, attention_mode="sdpa", lora=None, load_device="main_device"): + base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision] + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + if load_device == "main_device": + transformer_load_device = device + else: + transformer_load_device = offload_device + model_path = folder_paths.get_full_path_or_raise("diffusion_models", model) + model_config_path = os.path.join(script_directory, "transformer_config.json") + import json + with open(model_config_path, "r") as f: + config = json.load(f) + sd = load_torch_file(model_path, device=offload_device, safe_load=True) + model_weight_dtype = sd['single_transformer_blocks.0.attn.to_k.weight'].dtype + with init_empty_weights(): + transformer = HunyuanVideoTransformer3DModel(**config, attention_mode=attention_mode) + params_to_keep = {"norm", "bias", "time_in", "vector_in", "guidance_in", "txt_in", "img_in"} + if quantization == "fp8_e4m3fn" or quantization == "fp8_e4m3fn_fast" or quantization == "fp8_scaled": + dtype = torch.float8_e4m3fn + elif quantization == "fp8_e5m2": + dtype = torch.float8_e5m2 + else: + dtype = base_dtype + if lora is not None: + after_lora_dtype = dtype + dtype = base_dtype + print("Using accelerate to load and assign model weights to device...") + param_count = sum(1 for _ in transformer.named_parameters()) + for name, param in tqdm(transformer.named_parameters(), + desc=f"Loading transformer parameters to {transformer_load_device}", + total=param_count, + leave=True): + dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype + set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) + + if lora is not None: + adapter_list = [] + adapter_weights = [] + + for l in lora: + fuse = True if l["fuse_lora"] else False + lora_sd = load_torch_file(l["path"]) + + if "lora_unet_single_transformer_blocks_0_attn_to_k.lora_up.weight" in lora_sd: + from .utils import convert_to_diffusers + lora_sd = convert_to_diffusers("lora_unet_", lora_sd) + + if not "transformer.single_transformer_blocks.0.attn.to_k.lora_A.weight" in lora_sd: + log.info(f"Converting LoRA weights from {l['path']} to diffusers format...") + # Make a copy of the original state dict to avoid modifying it + state_dict_copy = {} + + # Remove scalar (0-dimensional) tensors that cause problems + for key, value in lora_sd.items(): + if isinstance(value, torch.Tensor): + if value.dim() == 0: + print(f"Skipping 0-dimensional tensor: {key}") + continue + state_dict_copy[key] = value + else: + print(f"Skipping non-tensor value: {key}") + + print(f"After filtering: {len(state_dict_copy)} valid keys") + + # Try the original conversion with the filtered state dict + try: + from diffusers.loaders.lora_conversion_utils import _convert_hunyuan_video_lora_to_diffusers + lora_sd = _convert_hunyuan_video_lora_to_diffusers(state_dict_copy) + print("Successfully converted LoRA weights") + except Exception as e: + print(f"Error in standard conversion: {e}") + # Fall back to empty dict if conversion fails + print("Conversion failed, returning empty state dict") + lora_sd = {} + + lora_rank = None + for key, val in lora_sd.items(): + if "lora_B" in key or "lora_up" in key: + lora_rank = val.shape[1] + break + if lora_rank is not None: + log.info(f"Merging rank {lora_rank} LoRA weights from {l['path']} with strength {l['strength']}") + adapter_name = l['path'].split("/")[-1].split(".")[0] + adapter_weight = l['strength'] + transformer.load_lora_adapter(lora_sd, weight_name=l['path'].split("/")[-1], lora_rank=lora_rank, adapter_name=adapter_name) + + adapter_list.append(adapter_name) + adapter_weights.append(adapter_weight) + + del lora_sd + mm.soft_empty_cache() + if adapter_list: + transformer.set_adapters(adapter_list, weights=adapter_weights) + if fuse: + if model_weight_dtype not in [torch.float32, torch.float16, torch.bfloat16]: + raise ValueError("Fusing LoRA doesn't work well with fp8 model weights. Please use a bf16 model file, or disable LoRA fusing.") + lora_scale = 1 + transformer.fuse_lora(lora_scale=lora_scale) + transformer.delete_adapters(adapter_list) + + if quantization == "fp8_e4m3fn" or quantization == "fp8_e4m3fn_fast" or quantization == "fp8_e5m2": + params_to_keep = {"norm", "bias", "time_in", "vector_in", "guidance_in", "txt_in", "img_in"} + for name, param in transformer.named_parameters(): + # Make sure to not cast the LoRA weights to fp8. + if not any(keyword in name for keyword in params_to_keep) and not 'lora' in name: + param.data = param.data.to(after_lora_dtype) + + if quantization == "fp8_e4m3fn_fast": + from .fp8_optimization import convert_fp8_linear + convert_fp8_linear(transformer, base_dtype, params_to_keep=params_to_keep) + + DynamicSwapInstaller.install_model(transformer, device=device) + + if compile_args is not None: + if compile_args["compile_single_blocks"]: + for i, block in enumerate(transformer.single_transformer_blocks): + transformer.single_transformer_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + if compile_args["compile_double_blocks"]: + for i, block in enumerate(transformer.transformer_blocks): + transformer.transformer_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + + pipe = { + "transformer": transformer.eval(), + "dtype": base_dtype, + } + return (pipe, ) + +class FramePackFindNearestBucket: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "image": ("IMAGE", {"tooltip": "Image to resize"}), + "base_resolution": ("INT", {"default": 640, "min": 64, "max": 2048, "step": 16, "tooltip": "Width of the image to encode"}), + }, + } + + RETURN_TYPES = ("INT", "INT", ) + RETURN_NAMES = ("width","height",) + FUNCTION = "process" + CATEGORY = "FramePackWrapper" + DESCRIPTION = "Finds the closes resolution bucket as defined in the orignal code" + + def process(self, image, base_resolution): + + H, W = image.shape[1], image.shape[2] + + new_height, new_width = find_nearest_bucket(H, W, resolution=base_resolution) + + return (new_width, new_height, ) + + +class FramePackSampler: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("FramePackMODEL",), + "positive": ("CONDITIONING",), + "negative": ("CONDITIONING",), + "start_latent": ("LATENT", {"tooltip": "init Latents to use for image2video"} ), + "steps": ("INT", {"default": 30, "min": 1}), + "use_teacache": ("BOOLEAN", {"default": True, "tooltip": "Use teacache for faster sampling."}), + "teacache_rel_l1_thresh": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "The threshold for the relative L1 loss."}), + "cfg": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 30.0, "step": 0.01}), + "guidance_scale": ("FLOAT", {"default": 10.0, "min": 0.0, "max": 32.0, "step": 0.01}), + "shift": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1000.0, "step": 0.01}), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + "latent_window_size": ("INT", {"default": 9, "min": 1, "max": 33, "step": 1, "tooltip": "The size of the latent window to use for sampling."}), + "total_second_length": ("FLOAT", {"default": 5, "min": 1, "max": 120, "step": 0.1, "tooltip": "The total length of the video in seconds."}), + "gpu_memory_preservation": ("FLOAT", {"default": 6.0, "min": 0.0, "max": 128.0, "step": 0.1, "tooltip": "The amount of GPU memory to preserve."}), + "sampler": (["unipc_bh1", "unipc_bh2"], + { + "default": 'unipc_bh1' + }), + }, + "optional": { + "image_embeds": ("CLIP_VISION_OUTPUT", ), + "end_latent": ("LATENT", {"tooltip": "end Latents to use for image2video"} ), + "end_image_embeds": ("CLIP_VISION_OUTPUT", {"tooltip": "end Image's clip embeds"} ), + "embed_interpolation": (["disabled", "weighted_average", "linear"], {"default": 'disabled', "tooltip": "Image embedding interpolation type. If linear, will smoothly interpolate with time, else it'll be weighted average with the specified weight."}), + "start_embed_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Weighted average constant for image embed interpolation. If end image is not set, the embed's strength won't be affected"}), + "initial_samples": ("LATENT", {"tooltip": "init Latents to use for video2video"} ), + "denoise_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), + } + } + + RETURN_TYPES = ("LATENT", ) + RETURN_NAMES = ("samples",) + FUNCTION = "process" + CATEGORY = "FramePackWrapper" + + def process(self, model, shift, positive, negative, latent_window_size, use_teacache, total_second_length, teacache_rel_l1_thresh, steps, cfg, + guidance_scale, seed, sampler, gpu_memory_preservation, start_latent=None, image_embeds=None, end_latent=None, end_image_embeds=None, embed_interpolation="linear", start_embed_strength=1.0, initial_samples=None, denoise_strength=1.0): + total_latent_sections = (total_second_length * 30) / (latent_window_size * 4) + total_latent_sections = int(max(round(total_latent_sections), 1)) + print("total_latent_sections: ", total_latent_sections) + + transformer = model["transformer"] + base_dtype = model["dtype"] + + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() + + if start_latent is not None: + start_latent = start_latent["samples"] * vae_scaling_factor + if initial_samples is not None: + initial_samples = initial_samples["samples"] * vae_scaling_factor + if end_latent is not None: + end_latent = end_latent["samples"] * vae_scaling_factor + has_end_image = end_latent is not None + print("start_latent", start_latent.shape) + B, C, T, H, W = start_latent.shape + + if image_embeds is not None: + start_image_encoder_last_hidden_state = image_embeds["last_hidden_state"].to(device, base_dtype) + + if has_end_image: + assert end_image_embeds is not None + end_image_encoder_last_hidden_state = end_image_embeds["last_hidden_state"].to(device, base_dtype) + else: + if image_embeds is not None: + end_image_encoder_last_hidden_state = torch.zeros_like(start_image_encoder_last_hidden_state) + + llama_vec = positive[0][0].to(device, base_dtype) + clip_l_pooler = positive[0][1]["pooled_output"].to(device, base_dtype) + + if not math.isclose(cfg, 1.0): + llama_vec_n = negative[0][0].to(device, base_dtype) + clip_l_pooler_n = negative[0][1]["pooled_output"].to(device, base_dtype) + else: + llama_vec_n = torch.zeros_like(llama_vec, device=device) + clip_l_pooler_n = torch.zeros_like(clip_l_pooler, device=device) + + llama_vec, llama_attention_mask = crop_or_pad_yield_mask(llama_vec, length=512) + llama_vec_n, llama_attention_mask_n = crop_or_pad_yield_mask(llama_vec_n, length=512) + + + # Sampling + + rnd = torch.Generator("cpu").manual_seed(seed) + + num_frames = latent_window_size * 4 - 3 + + history_latents = torch.zeros(size=(1, 16, 1 + 2 + 16, H, W), dtype=torch.float32).cpu() + + total_generated_latent_frames = 0 + + latent_paddings_list = list(reversed(range(total_latent_sections))) + latent_paddings = latent_paddings_list.copy() # Create a copy for iteration + + comfy_model = HyVideoModel( + HyVideoModelConfig(base_dtype), + model_type=comfy.model_base.ModelType.FLOW, + device=device, + ) + + patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, torch.device("cpu")) + from latent_preview import prepare_callback + callback = prepare_callback(patcher, steps) + + move_model_to_device_with_memory_preservation(transformer, target_device=device, preserved_memory_gb=gpu_memory_preservation) + + if total_latent_sections > 4: + # In theory the latent_paddings should follow the above sequence, but it seems that duplicating some + # items looks better than expanding it when total_latent_sections > 4 + # One can try to remove below trick and just + # use `latent_paddings = list(reversed(range(total_latent_sections)))` to compare + latent_paddings = [3] + [2] * (total_latent_sections - 3) + [1, 0] + latent_paddings_list = latent_paddings.copy() + + for i, latent_padding in enumerate(latent_paddings): + print(f"latent_padding: {latent_padding}") + is_last_section = latent_padding == 0 + is_first_section = latent_padding == latent_paddings[0] + latent_padding_size = latent_padding * latent_window_size + + if image_embeds is not None: + if embed_interpolation != "disabled": + if embed_interpolation == "linear": + if total_latent_sections <= 1: + frac = 1.0 # Handle case with only one section + else: + frac = 1 - i / (total_latent_sections - 1) # going backwards + else: + frac = start_embed_strength if has_end_image else 1.0 + + image_encoder_last_hidden_state = start_image_encoder_last_hidden_state * frac + (1 - frac) * end_image_encoder_last_hidden_state + else: + image_encoder_last_hidden_state = start_image_encoder_last_hidden_state * start_embed_strength + else: + image_encoder_last_hidden_state = None + + print(f'latent_padding_size = {latent_padding_size}, is_last_section = {is_last_section}, is_first_section = {is_first_section}') + + start_latent_frames = T # 0 or 1 + indices = torch.arange(0, sum([start_latent_frames, latent_padding_size, latent_window_size, 1, 2, 16])).unsqueeze(0) + clean_latent_indices_pre, blank_indices, latent_indices, clean_latent_indices_post, clean_latent_2x_indices, clean_latent_4x_indices = indices.split([start_latent_frames, latent_padding_size, latent_window_size, 1, 2, 16], dim=1) + clean_latent_indices = torch.cat([clean_latent_indices_pre, clean_latent_indices_post], dim=1) + + clean_latents_pre = start_latent.to(history_latents) + clean_latents_post, clean_latents_2x, clean_latents_4x = history_latents[:, :, :1 + 2 + 16, :, :].split([1, 2, 16], dim=2) + clean_latents = torch.cat([clean_latents_pre, clean_latents_post], dim=2) + + # Use end image latent for the first section if provided + if has_end_image and is_first_section: + clean_latents_post = end_latent.to(history_latents) + clean_latents = torch.cat([clean_latents_pre, clean_latents_post], dim=2) + + #vid2vid WIP + + if initial_samples is not None: + total_length = initial_samples.shape[2] + + # Get the max padding value for normalization + max_padding = max(latent_paddings_list) + + if is_last_section: + # Last section should capture the end of the sequence + start_idx = max(0, total_length - latent_window_size) + else: + # Calculate windows that distribute more evenly across the sequence + # This normalizes the padding values to create appropriate spacing + if max_padding > 0: # Avoid division by zero + progress = (max_padding - latent_padding) / max_padding + start_idx = int(progress * max(0, total_length - latent_window_size)) + else: + start_idx = 0 + + end_idx = min(start_idx + latent_window_size, total_length) + print(f"start_idx: {start_idx}, end_idx: {end_idx}, total_length: {total_length}") + input_init_latents = initial_samples[:, :, start_idx:end_idx, :, :].to(device) + + + if use_teacache: + transformer.initialize_teacache(enable_teacache=True, num_steps=steps, rel_l1_thresh=teacache_rel_l1_thresh) + else: + transformer.initialize_teacache(enable_teacache=False) + + with torch.autocast(device_type=mm.get_autocast_device(device), dtype=base_dtype, enabled=True): + generated_latents = sample_hunyuan( + transformer=transformer, + sampler=sampler, + initial_latent=input_init_latents if initial_samples is not None else None, + strength=denoise_strength, + width=W * 8, + height=H * 8, + frames=num_frames, + real_guidance_scale=cfg, + distilled_guidance_scale=guidance_scale, + guidance_rescale=0, + shift=shift if shift != 0 else None, + num_inference_steps=steps, + generator=rnd, + prompt_embeds=llama_vec, + prompt_embeds_mask=llama_attention_mask, + prompt_poolers=clip_l_pooler, + negative_prompt_embeds=llama_vec_n, + negative_prompt_embeds_mask=llama_attention_mask_n, + negative_prompt_poolers=clip_l_pooler_n, + device=device, + dtype=base_dtype, + image_embeddings=image_encoder_last_hidden_state, + latent_indices=latent_indices, + clean_latents=clean_latents, + clean_latent_indices=clean_latent_indices, + clean_latents_2x=clean_latents_2x, + clean_latent_2x_indices=clean_latent_2x_indices, + clean_latents_4x=clean_latents_4x, + clean_latent_4x_indices=clean_latent_4x_indices, + callback=callback, + ) + + if is_last_section: + generated_latents = torch.cat([start_latent.to(generated_latents), generated_latents], dim=2) + + total_generated_latent_frames += int(generated_latents.shape[2]) + history_latents = torch.cat([generated_latents.to(history_latents), history_latents], dim=2) + + real_history_latents = history_latents[:, :, :total_generated_latent_frames, :, :] + + if is_last_section: + break + + transformer.to(offload_device) + mm.soft_empty_cache() + + return {"samples": real_history_latents / vae_scaling_factor}, + +NODE_CLASS_MAPPINGS = { + "DownloadAndLoadFramePackModel": DownloadAndLoadFramePackModel, + "FramePackSampler": FramePackSampler, + "FramePackTorchCompileSettings": FramePackTorchCompileSettings, + "FramePackFindNearestBucket": FramePackFindNearestBucket, + "LoadFramePackModel": LoadFramePackModel, + "FramePackLoraSelect": FramePackLoraSelect, + } + +NODE_DISPLAY_NAME_MAPPINGS = { + "DownloadAndLoadFramePackModel": "(Down)Load FramePackModel", + "FramePackSampler": "FramePackSampler", + "FramePackTorchCompileSettings": "Torch Compile Settings", + "FramePackFindNearestBucket": "Find Nearest Bucket", + "LoadFramePackModel": "Load FramePackModel", + "FramePackLoraSelect": "Select Lora", + } \ No newline at end of file diff --git a/nodes_F1.py b/nodes_F1.py new file mode 100644 index 0000000..4663370 --- /dev/null +++ b/nodes_F1.py @@ -0,0 +1,635 @@ +import os +import torch +import math +import re + +import comfy.model_management as mm +import comfy.model_base +import comfy.model_patcher + +from .nodes import HyVideoModel, HyVideoModelConfig # Import the classes + +script_directory = os.path.dirname(os.path.abspath(__file__)) +vae_scaling_factor = 0.476986 + +from .diffusers_helper.models.hunyuan_video_packed import HunyuanVideoTransformer3DModel +from .diffusers_helper.memory import move_model_to_device_with_memory_preservation +from .diffusers_helper.pipelines.k_diffusion_hunyuan import sample_hunyuan +from .diffusers_helper.utils import crop_or_pad_yield_mask + +from dataclasses import dataclass +from typing import List, Optional, Tuple, Dict, Union # Add necessary types + +from latent_preview import prepare_callback + +# --- Helper Classes and Functions for Timestamped Prompts --- +@dataclass +class PromptSection: + prompt: str + start_time: float = 0.0 # in seconds + end_time: Optional[float] = None # in seconds, None means until the end + +def snap_to_section_boundaries(prompt_sections: List[PromptSection], latent_window_size: int, fps: int = 30) -> List[PromptSection]: + + section_frame_duration = latent_window_size * 4 - 3 + if section_frame_duration <= 0: section_frame_duration = 1 + section_duration_sec = section_frame_duration / float(fps) + if section_duration_sec <= 1e-5: section_duration_sec = 1.0 / fps # Avoid zero or near-zero duration + + aligned_sections = [] + for section in prompt_sections: + aligned_start = round(section.start_time / section_duration_sec) * section_duration_sec + aligned_end = None + if section.end_time is not None: + aligned_end = round(section.end_time / section_duration_sec) * section_duration_sec + if aligned_end <= aligned_start + 1e-5: # Ensure minimum duration + aligned_end = aligned_start + section_duration_sec + aligned_sections.append(PromptSection( + prompt=section.prompt, + start_time=aligned_start, + end_time=aligned_end + )) + return aligned_sections + +def parse_timestamped_prompt_f1(prompt_text: str, total_duration: float, latent_window_size: int = 9) -> List[PromptSection]: + + #Parse a prompt with timestamps like [0s: text], [1.5s-3s: text] for F1-style forward generation. + #Returns a list of PromptSection objects with timestamps aligned to section boundaries. + sections = [] + # Corrected Regex: Catches [Xs: text] or [Xs-Ys: text] + timestamp_pattern = r'\[\s*(\d+(?:\.\d+)?s)\s*(?:-\s*(\d+(?:\.\d+)?s)\s*)?:\s*(.*?)\s*\]' + matches = list(re.finditer(timestamp_pattern, prompt_text)) + last_end_index = 0 + + if not matches: + return [PromptSection(prompt=prompt_text.strip(), start_time=0.0, end_time=total_duration)] + + for match in matches: + plain_text_before = prompt_text[last_end_index:match.start()].strip() + current_start_time_str = match.group(1) + current_start_time = float(current_start_time_str.rstrip('s')) + if plain_text_before: + previous_end_time = sections[-1].end_time if sections and sections[-1].end_time is not None else (sections[-1].start_time if sections else 0.0) + if current_start_time > previous_end_time + 1e-5: + sections.append(PromptSection(prompt=plain_text_before, start_time=previous_end_time, end_time=current_start_time)) + elif not sections and current_start_time > 1e-5: # Plain text at the very beginning + sections.append(PromptSection(prompt=plain_text_before, start_time=0.0, end_time=current_start_time)) + + end_time_str = match.group(2) + section_text = match.group(3).strip() + start_time = current_start_time # Already parsed + end_time = float(end_time_str.rstrip('s')) if end_time_str else None + sections.append(PromptSection(prompt=section_text, start_time=start_time, end_time=end_time)) + last_end_index = match.end() + + plain_text_after = prompt_text[last_end_index:].strip() + if plain_text_after: + previous_end_time = sections[-1].end_time if sections and sections[-1].end_time is not None else sections[-1].start_time + if total_duration > previous_end_time + 1e-5: + sections.append(PromptSection(prompt=plain_text_after, start_time=previous_end_time, end_time=None)) + + if not sections: # Should not happen if regex matched, but safety + return [PromptSection(prompt=prompt_text.strip(), start_time=0.0, end_time=total_duration)] + + sections.sort(key=lambda x: x.start_time) + + # Sanitize and Fill Gaps/Set End Times + sanitized_sections = [] + current_time = 0.0 + for i, section in enumerate(sections): + section_start = max(current_time, section.start_time) # Ensure monotonic increase + section_start = min(section_start, total_duration) # Clamp to total duration + + # Fill gap if needed + if section_start > current_time + 1e-5: + filler_prompt = sanitized_sections[-1].prompt if sanitized_sections else "" # Use previous prompt + sanitized_sections.append(PromptSection(prompt=filler_prompt, start_time=current_time, end_time=section_start)) + + # Determine end time + section_end = section.end_time + if section_end is None: + if i + 1 < len(sections): + next_start = max(section_start, sections[i+1].start_time) # Ensure next start is after current start + section_end = min(next_start, total_duration) # End before next or at total duration + else: + section_end = total_duration # Last section ends at total duration + else: + section_end = min(max(section_start, section_end), total_duration) # Clamp user-defined end + + # Add the section if it has duration + if section_end > section_start + 1e-5: + sanitized_sections.append(PromptSection(prompt=section.prompt, start_time=section_start, end_time=section_end)) + current_time = section_end # Update current time marker + elif i == len(sections) - 1 and math.isclose(section_start, total_duration): # Allow point at the end? No, remove. + pass + + if not sanitized_sections: + return [PromptSection(prompt=prompt_text.strip(), start_time=0.0, end_time=total_duration)] + + # Snap timestamps to boundaries + aligned_sections = snap_to_section_boundaries(sanitized_sections, latent_window_size) + + # Merge identical consecutive prompts after snapping + merged_sections = [] + if not aligned_sections: return [PromptSection(prompt=prompt_text.strip(), start_time=0.0, end_time=total_duration)] + + current_merged = aligned_sections[0] + for i in range(1, len(aligned_sections)): + next_sec = aligned_sections[i] + # Merge if prompts are identical and sections are contiguous (or very close after snapping) + if next_sec.prompt == current_merged.prompt and abs(next_sec.start_time - current_merged.end_time) < 0.01: + current_merged.end_time = next_sec.end_time # Extend the end time + else: + current_merged.end_time = max(current_merged.start_time, current_merged.end_time) + if current_merged.start_time < current_merged.end_time - 1e-5: + merged_sections.append(current_merged) + current_merged = next_sec + + current_merged.end_time = max(current_merged.start_time, current_merged.end_time) + if current_merged.start_time < current_merged.end_time - 1e-5: + merged_sections.append(current_merged) + + if not merged_sections: return [PromptSection(prompt=prompt_text.strip(), start_time=0.0, end_time=total_duration)] + + print("Parsed Prompt Sections (F1):") + for sec in merged_sections: print(f" [{sec.start_time:.3f}s - {sec.end_time:.3f}s]: {sec.prompt}") + return merged_sections +# --- End Helper Code --- + +class FramePackSampler_F1: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("FramePackMODEL",), + "positive_timed_data": ("TIMED_CONDITIONING_WITH_METADATA", { "tooltip": "Output from FramePackTimestampedTextEncode. Dictionary containing sections, duration, and window size."}), + "negative": ("CONDITIONING",), + "steps": ("INT", {"default": 30, "min": 1}), + "use_teacache": ("BOOLEAN", {"default": True, "tooltip": "Use teacache for faster sampling."}), + "teacache_rel_l1_thresh": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "The threshold for the relative L1 loss."}), + "cfg": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 30.0, "step": 0.01}), + "guidance_scale": ("FLOAT", {"default": 10.0, "min": 0.0, "max": 32.0, "step": 0.01}), + "shift": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1000.0, "step": 0.01}), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + "gpu_memory_preservation": ("FLOAT", {"default": 6.0, "min": 0.0, "max": 128.0, "step": 0.1, "tooltip": "The amount of GPU memory to preserve."}), + "sampler": (["unipc_bh1", "unipc_bh2"], + { + "default": 'unipc_bh1' + }), + }, + "optional": { + "start_latent": ("LATENT", {"tooltip": "init Latents to use for image2video"} ), + "start_image_embeds": ("CLIP_VISION_OUTPUT", ), + "end_latent": ("LATENT", {"tooltip": "end Latents to use for image2video"} ), + "end_image_embeds": ("CLIP_VISION_OUTPUT", {"tooltip": "end Image's clip embeds"} ), + "embed_interpolation": (["disabled", "weighted_average", "linear"], {"default": 'disabled', "tooltip": "Image embedding interpolation type. If linear, will smoothly interpolate with time, else it'll be weighted average with the specified weight."}), + "start_embed_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Weighted average constant for image embed interpolation. If end image is not set, the embed's strength won't be affected"}), + "initial_samples": ("LATENT", {"tooltip": "init Latents to use for video2video"} ), + "denoise_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), + } + } + + RETURN_TYPES = ("LATENT", ) + RETURN_NAMES = ("samples",) + FUNCTION = "process" + CATEGORY = "FramePackWrapper" + + def process(self, model, positive_timed_data, negative, use_teacache, teacache_rel_l1_thresh, steps, cfg, + guidance_scale, shift, seed, sampler, gpu_memory_preservation, start_image_embeds=None, start_latent=None, end_latent=None, end_image_embeds=None, embed_interpolation="linear", start_embed_strength=1.0, initial_samples=None, denoise_strength=1.0): + + # --- Extract data from positive_timed_data --- + positive_timed_list = positive_timed_data["sections"] + total_second_length = positive_timed_data["total_duration"] + latent_window_size = positive_timed_data["window_size"] + prompt_blend_sections = positive_timed_data["blend_sections"] + print(f"Received - Total Duration: {total_second_length}s, Window Size: {latent_window_size}, Blend Sections: {prompt_blend_sections}") + + # --- F1 Model Type Assumption --- + # We assume the model loaded into this node is the F1 type. + + # Calculate total sections based on time and window size + section_frame_duration = latent_window_size * 4 - 3 + if section_frame_duration <= 0: section_frame_duration = 1 + fps = 30 # Assume 30 fps + section_duration_sec = section_frame_duration / float(fps) + if section_duration_sec <= 0: section_duration_sec = 1.0 / fps + + # Calculate total sections needed to cover the duration + total_latent_sections = int(math.ceil(total_second_length / section_duration_sec)) + total_latent_sections = max(total_latent_sections, 1) + print(f"Total latent sections calculated: {total_latent_sections} (Duration: {total_second_length}s, Section time: {section_duration_sec:.3f}s)") + + + transformer = model["transformer"] + base_dtype = model["dtype"] + + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() + + if start_latent is None: + # Handle case where start_latent is not provided (e.g., create default black latent) + # Get model's expected channel count (often 16 for FramePack) + latent_channels = getattr(transformer.config, 'in_channels', 16) + # Determine a default spatial size if not derivable (e.g., 64x64 or based on bucket?) + # Using a common default like 64x64 / 8 = 8x8 latent space, but this might need adjustment + H = W = 64 # Default spatial size assumption + print(f"Warning: start_latent not provided. Creating default black latent ({latent_channels}x1x{H}x{W}).") + start_latent_tensor = torch.zeros([1, latent_channels, 1, H, W], dtype=torch.float32) + else: + start_latent_tensor = start_latent["samples"] # Get tensor from dictionary + + # Get shape AFTER potentially creating the default + B, C, T, H, W = start_latent_tensor.shape + print(f"Latent dimensions: B={B}, C={C}, T={T}, H={H}, W={W}") + + start_latent_tensor = start_latent_tensor * vae_scaling_factor + + if initial_samples is not None: + initial_samples = initial_samples["samples"] * vae_scaling_factor + if end_latent is not None: + end_latent = end_latent["samples"] * vae_scaling_factor + has_end_image = end_latent is not None + + start_image_encoder_last_hidden_state = None # Initialize to None + if start_image_embeds is not None: + start_image_encoder_last_hidden_state = start_image_embeds["last_hidden_state"].to(base_dtype).to(device) + + end_image_encoder_last_hidden_state = None # Initialize to None + if has_end_image and embed_interpolation != "disabled" and end_image_embeds is not None: + end_image_encoder_last_hidden_state = end_image_embeds["last_hidden_state"].to(base_dtype).to(device) + elif start_image_encoder_last_hidden_state is not None: # Only create zeros if start exists + end_image_encoder_last_hidden_state = torch.zeros_like(start_image_encoder_last_hidden_state) + + # --- Conditioning Setup --- + # Negative conditioning + if not math.isclose(cfg, 1.0): + llama_vec_n = negative[0][0].to(dtype=base_dtype, device=device) + clip_l_pooler_n = negative[0][1]["pooled_output"].to(dtype=base_dtype, device=device) + llama_vec_n, llama_attention_mask_n = crop_or_pad_yield_mask(llama_vec_n, length=512) + else: + # Need dummy tensors with correct shape and device. Use shape from the first positive section. + if positive_timed_list: + try: + first_pos_cond = positive_timed_list[0][2][0][0].to(device=device) + first_pos_pooled = positive_timed_list[0][2][0][1]["pooled_output"].to(device=device) + llama_vec_n = torch.zeros_like(first_pos_cond) + clip_l_pooler_n = torch.zeros_like(first_pos_pooled) + # Still need to pad the zero tensor and get the mask + llama_vec_n, llama_attention_mask_n = crop_or_pad_yield_mask(llama_vec_n, length=512) + except Exception as e: + print(f"Error accessing positive_timed_list for negative shape when cfg=1.0: {e}. Creating fallback zero tensors.") + # Fallback zero tensors if list structure is unexpected or empty + llama_vec_n = torch.zeros((B, 512, 4096), dtype=base_dtype, device=device) # Guessing shape based on llama + llama_attention_mask_n = torch.ones((B, 512), dtype=torch.long, device=device) + clip_l_pooler_n = torch.zeros((B, 1280), dtype=base_dtype, device=device) # Guessing shape based on clip-l + else: + # This case remains the same - if no positive sections, create fallback zeros. + print("Warning: positive_timed_list is empty when cfg=1.0. Cannot determine negative shape. Creating fallback zero tensors.") + llama_vec_n = torch.zeros((B, 512, 4096), dtype=base_dtype, device=device) + llama_attention_mask_n = torch.ones((B, 512), dtype=torch.long, device=device) + clip_l_pooler_n = torch.zeros((B, 1280), dtype=base_dtype, device=device) + + # Positive conditioning: Handled inside the loop based on time. + # --- End Conditioning Setup --- + + # Sampling + rnd = torch.Generator("cpu").manual_seed(seed) + num_frames = latent_window_size * 4 - 3 # Frames generated per step + + # F1 History Latents Initialization + history_latents = torch.zeros(size=(B, 16, 16 + 2 + 1, H, W), dtype=torch.float32).cpu() + # F1: Start with the initial latent frame + history_latents = torch.cat([start_latent_tensor.to(history_latents)], dim=2) + total_generated_latent_frames = 1 # F1: Start count at 1, representing the initial frame + + # F1 Latent Paddings (determines number of generation steps) + latent_paddings = [1] * (total_latent_sections - 1) + [0] + latent_paddings_list = latent_paddings.copy() # For vid2vid indexing + + + comfy_model = HyVideoModel( + HyVideoModelConfig(base_dtype), + model_type=comfy.model_base.ModelType.FLOW, + device=device, + ) + + patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, torch.device("cpu")) + #from latent_preview import prepare_callback # Moved to top + callback = prepare_callback(patcher, steps) + + move_model_to_device_with_memory_preservation(transformer, target_device=device, preserved_memory_gb=gpu_memory_preservation) + + for i, latent_padding in enumerate(latent_paddings): + print(f"Sampling Section {i+1}/{total_latent_sections}, latent_padding: {latent_padding}") + is_last_section = latent_padding == 0 + + # F1 logic doesn't seem to use embed interpolation within the loop + # image_encoder_last_hidden_state = start_image_encoder_last_hidden_state * start_embed_strength + # ^-- This logic is removed as F1 doesn't use interpolation per step like the other sampler. + # We just pass the start_image_encoder_last_hidden_state directly to sample_hunyuan below. + # Handle case where image_embeds wasn't provided + current_image_embeds = start_image_encoder_last_hidden_state + + # --- Determine Current Positive Conditioning --- + # Calculate current time position based on the *start* of the section being generated + current_time_position = i * section_duration_sec + current_time_position = max(0.0, current_time_position) + print(f" Current time position: {current_time_position:.3f}s") + + active_section_index = -1 + if not positive_timed_list: + print("Error: positive_timed_list is empty! Cannot sample.") + # Handle error appropriately - maybe return black frames or raise exception? + # Returning empty/zeros for now + return {"samples": torch.zeros_like(start_latent_tensor) / vae_scaling_factor}, + + for idx, (start_sec, end_sec, _) in enumerate(positive_timed_list): + # Check if current_time_position falls within [start_sec, end_sec) + if start_sec <= current_time_position + 1e-4 and current_time_position < end_sec - 1e-4: + active_section_index = idx + # print(f" Found active prompt section index: {active_section_index} ({start_sec:.2f}s - {end_sec:.2f}s)") + break + else: + # If no section matches exactly, check edge cases + if math.isclose(current_time_position, positive_timed_list[-1][1], abs_tol=1e-4): + active_section_index = len(positive_timed_list) - 1 + # print(f" Time matches end of last section. Using index: {active_section_index}") + elif current_time_position >= positive_timed_list[-1][1] - 1e-4: + active_section_index = len(positive_timed_list) - 1 + # print(f" Time past end of last section. Using index: {active_section_index}") + elif current_time_position < positive_timed_list[0][0] + 1e-4: + active_section_index = 0 + # print(f" Time before first section. Using index: 0") + else: # Final fallback if list exists but no match (should be rare) + active_section_index = len(positive_timed_list) - 1 + print(f" Warning: No exact time match found, using last section index: {active_section_index}") + + print(f" Selected active prompt index: {active_section_index}") + + # --- Blending Logic --- + blend_alpha = 0.0 + prev_section_idx_for_blend = active_section_index + next_section_idx_for_blend = active_section_index + current_active_conditioning_tensor = positive_timed_list[active_section_index][2][0][0] + + # Find the index in the original list corresponding to the *start* of the next *different* conditioning + next_prompt_change_section_start_index = -1 + next_prompt_change_start_time = -1.0 + for k in range(active_section_index + 1, len(positive_timed_list)): + # Compare the actual conditioning data (tensors) + if not torch.equal(positive_timed_list[k][2][0][0], current_active_conditioning_tensor): + next_prompt_change_start_time = positive_timed_list[k][0] + next_prompt_change_section_start_index = int(round(next_prompt_change_start_time / section_duration_sec)) + prev_section_idx_for_blend = active_section_index # The prompt active before the change + next_section_idx_for_blend = k # The prompt active after the change + # print(f" Next prompt change detected at section index ~{next_prompt_change_section_start_index} (time {next_prompt_change_start_time:.2f}s)") + break + + # Check if we are within the blend window leading up to the change + if prompt_blend_sections > 0 and next_prompt_change_section_start_index != -1: + blend_start_section_idx = next_prompt_change_section_start_index - prompt_blend_sections + current_physical_section_idx = i # Use the actual loop iteration index + + if current_physical_section_idx >= blend_start_section_idx and current_physical_section_idx < next_prompt_change_section_start_index: + blend_progress = (current_physical_section_idx - blend_start_section_idx + 1) / float(prompt_blend_sections) + blend_alpha = max(0.0, min(1.0, blend_progress)) + print(f" Blending prompts: Section Index {current_physical_section_idx}, Blend Alpha: {blend_alpha:.3f}") + # No explicit 'else if >= next_prompt_change...' needed, blend_alpha remains 0 if not in window + + # --- End Blending Logic --- + + # Get the conditioning tensors + if blend_alpha > 0 and prev_section_idx_for_blend != next_section_idx_for_blend: + # Ensure indices are valid before accessing + if 0 <= prev_section_idx_for_blend < len(positive_timed_list) and 0 <= next_section_idx_for_blend < len(positive_timed_list): + cond_prev = positive_timed_list[prev_section_idx_for_blend][2][0][0].to(dtype=base_dtype, device=device) + pooled_prev = positive_timed_list[prev_section_idx_for_blend][2][0][1]['pooled_output'].to(dtype=base_dtype, device=device) + cond_next = positive_timed_list[next_section_idx_for_blend][2][0][0].to(dtype=base_dtype, device=device) + pooled_next = positive_timed_list[next_section_idx_for_blend][2][0][1]['pooled_output'].to(dtype=base_dtype, device=device) + + # Pad tensors before lerp + padded_cond_prev, mask_prev = crop_or_pad_yield_mask(cond_prev, length=512) + padded_cond_next, mask_next = crop_or_pad_yield_mask(cond_next, length=512) + + llama_vec = torch.lerp(padded_cond_prev, padded_cond_next, blend_alpha) + clip_l_pooler = torch.lerp(pooled_prev, pooled_next, blend_alpha) # Poolers assumed same shape + llama_attention_mask = mask_prev # Use mask from the first part of lerp + else: + print(f"Warning: Invalid blend indices ({prev_section_idx_for_blend}, {next_section_idx_for_blend}). Using non-blended active prompt.") + # Fallback to non-blended active prompt + selected_positive = positive_timed_list[active_section_index][2] + llama_vec = selected_positive[0][0].to(dtype=base_dtype, device=device) + clip_l_pooler = selected_positive[0][1]['pooled_output'].to(dtype=base_dtype, device=device) + llama_vec, llama_attention_mask = crop_or_pad_yield_mask(llama_vec, length=512) + else: + # Use the selected active conditioning directly + selected_positive = positive_timed_list[active_section_index][2] + llama_vec = selected_positive[0][0].to(dtype=base_dtype, device=device) + clip_l_pooler = selected_positive[0][1]['pooled_output'].to(dtype=base_dtype, device=device) + llama_vec, llama_attention_mask = crop_or_pad_yield_mask(llama_vec, length=512) + + # --- End Determine Current Positive Conditioning --- + + # F1 Indices Calculation + effective_window_size = int(latent_window_size) + indices = torch.arange(0, sum([1, 16, 2, 1, effective_window_size])).unsqueeze(0) + clean_latent_indices_start, clean_latent_4x_indices, clean_latent_2x_indices, clean_latent_1x_indices, latent_indices = indices.split([1, 16, 2, 1, effective_window_size], dim=1) + clean_latent_indices = torch.cat([clean_latent_indices_start, clean_latent_1x_indices], dim=1) + + # F1 Clean Latents Calculation + required_history_len = 16 + 2 + 1 # Need 19 previous frames + available_history_len = history_latents.shape[2] + + if available_history_len < required_history_len: + print(f"Warning: Not enough history frames ({available_history_len}) for clean latents (needed {required_history_len}). Padding with zeros.") + # Pad history_latents at the beginning with zeros to meet required length + padding_needed = required_history_len - available_history_len + padding_shape = list(history_latents.shape) + padding_shape[2] = padding_needed + zero_padding = torch.zeros(padding_shape, dtype=history_latents.dtype, device=history_latents.device) + padded_history = torch.cat([zero_padding, history_latents], dim=2) + clean_latents_4x, clean_latents_2x, clean_latents_1x = padded_history[:, :, -required_history_len:, :, :].split([16, 2, 1], dim=2) + else: + # Take the last 19 frames from history + clean_latents_4x, clean_latents_2x, clean_latents_1x = history_latents[:, :, -required_history_len:, :, :].split([16, 2, 1], dim=2) + + # Always prepend the original start_latent (frame 0) to clean_latents_1x (the most recent history frame) + clean_latents = torch.cat([start_latent_tensor.to(history_latents.device, dtype=history_latents.dtype), clean_latents_1x], dim=2) + + # vid2vid WIP (Using F1's method based on section index 'i') + input_init_latents = None + if initial_samples is not None: + total_length = initial_samples.shape[2] + # Use loop index 'i' for progress, mapping it to the vid2vid timeline + progress = i / (total_latent_sections - 1) if total_latent_sections > 1 else 0 + start_idx = int(progress * max(0, total_length - effective_window_size)) + end_idx = min(start_idx + effective_window_size, total_length) + # print(f"vid2vid (F1 logic) - Iteration {i}, Progress {progress:.2f}, Slice [{start_idx}:{end_idx}] of {total_length}") + if start_idx < end_idx: + input_init_latents = initial_samples[:, :, start_idx:end_idx, :, :].to(device) + else: + print("vid2vid - Warning: Calculated slice is empty.") + + if use_teacache: + transformer.initialize_teacache(enable_teacache=True, num_steps=steps, rel_l1_thresh=teacache_rel_l1_thresh) + else: + transformer.initialize_teacache(enable_teacache=False) + + with torch.autocast(device_type=mm.get_autocast_device(device), dtype=base_dtype, enabled=True): + generated_latents = sample_hunyuan( + transformer=transformer, + sampler=sampler, + initial_latent=input_init_latents, + strength=denoise_strength, + width=W * 8, + height=H * 8, + frames=num_frames, + real_guidance_scale=cfg, + distilled_guidance_scale=guidance_scale, + guidance_rescale=0, + shift=shift if shift != 0 else None, + num_inference_steps=steps, + generator=rnd, + prompt_embeds=llama_vec, + prompt_embeds_mask=llama_attention_mask, + prompt_poolers=clip_l_pooler, + negative_prompt_embeds=llama_vec_n, + negative_prompt_embeds_mask=llama_attention_mask_n, + negative_prompt_poolers=clip_l_pooler_n, + device=device, + dtype=base_dtype, + image_embeddings=current_image_embeds, + latent_indices=latent_indices, + clean_latents=clean_latents, + clean_latent_indices=clean_latent_indices, + clean_latents_2x=clean_latents_2x, + clean_latent_2x_indices=clean_latent_2x_indices, + clean_latents_4x=clean_latents_4x, + clean_latent_4x_indices=clean_latent_4x_indices, + callback=callback, + ) + + # F1 History Latents Update: Append new frames generated in this step + history_latents = torch.cat([history_latents, generated_latents.to(history_latents)], dim=2) + # Increment total frame count by the number of newly generated frames + total_generated_latent_frames += generated_latents.shape[2] + + # F1 Real History Latents Selection: Take from the end, ensuring we have `total_generated_latent_frames` count + real_history_latents = history_latents[:, :, -total_generated_latent_frames:, :, :] + + if is_last_section: + break + + transformer.to(offload_device) + mm.soft_empty_cache() + + # Ensure final output has the expected length (or close to it) + final_frame_count = real_history_latents.shape[2] + expected_latent_frames = total_generated_latent_frames # F1 should generate frame by frame + print(f"Final latent frames: {final_frame_count} (Expected based on generation: {expected_latent_frames})") + + return {"samples": real_history_latents / vae_scaling_factor}, + +class FramePackTimestampedTextEncode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "clip": ("CLIP", ), + "text": ("STRING", {"multiline": True, "dynamicPrompts": True, "tooltip": "Text prompt, use [Xs: prompt] or [Xs-Ys: prompt] for timed sections."}), + "negative_text": ("STRING", {"multiline": False, "default": "", "dynamicPrompts": False, "tooltip": "Single negative text prompt"}), + "total_second_length": ("FLOAT", {"default": 5.0, "min": 0.1, "max": 1200.0, "step": 0.1, "tooltip": "Expected total video duration in seconds for timestamp calculation."}), + "latent_window_size": ("INT", {"default": 9, "min": 1, "max": 33, "step": 1, "tooltip": "The latent window size used by the sampler for timestamp boundary snapping."}), + "prompt_blend_sections": ("INT", {"default": 0, "min": 0, "max": 10, "step": 1, "tooltip": "Number of latent sections (windows) over which to blend prompts when they change. 0 disables blending."}), + }, + } + RETURN_TYPES = ("TIMED_CONDITIONING_WITH_METADATA", "CONDITIONING",) + RETURN_NAMES = ("positive_timed_data", "negative",) + FUNCTION = "encode" + CATEGORY = "FramePackWrapper/experimental" + DESCRIPTION = """Encodes text prompts with optional timestamps for timed conditioning. + +Use format: [Xs: prompt] or [Xs-Ys: prompt] where X and Y are times in seconds (e.g., 0s, 1.5s, 10s). +- [Xs: prompt]: Prompt applies from time X until the next timestamp starts (or end of video). +- [Xs-Ys: prompt]: Prompt applies specifically between time X and time Y. + +Text before the first timestamp defaults to starting at 0s. +Gaps between specified timestamps are automatically filled, typically using the preceding prompt. +Timestamps are aligned to internal section boundaries based on latent_window_size. + +Outputs a dictionary containing: +- timed conditioning sections: List of (start_sec, end_sec, conditioning) tuples defining the prompt for each time segment. +- total duration: The overall video length in seconds, used for time calculations. +- latent window size: The sampler's processing window size, used for aligning timestamps. +- prompt blend sections: Number of sections over which to smoothly blend between changing prompts(if you want smoother visual transitions when your timed prompts change. A higher value gives a longer, more gradual blend). +""" + + def encode(self, clip, text, negative_text, total_second_length, latent_window_size, prompt_blend_sections): + prompt_sections = parse_timestamped_prompt_f1(text, total_second_length, latent_window_size) + unique_prompts = sorted(list(set(section.prompt for section in prompt_sections))) + encoded_prompts: Dict[str, List[List[Union[torch.Tensor, Dict[str, torch.Tensor]]]]] = {} + first_cond, first_pooled = None, None + + print(f"FramePackTimestampedTextEncode: Encoding {len(unique_prompts)} unique prompts.") + for i, prompt in enumerate(unique_prompts): + tokens = clip.tokenize(prompt) + cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True) + if i == 0: + first_cond, first_pooled = cond, pooled + encoded_prompts[prompt] = [[cond, {"pooled_output": pooled}]] + + positive_timed_list: List[Tuple[float, float, List[List[Union[torch.Tensor, Dict[str, torch.Tensor]]]]]] = [] + for section in prompt_sections: + if section.prompt in encoded_prompts: + encoded_cond = encoded_prompts[section.prompt] + positive_timed_list.append((section.start_time, section.end_time, encoded_cond)) + else: + print(f"Warning: Prompt '{section.prompt}' not found in encoded prompts. Skipping section.") + + if not positive_timed_list: + print("FramePackTimestampedTextEncode: Warning - No valid timed sections found. Creating a default empty section.") + tokens = clip.tokenize("") + cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True) + if first_cond is None: first_cond, first_pooled = cond, pooled # Store shape if needed + positive_timed_list.append((0.0, total_second_length, [[cond, {"pooled_output": pooled}]])) # Ensure list structure is maintained + + # --- Negative Conditioning --- + if negative_text: + tokens_neg = clip.tokenize(negative_text) + cond_neg, pooled_neg = clip.encode_from_tokens(tokens_neg, return_pooled=True) + negative = [[cond_neg, {"pooled_output": pooled_neg}]] + elif first_cond is not None: + negative = [[torch.zeros_like(first_cond), {"pooled_output": torch.zeros_like(first_pooled)}]] + else: + print("FramePackTimestampedTextEncode: Error - Cannot create empty negative conditioning, no positive prompts found and fallback failed.") + try: + tokens = clip.tokenize("") + cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True) + negative = [[torch.zeros_like(cond), {"pooled_output": torch.zeros_like(pooled)}]] + except Exception as e: + print(f"Fallback negative shape guess failed: {e}") + # Minimal fallback guess + negative = [[torch.zeros((1, 77, 768)), {"pooled_output": torch.zeros((1, 768))}]] + + # Package results into a dictionary + timed_data = { + "sections": positive_timed_list, + "total_duration": total_second_length, + "window_size": latent_window_size, + "blend_sections": prompt_blend_sections + } + return (timed_data, negative) + +NODE_CLASS_MAPPINGS = { + "FramePackSampler_F1": FramePackSampler_F1, + "FramePackTimestampedTextEncode": FramePackTimestampedTextEncode, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "FramePackSampler_F1": "FramePackSampler (F1)", + "FramePackTimestampedTextEncode": "FramePack Text Encode (Timestamped)", +} \ No newline at end of file diff --git a/transformer_config.json b/transformer_config.json new file mode 100644 index 0000000..844039f --- /dev/null +++ b/transformer_config.json @@ -0,0 +1,28 @@ +{ + "_class_name": "HunyuanVideoTransformer3DModelPacked", + "_diffusers_version": "0.33.0.dev0", + "_name_or_path": "hunyuanvideo-community/HunyuanVideo", + "attention_head_dim": 128, + "guidance_embeds": true, + "has_clean_x_embedder": true, + "has_image_proj": true, + "image_proj_dim": 1152, + "in_channels": 16, + "mlp_ratio": 4.0, + "num_attention_heads": 24, + "num_layers": 20, + "num_refiner_layers": 2, + "num_single_layers": 40, + "out_channels": 16, + "patch_size": 2, + "patch_size_t": 1, + "pooled_projection_dim": 768, + "qk_norm": "rms_norm", + "rope_axes_dim": [ + 16, + 56, + 56 + ], + "rope_theta": 256.0, + "text_embed_dim": 4096 +} diff --git a/utils.py b/utils.py new file mode 100644 index 0000000..b981dff --- /dev/null +++ b/utils.py @@ -0,0 +1,93 @@ +import importlib.metadata +import torch +import logging +logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') +log = logging.getLogger(__name__) + +def check_diffusers_version(): + try: + version = importlib.metadata.version('diffusers') + required_version = '0.31.0' + if version < required_version: + raise AssertionError(f"diffusers version {version} is installed, but version {required_version} or higher is required.") + except importlib.metadata.PackageNotFoundError: + raise AssertionError("diffusers is not installed.") + +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}") + +def convert_to_diffusers(prefix, weights_sd): + # convert from default LoRA to diffusers + # https://github.com/kohya-ss/musubi-tuner/blob/main/convert_lora.py + + # get alphas + lora_alphas = {} + for key, weight in weights_sd.items(): + if key.startswith(prefix): + lora_name = key.split(".", 1)[0] # before first dot + if lora_name not in lora_alphas and "alpha" in key: + lora_alphas[lora_name] = weight + + new_weights_sd = {} + for key, weight in weights_sd.items(): + if key.startswith(prefix): + if "alpha" in key: + continue + + lora_name = key.split(".", 1)[0] # before first dot + + module_name = lora_name[len(prefix) :] # remove "lora_unet_" + module_name = module_name.replace("_", ".") # replace "_" with "." + + # HunyuanVideo lora name to module name: ugly but works + #module_name = module_name.replace("double.blocks.", "double_blocks.") # fix double blocks + module_name = module_name.replace("single.transformer.blocks.", "single_transformer_blocks.") # fix single blocks + module_name = module_name.replace("transformer.blocks.", "transformer_blocks.") # fix double blocks + + module_name = module_name.replace("img.", "img_") # fix img + module_name = module_name.replace("txt.", "txt_") # fix txt + module_name = module_name.replace("to.q", "to_q") # fix attn + module_name = module_name.replace("to.k", "to_k") + module_name = module_name.replace("to.v", "to_v") + module_name = module_name.replace("to.add.out", "to_add_out") + module_name = module_name.replace("add.k.proj", "add_k_proj") + module_name = module_name.replace("add.q.proj", "add_q_proj") + module_name = module_name.replace("add.v.proj", "add_v_proj") + module_name = module_name.replace("add.out.proj", "add_out_proj") + module_name = module_name.replace("proj.", "proj_") # fix proj + module_name = module_name.replace("to.out", "to_out") # fix to_out + module_name = module_name.replace("ff.context", "ff_context") # fix ff context + + diffusers_prefix = "transformer" + if "lora_down" in key: + new_key = f"{diffusers_prefix}.{module_name}.lora_A.weight" + dim = weight.shape[0] + elif "lora_up" in key: + new_key = f"{diffusers_prefix}.{module_name}.lora_B.weight" + dim = weight.shape[1] + else: + log.warning(f"unexpected key: {key} in default LoRA format") + continue + + # scale weight by alpha + if lora_name in lora_alphas: + # we scale both down and up, so scale is sqrt + scale = lora_alphas[lora_name] / dim + scale = scale.sqrt() + weight = weight * scale + else: + log.warning(f"missing alpha for {lora_name}") + + new_weights_sd[new_key] = weight + + return new_weights_sd +