This commit is contained in:
Shmuel Ronen
2025-05-17 17:42:21 +03:00
committed by GitHub
parent e928f4a02a
commit 9c6c7873d6
6 changed files with 1740 additions and 0 deletions
+8
View File
@@ -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']
+39
View File
@@ -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))
+937
View File
@@ -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",
}
+635
View File
@@ -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)",
}
+28
View File
@@ -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
}
+93
View File
@@ -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