# -*- coding: utf-8 -*- # Copyright (c) 2025 Salvador E. Tropea # Copyright (c) 2025 Instituto Nacional de TecnologĂ­a Industrial # License: CC BY-NC-SA 4.0 # Project: ComfyUI-Float_Optimized import contextlib # For context manager import logging import torch try: import comfy.model_management as mm with_comfy = True except Exception: with_comfy = False from .misc import NODES_NAME logger = logging.getLogger(f"{NODES_NAME}.torch") def get_torch_device_options(): # We always have CPU default = "cpu" options = [default] # Do we have CUDA? if torch.cuda.is_available(): default = "cuda" options.append(default) for i in range(torch.cuda.device_count()): options.append(f"cuda:{i}") # Specific CUDA devices # Is this a Mac? if torch.backends.mps.is_available() and torch.backends.mps.is_built(): options.append("mps") if default == "cpu": default = "mps" return options, default # ################################################################################## # # Helper for inference (Target device, offload, eval, no_grad and cuDNN Benchmark) # ################################################################################## @contextlib.contextmanager def model_to_target(model): """ Consolidated context manager for model device placement and inference state. - Moves the model to its designated `model.target_device`. - Sets `torch.backends.cudnn.benchmark` based on `model.cudnn_benchmark_setting` if available. - Sets the model to `eval()` mode. - Wraps the operation in a `torch.no_grad()` context. - Offloads the model to the CPU (`mm.unet_offload_device()`) afterwards. """ if not isinstance(model, torch.nn.Module): with torch.no_grad(): yield # The code inside the 'with' statement runs here return # 1. Determine target device from the model object try: target_device = model.target_device assert isinstance(target_device, torch.device) except Exception as e: logger.warning(f"model_to_target: Could not get 'target_device' from model ({e}). " "Defaulting to model's current device.") target_device = next(model.parameters()).device # 2. Get CUDNN benchmark setting from the model object (optional) # Use hasattr as this is an optional setting that not all models might have. cudnn_benchmark_enabled = None # Default is to keep the current setting if hasattr(model, 'cudnn_benchmark_setting'): cudnn_benchmark_enabled = model.cudnn_benchmark_setting original_device = next(model.parameters()).device original_cudnn_benchmark_state = None is_cuda_target = target_device.type == 'cuda' try: # 3. Manage cuDNN benchmark state if (cudnn_benchmark_enabled is not None and is_cuda_target and hasattr(torch.backends, 'cudnn') and torch.backends.cudnn.is_available()): if torch.backends.cudnn.benchmark != cudnn_benchmark_enabled: original_cudnn_benchmark_state = torch.backends.cudnn.benchmark torch.backends.cudnn.benchmark = cudnn_benchmark_enabled logger.debug(f"Temporarily set cuDNN benchmark to {torch.backends.cudnn.benchmark}") # 4. Move model to target device if not already there if original_device != target_device: logger.debug(f"Moving model from `{original_device}` to target device `{target_device}` for inference.") model.to(target_device) # 5. Set to eval mode and disable gradients for the operation model.eval() with torch.no_grad(): yield # The code inside the 'with' statement runs here finally: # 6. Restore original cuDNN benchmark state if original_cudnn_benchmark_state is not None: # This check is sufficient because it will only be not None if we set it inside the try block torch.backends.cudnn.benchmark = original_cudnn_benchmark_state logger.debug(f"Restored cuDNN benchmark to {original_cudnn_benchmark_state}") # 7. Offload model back to CPU if with_comfy: offload_device = mm.unet_offload_device() current_device_after_yield = next(model.parameters()).device if current_device_after_yield != offload_device: logger.debug(f"Offloading model from `{current_device_after_yield}` to offload device `{offload_device}`.") model.to(offload_device) # Clear cache if we were on a CUDA device if 'cuda' in str(current_device_after_yield): torch.cuda.empty_cache()