custom block_swap to reduce VRAM use

thanks to 2kpr for the initial block swap code!
This commit is contained in:
kijai
2024-12-04 02:13:48 +02:00
parent 7b2fa72bc5
commit 1073ef9156
3 changed files with 512 additions and 17 deletions
@@ -0,0 +1,444 @@
{
"last_node_id": 35,
"last_link_id": 43,
"nodes": [
{
"id": 30,
"type": "HyVideoTextEncode",
"pos": [
203,
247
],
"size": [
400,
200
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "text_encoders",
"type": "HYVIDTEXTENCODER",
"link": 35
}
],
"outputs": [
{
"name": "hyvid_embeds",
"type": "HYVIDEMBEDS",
"links": [
36
]
}
],
"properties": {
"Node name for S&R": "HyVideoTextEncode"
},
"widgets_values": [
"high quality anime style movie featuring a wolf in a forest",
"bad quality video",
true
]
},
{
"id": 16,
"type": "DownloadAndLoadHyVideoTextEncoder",
"pos": [
-310,
248
],
"size": [
441,
106
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "hyvid_text_encoder",
"type": "HYVIDTEXTENCODER",
"links": [
35
]
}
],
"properties": {
"Node name for S&R": "DownloadAndLoadHyVideoTextEncoder"
},
"widgets_values": [
"Kijai/llava-llama-3-8b-text-encoder-tokenizer",
"openai/clip-vit-large-patch14",
"fp16"
]
},
{
"id": 5,
"type": "HyVideoDecode",
"pos": [
920,
-279
],
"size": [
345.4285888671875,
102
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "vae",
"type": "VAE",
"link": 6
},
{
"name": "samples",
"type": "LATENT",
"link": 4
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
42
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "HyVideoDecode"
},
"widgets_values": [
true,
16
]
},
{
"id": 7,
"type": "HyVideoVAELoader",
"pos": [
442,
-282
],
"size": [
379.166748046875,
82
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [
{
"name": "compile_args",
"type": "COMPILEARGS",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "vae",
"type": "VAE",
"links": [
6
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "HyVideoVAELoader"
},
"widgets_values": [
"hyvid\\hunyuan_video_vae_bf16.safetensors",
"fp16"
]
},
{
"id": 1,
"type": "HyVideoModelLoader",
"pos": [
24,
-63
],
"size": [
509.7506103515625,
178
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "compile_args",
"type": "COMPILEARGS",
"link": null,
"shape": 7
},
{
"name": "block_swap_args",
"type": "BLOCKSWAPARGS",
"link": 43,
"shape": 7
}
],
"outputs": [
{
"name": "model",
"type": "HYVIDEOMODEL",
"links": [
2
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "HyVideoModelLoader"
},
"widgets_values": [
"hyvideo\\hunyuan_video_720_fp8_e4m3fn.safetensors",
"bf16",
"fp8_e4m3fn",
"main_device",
"sageattn_varlen"
]
},
{
"id": 35,
"type": "HyVideoBlockSwap",
"pos": [
-351,
-44
],
"size": [
315,
82
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "block_swap_args",
"type": "BLOCKSWAPARGS",
"links": [
43
]
}
],
"properties": {
"Node name for S&R": "HyVideoBlockSwap"
},
"widgets_values": [
20,
0
]
},
{
"id": 34,
"type": "VHS_VideoCombine",
"pos": [
1367,
-275
],
"size": [
371.7926940917969,
675.792724609375
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 42
},
{
"name": "audio",
"type": "AUDIO",
"link": null,
"shape": 7
},
{
"name": "meta_batch",
"type": "VHS_BatchManager",
"link": null,
"shape": 7
},
{
"name": "vae",
"type": "VAE",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 16,
"loop_count": 0,
"filename_prefix": "HunyuanVideo",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 19,
"save_metadata": true,
"pingpong": false,
"save_output": false,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "HunyuanVideo_00059.mp4",
"subfolder": "",
"type": "temp",
"format": "video/h264-mp4",
"frame_rate": 16
},
"muted": false
}
}
},
{
"id": 3,
"type": "HyVideoSampler",
"pos": [
668,
-62
],
"size": [
315,
314
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "HYVIDEOMODEL",
"link": 2
},
{
"name": "hyvid_embeds",
"type": "HYVIDEMBEDS",
"link": 36
},
{
"name": "samples",
"type": "LATENT",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "samples",
"type": "LATENT",
"links": [
4
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "HyVideoSampler"
},
"widgets_values": [
512,
512,
33,
20,
2.5,
9,
2,
"fixed",
1,
1
]
}
],
"links": [
[
2,
1,
0,
3,
0,
"HYVIDEOMODEL"
],
[
4,
3,
0,
5,
1,
"LATENT"
],
[
6,
7,
0,
5,
0,
"VAE"
],
[
35,
16,
0,
30,
0,
"HYVIDTEXTENCODER"
],
[
36,
30,
0,
3,
1,
"HYVIDEMBEDS"
],
[
42,
5,
0,
34,
0,
"IMAGE"
],
[
43,
35,
0,
1,
1,
"BLOCKSWAPARGS"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.9090909090909091,
"offset": [
488.63705586672444,
501.29751842889374
]
}
},
"version": 0.4
}
+35 -3
View File
@@ -16,7 +16,7 @@ from .posemb_layers import apply_rotary_emb
from .mlp_layers import MLP, MLPEmbedder, FinalLayer
from .modulate_layers import ModulateDiT, modulate, apply_gate
from .token_refiner import SingleTokenRefiner
import comfy.model_management as mm
class MMDoubleStreamBlock(nn.Module):
"""
@@ -443,6 +443,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
text_states_dim_2: int = 768,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
offload_device: Optional[torch.device] = None,
attention_mode: str = "flash_attn",
):
factory_kwargs = {"device": device, "dtype": dtype}
@@ -455,6 +456,9 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
self.guidance_embed = guidance_embed
self.rope_dim_list = rope_dim_list
self.main_device = device
self.offload_device = offload_device
# Text projection. Default to linear projection.
# Alternative: TokenRefiner. See more details (LI-DiT): http://arxiv.org/abs/2406.11831
self.use_attention_mask = use_attention_mask
@@ -558,6 +562,22 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
get_activation_layer("silu"),
**factory_kwargs,
)
self.double_blocks_to_swap = 0
self.single_blocks_to_swap = 0
# thanks @2kpr for the initial block swap code!
def block_swap(self, double_blocks_to_swap, single_blocks_to_swap):
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
for b, block in enumerate(self.double_blocks):
if b < 0 or b > self.double_blocks_to_swap:
mm.soft_empty_cache()
block.to(self.main_device)
for b, block in enumerate(self.single_blocks):
if b < 0 or b > self.single_blocks_to_swap:
mm.soft_empty_cache()
block.to(self.main_device)
def enable_deterministic(self):
for block in self.double_blocks:
@@ -632,7 +652,10 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# --------------------- Pass through DiT blocks ------------------------
for _, block in enumerate(self.double_blocks):
for b, block in enumerate(self.double_blocks):
if b >= 0 and b <= self.double_blocks_to_swap:
mm.soft_empty_cache()
block.to(self.main_device)
double_block_args = [
img,
txt,
@@ -645,11 +668,17 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
]
img, txt = block(*double_block_args)
if b >= 0 and b <= self.double_blocks_to_swap:
mm.soft_empty_cache()
block.to(self.offload_device, non_blocking=True)
# Merge txt and img to pass through single stream blocks.
x = torch.cat((img, txt), 1)
if len(self.single_blocks) > 0:
for _, block in enumerate(self.single_blocks):
for b, block in enumerate(self.single_blocks):
if b >= 0 and b <= self.single_blocks_to_swap:
mm.soft_empty_cache()
block.to(self.main_device)
single_block_args = [
x,
vec,
@@ -662,6 +691,9 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
]
x = block(*single_block_args)
if b >= 0 and b <= self.single_blocks_to_swap:
mm.soft_empty_cache()
block.to(self.offload_device, non_blocking=True)
img = x[:, :img_seq_len, ...]
+33 -14
View File
@@ -74,6 +74,24 @@ def get_rotary_pos_embed(transformer, video_length, height, width):
)
return freqs_cos, freqs_sin
class HyVideoBlockSwap:
@classmethod
def INPUT_TYPES(s):
return {
"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"}),
},
}
RETURN_TYPES = ("BLOCKSWAPARGS",)
RETURN_NAMES = ("block_swap_args",)
FUNCTION = "setargs"
CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "Settings for block swapping, reduces VRAM use by swapping blocks to CPU memory"
def setargs(self, **kwargs):
return (kwargs, )
#region Model loading
class HyVideoModelLoader:
@classmethod
@@ -85,7 +103,6 @@ class HyVideoModelLoader:
"base_precision": (["fp16", "fp32", "bf16"], {"default": "bf16"}),
"quantization": (['disabled', 'fp8_e4m3fn', 'torchao_fp8dq', "torchao_fp8dqrow", "torchao_int8dq", "torchao_fp6"], {"default": 'disabled', "tooltip": "optional quantization method"}),
"load_device": (["main_device", "offload_device"], {"default": "main_device"}),
"enable_sequential_cpu_offload": ("BOOLEAN", {"default": False, "tooltip": "significantly reducing memory usage and slows down the inference"}),
},
"optional": {
"attention_mode": ([
@@ -94,6 +111,7 @@ class HyVideoModelLoader:
"sageattn_varlen",
], {"default": "flash_attn"}),
"compile_args": ("COMPILEARGS", ),
"block_swap_args": ("BLOCKSWAPARGS", ),
}
}
@@ -103,7 +121,7 @@ class HyVideoModelLoader:
CATEGORY = "HunyuanVideoWrapper"
def loadmodel(self, model, base_precision, load_device, quantization,
compile_args=None, attention_mode="sdpa", enable_sequential_cpu_offload=False):
compile_args=None, attention_mode="sdpa", enable_sequential_cpu_offload=False, block_swap_args=None):
transformer = None
manual_offloading = True
if "sage" in attention_mode:
@@ -124,7 +142,7 @@ class HyVideoModelLoader:
sd = load_torch_file(model_path, device=transformer_load_device)
in_channels = out_channels = 16
factor_kwargs = {"device": device, "dtype": base_dtype}
factor_kwargs = {"device": transformer_load_device, "dtype": base_dtype}
HUNYUAN_VIDEO_CONFIG = {
"mm_double_blocks_depth": 20,
"mm_single_blocks_depth": 40,
@@ -139,6 +157,7 @@ class HyVideoModelLoader:
in_channels=in_channels,
out_channels=out_channels,
attention_mode=attention_mode,
offload_device=offload_device,
**HUNYUAN_VIDEO_CONFIG,
**factor_kwargs
)
@@ -195,9 +214,9 @@ class HyVideoModelLoader:
log.info(f"Quantized transformer blocks to {quantization}")
scheduler = FlowMatchDiscreteScheduler(
shift=9.0, #this is not even used?
reverse=True, #has to be true or noise
solver="euler", #has to be euler
shift=9.0,
reverse=True,
solver="euler",
)
pipe = HunyuanVideoPipeline(
@@ -205,19 +224,15 @@ class HyVideoModelLoader:
scheduler=scheduler,
progress_bar_config=None
)
if enable_sequential_cpu_offload:
pipe.enable_sequential_cpu_offload()
manual_offloading = False
pipeline = {
"pipe": pipe,
"dtype": base_dtype,
"base_path": model_path,
"cpu_offloading": enable_sequential_cpu_offload,
"model_name": model,
"manual_offloading": manual_offloading,
"quantization": "disabled",
"block_swap_args": block_swap_args
}
return (pipeline,)
@@ -630,9 +645,11 @@ class HyVideoSampler:
# mm.get_autocast_device(device), dtype=dtype
# ) if any(q in model["quantization"] for q in ("e4m3fn", "GGUF")) else nullcontext()
#with autocast_context:
if not model["cpu_offloading"] and model["manual_offloading"]:
if model["block_swap_args"] is not None:
model["pipe"].transformer.to(device)
model["pipe"].transformer.block_swap(20, 0)
elif not model["manual_offloading"]:
model["pipe"].transformer.to(device)
out_latents = model["pipe"](
num_inference_steps=steps,
@@ -656,7 +673,7 @@ class HyVideoSampler:
pass
if force_offload:
if not model["cpu_offloading"] and model["manual_offloading"]:
if model["manual_offloading"]:
model["pipe"].transformer.to(offload_device)
mm.soft_empty_cache()
@@ -834,6 +851,7 @@ NODE_CLASS_MAPPINGS = {
"HyVideoVAELoader": HyVideoVAELoader,
"DownloadAndLoadHyVideoTextEncoder": DownloadAndLoadHyVideoTextEncoder,
"HyVideoEncode": HyVideoEncode,
"HyVideoBlockSwap": HyVideoBlockSwap,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"HyVideoSampler": "HunyuanVideo Sampler",
@@ -843,4 +861,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"HyVideoVAELoader": "HunyuanVideo VAE Loader",
"DownloadAndLoadHyVideoTextEncoder": "(Down)Load HunyuanVideo TextEncoder",
"HyVideoEncode": "HunyuanVideo Encode",
"HyVideoBlockSwap": "HunyuanVideo BlockSwap",
}