block_swap tweaks

This commit is contained in:
kijai
2024-12-08 01:44:16 +02:00
parent c64db95155
commit e9e058436e
3 changed files with 62 additions and 52 deletions
@@ -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
+26 -7
View File
@@ -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)
+25 -7
View File
@@ -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)