diff --git a/__init__.py b/__init__.py index 4efb858..6e04899 100644 --- a/__init__.py +++ b/__init__.py @@ -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 diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 5a53877..79f7f42 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -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 diff --git a/nodes_sampler.py b/nodes_sampler.py index 489e216..9f415a5 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -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) diff --git a/utils.py b/utils.py index 243e5a6..4636dee 100644 --- a/utils.py +++ b/utils.py @@ -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/ diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 27d5384..887cbe1 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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, diff --git a/wanvideo/wan_video_vae.py b/wanvideo/wan_video_vae.py index 757a77f..38c0b4d 100644 --- a/wanvideo/wan_video_vae.py +++ b/wanvideo/wan_video_vae.py @@ -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 \ No newline at end of file + 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