Make VAE memory reporting optional to reduce log spam, and other logging updates
This commit is contained in:
+3
-3
@@ -1,11 +1,11 @@
|
||||
try:
|
||||
from .utils import check_duplicate_nodes, log
|
||||
from .utils import check_duplicate_nodes, log, color_text
|
||||
duplicate_dirs = check_duplicate_nodes()
|
||||
if duplicate_dirs:
|
||||
warning_msg = f"WARNING: Found {len(duplicate_dirs)} other WanVideoWrapper directories:\n"
|
||||
for dir_path in duplicate_dirs:
|
||||
warning_msg += f" - {dir_path}\n"
|
||||
log.warning(warning_msg + "Please remove duplicates to avoid possible conflicts.")
|
||||
warning_msg += f" - {color_text(dir_path, 'yellow')}\n"
|
||||
log.warning(color_text(warning_msg + "Please remove duplicates to avoid possible conflicts.", "red"))
|
||||
except:
|
||||
pass
|
||||
|
||||
|
||||
+11
-5
@@ -13,7 +13,7 @@ from .wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38
|
||||
from .custom_linear import _replace_linear
|
||||
|
||||
from accelerate import init_empty_weights
|
||||
from .utils import set_module_tensor_to_device
|
||||
from .utils import set_module_tensor_to_device, get_module_memory_mb_per_device
|
||||
|
||||
import folder_paths
|
||||
import comfy.model_management as mm
|
||||
@@ -864,7 +864,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
)
|
||||
transformer.gguf_patched = True
|
||||
else:
|
||||
log.info("Using accelerate to load and assign model weights to device...")
|
||||
log.info("Loading and assigning model weights to device...")
|
||||
named_params = transformer.named_parameters()
|
||||
|
||||
for name, param in tqdm(named_params,
|
||||
@@ -925,6 +925,11 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
pbar.update(100)
|
||||
|
||||
#[print(name, param.device, param.dtype) for name, param in transformer.named_parameters()]
|
||||
memory_on_device = get_module_memory_mb_per_device(transformer)
|
||||
log.info("-" * 25)
|
||||
log.info("Transformer weights loaded:")
|
||||
for dev, mem_mb in memory_on_device.items():
|
||||
log.info(f"Device: {dev:8s} | Memory: {mem_mb:,.2f} MB")
|
||||
|
||||
pbar.update_absolute(0)
|
||||
|
||||
@@ -1856,6 +1861,7 @@ class WanVideoVAELoader:
|
||||
),
|
||||
"compile_args": ("WANCOMPILEARGS", ),
|
||||
"use_cpu_cache": ("BOOLEAN", {"default": False, "tooltip": "Reduces VRAM usage, but slows the VAE down a lot"}),
|
||||
"verbose": ("BOOLEAN", {"default": False, "tooltip": "Enables memory usage logging when using the model"}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1865,7 +1871,7 @@ class WanVideoVAELoader:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Loads Wan VAE model from 'ComfyUI/models/vae'"
|
||||
|
||||
def loadmodel(self, model_name, precision, compile_args=None, use_cpu_cache=False):
|
||||
def loadmodel(self, model_name, precision, compile_args=None, use_cpu_cache=False, verbose=False):
|
||||
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||
model_path = folder_paths.get_full_path_or_raise("vae", model_name)
|
||||
vae_sd = load_torch_file(model_path, safe_load=True)
|
||||
@@ -1882,9 +1888,9 @@ class WanVideoVAELoader:
|
||||
pruning_rate = 0.0
|
||||
|
||||
if vae_sd["model.conv2.weight"].shape[0] == 16:
|
||||
vae = WanVideoVAE(dtype=dtype, pruning_rate=pruning_rate, cpu_cache=use_cpu_cache)
|
||||
vae = WanVideoVAE(dtype=dtype, pruning_rate=pruning_rate, cpu_cache=use_cpu_cache, verbose=verbose)
|
||||
elif vae_sd["model.conv2.weight"].shape[0] == 48:
|
||||
vae = WanVideoVAE38(dtype=dtype, pruning_rate=pruning_rate, cpu_cache=use_cpu_cache)
|
||||
vae = WanVideoVAE38(dtype=dtype, pruning_rate=pruning_rate, cpu_cache=use_cpu_cache, verbose=verbose)
|
||||
|
||||
vae.load_state_dict(vae_sd)
|
||||
del vae_sd
|
||||
|
||||
+4
-5
@@ -177,7 +177,6 @@ class WanVideoSampler:
|
||||
start_step = scheduler.get("start_step", start_step)
|
||||
elif scheduler != "multitalk":
|
||||
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, denoise_strength, sigmas=sigmas, log_timesteps=True)
|
||||
log.info(f"sigmas: {sample_scheduler.sigmas}")
|
||||
else:
|
||||
timesteps = torch.tensor([1000, 750, 500, 250], device=device)
|
||||
|
||||
@@ -241,8 +240,6 @@ class WanVideoSampler:
|
||||
else:
|
||||
image_cond[:, 1:] = 0
|
||||
|
||||
log.info(f"image_cond shape: {image_cond.shape}")
|
||||
|
||||
#ATI tracks
|
||||
if transformer_options is not None:
|
||||
ATI_tracks = transformer_options.get("ati_tracks", None)
|
||||
@@ -1679,8 +1676,8 @@ class WanVideoSampler:
|
||||
callback = prepare_callback(patcher, len(timesteps))
|
||||
|
||||
if not multitalk_sampling and not framepack and not wananimate_loop:
|
||||
log.info(f"Input sequence length: {seq_len}")
|
||||
log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps-ttm_start_step} steps")
|
||||
log.info("-" * 10 + " Sampling start " + "-" * 10)
|
||||
log.info(f"{(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} (Input sequence length: {seq_len}) with {steps-ttm_start_step} steps")
|
||||
|
||||
|
||||
# Differential diffusion prep
|
||||
@@ -2573,6 +2570,8 @@ class WanVideoSampler:
|
||||
if story_mem_latents is not None:
|
||||
latent = latent[:, story_mem_latents.shape[1]:]
|
||||
|
||||
log.info("-" * 10 + " Sampling end " + "-" * 12)
|
||||
|
||||
cache_states = None
|
||||
if cache_args is not None:
|
||||
cache_report(transformer, cache_args)
|
||||
|
||||
@@ -25,6 +25,23 @@ try:
|
||||
except:
|
||||
pass
|
||||
|
||||
COLOR_CODES = {
|
||||
"reset": "\033[0m",
|
||||
"red": "\033[31m",
|
||||
"green": "\033[32m",
|
||||
"yellow": "\033[33m",
|
||||
"blue": "\033[34m",
|
||||
"magenta": "\033[35m",
|
||||
"cyan": "\033[36m",
|
||||
"white": "\033[37m",
|
||||
}
|
||||
|
||||
def color_text(text, color):
|
||||
try:
|
||||
return f"{COLOR_CODES.get(color, COLOR_CODES['reset'])}{text}{COLOR_CODES['reset']}"
|
||||
except Exception:
|
||||
return text
|
||||
|
||||
class MetaParameter(torch.nn.Parameter):
|
||||
def __new__(cls, dtype, quant_type=None):
|
||||
data = torch.empty(0, dtype=dtype)
|
||||
@@ -191,10 +208,8 @@ def check_diffusers_version():
|
||||
raise AssertionError("diffusers is not installed.")
|
||||
|
||||
def print_memory(device, process="Sampling"):
|
||||
memory = torch.cuda.memory_allocated(device) / 1024**3
|
||||
max_memory = torch.cuda.max_memory_allocated(device) / 1024**3
|
||||
max_reserved = torch.cuda.max_memory_reserved(device) / 1024**3
|
||||
log.info(f"[{process}] Allocated memory: {memory=:.3f} GB")
|
||||
log.info(f"[{process}] Max allocated memory: {max_memory=:.3f} GB")
|
||||
log.info(f"[{process}] Max reserved memory: {max_reserved=:.3f} GB")
|
||||
#memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False)
|
||||
@@ -207,6 +222,18 @@ def get_module_memory_mb(module):
|
||||
memory += param.nelement() * param.element_size()
|
||||
return memory / (1024 * 1024) # Convert to MB
|
||||
|
||||
def get_module_memory_mb_per_device(module):
|
||||
memory_per_device = {}
|
||||
memory = 0
|
||||
for param in module.parameters():
|
||||
if param.data is not None:
|
||||
device = str(param.device)
|
||||
memory += param.nelement() * param.element_size()
|
||||
memory_per_device[device] = memory_per_device.get(device, 0) + memory
|
||||
|
||||
memory_per_device = {dev: mem / (1024 * 1024) for dev, mem in memory_per_device.items()}
|
||||
return memory_per_device
|
||||
|
||||
def get_tensor_memory(tensor):
|
||||
memory_bytes = tensor.element_size() * tensor.nelement()
|
||||
return f"{memory_bytes / (1024 * 1024):.2f} MB"
|
||||
@@ -666,9 +693,9 @@ def check_duplicate_nodes():
|
||||
"""Check ComfyUI custom_nodes directory for duplicate installations"""
|
||||
custom_nodes_dir = Path(folder_paths.folder_names_and_paths["custom_nodes"][0][0])
|
||||
current_path = Path(__file__).parent
|
||||
|
||||
|
||||
wanvideo_dirs = []
|
||||
|
||||
|
||||
# Check all directories in custom_nodes
|
||||
for path in custom_nodes_dir.iterdir():
|
||||
if (path.is_dir() and
|
||||
@@ -676,7 +703,7 @@ def check_duplicate_nodes():
|
||||
'wanvideo' in path.name.lower() and
|
||||
'wrapper' in path.name.lower()):
|
||||
wanvideo_dirs.append(str(path))
|
||||
|
||||
|
||||
return wanvideo_dirs
|
||||
|
||||
#https://github.com/temporalscorerescaling/TSR/
|
||||
|
||||
+10
-10
@@ -2019,24 +2019,24 @@ class WanModel(torch.nn.Module):
|
||||
def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False, vace_blocks_to_swap=None, prefetch_blocks=0, block_swap_debug=False):
|
||||
# Clamp blocks_to_swap to valid range
|
||||
blocks_to_swap = max(0, min(blocks_to_swap, len(self.blocks)))
|
||||
|
||||
|
||||
log.info(f"Swapping {blocks_to_swap} transformer blocks")
|
||||
self.blocks_to_swap = blocks_to_swap
|
||||
self.prefetch_blocks = prefetch_blocks
|
||||
self.block_swap_debug = block_swap_debug
|
||||
|
||||
|
||||
self.offload_img_emb = offload_img_emb
|
||||
self.offload_txt_emb = offload_txt_emb
|
||||
|
||||
total_offload_memory = 0
|
||||
total_main_memory = 0
|
||||
|
||||
|
||||
# Calculate the index where swapping starts
|
||||
swap_start_idx = len(self.blocks) - blocks_to_swap
|
||||
|
||||
|
||||
for b, block in tqdm(enumerate(self.blocks), total=len(self.blocks), desc="Initializing block swap"):
|
||||
block_memory = get_module_memory_mb(block)
|
||||
|
||||
|
||||
if b < swap_start_idx:
|
||||
block.to(self.main_device)
|
||||
total_main_memory += block_memory
|
||||
@@ -2051,13 +2051,13 @@ class WanModel(torch.nn.Module):
|
||||
# Clamp vace_blocks_to_swap to valid range
|
||||
vace_blocks_to_swap = max(0, min(vace_blocks_to_swap, len(self.vace_blocks)))
|
||||
self.vace_blocks_to_swap = vace_blocks_to_swap
|
||||
|
||||
|
||||
# Calculate the index where VACE swapping starts
|
||||
vace_swap_start_idx = len(self.vace_blocks) - vace_blocks_to_swap
|
||||
|
||||
for b, block in tqdm(enumerate(self.vace_blocks), total=len(self.vace_blocks), desc="Initializing vace block swap"):
|
||||
block_memory = get_module_memory_mb(block)
|
||||
|
||||
|
||||
if b < vace_swap_start_idx:
|
||||
block.to(self.main_device)
|
||||
total_main_memory += block_memory
|
||||
@@ -2068,13 +2068,13 @@ class WanModel(torch.nn.Module):
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
log.info("----------------------")
|
||||
log.info(f"Block swap memory summary:")
|
||||
log.info("-" * 25)
|
||||
log.info("Block swap memory summary:")
|
||||
log.info(f"Transformer blocks on {self.offload_device}: {total_offload_memory:.2f}MB")
|
||||
log.info(f"Transformer blocks on {self.main_device}: {total_main_memory:.2f}MB")
|
||||
log.info(f"Total memory used by transformer blocks: {(total_offload_memory + total_main_memory):.2f}MB")
|
||||
log.info(f"Non-blocking memory transfer: {self.use_non_blocking}")
|
||||
log.info("----------------------")
|
||||
log.info("-" * 25)
|
||||
|
||||
def forward_vace(
|
||||
self,
|
||||
|
||||
+41
-32
@@ -989,7 +989,8 @@ class VideoVAE_(nn.Module):
|
||||
mean=None,
|
||||
inv_std=None,
|
||||
pruning_rate=0.0,
|
||||
cpu_cache=False):
|
||||
cpu_cache=False,
|
||||
verbose=False):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
@@ -1000,6 +1001,7 @@ class VideoVAE_(nn.Module):
|
||||
self.temperal_upsample = temperal_downsample[::-1]
|
||||
self.mean = mean
|
||||
self.inv_std = inv_std
|
||||
self.verbose = verbose
|
||||
|
||||
# modules
|
||||
self.encoder = Encoder3d(dim, z_dim * 2, dim_mult, num_res_blocks,
|
||||
@@ -1085,12 +1087,13 @@ class VideoVAE_(nn.Module):
|
||||
std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0))
|
||||
eps = torch.randn_like(std)
|
||||
return mu + std * eps
|
||||
try:
|
||||
log.info(f"WanVAE encoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE encode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
if self.verbose:
|
||||
try:
|
||||
log.info(f"WanVAE encoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE encode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return mu
|
||||
|
||||
|
||||
@@ -1154,12 +1157,13 @@ class VideoVAE_(nn.Module):
|
||||
if pbar:
|
||||
pbar.update_absolute(0)
|
||||
self.clear_cache()
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
if self.verbose:
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return out
|
||||
|
||||
def reparameterize(self, mu, log_var):
|
||||
@@ -1186,12 +1190,12 @@ class VideoVAE_(nn.Module):
|
||||
|
||||
class WanVideoVAE(nn.Module):
|
||||
|
||||
def __init__(self, z_dim=16, dtype=torch.float32, pruning_rate=0.0, cpu_cache=False):
|
||||
def __init__(self, z_dim=16, dtype=torch.float32, pruning_rate=0.0, cpu_cache=False, verbose=False):
|
||||
super().__init__()
|
||||
|
||||
self.dtype = dtype
|
||||
self.cpu_cache = cpu_cache
|
||||
|
||||
self.verbose = verbose
|
||||
mean = [
|
||||
-0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508,
|
||||
0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921
|
||||
@@ -1205,7 +1209,7 @@ class WanVideoVAE(nn.Module):
|
||||
self.z_dim = z_dim
|
||||
|
||||
# init model
|
||||
self.model = VideoVAE_(z_dim=z_dim, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=self.cpu_cache).eval().requires_grad_(False)
|
||||
self.model = VideoVAE_(z_dim=z_dim, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=self.cpu_cache, verbose=self.verbose).eval().requires_grad_(False)
|
||||
self.upsampling_factor = 8
|
||||
|
||||
|
||||
@@ -1431,7 +1435,8 @@ class VideoVAE38_(VideoVAE_):
|
||||
mean=None,
|
||||
inv_std=None,
|
||||
pruning_rate=0.0,
|
||||
cpu_cache=False):
|
||||
cpu_cache=False,
|
||||
verbose=False):
|
||||
super(VideoVAE_, self).__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
@@ -1444,6 +1449,7 @@ class VideoVAE38_(VideoVAE_):
|
||||
self.mean = mean
|
||||
self.inv_std = inv_std
|
||||
self.cpu_cache = cpu_cache
|
||||
self.verbose = verbose
|
||||
|
||||
# modules
|
||||
self.encoder = Encoder3d_38(dim, z_dim * 2, dim_mult, num_res_blocks,
|
||||
@@ -1481,12 +1487,13 @@ class VideoVAE38_(VideoVAE_):
|
||||
mu = self.conv1(out).chunk(2, dim=1)[0]
|
||||
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
|
||||
self.clear_cache()
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
if self.verbose:
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return mu
|
||||
|
||||
|
||||
@@ -1519,18 +1526,19 @@ class VideoVAE38_(VideoVAE_):
|
||||
pbar.update(1)
|
||||
out = unpatchify(out, patch_size=2)
|
||||
self.clear_cache()
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
if self.verbose:
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return out
|
||||
|
||||
|
||||
class WanVideoVAE38(WanVideoVAE):
|
||||
|
||||
def __init__(self, z_dim=48, dim=160, dtype=torch.bfloat16, pruning_rate=0.0, cpu_cache=False):
|
||||
def __init__(self, z_dim=48, dim=160, dtype=torch.bfloat16, pruning_rate=0.0, cpu_cache=False, verbose=False):
|
||||
super(WanVideoVAE, self).__init__()
|
||||
|
||||
mean = [
|
||||
@@ -1554,7 +1562,8 @@ class WanVideoVAE38(WanVideoVAE):
|
||||
self.dtype = dtype
|
||||
self.z_dim = z_dim
|
||||
self.cpu_cache = cpu_cache
|
||||
self.verbose = verbose
|
||||
|
||||
# init model
|
||||
self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=cpu_cache).eval().requires_grad_(False)
|
||||
self.upsampling_factor = 16
|
||||
self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=cpu_cache, verbose=verbose).eval().requires_grad_(False)
|
||||
self.upsampling_factor = 16
|
||||
|
||||
Reference in New Issue
Block a user