block_swap tweaks
This commit is contained in:
@@ -38,25 +38,6 @@ logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
EXAMPLE_DOC_STRING = """"""
|
||||
|
||||
|
||||
def rescale_noise_cfg(noise_cfg, noise_pred_text, guidance_rescale=0.0):
|
||||
"""
|
||||
Rescale `noise_cfg` according to `guidance_rescale`. Based on findings of [Common Diffusion Noise Schedules and
|
||||
Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf). See Section 3.4
|
||||
"""
|
||||
std_text = noise_pred_text.std(
|
||||
dim=list(range(1, noise_pred_text.ndim)), keepdim=True
|
||||
)
|
||||
std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True)
|
||||
# rescale the results from guidance (fixes overexposure)
|
||||
noise_pred_rescaled = noise_cfg * (std_text / std_cfg)
|
||||
# mix with the original results from guidance by factor guidance_rescale to avoid "plain looking" images
|
||||
noise_cfg = (
|
||||
guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg
|
||||
)
|
||||
return noise_cfg
|
||||
|
||||
|
||||
def retrieve_timesteps(
|
||||
scheduler,
|
||||
num_inference_steps: Optional[int] = None,
|
||||
@@ -446,8 +427,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
negative_prompt_mask = prompt_embed_dict["negative_attention_mask"]
|
||||
prompt_embeds_2 = prompt_embed_dict["prompt_embeds_2"]
|
||||
negative_prompt_embeds_2 = prompt_embed_dict["negative_prompt_embeds_2"]
|
||||
prompt_mask_2 = prompt_embed_dict["attention_mask_2"]
|
||||
negative_prompt_mask_2 = prompt_embed_dict["negative_attention_mask_2"]
|
||||
#prompt_mask_2 = prompt_embed_dict["attention_mask_2"]
|
||||
#negative_prompt_mask_2 = prompt_embed_dict["negative_attention_mask_2"]
|
||||
|
||||
# For classifier free guidance, we need to do two forward passes.
|
||||
# Here we concatenate the unconditional and text embeddings into a single batch
|
||||
@@ -458,8 +439,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
prompt_mask = torch.cat([negative_prompt_mask, prompt_mask])
|
||||
if prompt_embeds_2 is not None:
|
||||
prompt_embeds_2 = torch.cat([negative_prompt_embeds_2, prompt_embeds_2])
|
||||
if prompt_mask_2 is not None:
|
||||
prompt_mask_2 = torch.cat([negative_prompt_mask_2, prompt_mask_2])
|
||||
#if prompt_mask_2 is not None:
|
||||
# prompt_mask_2 = torch.cat([negative_prompt_mask_2, prompt_mask_2])
|
||||
elif self.do_classifier_free_guidance and self.do_spatio_temporal_guidance:
|
||||
prompt_embeds = torch.cat(
|
||||
[negative_prompt_embeds, prompt_embeds, prompt_embeds]
|
||||
@@ -470,18 +451,18 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
prompt_embeds_2 = torch.cat(
|
||||
[negative_prompt_embeds_2, prompt_embeds_2, prompt_embeds_2]
|
||||
)
|
||||
if prompt_mask_2 is not None:
|
||||
prompt_mask_2 = torch.cat(
|
||||
[negative_prompt_mask_2, prompt_mask_2, prompt_mask_2]
|
||||
)
|
||||
#if prompt_mask_2 is not None:
|
||||
# prompt_mask_2 = torch.cat(
|
||||
# [negative_prompt_mask_2, prompt_mask_2, prompt_mask_2]
|
||||
# )
|
||||
elif self.do_spatio_temporal_guidance:
|
||||
prompt_embeds = torch.cat([prompt_embeds, prompt_embeds])
|
||||
if prompt_mask is not None:
|
||||
prompt_mask = torch.cat([prompt_mask, prompt_mask])
|
||||
if prompt_embeds_2 is not None:
|
||||
prompt_embeds_2 = torch.cat([prompt_embeds_2, prompt_embeds_2])
|
||||
if prompt_mask_2 is not None:
|
||||
prompt_mask_2 = torch.cat([prompt_mask_2, prompt_mask_2])
|
||||
#if prompt_mask_2 is not None:
|
||||
# prompt_mask_2 = torch.cat([prompt_mask_2, prompt_mask_2])
|
||||
|
||||
|
||||
# 4. Prepare timesteps
|
||||
@@ -617,14 +598,6 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
noise_pred_text - noise_pred_perturb
|
||||
)
|
||||
|
||||
if self.do_classifier_free_guidance and self.guidance_rescale > 0.0:
|
||||
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
|
||||
noise_pred = rescale_noise_cfg(
|
||||
noise_pred,
|
||||
noise_pred_text,
|
||||
guidance_rescale=self.guidance_rescale,
|
||||
)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(
|
||||
noise_pred, t, latents, **extra_step_kwargs, return_dict=False
|
||||
@@ -656,6 +629,6 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
#latents = (latents / 2 + 0.5).clamp(0, 1).cpu()
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
#self.maybe_free_model_hooks()
|
||||
|
||||
return latents
|
||||
@@ -614,20 +614,28 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
)
|
||||
self.double_blocks_to_swap = 0
|
||||
self.single_blocks_to_swap = 0
|
||||
self.offload_txt_in = False
|
||||
self.offload_img_in = False
|
||||
|
||||
# thanks @2kpr for the initial block swap code!
|
||||
def block_swap(self, double_blocks_to_swap, single_blocks_to_swap):
|
||||
def block_swap(self, double_blocks_to_swap, single_blocks_to_swap, offload_txt_in=False, offload_img_in=False):
|
||||
print(f"Swapping {double_blocks_to_swap} double blocks and {single_blocks_to_swap} single blocks")
|
||||
self.double_blocks_to_swap = double_blocks_to_swap
|
||||
self.single_blocks_to_swap = single_blocks_to_swap
|
||||
self.offload_txt_in = offload_txt_in
|
||||
self.offload_img_in = offload_img_in
|
||||
for b, block in enumerate(self.double_blocks):
|
||||
if b < 0 or b > self.double_blocks_to_swap:
|
||||
#mm.soft_empty_cache()
|
||||
if b > self.double_blocks_to_swap:
|
||||
print(f"Moving double_block {b} to main device")
|
||||
block.to(self.main_device)
|
||||
else:
|
||||
print(f"Moving double_block {b} to offload_device")
|
||||
block.to(self.offload_device)
|
||||
for b, block in enumerate(self.single_blocks):
|
||||
if b < 0 or b > self.single_blocks_to_swap:
|
||||
#mm.soft_empty_cache()
|
||||
if b > self.single_blocks_to_swap:
|
||||
block.to(self.main_device)
|
||||
else:
|
||||
block.to(self.offload_device)
|
||||
|
||||
def enable_deterministic(self):
|
||||
for block in self.double_blocks:
|
||||
@@ -683,6 +691,11 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
vec = vec + self.guidance_in(guidance)
|
||||
|
||||
# Embed image and text.
|
||||
if self.offload_txt_in:
|
||||
self.txt_in.to(self.main_device)
|
||||
if self.offload_img_in:
|
||||
self.img_in.to(self.main_device)
|
||||
|
||||
img = self.img_in(img)
|
||||
if self.text_projection == "linear":
|
||||
txt = self.txt_in(txt)
|
||||
@@ -692,6 +705,10 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
raise NotImplementedError(
|
||||
f"Unsupported text_projection: {self.text_projection}"
|
||||
)
|
||||
if self.offload_txt_in:
|
||||
self.txt_in.to(self.offload_device)
|
||||
if self.offload_img_in:
|
||||
self.img_in.to(self.offload_device)
|
||||
|
||||
txt_seq_len = txt.shape[1]
|
||||
img_seq_len = img.shape[1]
|
||||
@@ -718,7 +735,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
# --------------------- Pass through DiT blocks ------------------------
|
||||
for b, block in enumerate(self.double_blocks):
|
||||
if b <= self.double_blocks_to_swap and self.double_blocks_to_swap > 0:
|
||||
#mm.soft_empty_cache()
|
||||
#print(f"Moving double_block {b} to main device")
|
||||
block.to(self.main_device)
|
||||
double_block_args = [
|
||||
img,
|
||||
@@ -734,7 +751,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
|
||||
img, txt = block(*double_block_args)
|
||||
if b <= self.double_blocks_to_swap and self.double_blocks_to_swap > 0:
|
||||
#mm.soft_empty_cache()
|
||||
#print(f"Moving double_block {b} to offload device")
|
||||
block.to(self.offload_device, non_blocking=True)
|
||||
|
||||
# Merge txt and img to pass through single stream blocks.
|
||||
@@ -742,6 +759,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
if len(self.single_blocks) > 0:
|
||||
for b, block in enumerate(self.single_blocks):
|
||||
if b <= self.single_blocks_to_swap and self.single_blocks_to_swap > 0:
|
||||
#print(f"Moving single_block {b} to main device")
|
||||
#mm.soft_empty_cache()
|
||||
block.to(self.main_device)
|
||||
curr_stg_mode = stg_mode if b == stg_block_idx else None
|
||||
@@ -760,6 +778,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
|
||||
x = block(*single_block_args)
|
||||
if b <= self.single_blocks_to_swap and self.single_blocks_to_swap > 0:
|
||||
#print(f"Moving single_block {b} to offload device")
|
||||
#mm.soft_empty_cache()
|
||||
block.to(self.offload_device, non_blocking=True)
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ import os
|
||||
import torch
|
||||
import json
|
||||
from typing import List
|
||||
import gc
|
||||
from .utils import log, check_diffusers_version, print_memory
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
|
||||
@@ -77,6 +78,8 @@ class HyVideoBlockSwap:
|
||||
"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"}),
|
||||
"offload_img_in": ("BOOLEAN", {"default": False, "tooltip": "Offload img_in layer"}),
|
||||
},
|
||||
}
|
||||
RETURN_TYPES = ("BLOCKSWAPARGS",)
|
||||
@@ -275,7 +278,9 @@ class HyVideoModelLoader:
|
||||
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,
|
||||
@@ -333,10 +338,13 @@ class HyVideoVAELoader:
|
||||
model_path = folder_paths.get_full_path("vae", model_name)
|
||||
vae_sd = load_torch_file(model_path)
|
||||
|
||||
vae = AutoencoderKLCausal3D.from_config(vae_config).to(dtype).to(offload_device)
|
||||
vae = AutoencoderKLCausal3D.from_config(vae_config)
|
||||
vae.load_state_dict(vae_sd)
|
||||
del vae_sd
|
||||
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"]
|
||||
@@ -757,13 +765,21 @@ class HyVideoSampler:
|
||||
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"] , model["block_swap_args"]["single_blocks_to_swap"])
|
||||
# for name, param in model["pipe"].transformer.named_parameters():
|
||||
# print(name, param.data.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"],
|
||||
)
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
elif model["manual_offloading"]:
|
||||
transformer.to(device)
|
||||
|
||||
#for name, param in transformer.named_parameters():
|
||||
# print(name, param.data.device)
|
||||
|
||||
out_latents = model["pipe"](
|
||||
num_inference_steps=steps,
|
||||
height = target_height,
|
||||
@@ -789,11 +805,12 @@ class HyVideoSampler:
|
||||
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": out_latents
|
||||
@@ -820,6 +837,7 @@ class HyVideoDecode:
|
||||
def decode(self, vae, samples, enable_vae_tiling, temporal_tiling_sample_size):
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
mm.soft_empty_cache()
|
||||
latents = samples["samples"]
|
||||
generator = torch.Generator(device=torch.device("cpu"))#.manual_seed(seed)
|
||||
vae.to(device)
|
||||
|
||||
Reference in New Issue
Block a user