initial RF-inversion for testing
This commit is contained in:
@@ -7,5 +7,4 @@ master_ip
|
||||
logs/
|
||||
*.DS_Store
|
||||
.idea
|
||||
*.pt
|
||||
tools/
|
||||
+4
-1
@@ -1,4 +1,7 @@
|
||||
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .nodes import NODE_CLASS_MAPPINGS as NODES_CLASS, NODE_DISPLAY_NAME_MAPPINGS as NODES_DISPLAY
|
||||
from .nodes_rf_inversion import NODE_CLASS_MAPPINGS as NODE_CLASS_MAPPINGS_RF_INVERSION, NODE_DISPLAY_NAME_MAPPINGS as NODE_DISPLAY_NAME_MAPPINGS_RF_INVERSION
|
||||
|
||||
NODE_CLASS_MAPPINGS = {**NODES_CLASS, **NODE_CLASS_MAPPINGS_RF_INVERSION}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {**NODES_DISPLAY, **NODE_DISPLAY_NAME_MAPPINGS_RF_INVERSION}
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@@ -509,6 +509,10 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
# 8. Preview callback
|
||||
from ....latent_preview import prepare_callback
|
||||
callback = prepare_callback(self.transformer, num_inference_steps)
|
||||
|
||||
|
||||
logger.info(f"Sampling {video_length} frames in {latents.shape[2]} latents at {width}x{height} with {len(timesteps)} inference steps")
|
||||
comfy_pbar = ProgressBar(len(timesteps))
|
||||
@@ -619,10 +623,10 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
):
|
||||
if progress_bar is not None:
|
||||
progress_bar.update()
|
||||
if callback is not None:
|
||||
callback(i, latents.detach()[-1].permute(1,0,2,3), None, num_inference_steps)
|
||||
else:
|
||||
comfy_pbar.update(1)
|
||||
if callback is not None and i % callback_steps == 0:
|
||||
step_idx = i // getattr(self.scheduler, "order", 1)
|
||||
callback(step_idx, t, latents)
|
||||
|
||||
#latents = (latents / 2 + 0.5).clamp(0, 1).cpu()
|
||||
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from comfy.cli_args import args, LatentPreviewMethod
|
||||
import comfy.model_management
|
||||
import comfy.utils
|
||||
|
||||
MAX_PREVIEW_RESOLUTION = args.preview_size
|
||||
|
||||
def preview_to_image(latent_image):
|
||||
latents_ubyte = (((latent_image + 1.0) / 2.0).clamp(0, 1) # change scale from -1..1 to 0..1
|
||||
.mul(0xFF) # to 0..255
|
||||
).to(device="cpu", dtype=torch.uint8, non_blocking=comfy.model_management.device_supports_non_blocking(latent_image.device))
|
||||
|
||||
return Image.fromarray(latents_ubyte.numpy())
|
||||
|
||||
class LatentPreviewer:
|
||||
def decode_latent_to_preview(self, x0):
|
||||
pass
|
||||
|
||||
def decode_latent_to_preview_image(self, preview_format, x0):
|
||||
preview_image = self.decode_latent_to_preview(x0)
|
||||
return ("GIF", preview_image, MAX_PREVIEW_RESOLUTION)
|
||||
|
||||
class Latent2RGBPreviewer(LatentPreviewer):
|
||||
def __init__(self):
|
||||
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]]
|
||||
self.latent_rgb_factors = torch.tensor(latent_rgb_factors, device="cpu").transpose(0, 1)
|
||||
self.latent_rgb_factors_bias = torch.tensor([0.138, 0.025, -0.299], device="cpu")
|
||||
|
||||
def decode_latent_to_preview(self, x0):
|
||||
self.latent_rgb_factors = self.latent_rgb_factors.to(dtype=x0.dtype, device=x0.device)
|
||||
if self.latent_rgb_factors_bias is not None:
|
||||
self.latent_rgb_factors_bias = self.latent_rgb_factors_bias.to(dtype=x0.dtype, device=x0.device)
|
||||
|
||||
latent_image = torch.nn.functional.linear(x0[0].permute(1, 2, 0), self.latent_rgb_factors,
|
||||
bias=self.latent_rgb_factors_bias)
|
||||
return preview_to_image(latent_image)
|
||||
|
||||
|
||||
def get_previewer():
|
||||
previewer = None
|
||||
method = args.preview_method
|
||||
if method != LatentPreviewMethod.NoPreviews:
|
||||
# TODO previewer method
|
||||
|
||||
if method == LatentPreviewMethod.Auto:
|
||||
method = LatentPreviewMethod.Latent2RGB
|
||||
|
||||
if previewer is None:
|
||||
previewer = Latent2RGBPreviewer()
|
||||
return previewer
|
||||
|
||||
def prepare_callback(model, steps, x0_output_dict=None):
|
||||
preview_format = "JPEG"
|
||||
if preview_format not in ["JPEG", "PNG"]:
|
||||
preview_format = "JPEG"
|
||||
|
||||
previewer = get_previewer()
|
||||
|
||||
pbar = comfy.utils.ProgressBar(steps)
|
||||
def callback(step, x0, x, total_steps):
|
||||
if x0_output_dict is not None:
|
||||
x0_output_dict["x0"] = x0
|
||||
preview_bytes = None
|
||||
if previewer:
|
||||
preview_bytes = previewer.decode_latent_to_preview_image(preview_format, x0)
|
||||
pbar.update_absolute(step + 1, total_steps, preview_bytes)
|
||||
return callback
|
||||
|
||||
@@ -855,7 +855,7 @@ class HyVideoSampler:
|
||||
return ({
|
||||
"samples": out_latents
|
||||
},)
|
||||
|
||||
|
||||
#region VideoDecode
|
||||
class HyVideoDecode:
|
||||
@classmethod
|
||||
@@ -990,7 +990,7 @@ class HyVideoEncode:
|
||||
|
||||
return ({"samples": latents},)
|
||||
|
||||
class CogVideoLatentPreview:
|
||||
class HyVideoLatentPreview:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
@@ -1008,25 +1008,39 @@ class CogVideoLatentPreview:
|
||||
RETURN_TYPES = ("IMAGE", "STRING", )
|
||||
RETURN_NAMES = ("images", "latent_rgb_factors",)
|
||||
FUNCTION = "sample"
|
||||
CATEGORY = "PyramidFlowWrapper"
|
||||
CATEGORY = "HunyuanVideoWrapper"
|
||||
|
||||
def sample(self, samples, seed, min_val, max_val, r_bias, g_bias, b_bias):
|
||||
mm.soft_empty_cache()
|
||||
|
||||
latents = samples["samples"].clone()
|
||||
print("in sample", latents.shape)
|
||||
latents = latents.permute(0, 2, 1, 3, 4) # [batch_size, num_channels, num_frames, height, width]
|
||||
|
||||
#[[0.0658900170023352, 0.04687556512203313, -0.056971557475649186], [-0.01265770449940036, -0.02814809569100843, -0.0768912512529372], [0.061456544746314665, 0.0005511617552452358, -0.0652574975291287], [-0.09020669168815276, -0.004755440180558637, -0.023763970904494294], [0.031766964513999865, -0.030959599938418375, 0.08654669098083616], [-0.005981764690055846, -0.08809119252349802, -0.06439852368217663], [-0.0212114426433989, 0.08894281999597677, 0.05155629477559985], [-0.013947446911030725, -0.08987475069900677, -0.08923124751217484], [-0.08235967967978511, 0.07268025379974379, 0.08830486164536037], [-0.08052049179735378, -0.050116143175332195, 0.02023752569687405], [-0.07607527759162447, 0.06827156419895981, 0.08678111754261035], [-0.04689089232553825, 0.017294986041038893, -0.10280492336438908], [-0.06105783150270304, 0.07311850680875913, 0.019995735372550075], [-0.09232589996527711, -0.012869815059053047, -0.04355587834255975], [-0.06679931010802251, 0.018399815879067458, 0.06802404982033876], [-0.013062632927118165, -0.04292991477896661, 0.07476243356192845]]
|
||||
latent_rgb_factors =[[0.11945946736445662, 0.09919175788574555, -0.004832707433877734], [-0.0011977028264356232, 0.05496505130267682, 0.021321622433638193], [-0.014088548986590666, -0.008701477861945644, -0.020991313281459367], [0.03063921972519621, 0.12186477097625073, 0.0139593690235148], [0.0927403067854673, 0.030293187650929136, 0.05083134241694003], [0.0379112441305742, 0.04935199882777209, 0.058562766246777774], [0.017749911959153715, 0.008839453404921545, 0.036005638019226294], [0.10610119248526109, 0.02339855688237826, 0.057154257614084596], [0.1273639464837117, -0.010959856130713416, 0.043268631260428896], [-0.01873510946881321, 0.08220930648486932, 0.10613256772247093], [0.008429116376722327, 0.07623856561000408, 0.09295712117576727], [0.12938137079617007, 0.12360403483892413, 0.04478930933220116], [0.04565908794779364, 0.041064156741596365, -0.017695041535528512], [0.00019003240570281826, -0.013965147883381978, 0.05329669529635849], [0.08082391586738358, 0.11548306825496074, -0.021464170006615893], [-0.01517932393230994, -0.0057985555313003236, 0.07216646476618871]]
|
||||
#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],
|
||||
[-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)]
|
||||
#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.085, 0.137, 0.158]
|
||||
#latent_rgb_factors_bias = [r_bias, g_bias, b_bias]
|
||||
#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)
|
||||
@@ -1063,7 +1077,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"HyVideoTorchCompileSettings": HyVideoTorchCompileSettings,
|
||||
"HyVideoSTG": HyVideoSTG,
|
||||
"HyVideoCustomPromptTemplate": HyVideoCustomPromptTemplate,
|
||||
}
|
||||
"HyVideoLatentPreview": HyVideoLatentPreview,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"HyVideoSampler": "HunyuanVideo Sampler",
|
||||
"HyVideoDecode": "HunyuanVideo Decode",
|
||||
@@ -1076,4 +1091,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"HyVideoTorchCompileSettings": "HunyuanVideo Torch Compile Settings",
|
||||
"HyVideoSTG": "HunyuanVideo STG",
|
||||
"HyVideoCustomPromptTemplate": "HunyuanVideo Custom Prompt Template",
|
||||
"HyVideoLatentPreview": "HunyuanVideo Latent Preview",
|
||||
}
|
||||
|
||||
@@ -0,0 +1,437 @@
|
||||
#based on https://github.com/DarkMnDragon/rf-inversion-diffuser/blob/main/inversion_editing_cli.py
|
||||
import torch
|
||||
import gc
|
||||
import os
|
||||
from .utils import log, print_memory
|
||||
|
||||
from .hyvideo.utils.data_utils import align_to
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
import comfy.model_management as mm
|
||||
from .nodes import get_rotary_pos_embed
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
def generate_eta_values(
|
||||
timesteps,
|
||||
start_step,
|
||||
end_step,
|
||||
eta,
|
||||
eta_trend,
|
||||
):
|
||||
assert start_step < end_step and start_step >= 0 and end_step <= len(timesteps), "Invalid start_step and end_step"
|
||||
# timesteps are monotonically decreasing, from 1.0 to 0.0
|
||||
print("eta timesteps", timesteps)
|
||||
eta_values = [0.0] * (len(timesteps) - 1)
|
||||
|
||||
if eta_trend == 'constant':
|
||||
for i in range(start_step, end_step):
|
||||
eta_values[i] = eta
|
||||
elif eta_trend == 'linear_increase':
|
||||
total_time = timesteps[start_step] - timesteps[end_step - 1]
|
||||
for i in range(start_step, end_step):
|
||||
eta_values[i] = eta * (timesteps[start_step] - timesteps[i]) / total_time
|
||||
elif eta_trend == 'linear_decrease':
|
||||
total_time = timesteps[start_step] - timesteps[end_step - 1]
|
||||
for i in range(start_step, end_step):
|
||||
eta_values[i] = eta * (timesteps[i] - timesteps[end_step - 1]) / total_time
|
||||
else:
|
||||
raise NotImplementedError(f"Unsupported eta_trend: {eta_trend}")
|
||||
|
||||
return eta_values
|
||||
|
||||
class HyVideoEmptyTextEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("HYVIDEMBEDS", )
|
||||
RETURN_NAMES = ("hyvid_embeds",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "HunyuanVideoWrapper"
|
||||
DESCRIPTION = "Empty Text Embeds for HunyuanVideoWrapper, to avoid having to encode prompts for inverse sampling"
|
||||
|
||||
def process(self):
|
||||
device = mm.text_encoder_device()
|
||||
offload_device = mm.text_encoder_offload_device()
|
||||
|
||||
prompt_embeds_dict = torch.load(os.path.join(script_directory, "hunyuan_empty_prompt_embeds_dict.pt"))
|
||||
return (prompt_embeds_dict,)
|
||||
|
||||
class HyVideoInverseSampler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("HYVIDEOMODEL",),
|
||||
"hyvid_embeds": ("HYVIDEMBEDS", ),
|
||||
"samples": ("LATENT", {"tooltip": "init Latents to use for video2video process"} ),
|
||||
"steps": ("INT", {"default": 30, "min": 1}),
|
||||
"embedded_guidance_scale": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 30.0, "step": 0.01}),
|
||||
"flow_shift": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 30.0, "step": 0.01}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"force_offload": ("BOOLEAN", {"default": True}),
|
||||
"gamma": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
RETURN_NAMES = ("samples",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "HunyuanVideoWrapper"
|
||||
|
||||
def process(self, model, hyvid_embeds, flow_shift, steps, embedded_guidance_scale, seed, samples, gamma, force_offload):
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
dtype = model["dtype"]
|
||||
transformer = model["pipe"].transformer
|
||||
pipeline = model["pipe"]
|
||||
|
||||
generator = torch.Generator(device=torch.device("cpu")).manual_seed(seed)
|
||||
|
||||
latents = samples["samples"] if samples is not None else None
|
||||
batch_size, num_channels_latents, latent_num_frames, latent_height, latent_width = latents.shape
|
||||
height = latent_height * pipeline.vae_scale_factor
|
||||
width = latent_width * pipeline.vae_scale_factor
|
||||
num_frames = (latent_num_frames - 1) * 4 + 1
|
||||
|
||||
|
||||
if width <= 0 or height <= 0 or num_frames <= 0:
|
||||
raise ValueError(
|
||||
f"`height` and `width` and `video_length` must be positive integers, got height={height}, width={width}, video_length={num_frames}"
|
||||
)
|
||||
if (num_frames - 1) % 4 != 0:
|
||||
raise ValueError(
|
||||
f"`video_length-1` must be a multiple of 4, got {num_frames}"
|
||||
)
|
||||
|
||||
log.info(
|
||||
f"Input (height, width, video_length) = ({height}, {width}, {num_frames})"
|
||||
)
|
||||
|
||||
freqs_cos, freqs_sin = get_rotary_pos_embed(transformer, num_frames, height, width)
|
||||
|
||||
pipeline.scheduler.shift = flow_shift
|
||||
|
||||
if model["block_swap_args"] is not None:
|
||||
for name, param in transformer.named_parameters():
|
||||
#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"]["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"],
|
||||
)
|
||||
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
|
||||
|
||||
pipeline.scheduler.set_timesteps(steps, device=device)
|
||||
timesteps = pipeline.scheduler.timesteps
|
||||
timesteps = timesteps.flip(0)
|
||||
print("timesteps", timesteps)
|
||||
print("pipeline.scheduler.order", pipeline.scheduler.order)
|
||||
print("len(timesteps)", len(timesteps))
|
||||
|
||||
latent_video_length = (num_frames - 1) // 4 + 1
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = transformer.config.in_channels
|
||||
|
||||
|
||||
latents = latents.to(device)
|
||||
|
||||
shape = (
|
||||
1,
|
||||
num_channels_latents,
|
||||
latent_video_length,
|
||||
int(height) // pipeline.vae_scale_factor,
|
||||
int(width) // pipeline.vae_scale_factor,
|
||||
)
|
||||
noise = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
|
||||
frames_needed = noise.shape[1]
|
||||
current_frames = latents.shape[1]
|
||||
|
||||
if frames_needed > current_frames:
|
||||
repeat_factor = frames_needed - current_frames
|
||||
additional_frame = torch.randn((latents.size(0), repeat_factor, latents.size(2), latents.size(3), latents.size(4)), dtype=latents.dtype, device=latents.device)
|
||||
latents = torch.cat((additional_frame, latents), dim=1)
|
||||
self.additional_frames = repeat_factor
|
||||
elif frames_needed < current_frames:
|
||||
latents = latents[:, :frames_needed, :, :, :]
|
||||
|
||||
|
||||
|
||||
# 7. Denoising loop
|
||||
num_warmup_steps = len(timesteps) - steps * pipeline.scheduler.order
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
from .latent_preview import prepare_callback
|
||||
callback = prepare_callback(transformer, steps)
|
||||
|
||||
from comfy.utils import ProgressBar
|
||||
from tqdm import tqdm
|
||||
log.info(f"Sampling {num_frames} frames in {latents.shape[2]} latents at {width}x{height} with {len(timesteps)} inference steps")
|
||||
comfy_pbar = ProgressBar(len(timesteps))
|
||||
with tqdm(total=len(timesteps)) as progress_bar:
|
||||
for idx, (t, t_prev) in enumerate(zip(timesteps[:-1], timesteps[1:])):
|
||||
latent_model_input = latents
|
||||
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
guidance_expand = (
|
||||
torch.tensor(
|
||||
[embedded_guidance_scale] * latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
).to(pipeline.base_dtype)
|
||||
* 1000.0
|
||||
if embedded_guidance_scale is not None
|
||||
else None
|
||||
)
|
||||
|
||||
# predict the noise residual
|
||||
with torch.autocast(
|
||||
device_type="cuda", dtype=pipeline.base_dtype, enabled=True
|
||||
):
|
||||
noise_pred = transformer( # For an input image (129, 192, 336) (1, 256, 256)
|
||||
latent_model_input, # [2, 16, 33, 24, 42]
|
||||
t_expand, # [2]
|
||||
text_states=hyvid_embeds["prompt_embeds"], # [2, 256, 4096]
|
||||
text_mask=hyvid_embeds["attention_mask"], # [2, 256]
|
||||
text_states_2=hyvid_embeds["prompt_embeds_2"], # [2, 768]
|
||||
freqs_cos=freqs_cos, # [seqlen, head_dim]
|
||||
freqs_sin=freqs_sin, # [seqlen, head_dim]
|
||||
guidance=guidance_expand,
|
||||
stg_block_idx=-1,
|
||||
stg_mode=None,
|
||||
return_dict=True,
|
||||
)["x"]
|
||||
sigma = t / 1000.0
|
||||
sigma_prev = t_prev / 1000.0
|
||||
target_noise_velocity = (noise - latents) / (1.0 - sigma)
|
||||
interpolated_velocity = gamma * target_noise_velocity + (1 - gamma) * noise_pred
|
||||
|
||||
latents = latents + (sigma_prev - sigma) * interpolated_velocity
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
#latents = pipeline.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
|
||||
|
||||
|
||||
progress_bar.update()
|
||||
if callback is not None:
|
||||
print("callback", latents.shape)
|
||||
callback(idx, latents.detach()[-1].permute(1,0,2,3), None, steps)
|
||||
else:
|
||||
comfy_pbar.update(1)
|
||||
|
||||
|
||||
print_memory(device)
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
|
||||
if force_offload:
|
||||
if model["manual_offloading"]:
|
||||
transformer.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
return ({
|
||||
"samples": latents
|
||||
},)
|
||||
|
||||
class HyVideoReSampler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("HYVIDEOMODEL",),
|
||||
"hyvid_embeds": ("HYVIDEMBEDS", ),
|
||||
"samples": ("LATENT", {"tooltip": "init Latents to use for video2video process"} ),
|
||||
"inversed_latents": ("LATENT", {"tooltip": "inversed latents from HyVideoInverseSampler"} ),
|
||||
"steps": ("INT", {"default": 30, "min": 1}),
|
||||
"embedded_guidance_scale": ("FLOAT", {"default": 6.0, "min": 0.0, "max": 30.0, "step": 0.01}),
|
||||
"flow_shift": ("FLOAT", {"default": 1.0, "min": 1.0, "max": 30.0, "step": 0.01}),
|
||||
"force_offload": ("BOOLEAN", {"default": True}),
|
||||
"start_step": ("INT", {"default": 0, "min": 0}),
|
||||
"end_step": ("INT", {"default": 18, "min": 0}),
|
||||
"eta_base": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"eta_trend": (['constant', 'linear_increase', 'linear_decrease'], {"default": "constant"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
RETURN_NAMES = ("samples",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "HunyuanVideoWrapper"
|
||||
|
||||
def process(self, model, hyvid_embeds, flow_shift, steps, embedded_guidance_scale,
|
||||
samples, inversed_latents, force_offload, start_step, end_step, eta_base, eta_trend):
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
dtype = model["dtype"]
|
||||
transformer = model["pipe"].transformer
|
||||
pipeline = model["pipe"]
|
||||
|
||||
target_latents = samples["samples"]
|
||||
|
||||
batch_size, num_channels_latents, latent_num_frames, latent_height, latent_width = target_latents.shape
|
||||
height = latent_height * pipeline.vae_scale_factor
|
||||
width = latent_width * pipeline.vae_scale_factor
|
||||
num_frames = (latent_num_frames - 1) * 4 + 1
|
||||
|
||||
if width <= 0 or height <= 0 or num_frames <= 0:
|
||||
raise ValueError(
|
||||
f"`height` and `width` and `video_length` must be positive integers, got height={height}, width={width}, video_length={num_frames}"
|
||||
)
|
||||
if (num_frames - 1) % 4 != 0:
|
||||
raise ValueError(
|
||||
f"`video_length-1` must be a multiple of 4, got {num_frames}"
|
||||
)
|
||||
|
||||
log.info(
|
||||
f"Input (height, width, video_length) = ({height}, {width}, {num_frames})"
|
||||
)
|
||||
|
||||
freqs_cos, freqs_sin = get_rotary_pos_embed(transformer, num_frames, height, width)
|
||||
|
||||
pipeline.scheduler.shift = flow_shift
|
||||
|
||||
if model["block_swap_args"] is not None:
|
||||
for name, param in transformer.named_parameters():
|
||||
#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"]["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"],
|
||||
)
|
||||
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
|
||||
|
||||
pipeline.scheduler.set_timesteps(steps, device=device)
|
||||
timesteps = pipeline.scheduler.timesteps
|
||||
|
||||
eta_values = generate_eta_values(timesteps / 1000, start_step, end_step, eta_base, eta_trend)
|
||||
|
||||
|
||||
target_latents = target_latents.to(device)
|
||||
latents = inversed_latents["samples"]
|
||||
|
||||
# 7. Denoising loop
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
from .latent_preview import prepare_callback
|
||||
callback = prepare_callback(transformer, steps)
|
||||
|
||||
from comfy.utils import ProgressBar
|
||||
from tqdm import tqdm
|
||||
log.info(f"Sampling {num_frames} frames in {latents.shape[2]} latents at {width}x{height} with {len(timesteps)} inference steps")
|
||||
comfy_pbar = ProgressBar(len(timesteps))
|
||||
|
||||
with tqdm(total=len(timesteps)) as progress_bar:
|
||||
for idx, (t, t_prev) in enumerate(zip(timesteps[:-1], timesteps[1:])):
|
||||
|
||||
latent_model_input = latents
|
||||
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
guidance_expand = (
|
||||
torch.tensor(
|
||||
[embedded_guidance_scale] * latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
).to(pipeline.base_dtype)
|
||||
* 1000.0
|
||||
if embedded_guidance_scale is not None
|
||||
else None
|
||||
)
|
||||
|
||||
# predict the noise residual
|
||||
with torch.autocast(
|
||||
device_type="cuda", dtype=pipeline.base_dtype, enabled=True
|
||||
):
|
||||
noise_pred = transformer( # For an input image (129, 192, 336) (1, 256, 256)
|
||||
latent_model_input, # [2, 16, 33, 24, 42]
|
||||
t_expand, # [2]
|
||||
text_states=hyvid_embeds["prompt_embeds"], # [2, 256, 4096]
|
||||
text_mask=hyvid_embeds["attention_mask"], # [2, 256]
|
||||
text_states_2=hyvid_embeds["prompt_embeds_2"], # [2, 768]
|
||||
freqs_cos=freqs_cos, # [seqlen, head_dim]
|
||||
freqs_sin=freqs_sin, # [seqlen, head_dim]
|
||||
guidance=guidance_expand,
|
||||
stg_block_idx=-1,
|
||||
stg_mode=None,
|
||||
return_dict=True,
|
||||
)["x"]
|
||||
sigma = t / 1000.0
|
||||
sigma_prev = t_prev / 1000.0
|
||||
noise_pred = noise_pred.to(torch.float32)
|
||||
latents = latents.to(torch.float32)
|
||||
target_latents = target_latents.to(torch.float32)
|
||||
target_img_velocity = -(target_latents - latents) / sigma
|
||||
|
||||
# interpolated velocity
|
||||
eta = eta_values[idx]
|
||||
interpolated_velocity = eta * target_img_velocity + (1 - eta) * noise_pred
|
||||
latents = latents + (sigma_prev - sigma) * interpolated_velocity
|
||||
|
||||
print(f"X_{sigma_prev:.3f} = X_{sigma:.3f} + {sigma_prev - sigma:.3f} * ({eta:.3f} * target_img_velocity + {1 - eta:.3f} * noise_pred)")
|
||||
latents = latents.to(torch.bfloat16)
|
||||
|
||||
if callback is not None:
|
||||
callback(idx, latents.detach()[-1].permute(1,0,2,3), None, steps)
|
||||
else:
|
||||
comfy_pbar.update(1)
|
||||
|
||||
print_memory(device)
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
|
||||
if force_offload:
|
||||
if model["manual_offloading"]:
|
||||
transformer.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
return ({
|
||||
"samples": latents
|
||||
},)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"HyVideoInverseSampler": HyVideoInverseSampler,
|
||||
"HyVideoReSampler": HyVideoReSampler,
|
||||
"HyVideoEmptyTextEmbeds": HyVideoEmptyTextEmbeds
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"HyVideoInverseSampler": "HunyuanVideo Inverse Sampler",
|
||||
"HyVideoReSampler": "HunyuanVideo ReSampler",
|
||||
"HyVideoEmptyTextEmbeds": "HunyuanVideo Empty Text Embeds"
|
||||
}
|
||||
Reference in New Issue
Block a user