From 2fd0a1296f9223b4b045759b631115787513a17b Mon Sep 17 00:00:00 2001 From: loxotron Date: Sat, 17 May 2025 22:22:55 +0300 Subject: [PATCH] Fixes for DirectML detection in fast_sampler.py, bugfixes in VRAM management for DirectML devices, improved logging for better debugging and profiling --- comfy/model_management.py | 71 ++++++++++++++++--- fast_sampler.py | 139 ++++++++++++++++++++++++++------------ nodes.py | 2 +- 3 files changed, 158 insertions(+), 54 deletions(-) diff --git a/comfy/model_management.py b/comfy/model_management.py index dac5985dc..27d79545d 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -87,6 +87,9 @@ _directml_vram_cache = {} # Cache for active models memory in DirectML _directml_active_memory_cache = {} +# Model management +current_loaded_models = [] + def cpu_mode(): """Check if system is in CPU mode.""" global cpu_state @@ -177,7 +180,7 @@ def get_directml_vram(dev): # Try torch_directml heuristic if _torch_directml_available: try: - device_index = dev.index if hasattr(dev, 'index') else 0 + device_index = dev.index if hasattr(dev, 'index') and dev.index is not None else 0 device_name = torch_directml.device_name(device_index).lower() vram_map = { 'gtx 1660': 6 * 1024 * 1024 * 1024, @@ -188,6 +191,7 @@ def get_directml_vram(dev): 'rx 580': 8 * 1024 * 1024 * 1024, 'rx 570': 8 * 1024 * 1024 * 1024, 'rx 6700': 12 * 1024 * 1024 * 1024, + 'rx 6800': 16 * 1024 * 1024 * 1024, 'arc a770': 16 * 1024 * 1024 * 1024, } vram = 6 * 1024 * 1024 * 1024 @@ -579,10 +583,55 @@ def get_free_memory(dev=None, torch_free_too=False): if directml_enabled: total_vram = get_directml_vram(dev) cache_key = (dev, 'active_models') - if cache_key not in _directml_active_memory_cache: - active_models = sum(m.model_loaded_memory() for m in current_loaded_models if m.device == dev) - _directml_active_memory_cache[cache_key] = active_models - active_models = _directml_active_memory_cache[cache_key] + # Invalidate cache if models list has changed + current_models_hash = hash(tuple((id(m), m.model_loaded_memory() if not m.is_dead() else 0) for m in current_loaded_models)) + if cache_key in _directml_active_memory_cache: + cached_hash, cached_active_models = _directml_active_memory_cache[cache_key] + if cached_hash == current_models_hash: + active_models = cached_active_models + if DEBUG_ENABLED: + logging.debug(f"Using cached active_models={active_models / (1024**3):.2f} GB for device {dev}") + else: + if DEBUG_ENABLED: + logging.debug(f"Cache invalidated for {dev}: models list changed") + active_models = None + else: + active_models = None + + if active_models is None: + active_models = 0 + try: + if DEBUG_ENABLED: + logging.debug(f"Processing {len(current_loaded_models)} models in get_free_memory for device {dev}") + for m in current_loaded_models: + model_name = m.model.__class__.__name__ if m.model else "Unknown" + if m.device != dev: + if DEBUG_ENABLED: + logging.debug(f"Skipping model {model_name}: device mismatch (model on {m.device}, expected {dev})") + continue + if m.is_dead(): + if DEBUG_ENABLED: + logging.debug(f"Skipping model {model_name}: model is dead") + continue + try: + mem = m.model_loaded_memory() + if DEBUG_ENABLED: + logging.debug(f"Loaded model {model_name} on device {m.device}, memory={mem / (1024**3):.2f} GB, is_dead={m.is_dead()}") + if mem <= 0: + logging.warning(f"Model {model_name} returned invalid memory: {mem}. Skipping.") + continue + active_models += mem + if DEBUG_ENABLED: + logging.debug(f"Model {model_name} on {dev}: loaded_memory={mem / (1024**3):.2f} GB") + except Exception as e: + logging.warning(f"Failed to calculate memory for model {model_name}: {str(e)}") + # Update cache + _directml_active_memory_cache[cache_key] = (current_models_hash, active_models) + except NameError: + logging.warning("current_loaded_models not defined yet in get_free_memory") + _directml_active_memory_cache[cache_key] = (current_models_hash, 0) + + # Apply safety margin (1.2x) and ensure at least 1 GB free mem_free_total = max(1024 * 1024 * 1024, total_vram - active_models * 1.2) mem_free_torch = mem_free_total if DEBUG_ENABLED: @@ -776,9 +825,6 @@ def register_vram_optimizer(optimizer): """Register a VRAM optimizer.""" _vram_optimizers.append(optimizer) -# Model management -current_loaded_models = [] - class LoadedModel: def __init__(self, model): self._set_model(model) @@ -972,6 +1018,15 @@ def module_size(model, shape=None, dtype=None): """ from diffusers import AutoencoderKL + # Early check for None model to avoid unnecessary processing + if model is None: + if DEBUG_ENABLED: + logging.warning( + f"module_size: Received None model. Assuming minimal memory (1 MB). " + f"Call stack: {''.join(traceback.format_stack(limit=5))}" + ) + return 1024 * 1024 # Minimal memory assumption for None model + module_mem = 0 if shape is not None and dtype is not None and isinstance(model, AutoencoderKL): try: diff --git a/fast_sampler.py b/fast_sampler.py index a821d7895..ae752cbb9 100644 --- a/fast_sampler.py +++ b/fast_sampler.py @@ -33,9 +33,9 @@ def profile_section(name): else: yield -def profile_cuda_sync(is_gpu, message="CUDA sync"): +def profile_cuda_sync(is_gpu, device, message="CUDA sync"): """Profile CUDA synchronization time if GPU is used.""" - if PROFILING_ENABLED and is_gpu: + if PROFILING_ENABLED and is_gpu and device.type == 'cuda': logging.debug(f"{message} started") sync_start = time.time() torch.cuda.synchronize() @@ -64,7 +64,7 @@ def initialize_device_and_dtype(model, device=None): if device is None: device = get_torch_device() dtype = getattr(model, 'dtype', torch.float32) - is_gpu = device.type == 'cuda' and torch.cuda.is_available() + is_gpu = (device.type == 'cuda' and torch.cuda.is_available()) or (device.type == 'privateuseone') return device, dtype, is_gpu def clear_vram(device, threshold=0.5, min_free=1.5): @@ -87,6 +87,13 @@ def clear_vram(device, threshold=0.5, min_free=1.5): if PROFILING_ENABLED: logging.debug(f"VRAM not cleared: {mem_allocated:.2f} GB / {mem_total:.2f} GB, sufficient free memory") return mem_allocated, mem_total + elif device.type == 'privateuseone': + # For DirectML just clearing cache, cause mem_get_info not working + if PROFILING_ENABLED: + logging.debug("Clearing VRAM for DirectML (mem_get_info unavailable)") + torch.cuda.empty_cache() + return 0, 0 + return 0, 0 def preload_model(model, device, is_vae=False): """Preload model or VAE to device, avoiding unnecessary unloading.""" @@ -129,8 +136,13 @@ def preload_model(model, device, is_vae=False): comfy.model_management.load_model_gpu(model) model._loaded_to_device = device if PROFILING_ENABLED: - free_mem = (torch.cuda.get_device_properties(device).total_memory - torch.cuda.memory_allocated(device)) / 1024**3 - logging.debug(f"U-Net loaded to {device}, VRAM free: {free_mem:.2f} GB") + if device.type == 'cuda': + free_mem = (torch.cuda.get_device_properties(device).total_memory - torch.cuda.memory_allocated(device)) / 1024**3 + logging.debug(f"U-Net loaded to {device}, VRAM free: {free_mem:.2f} GB") + elif device.type == 'privateuseone': + logging.debug(f"U-Net loaded to {device}, VRAM info unavailable") + else: + logging.debug(f"U-Net loaded to {device}, no VRAM info for non-GPU device") def optimized_transfer(tensor, device, dtype): """Synchronous tensor transfer to device.""" @@ -150,11 +162,11 @@ def finalize_images(images, device): """Process and finalize output images.""" if len(images.shape) == 5: # Combine batches images = images.reshape(-1, images.shape[-3], images.shape[-2], images.shape[-1]) - # Apply channels_last only for CUDA devices if force_channels_last is enabled - is_gpu = device.type == 'cuda' and torch.cuda.is_available() + # Apply channels_last for CUDA or DirectML devices if force_channels_last is enabled + is_gpu = (device.type == 'cuda' and torch.cuda.is_available()) or (device.type == 'privateuseone') memory_format = torch.channels_last if (is_gpu and not directml_enabled and force_channels_last()) else torch.contiguous_format if DEBUG_ENABLED: - logging.debug(f"finalize_images: Using memory_format={memory_format} for device={device}, directml_enabled={directml_enabled}") + logging.debug(f"finalize_images: Using memory_format={memory_format} for device={device}, directml_enabled={directml_enabled}, is_gpu={is_gpu}") return images.to(device=device, memory_format=memory_format) def fast_sample(model, noise, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, @@ -195,9 +207,11 @@ def fast_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, ne if DEBUG_ENABLED: if model is None: logging.warning("fast_ksampler: model is None") - + logging.debug(f"Starting fast_ksampler, device={device}, is_gpu={is_gpu}") if device is None or dtype is None or is_gpu is None: device, dtype, is_gpu = initialize_device_and_dtype(model.model) + if DEBUG_ENABLED: + logging.debug(f"Initialized device: {device}, dtype: {dtype}, is_gpu: {is_gpu}") try: # Enable cuDNN benchmarking if requested @@ -270,14 +284,17 @@ def fast_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, ne # Log VRAM state after sampling if is_gpu and PROFILING_ENABLED: - mem_total = torch.cuda.get_device_properties(device).total_memory / 1024**3 - mem_allocated = torch.cuda.memory_allocated(device) / 1024**3 - logging.debug(f"VRAM after sampling: {mem_allocated:.2f} GB / {mem_total:.2f} GB") + if device.type == 'cuda': + mem_total = torch.cuda.get_device_properties(device).total_memory / 1024**3 + mem_allocated = torch.cuda.memory_allocated(device) / 1024**3 + logging.debug(f"VRAM after sampling: {mem_allocated:.2f} GB / {mem_total:.2f} GB") + elif device.type == 'privateuseone': + logging.debug("VRAM info unavailable for DirectML after sampling") # Log completion of sampling if PROFILING_ENABLED: logging.debug(f"Sampling completed, preparing for VAE") - profile_cuda_sync(is_gpu) + profile_cuda_sync(is_gpu, device) # Clear VRAM after sampling if is_gpu: @@ -285,8 +302,12 @@ def fast_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, ne clear_vram(device, threshold=0.5, min_free=1.5) else: clear_start = time.time() - mem_allocated, mem_total = clear_vram(device, threshold=0.5, min_free=1.5) - logging.debug(f"VRAM after sampling: {mem_allocated:.2f} GB / {mem_total:.2f} GB, clear took {time.time() - clear_start:.3f} s") + if device.type == 'cuda': + mem_allocated, mem_total = clear_vram(device, threshold=0.5, min_free=1.5) + logging.debug(f"VRAM after sampling: {mem_allocated:.2f} GB / {mem_total:.2f} GB, clear took {time.time() - clear_start:.3f} s") + elif device.type == 'privateuseone': + clear_vram(device, threshold=0.5, min_free=1.5) + logging.debug(f"VRAM clear after sampling took {time.time() - clear_start:.3f} s (VRAM info unavailable for DirectML)") logging.debug(f"Post-VRAM checkpoint: {time.time()}") out = latent.copy() @@ -307,11 +328,13 @@ def fast_vae_decode(vae, samples): """ device = get_torch_device() vae_dtype_val = vae_dtype(device=device) - is_gpu = device.type == 'cuda' and torch.cuda.is_available() + is_gpu = (device.type == 'cuda' and torch.cuda.is_available()) or (device.type == 'privateuseone') if DEBUG_ENABLED: logging.debug(f"VAE dtype: {vae_dtype_val}") logging.debug(f"Pre-VAE checkpoint: {time.time()}") + logging.debug(f"Starting fast_vae_decode, device={device}, dtype={vae_dtype_val}, is_gpu={is_gpu}") + logging.debug(f"Latent samples shape: {samples['samples'].shape}") try: # Disable cuDNN benchmark for VAE stability if enabled @@ -320,21 +343,25 @@ def fast_vae_decode(vae, samples): # Prepare VRAM for VAE if is_gpu: - mem_total = torch.cuda.get_device_properties(device).total_memory / 1024**3 - latent_size = samples["samples"].shape - model_for_memory = getattr(vae, 'first_stage_model', vae) - vae_memory_required = estimate_vae_decode_memory(model_for_memory, latent_size, vae_dtype_val) / 1024**3 - vram_threshold = 1.0 if mem_total < 5.9 else 1.1 - vae_memory_required *= vram_threshold - if PROFILING_ENABLED: - logging.debug(f"Estimated VAE memory: {vae_memory_required:.2f} GB") - mem_allocated = torch.cuda.memory_allocated(device) / 1024**3 - free_mem = mem_total - mem_allocated - if free_mem < vae_memory_required: - #free_memory(vae_memory_required) - mem_allocated, mem_total = clear_vram(device, threshold=0.4, min_free=2.0) + if device.type == 'cuda': + mem_total = torch.cuda.get_device_properties(device).total_memory / 1024**3 + latent_size = samples["samples"].shape + model_for_memory = getattr(vae, 'first_stage_model', vae) + vae_memory_required = estimate_vae_decode_memory(model_for_memory, latent_size, vae_dtype_val) / 1024**3 + vram_threshold = 1.0 if mem_total < 5.9 else 1.1 + vae_memory_required *= vram_threshold if PROFILING_ENABLED: - logging.debug(f"VRAM after free_memory: {mem_allocated:.2f} GB / {mem_total:.2f} GB") + logging.debug(f"Estimated VAE memory: {vae_memory_required:.2f} GB") + mem_allocated = torch.cuda.memory_allocated(device) / 1024**3 + free_mem = mem_total - mem_allocated + if free_mem < vae_memory_required: + #free_memory(vae_memory_required) + mem_allocated, mem_total = clear_vram(device, threshold=0.4, min_free=2.0) + if PROFILING_ENABLED: + logging.debug(f"VRAM after free_memory: {mem_allocated:.2f} GB / {mem_total:.2f} GB") + elif device.type == 'privateuseone': + if PROFILING_ENABLED: + logging.debug("Memory info unavailable for DirectML, skipping VRAM check") # Preload VAE to device preload_model(vae, device, is_vae=True) @@ -384,6 +411,8 @@ def fast_vae_tiled_decode(vae, samples, tile_size=512, overlap=64, temporal_size if DEBUG_ENABLED: logging.debug(f"VAE dtype: {vae_dtype_val}") logging.debug(f"Pre-VAE checkpoint: {time.time()}") + logging.debug(f"Starting fast_vae_tiled_decode, device={device}, dtype={vae_dtype_val}, is_gpu={is_gpu}") + logging.debug(f"Latent samples shape: {samples['samples'].shape}, tile_size={tile_size}, overlap={overlap}") try: # Disable cuDNN benchmark for tiled decoding stability if enabled @@ -392,25 +421,42 @@ def fast_vae_tiled_decode(vae, samples, tile_size=512, overlap=64, temporal_size # Clear VRAM before VAE if is_gpu: - mem_total = torch.cuda.get_device_properties(device).total_memory / 1024**3 - mem_allocated = torch.cuda.memory_allocated(device) / 1024**3 - free_mem = mem_total - mem_allocated + if device.type == 'cuda': + mem_total = torch.cuda.get_device_properties(device).total_memory / 1024**3 + mem_allocated = torch.cuda.memory_allocated(device) / 1024**3 + free_mem = mem_total - mem_allocated + elif device.type == 'privateuseone': + if PROFILING_ENABLED: + logging.debug("Memory info unavailable for DirectML, skipping VRAM check") + # Estimate memory for tiled decoding (conservative, ~50% of full decode) vae_memory_required = (vae.memory_used_decode(samples["samples"].shape, vae_dtype_val) / 1024**3 * 0.5 if hasattr(vae, 'memory_used_decode') else 0.75) if PROFILING_ENABLED: - logging.debug(f"VRAM before tiled VAE: {mem_allocated:.2f} GB / {mem_total:.2f} GB") + if device.type == 'cuda': + logging.debug(f"VRAM before tiled VAE: {mem_allocated:.2f} GB / {mem_total:.2f} GB") + else: + logging.debug("VRAM before tiled VAE: unavailable for DirectML") logging.debug(f"Estimated tiled VAE memory: {vae_memory_required:.2f} GB") # Skip VRAM cleanup if VAE is already loaded and memory is sufficient - if (hasattr(vae, '_loaded_to_device') and vae._loaded_to_device == device and - free_mem >= vae_memory_required * 1.1): - if PROFILING_ENABLED: - logging.debug(f"VAE already loaded, sufficient memory: {free_mem:.2f} GB") - elif mem_allocated > 0.4 * mem_total or free_mem < vae_memory_required: - if PROFILING_ENABLED: - logging.debug(f"Clearing VRAM: {mem_allocated:.2f} GB used of {mem_total:.2f} GB") - mem_allocated, mem_total = clear_vram(device, threshold=0.4, min_free=0.75) + if device.type == 'cuda': + if (hasattr(vae, '_loaded_to_device') and vae._loaded_to_device == device and + free_mem >= vae_memory_required * 1.1): + if PROFILING_ENABLED: + logging.debug(f"VAE already loaded, sufficient memory: {free_mem:.2f} GB") + elif mem_allocated > 0.4 * mem_total or free_mem < vae_memory_required: + if PROFILING_ENABLED: + logging.debug(f"Clearing VRAM: {mem_allocated:.2f} GB used of {mem_total:.2f} GB") + mem_allocated, mem_total = clear_vram(device, threshold=0.4, min_free=0.75) + else: + if (hasattr(vae, '_loaded_to_device') and vae._loaded_to_device == device): + if PROFILING_ENABLED: + logging.debug("VAE already loaded on DirectML, skipping VRAM cleanup") + else: + if PROFILING_ENABLED: + logging.debug("Clearing VRAM for DirectML") + clear_vram(device, threshold=0.4, min_free=0.75) # Preload VAE if not PROFILING_ENABLED: @@ -481,9 +527,12 @@ def fast_vae_tiled_decode(vae, samples, tile_size=512, overlap=64, temporal_size images = finalize_images(images, device) if is_gpu and PROFILING_ENABLED: - mem_allocated = torch.cuda.memory_allocated(device) / 1024**3 - mem_total = torch.cuda.get_device_properties(device).total_memory / 1024**3 - logging.debug(f"VRAM after tiled decoding: {mem_allocated:.2f} GB / {mem_total:.2f} GB") + if device.type == 'cuda': + mem_allocated = torch.cuda.memory_allocated(device) / 1024**3 + mem_total = torch.cuda.get_device_properties(device).total_memory / 1024**3 + logging.debug(f"VRAM after tiled decoding: {mem_allocated:.2f} GB / {mem_total:.2f} GB") + else: + logging.debug("VRAM after tiled decoding: unavailable for DirectML") logging.debug(f"Post-decode checkpoint: {time.time()}") if PROFILING_ENABLED: diff --git a/nodes.py b/nodes.py index 28a38035b..e3e3d961a 100644 --- a/nodes.py +++ b/nodes.py @@ -1468,7 +1468,7 @@ def common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, # Get device and dtype device = comfy.model_management.get_torch_device() dtype = getattr(model.model, 'dtype', torch.float32) - is_gpu = device.type == 'cuda' and torch.cuda.is_available() + is_gpu = (device.type == 'cuda' and torch.cuda.is_available()) or (device.type == 'privateuseone') # Prepare latent image latent_image = latent["samples"]