v1
This commit is contained in:
@@ -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']
|
||||
@@ -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))
|
||||
@@ -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
@@ -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)",
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user