Merge branch 'pr/72'
This commit is contained in:
@@ -1,9 +1,8 @@
|
||||
import os
|
||||
import torch
|
||||
import json
|
||||
from typing import List
|
||||
import gc
|
||||
from .utils import log, check_diffusers_version, print_memory
|
||||
from .utils import log, print_memory
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
|
||||
from .hyvideo.constants import PROMPT_TEMPLATE
|
||||
@@ -20,6 +19,8 @@ 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
|
||||
import comfy.model_base
|
||||
import comfy.latent_formats
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
@@ -56,7 +57,7 @@ def get_rotary_pos_embed(transformer, video_length, height, width):
|
||||
|
||||
if len(rope_sizes) != target_ndim:
|
||||
rope_sizes = [1] * (target_ndim - len(rope_sizes)) + rope_sizes # time axis
|
||||
|
||||
|
||||
if rope_dim_list is None:
|
||||
rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)]
|
||||
assert (
|
||||
@@ -71,11 +72,89 @@ def get_rotary_pos_embed(transformer, video_length, height, width):
|
||||
)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
def filter_state_dict_by_blocks(state_dict, blocks_mapping):
|
||||
filtered_dict = {}
|
||||
|
||||
for key in state_dict:
|
||||
if 'double_blocks.' in key or 'single_blocks.' in key:
|
||||
block_pattern = key.split('diffusion_model.')[1].split('.', 2)[0:2]
|
||||
block_key = f'{block_pattern[0]}.{block_pattern[1]}.'
|
||||
|
||||
if block_key in blocks_mapping:
|
||||
filtered_dict[key] = state_dict[key]
|
||||
|
||||
return filtered_dict
|
||||
|
||||
class HyVideoLoraBlockEdit:
|
||||
def __init__(self):
|
||||
self.loaded_lora = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
arg_dict = {}
|
||||
argument = ("BOOLEAN", {"default": True})
|
||||
|
||||
for i in range(20):
|
||||
arg_dict["double_blocks.{}.".format(i)] = argument
|
||||
|
||||
for i in range(40):
|
||||
arg_dict["single_blocks.{}.".format(i)] = argument
|
||||
|
||||
return {"required": arg_dict}
|
||||
|
||||
RETURN_TYPES = ("SELECTEDBLOCKS", )
|
||||
RETURN_NAMES = ("blocks", )
|
||||
OUTPUT_TOOLTIPS = ("The modified diffusion model.",)
|
||||
FUNCTION = "select"
|
||||
|
||||
CATEGORY = "HunyuanVideoWrapper"
|
||||
|
||||
def select(self, **kwargs):
|
||||
selected_blocks = {k: v for k, v in kwargs.items() if v is True}
|
||||
print("Selected blocks: ", selected_blocks)
|
||||
return (selected_blocks,)
|
||||
class HyVideoLoraSelect:
|
||||
@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"}),
|
||||
},
|
||||
"optional": {
|
||||
"prev_lora":("HYVIDLORA", {"default": None, "tooltip": "For loading multiple LoRAs"}),
|
||||
"blocks":("SELECTEDBLOCKS", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("HYVIDLORA",)
|
||||
RETURN_NAMES = ("lora", )
|
||||
FUNCTION = "getlorapath"
|
||||
CATEGORY = "HunyuanVideoWrapper"
|
||||
DESCRIPTION = "Select a LoRA model from ComfyUI/models/loras"
|
||||
|
||||
def getlorapath(self, lora, strength, blocks=None, prev_lora=None, fuse_lora=False):
|
||||
loras_list = []
|
||||
|
||||
lora = {
|
||||
"path": folder_paths.get_full_path("loras", lora),
|
||||
"strength": strength,
|
||||
"name": lora.split(".")[0],
|
||||
"fuse_lora": fuse_lora,
|
||||
"blocks": blocks
|
||||
}
|
||||
if prev_lora is not None:
|
||||
loras_list.extend(prev_lora)
|
||||
|
||||
loras_list.append(lora)
|
||||
return (loras_list,)
|
||||
|
||||
class HyVideoBlockSwap:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"required": {
|
||||
"double_blocks_to_swap": ("INT", {"default": 20, "min": 0, "max": 20, "step": 1, "tooltip": "Number of double blocks to swap"}),
|
||||
"single_blocks_to_swap": ("INT", {"default": 0, "min": 0, "max": 40, "step": 1, "tooltip": "Number of single blocks to swap"}),
|
||||
"offload_txt_in": ("BOOLEAN", {"default": False, "tooltip": "Offload txt_in layer"}),
|
||||
@@ -95,7 +174,7 @@ class HyVideoSTG:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"required": {
|
||||
"stg_mode": (["STG-A", "STG-R"],),
|
||||
"stg_block_idx": ("INT", {"default": 0, "min": -1, "max": 39, "step": 1, "tooltip": "Block index to apply STG"}),
|
||||
"stg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Recommended values are ≤2.0"}),
|
||||
@@ -112,6 +191,33 @@ class HyVideoSTG:
|
||||
def setargs(self, **kwargs):
|
||||
return (kwargs, )
|
||||
|
||||
|
||||
class HyVideoModel(comfy.model_base.BaseModel):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.pipeline = {}
|
||||
|
||||
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.LatentFormat()
|
||||
self.latent_format.latent_channels = 16
|
||||
self.manual_cast_dtype = dtype
|
||||
self.sampling_settings = {"multiplier": 1.0}
|
||||
# Don't know what this is. Value taken from ComfyUI Mochi model.
|
||||
self.memory_usage_factor = 2.0
|
||||
# denoiser is handled by extension
|
||||
self.unet_config["disable_unet_model_creation"] = True
|
||||
|
||||
|
||||
#region Model loading
|
||||
class HyVideoModelLoader:
|
||||
@classmethod
|
||||
@@ -119,8 +225,8 @@ class HyVideoModelLoader:
|
||||
return {
|
||||
"required": {
|
||||
"model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}),
|
||||
|
||||
"base_precision": (["fp16", "fp32", "bf16"], {"default": "bf16"}),
|
||||
|
||||
"base_precision": (["fp32", "bf16"], {"default": "bf16"}),
|
||||
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'torchao_fp8dq', "torchao_fp8dqrow", "torchao_int8dq", "torchao_fp6", "torchao_int4", "torchao_int8"], {"default": 'disabled', "tooltip": "optional quantization method"}),
|
||||
"load_device": (["main_device", "offload_device"], {"default": "main_device"}),
|
||||
},
|
||||
@@ -133,6 +239,7 @@ class HyVideoModelLoader:
|
||||
], {"default": "flash_attn"}),
|
||||
"compile_args": ("COMPILEARGS", ),
|
||||
"block_swap_args": ("BLOCKSWAPARGS", ),
|
||||
"lora": ("HYVIDLORA", {"default": None}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -142,7 +249,7 @@ class HyVideoModelLoader:
|
||||
CATEGORY = "HunyuanVideoWrapper"
|
||||
|
||||
def loadmodel(self, model, base_precision, load_device, quantization,
|
||||
compile_args=None, attention_mode="sdpa", block_swap_args=None):
|
||||
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None):
|
||||
transformer = None
|
||||
manual_offloading = True
|
||||
if "sage" in attention_mode:
|
||||
@@ -154,7 +261,7 @@ class HyVideoModelLoader:
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
manual_offloading = True
|
||||
transformer_load_device = device if load_device == "main_device" else offload_device
|
||||
transformer_load_device = device if load_device == "main_device" or lora is None else offload_device
|
||||
mm.soft_empty_cache()
|
||||
|
||||
base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[base_precision]
|
||||
@@ -184,7 +291,74 @@ class HyVideoModelLoader:
|
||||
**factor_kwargs
|
||||
)
|
||||
transformer.eval()
|
||||
if "torchao" in quantization:
|
||||
|
||||
comfy_model = HyVideoModel(
|
||||
HyVideoModelConfig(base_dtype),
|
||||
model_type=comfy.model_base.ModelType.FLOW,
|
||||
device=device,
|
||||
)
|
||||
scheduler = FlowMatchDiscreteScheduler(
|
||||
shift=9.0,
|
||||
reverse=True,
|
||||
solver="euler",
|
||||
)
|
||||
pipe = HunyuanVideoPipeline(
|
||||
transformer=transformer,
|
||||
scheduler=scheduler,
|
||||
progress_bar_config=None,
|
||||
base_dtype=base_dtype
|
||||
)
|
||||
|
||||
if not "torchao" in quantization:
|
||||
log.info("Using accelerate to load and assign model weights to device...")
|
||||
if quantization == "fp8_e4m3fn" or quantization == "fp8_e4m3fn_fast":
|
||||
dtype = torch.float8_e4m3fn
|
||||
else:
|
||||
dtype = base_dtype
|
||||
params_to_keep = {"norm", "bias", "time_in", "vector_in", "guidance_in", "txt_in", "img_in"}
|
||||
for name, param in transformer.named_parameters():
|
||||
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])
|
||||
|
||||
comfy_model.diffusion_model = transformer
|
||||
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
|
||||
|
||||
if lora is not None:
|
||||
from comfy.sd import load_lora_for_models
|
||||
for l in lora:
|
||||
lora_path = l["path"]
|
||||
lora_strength = l["strength"]
|
||||
lora_sd = load_torch_file(lora_path, safe_load=True)
|
||||
if l["blocks"]:
|
||||
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"])
|
||||
|
||||
#for k in lora_sd.keys():
|
||||
# print(k)
|
||||
|
||||
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
|
||||
|
||||
comfy.model_management.load_models_gpu([patcher])
|
||||
|
||||
if quantization == "fp8_e4m3fn_fast":
|
||||
from .fp8_optimization import convert_fp8_linear
|
||||
convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep)
|
||||
|
||||
#compile
|
||||
if compile_args is not None:
|
||||
torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"]
|
||||
if compile_args["compile_single_blocks"]:
|
||||
for i, block in enumerate(patcher.model.diffusion_model.single_blocks):
|
||||
patcher.model.diffusion_model.single_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(patcher.model.diffusion_model.double_blocks):
|
||||
patcher.model.diffusion_model.double_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_txt_in"]:
|
||||
patcher.model.diffusion_model.txt_in = torch.compile(patcher.model.diffusion_model.txt_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
if compile_args["compile_vector_in"]:
|
||||
patcher.model.diffusion_model.vector_in = torch.compile(patcher.model.diffusion_model.vector_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
if compile_args["compile_final_layer"]:
|
||||
patcher.model.diffusion_model.final_layer = torch.compile(patcher.model.diffusion_model.final_layer, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
elif "torchao" in quantization:
|
||||
try:
|
||||
from torchao.quantization import (
|
||||
quantize_,
|
||||
@@ -195,14 +369,14 @@ class HyVideoModelLoader:
|
||||
int4_weight_only
|
||||
)
|
||||
except:
|
||||
raise ImportError("torchao is not installed, please install torchao to use fp8dq")
|
||||
raise ImportError("torchao is not installed")
|
||||
|
||||
# def filter_fn(module: nn.Module, fqn: str) -> bool:
|
||||
# target_submodules = {'attn1', 'ff'} # avoid norm layers, 1.5 at least won't work with quantized norm1 #todo: test other models
|
||||
# if any(sub in fqn for sub in target_submodules):
|
||||
# return isinstance(module, nn.Linear)
|
||||
# return False
|
||||
|
||||
|
||||
if "fp6" in quantization:
|
||||
quant_func = fpx_weight_only(3, 2)
|
||||
elif "int4" in quantization:
|
||||
@@ -218,93 +392,60 @@ class HyVideoModelLoader:
|
||||
quant_func = int8_dynamic_activation_int8_weight()
|
||||
|
||||
log.info(f"Quantizing model with {quant_func}")
|
||||
comfy_model.diffusion_model = transformer
|
||||
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
|
||||
|
||||
for i, block in enumerate(transformer.single_blocks):
|
||||
if lora is not None:
|
||||
from comfy.sd import load_lora_for_models
|
||||
for l in lora:
|
||||
lora_path = l["path"]
|
||||
lora_strength = l["strength"]
|
||||
lora_sd = load_torch_file(lora_path, safe_load=True)
|
||||
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
|
||||
|
||||
comfy.model_management.load_models_gpu([patcher])
|
||||
|
||||
for i, block in enumerate(patcher.model.diffusion_model.single_blocks):
|
||||
log.info(f"Quantizing single_block {i}")
|
||||
for name, _ in block.named_parameters(prefix=f"single_blocks.{i}"):
|
||||
#print(f"Parameter name: {name}")
|
||||
set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=base_dtype, value=sd[name])
|
||||
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name])
|
||||
if compile_args is not None:
|
||||
transformer.single_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
patcher.model.diffusion_model.single_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
quantize_(block, quant_func)
|
||||
print(block)
|
||||
block.to(offload_device)
|
||||
for i, block in enumerate(transformer.double_blocks):
|
||||
for i, block in enumerate(patcher.model.diffusion_model.double_blocks):
|
||||
log.info(f"Quantizing double_block {i}")
|
||||
for name, _ in block.named_parameters(prefix=f"double_blocks.{i}"):
|
||||
#print(f"Parameter name: {name}")
|
||||
set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=base_dtype, value=sd[name])
|
||||
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name])
|
||||
if compile_args is not None:
|
||||
transformer.double_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
patcher.model.diffusion_model.double_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
quantize_(block, quant_func)
|
||||
for name, param in transformer.named_parameters():
|
||||
for name, param in patcher.model.diffusion_model.named_parameters():
|
||||
if "single_blocks" not in name and "double_blocks" not in name:
|
||||
set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=base_dtype, value=sd[name])
|
||||
|
||||
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name])
|
||||
|
||||
manual_offloading = False # to disable manual .to(device) calls
|
||||
log.info(f"Quantized transformer blocks to {quantization}")
|
||||
for name, param in transformer.named_parameters():
|
||||
for name, param in patcher.model.diffusion_model.named_parameters():
|
||||
print(name, param.dtype)
|
||||
#param.data = param.data.to(self.vae_dtype).to(device)
|
||||
else:
|
||||
log.info("Using accelerate to load and assign model weights to device...")
|
||||
if quantization == "fp8_e4m3fn" or quantization == "fp8_e4m3fn_fast":
|
||||
dtype = torch.float8_e4m3fn
|
||||
else:
|
||||
dtype = base_dtype
|
||||
params_to_keep = {"norm", "bias", "time_in", "vector_in", "guidance_in", "txt_in", "img_in"}
|
||||
for name, param in transformer.named_parameters():
|
||||
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 quantization == "fp8_e4m3fn_fast":
|
||||
from .fp8_optimization import convert_fp8_linear
|
||||
convert_fp8_linear(transformer, base_dtype, params_to_keep=params_to_keep)
|
||||
|
||||
#compile
|
||||
if compile_args is not None:
|
||||
torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"]
|
||||
if compile_args["compile_single_blocks"]:
|
||||
for i, block in enumerate(transformer.single_blocks):
|
||||
transformer.single_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.double_blocks):
|
||||
transformer.double_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_txt_in"]:
|
||||
transformer.txt_in = torch.compile(transformer.txt_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
if compile_args["compile_vector_in"]:
|
||||
transformer.vector_in = torch.compile(transformer.vector_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
if compile_args["compile_final_layer"]:
|
||||
transformer.final_layer = torch.compile(transformer.final_layer, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
|
||||
del sd
|
||||
mm.soft_empty_cache()
|
||||
|
||||
scheduler = FlowMatchDiscreteScheduler(
|
||||
shift=9.0,
|
||||
reverse=True,
|
||||
solver="euler",
|
||||
)
|
||||
|
||||
pipe = HunyuanVideoPipeline(
|
||||
transformer=transformer,
|
||||
scheduler=scheduler,
|
||||
progress_bar_config=None,
|
||||
base_dtype=base_dtype
|
||||
)
|
||||
|
||||
pipeline = {
|
||||
"pipe": pipe,
|
||||
"dtype": base_dtype,
|
||||
"base_path": model_path,
|
||||
"model_name": model,
|
||||
"manual_offloading": manual_offloading,
|
||||
"quantization": "disabled",
|
||||
"block_swap_args": block_swap_args,
|
||||
}
|
||||
return (pipeline,)
|
||||
|
||||
patcher.model["pipe"] = pipe
|
||||
patcher.model["dtype"] = base_dtype
|
||||
patcher.model["base_path"] = model_path
|
||||
patcher.model["model_name"] = model
|
||||
patcher.model["manual_offloading"] = manual_offloading
|
||||
patcher.model["quantization"] = "disabled"
|
||||
patcher.model["block_swap_args"] = block_swap_args
|
||||
|
||||
return (patcher,)
|
||||
|
||||
#region load VAE
|
||||
|
||||
class HyVideoVAELoader:
|
||||
@@ -329,7 +470,7 @@ class HyVideoVAELoader:
|
||||
DESCRIPTION = "Loads Hunyuan VAE model from 'ComfyUI/models/vae'"
|
||||
|
||||
def loadmodel(self, model_name, precision, compile_args=None):
|
||||
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
@@ -345,21 +486,21 @@ class HyVideoVAELoader:
|
||||
vae.requires_grad_(False)
|
||||
vae.eval()
|
||||
vae.to(device = device, dtype = dtype)
|
||||
|
||||
|
||||
#compile
|
||||
if compile_args is not None:
|
||||
torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"]
|
||||
vae = torch.compile(vae, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
|
||||
return (vae,)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
class HyVideoTorchCompileSettings:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"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"}),
|
||||
@@ -395,7 +536,7 @@ class HyVideoTorchCompileSettings:
|
||||
}
|
||||
|
||||
return (compile_args, )
|
||||
|
||||
|
||||
#region TextEncode
|
||||
|
||||
class DownloadAndLoadHyVideoTextEncoder:
|
||||
@@ -405,7 +546,7 @@ class DownloadAndLoadHyVideoTextEncoder:
|
||||
"required": {
|
||||
"llm_model": (["Kijai/llava-llama-3-8b-text-encoder-tokenizer",],),
|
||||
"clip_model": (["disabled","openai/clip-vit-large-patch14",],),
|
||||
|
||||
|
||||
"precision": (["fp16", "fp32", "bf16"],
|
||||
{"default": "bf16"}
|
||||
),
|
||||
@@ -424,7 +565,7 @@ class DownloadAndLoadHyVideoTextEncoder:
|
||||
DESCRIPTION = "Loads Hunyuan text_encoder model from 'ComfyUI/models/LLM'"
|
||||
|
||||
def loadmodel(self, llm_model, clip_model, precision, apply_final_norm=False, hidden_state_skip_layer=2, quantization="disabled"):
|
||||
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||
@@ -473,7 +614,7 @@ class DownloadAndLoadHyVideoTextEncoder:
|
||||
local_dir=base_path,
|
||||
local_dir_use_symlinks=False,
|
||||
)
|
||||
|
||||
|
||||
text_encoder = TextEncoder(
|
||||
text_encoder_path=base_path,
|
||||
text_encoder_type="llm",
|
||||
@@ -487,13 +628,13 @@ class DownloadAndLoadHyVideoTextEncoder:
|
||||
dtype=dtype,
|
||||
quantization_config=quantization_config
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
hyvid_text_encoders = {
|
||||
"text_encoder": text_encoder,
|
||||
"text_encoder_2": text_encoder_2,
|
||||
}
|
||||
|
||||
|
||||
return (hyvid_text_encoders,)
|
||||
|
||||
class HyVideoCustomPromptTemplate:
|
||||
@@ -516,7 +657,7 @@ class HyVideoCustomPromptTemplate:
|
||||
"crop_start": crop_start,
|
||||
}
|
||||
return (prompt_template_dict,)
|
||||
|
||||
|
||||
class HyVideoTextEncode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -549,7 +690,7 @@ class HyVideoTextEncode:
|
||||
text_encoder_2 = None
|
||||
|
||||
negative_prompt = None
|
||||
|
||||
|
||||
if prompt_template != "disabled":
|
||||
if prompt_template == "custom":
|
||||
prompt_template_dict = custom_prompt_template
|
||||
@@ -574,12 +715,12 @@ class HyVideoTextEncode:
|
||||
batch_size = 1
|
||||
num_videos_per_prompt = 1
|
||||
do_classifier_free_guidance = False # not implemented, for now we only have cfg distilled model
|
||||
|
||||
|
||||
text_inputs = text_encoder.text2tokens(prompt, prompt_template=prompt_template_dict)
|
||||
|
||||
prompt_outputs = text_encoder.encode(text_inputs, prompt_template=prompt_template_dict, device=device)
|
||||
prompt_embeds = prompt_outputs.hidden_state
|
||||
|
||||
|
||||
attention_mask = prompt_outputs.attention_mask
|
||||
log.info(f"{text_encoder.text_encoder_type} prompt attention_mask shape: {attention_mask.shape}, masked tokens: {attention_mask[0].sum().item()}")
|
||||
if attention_mask is not None:
|
||||
@@ -696,7 +837,7 @@ class HyVideoTextEncode:
|
||||
prompt_embeds_2 = prompt_embeds_2.to(device=device)
|
||||
|
||||
negative_prompt_embeds_2, attention_mask_2, negative_attention_mask_2 = None, None, None
|
||||
|
||||
|
||||
if force_offload:
|
||||
clip_l.cond_stage_model.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
@@ -705,7 +846,7 @@ class HyVideoTextEncode:
|
||||
negative_prompt_embeds_2 = None
|
||||
attention_mask_2 = None
|
||||
negative_attention_mask_2 = None
|
||||
|
||||
|
||||
prompt_embeds_dict = {
|
||||
"prompt_embeds": prompt_embeds,
|
||||
"negative_prompt_embeds": negative_prompt_embeds,
|
||||
@@ -717,9 +858,9 @@ class HyVideoTextEncode:
|
||||
"negative_attention_mask_2": negative_attention_mask_2,
|
||||
}
|
||||
return (prompt_embeds_dict,)
|
||||
|
||||
|
||||
#region Sampler
|
||||
|
||||
#region Sampler
|
||||
class HyVideoSampler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -735,7 +876,7 @@ class HyVideoSampler:
|
||||
"flow_shift": ("FLOAT", {"default": 9.0, "min": 0.0, "max": 30.0, "step": 0.01}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"force_offload": ("BOOLEAN", {"default": True}),
|
||||
|
||||
|
||||
},
|
||||
"optional": {
|
||||
"samples": ("LATENT", {"tooltip": "init Latents to use for video2video process"} ),
|
||||
@@ -750,6 +891,8 @@ class HyVideoSampler:
|
||||
CATEGORY = "HunyuanVideoWrapper"
|
||||
|
||||
def process(self, model, hyvid_embeds, flow_shift, steps, embedded_guidance_scale, seed, width, height, num_frames, samples=None, denoise_strength=1.0, force_offload=True, stg_args=None):
|
||||
model = model.model
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
dtype = model["dtype"]
|
||||
@@ -760,7 +903,7 @@ class HyVideoSampler:
|
||||
raise ValueError(
|
||||
f"STG-A requires attention_mode to be 'sdpa', but got {transformer.attention_mode}."
|
||||
)
|
||||
|
||||
|
||||
generator = torch.Generator(device=torch.device("cpu")).manual_seed(seed)
|
||||
|
||||
if width <= 0 or height <= 0 or num_frames <= 0:
|
||||
@@ -785,7 +928,7 @@ class HyVideoSampler:
|
||||
n_tokens = freqs_cos.shape[0]
|
||||
|
||||
model["pipe"].scheduler.shift = flow_shift
|
||||
|
||||
|
||||
# autocast_context = torch.autocast(
|
||||
# mm.get_autocast_device(device), dtype=dtype
|
||||
# ) if any(q in model["quantization"] for q in ("e4m3fn", "GGUF")) else nullcontext()
|
||||
@@ -795,9 +938,9 @@ class HyVideoSampler:
|
||||
#print(name, param.data.device)
|
||||
if "single" not in name and "double" not in name:
|
||||
param.data = param.data.to(device)
|
||||
|
||||
|
||||
transformer.block_swap(
|
||||
model["block_swap_args"]["double_blocks_to_swap"] - 1 ,
|
||||
model["block_swap_args"]["double_blocks_to_swap"] - 1 ,
|
||||
model["block_swap_args"]["single_blocks_to_swap"] - 1,
|
||||
offload_txt_in = model["block_swap_args"]["offload_txt_in"],
|
||||
offload_img_in = model["block_swap_args"]["offload_img_in"],
|
||||
@@ -808,18 +951,17 @@ class HyVideoSampler:
|
||||
elif model["manual_offloading"]:
|
||||
transformer.to(device)
|
||||
|
||||
mm.unload_all_models()
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
|
||||
|
||||
#for name, param in transformer.named_parameters():
|
||||
# print(name, param.data.device)
|
||||
|
||||
|
||||
out_latents = model["pipe"](
|
||||
num_inference_steps=steps,
|
||||
height = target_height,
|
||||
@@ -845,7 +987,7 @@ class HyVideoSampler:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
|
||||
|
||||
if force_offload:
|
||||
if model["manual_offloading"]:
|
||||
transformer.to(offload_device)
|
||||
@@ -855,8 +997,8 @@ class HyVideoSampler:
|
||||
return ({
|
||||
"samples": out_latents
|
||||
},)
|
||||
|
||||
#region VideoDecode
|
||||
|
||||
#region VideoDecode
|
||||
class HyVideoDecode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -867,7 +1009,7 @@ class HyVideoDecode:
|
||||
"temporal_tiling_sample_size": ("INT", {"default": 64, "min": 4, "max": 256, "tooltip": "Smaller values use less VRAM, model default is 64, any other value will cause stutter"}),
|
||||
"spatial_tile_sample_min_size": ("INT", {"default": 256, "min": 32, "max": 2048, "step": 32, "tooltip": "Spatial tile minimum size in pixels, smaller values use less VRAM, may introduce more seams"}),
|
||||
"auto_tile_size": ("BOOLEAN", {"default": True, "tooltip": "Automatically set tile size based on defaults, above settings are ignored"}),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
@@ -895,8 +1037,8 @@ class HyVideoDecode:
|
||||
vae.tile_latent_min_tsize = 16
|
||||
vae.tile_sample_min_size = 256
|
||||
vae.tile_latent_min_size = 32
|
||||
|
||||
|
||||
|
||||
|
||||
expand_temporal_dim = False
|
||||
if len(latents.shape) == 4:
|
||||
if isinstance(vae, AutoencoderKLCausal3D):
|
||||
@@ -940,7 +1082,7 @@ class HyVideoDecode:
|
||||
|
||||
return (out,)
|
||||
|
||||
#region VideoEncode
|
||||
#region VideoEncode
|
||||
class HyVideoEncode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -951,7 +1093,7 @@ class HyVideoEncode:
|
||||
"temporal_tiling_sample_size": ("INT", {"default": 64, "min": 4, "max": 256, "tooltip": "Smaller values use less VRAM, model default is 64, any other value will cause stutter"}),
|
||||
"spatial_tile_sample_min_size": ("INT", {"default": 256, "min": 32, "max": 2048, "step": 32, "tooltip": "Spatial tile minimum size in pixels, smaller values use less VRAM, may introduce more seams"}),
|
||||
"auto_tile_size": ("BOOLEAN", {"default": True, "tooltip": "Automatically set tile size based on defaults, above settings are ignored"}),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
@@ -962,7 +1104,7 @@ class HyVideoEncode:
|
||||
def encode(self, vae, image, enable_vae_tiling, temporal_tiling_sample_size, auto_tile_size, spatial_tile_sample_min_size):
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
|
||||
generator = torch.Generator(device=torch.device("cpu"))#.manual_seed(seed)
|
||||
vae.to(device)
|
||||
if not auto_tile_size:
|
||||
@@ -985,8 +1127,8 @@ class HyVideoEncode:
|
||||
latents = vae.encode(image).latent_dist.sample(generator)
|
||||
latents = latents * vae.config.scaling_factor
|
||||
vae.to(offload_device)
|
||||
print("encoded latents shape",latents.shape)
|
||||
|
||||
print("encoded latents shape",latents.shape)
|
||||
|
||||
|
||||
return ({"samples": latents},)
|
||||
|
||||
@@ -1016,32 +1158,32 @@ class HyVideoLatentPreview:
|
||||
latents = samples["samples"].clone()
|
||||
print("in sample", latents.shape)
|
||||
#latent_rgb_factors =[[-0.02531045419704009, -0.00504800612542497, 0.13293717293982546], [-0.03421835830845858, 0.13996708548892614, -0.07081038680118075], [0.011091819063647063, -0.03372949685846012, -0.0698232210116172], [-0.06276524604742019, -0.09322986677909442, 0.01826383612148913], [0.021290659938126788, -0.07719530444034409, -0.08247812477766273], [0.04401102991215147, -0.0026401932105894754, -0.01410913586718443], [0.08979717602613707, 0.05361221258740831, 0.11501425309699129], [0.04695121980405198, -0.13053491609675175, 0.05025986885867986], [-0.09704684176098193, 0.03397687417738002, -0.1105886644677771], [0.14694697234804935, -0.12316902186157716, 0.04210404546699645], [0.14432470831243552, -0.002580008133591355, -0.08490676947390643], [0.051502750076553944, -0.10071695490292451, -0.01786223610178095], [-0.12503276881774464, 0.08877830923879379, 0.1076584501927316], [-0.020191205513213406, -0.1493425056303128, -0.14289740371758308], [-0.06470138952271293, -0.07410426095060325, 0.00980804676890873], [0.11747671720735695, 0.10916082743849789, -0.12235599365235904]]
|
||||
latent_rgb_factors = [[-0.41, -0.25, -0.26],
|
||||
[-0.26, -0.49, -0.24],
|
||||
[-0.37, -0.54, -0.3],
|
||||
[-0.04, -0.29, -0.29],
|
||||
[-0.52, -0.59, -0.39],
|
||||
[-0.56, -0.6, -0.02],
|
||||
[-0.53, -0.06, -0.48],
|
||||
[-0.51, -0.28, -0.18],
|
||||
[-0.59, -0.1, -0.33],
|
||||
[-0.56, -0.54, -0.41],
|
||||
[-0.61, -0.19, -0.5],
|
||||
[-0.05, -0.25, -0.17],
|
||||
[-0.23, -0.04, -0.22],
|
||||
[-0.51, -0.56, -0.43],
|
||||
[-0.13, -0.4, -0.05],
|
||||
latent_rgb_factors = [[-0.41, -0.25, -0.26],
|
||||
[-0.26, -0.49, -0.24],
|
||||
[-0.37, -0.54, -0.3],
|
||||
[-0.04, -0.29, -0.29],
|
||||
[-0.52, -0.59, -0.39],
|
||||
[-0.56, -0.6, -0.02],
|
||||
[-0.53, -0.06, -0.48],
|
||||
[-0.51, -0.28, -0.18],
|
||||
[-0.59, -0.1, -0.33],
|
||||
[-0.56, -0.54, -0.41],
|
||||
[-0.61, -0.19, -0.5],
|
||||
[-0.05, -0.25, -0.17],
|
||||
[-0.23, -0.04, -0.22],
|
||||
[-0.51, -0.56, -0.43],
|
||||
[-0.13, -0.4, -0.05],
|
||||
[-0.01, -0.01, -0.48]]
|
||||
|
||||
|
||||
import random
|
||||
random.seed(seed)
|
||||
#latent_rgb_factors = [[random.uniform(min_val, max_val) for _ in range(3)] for _ in range(16)]
|
||||
out_factors = latent_rgb_factors
|
||||
print(latent_rgb_factors)
|
||||
|
||||
|
||||
#latent_rgb_factors_bias = [0.138, 0.025, -0.299]
|
||||
latent_rgb_factors_bias = [r_bias, g_bias, b_bias]
|
||||
|
||||
|
||||
latent_rgb_factors = torch.tensor(latent_rgb_factors, device=latents.device, dtype=latents.dtype).transpose(0, 1)
|
||||
latent_rgb_factors_bias = torch.tensor(latent_rgb_factors_bias, device=latents.device, dtype=latents.dtype)
|
||||
|
||||
@@ -1062,9 +1204,9 @@ class HyVideoLatentPreview:
|
||||
latent_images_min = latent_images.min()
|
||||
latent_images_max = latent_images.max()
|
||||
latent_images = (latent_images - latent_images_min) / (latent_images_max - latent_images_min)
|
||||
|
||||
|
||||
return (latent_images.float().cpu(), out_factors)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"HyVideoSampler": HyVideoSampler,
|
||||
"HyVideoDecode": HyVideoDecode,
|
||||
@@ -1078,6 +1220,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"HyVideoSTG": HyVideoSTG,
|
||||
"HyVideoCustomPromptTemplate": HyVideoCustomPromptTemplate,
|
||||
"HyVideoLatentPreview": HyVideoLatentPreview,
|
||||
"HyVideoLoraSelect": HyVideoLoraSelect,
|
||||
"HyVideoLoraBlockEdit": HyVideoLoraBlockEdit,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"HyVideoSampler": "HunyuanVideo Sampler",
|
||||
@@ -1092,4 +1236,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"HyVideoSTG": "HunyuanVideo STG",
|
||||
"HyVideoCustomPromptTemplate": "HunyuanVideo Custom Prompt Template",
|
||||
"HyVideoLatentPreview": "HunyuanVideo Latent Preview",
|
||||
"HyVideoLoraSelect": "HunyuanVideo Lora Select",
|
||||
"HyVideoLoraBlockEdit": "HunyuanVideo Lora Block Edit",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user