2024-08-08 03:27:37 -04:00
"""
This file is part of ComfyUI.
Copyright (C) 2024 Comfy
This program is free software: you can redistribute it and/or modify
it under the terms of the GNU General Public License as published by
the Free Software Foundation, either version 3 of the License, or
(at your option) any later version.
This program is distributed in the hope that it will be useful,
but WITHOUT ANY WARRANTY; without even the implied warranty of
MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
GNU General Public License for more details.
You should have received a copy of the GNU General Public License
along with this program. If not, see <https://www.gnu.org/licenses/>.
"""
2026-05-25 18:26:40 -07:00
from __future__ import annotations
2024-08-08 03:27:37 -04:00
2023-04-05 23:41:23 -04:00
import psutil
2024-03-10 11:37:08 -04:00
import logging
2023-04-05 23:41:23 -04:00
from enum import Enum
2026-02-09 13:16:08 -08:00
from comfy . cli_args import args , PerformanceFeature
2026-02-02 14:35:20 -08:00
import threading
2023-06-02 15:05:25 -04:00
import torch
2023-08-17 01:06:34 -04:00
import sys
2024-05-21 16:56:33 -04:00
import platform
2024-12-02 14:39:34 -05:00
import weakref
import gc
2025-12-19 15:01:50 -07:00
import os
2026-05-25 18:26:40 -07:00
from contextlib import contextmanager , nullcontext
2026-01-31 22:01:11 -08:00
import comfy . memory_management
2026-08-27 18:57:31 +04:00
import comfy . system_memory
2026-01-31 22:01:11 -08:00
import comfy . utils
import comfy . quant_ops
2026-05-21 10:03:58 +10:00
import comfy_aimdo . host_buffer
2026-05-03 09:23:24 +10:00
import comfy_aimdo . vram_buffer
2026-08-02 22:20:27 +02:00
from comfy . internal_logging import detail
2026-01-31 22:01:11 -08:00
2026-05-25 18:26:40 -07:00
from typing import TYPE_CHECKING
if TYPE_CHECKING :
from comfy . model_patcher import ModelPatcher
2023-04-05 23:41:23 -04:00
class VRAMState ( Enum ) :
2023-06-04 17:51:04 -04:00
DISABLED = 0 #No vram present: no need to move models to vram
NO_VRAM = 1 #Very low vram: enable all the options to save vram
2023-04-05 23:41:23 -04:00
LOW_VRAM = 2
NORMAL_VRAM = 3
HIGH_VRAM = 4
2023-06-04 17:51:04 -04:00
SHARED = 5 #No dedicated vram: memory shared between CPU and GPU but models still need to be moved between both.
2023-06-03 11:05:37 -04:00
class CPUState ( Enum ) :
GPU = 0
CPU = 1
MPS = 2
2023-02-08 11:37:10 -05:00
2023-04-05 23:41:23 -04:00
# Determine VRAM State
vram_state = VRAMState . NORMAL_VRAM
set_vram_to = VRAMState . NORMAL_VRAM
2023-06-03 11:05:37 -04:00
cpu_state = CPUState . GPU
2023-02-08 11:37:10 -05:00
2023-02-08 14:05:31 -05:00
total_vram = 0
2023-02-08 11:42:37 -05:00
2026-02-11 10:45:19 +08:00
# Training Related State
in_training = False
2026-03-25 08:39:04 +08:00
training_fp8_bwd = False
2026-02-11 10:45:19 +08:00
2025-03-25 05:23:49 -04:00
def get_supported_float8_types ( ) :
float8_types = [ ]
try :
float8_types . append ( torch . float8_e4m3fn )
except :
pass
try :
float8_types . append ( torch . float8_e4m3fnuz )
except :
pass
try :
float8_types . append ( torch . float8_e5m2 )
except :
pass
try :
float8_types . append ( torch . float8_e5m2fnuz )
except :
pass
try :
float8_types . append ( torch . float8_e8m0fnu )
except :
pass
return float8_types
FLOAT8_TYPES = get_supported_float8_types ( )
2024-08-23 04:04:55 -04:00
xpu_available = False
2024-08-30 12:48:42 -04:00
torch_version = " "
2024-08-23 04:04:55 -04:00
try :
torch_version = torch . version . __version__
2025-02-17 04:36:45 -05:00
temp = torch_version . split ( " . " )
torch_version_numeric = ( int ( temp [ 0 ] ) , int ( temp [ 1 ] ) )
2024-08-23 04:04:55 -04:00
except :
pass
2024-08-23 00:59:57 -07:00
2023-05-30 12:36:41 -04:00
lowvram_available = True
2023-12-17 16:59:21 -05:00
if args . deterministic :
2024-03-11 13:54:56 -04:00
logging . info ( " Using deterministic algorithms for pytorch " )
2023-12-17 16:59:21 -05:00
torch . use_deterministic_algorithms ( True , warn_only = True )
2023-04-28 14:28:57 -04:00
directml_enabled = False
2023-04-28 16:51:35 -04:00
if args . directml is not None :
2025-10-25 17:05:22 -07:00
logging . warning ( " WARNING: torch-directml barely works, is very slow, has not been updated in over 1 year and might be removed soon, please don ' t use it, there are better options. " )
2023-04-28 14:28:57 -04:00
import torch_directml
directml_enabled = True
2023-04-28 16:51:35 -04:00
device_index = args . directml
if device_index < 0 :
directml_device = torch_directml . device ( )
else :
directml_device = torch_directml . device ( device_index )
2024-03-11 13:54:56 -04:00
logging . info ( " Using directml with device: {} " . format ( torch_directml . device_name ( device_index ) ) )
2023-04-28 14:28:57 -04:00
# torch_directml.disable_tiled_resources(True)
2023-05-30 12:36:41 -04:00
lowvram_available = False #TODO: need to find a way to get free memory in directml before this can be enabled by default.
2023-04-28 14:28:57 -04:00
2025-08-13 16:13:35 -07:00
try :
2024-08-23 00:59:57 -07:00
_ = torch . xpu . device_count ( )
2025-08-13 16:13:35 -07:00
xpu_available = torch . xpu . is_available ( )
2023-02-08 14:05:31 -05:00
except :
2025-08-13 16:13:35 -07:00
xpu_available = False
2023-02-08 14:05:31 -05:00
2023-06-03 11:05:37 -04:00
try :
if torch . backends . mps . is_available ( ) :
cpu_state = CPUState . MPS
2023-07-12 10:06:34 +08:00
import torch . mps
2023-06-03 11:05:37 -04:00
except :
pass
2024-12-27 08:36:50 +08:00
try :
2024-12-26 20:05:54 -05:00
import torch_npu # noqa: F401
2024-12-27 08:36:50 +08:00
_ = torch . npu . device_count ( )
npu_available = torch . npu . is_available ( )
except :
npu_available = False
2025-02-27 09:45:13 +08:00
try :
import torch_mlu # noqa: F401
_ = torch . mlu . device_count ( )
mlu_available = torch . mlu . is_available ( )
except :
mlu_available = False
2025-07-25 01:57:36 +08:00
try :
ixuca_available = hasattr ( torch , " corex " )
except :
ixuca_available = False
2023-06-03 11:05:37 -04:00
if args . cpu :
cpu_state = CPUState . CPU
2023-09-02 18:22:10 -07:00
def is_intel_xpu ( ) :
global cpu_state
2023-06-02 15:05:25 -04:00
global xpu_available
2023-09-02 18:22:10 -07:00
if cpu_state == CPUState . GPU :
if xpu_available :
return True
return False
2024-12-27 08:36:50 +08:00
def is_ascend_npu ( ) :
global npu_available
if npu_available :
return True
return False
2025-02-27 09:45:13 +08:00
def is_mlu ( ) :
global mlu_available
if mlu_available :
return True
return False
2025-07-25 01:57:36 +08:00
def is_ixuca ( ) :
global ixuca_available
if ixuca_available :
return True
return False
2026-02-28 19:23:28 -08:00
def is_wsl ( ) :
version = platform . uname ( ) . release
if version . endswith ( " -Microsoft " ) :
return True
elif version . endswith ( " microsoft-standard-WSL2 " ) :
return True
return False
2023-09-02 18:22:10 -07:00
def get_torch_device ( ) :
2023-06-02 15:05:25 -04:00
global directml_enabled
2023-06-03 11:05:37 -04:00
global cpu_state
2023-06-02 15:05:25 -04:00
if directml_enabled :
global directml_device
return directml_device
2023-06-03 11:05:37 -04:00
if cpu_state == CPUState . MPS :
2023-06-02 15:05:25 -04:00
return torch . device ( " mps " )
2023-06-03 11:05:37 -04:00
if cpu_state == CPUState . CPU :
2023-06-02 15:05:25 -04:00
return torch . device ( " cpu " )
else :
2023-09-02 18:22:10 -07:00
if is_intel_xpu ( ) :
2024-05-02 00:26:50 -07:00
return torch . device ( " xpu " , torch . xpu . current_device ( ) )
2024-12-27 08:36:50 +08:00
elif is_ascend_npu ( ) :
return torch . device ( " npu " , torch . npu . current_device ( ) )
2025-02-27 09:45:13 +08:00
elif is_mlu ( ) :
return torch . device ( " mlu " , torch . mlu . current_device ( ) )
2023-06-02 15:05:25 -04:00
else :
return torch . device ( torch . cuda . current_device ( ) )
2026-05-25 18:26:40 -07:00
def get_all_torch_devices ( exclude_current = False ) :
global cpu_state
devices = [ ]
if cpu_state == CPUState . GPU :
# NVIDIA + AMD/ROCm both expose their GPUs through torch.cuda.*;
# without the AMD arm, single-GPU ROCm users get an empty list
# which silently turns unload_all_models() into a no-op.
if is_nvidia ( ) or is_amd ( ) :
for i in range ( torch . cuda . device_count ( ) ) :
devices . append ( torch . device ( " cuda " , i ) )
elif is_intel_xpu ( ) :
for i in range ( torch . xpu . device_count ( ) ) :
devices . append ( torch . device ( " xpu " , i ) )
elif is_ascend_npu ( ) :
for i in range ( torch . npu . device_count ( ) ) :
devices . append ( torch . device ( " npu " , i ) )
elif is_mlu ( ) :
for i in range ( torch . mlu . device_count ( ) ) :
devices . append ( torch . device ( " mlu " , i ) )
else :
# Fallback for unhandled GPU backends (e.g. DirectML): at least
# report the current device so callers like unload_all_models()
# do not silently no-op.
devices . append ( get_torch_device ( ) )
else :
devices . append ( get_torch_device ( ) )
if exclude_current :
current = get_torch_device ( )
if current in devices :
devices . remove ( current )
return devices
def get_gpu_device_options ( ) :
""" Return list of device option strings for node widgets.
Always includes " default " and " cpu " . When multiple GPUs are present,
adds " gpu:0 " , " gpu:1 " , etc. (vendor-agnostic labels).
"""
options = [ " default " , " cpu " ]
devices = get_all_torch_devices ( )
if len ( devices ) > 1 :
for i in range ( len ( devices ) ) :
options . append ( f " gpu: { i } " )
return options
def get_gpu_device_options_no_cpu ( ) :
""" Variant of get_gpu_device_options that omits " cpu " .
Intended for components like the VAE selector where running on CPU
is impractical and should not be offered as a choice.
"""
return [ o for o in get_gpu_device_options ( ) if o != " cpu " ]
def resolve_gpu_device_option ( option : str ) :
""" Resolve a device option string to a torch.device.
Returns None for " default " (let the caller use its normal default).
Returns torch.device( " cpu " ) for " cpu " .
For " gpu:N " , returns the Nth torch device. Returns None if the
index is out of range, the option string is malformed, or
unrecognized (callers are expected to log their own context-rich
message before falling back to the default device).
"""
if option is None or option == " default " :
return None
if option == " cpu " :
return torch . device ( " cpu " )
if option . startswith ( " gpu: " ) :
try :
idx = int ( option [ 4 : ] )
except ValueError :
return None
devices = get_all_torch_devices ( )
if 0 < = idx < len ( devices ) :
return devices [ idx ]
return None
@contextmanager
def cuda_device_context ( device ) :
""" Context manager that sets torch.cuda.current_device to match *device*.
Used when running operations on a non-default CUDA device so that custom
CUDA kernels (e.g. comfy_kitchen fp8 quantization) pick up the correct
device index. The previous device is restored on exit.
No-op when *device* is not CUDA, has no explicit index, or already matches
the current device.
"""
prev = None
if device . type == " cuda " and device . index is not None :
prev = torch . cuda . current_device ( )
if prev != device . index :
torch . cuda . set_device ( device )
else :
prev = None
try :
yield
finally :
if prev is not None :
torch . cuda . set_device ( prev )
2023-06-02 15:05:25 -04:00
def get_total_memory ( dev = None , torch_total_too = False ) :
global directml_enabled
if dev is None :
dev = get_torch_device ( )
if hasattr ( dev , ' type ' ) and ( dev . type == ' cpu ' or dev . type == ' mps ' ) :
2026-08-27 18:57:31 +04:00
mem_total = comfy . system_memory . virtual_memory_total ( )
2023-06-02 15:05:25 -04:00
mem_total_torch = mem_total
else :
if directml_enabled :
mem_total = 1024 * 1024 * 1024 #TODO
mem_total_torch = mem_total
2023-09-02 18:22:10 -07:00
elif is_intel_xpu ( ) :
2023-08-17 03:12:17 -07:00
stats = torch . xpu . memory_stats ( dev )
mem_reserved = stats [ ' reserved_bytes.all.current ' ]
2025-07-23 15:10:59 -07:00
mem_total_xpu = torch . xpu . get_device_properties ( dev ) . total_memory
2023-08-17 03:12:17 -07:00
mem_total_torch = mem_reserved
2025-07-22 12:20:09 -07:00
mem_total = mem_total_xpu
2024-12-27 08:36:50 +08:00
elif is_ascend_npu ( ) :
stats = torch . npu . memory_stats ( dev )
mem_reserved = stats [ ' reserved_bytes.all.current ' ]
_ , mem_total_npu = torch . npu . mem_get_info ( dev )
mem_total_torch = mem_reserved
mem_total = mem_total_npu
2025-02-27 09:45:13 +08:00
elif is_mlu ( ) :
stats = torch . mlu . memory_stats ( dev )
mem_reserved = stats [ ' reserved_bytes.all.current ' ]
_ , mem_total_mlu = torch . mlu . mem_get_info ( dev )
mem_total_torch = mem_reserved
mem_total = mem_total_mlu
2023-06-02 15:05:25 -04:00
else :
stats = torch . cuda . memory_stats ( dev )
mem_reserved = stats [ ' reserved_bytes.all.current ' ]
_ , mem_total_cuda = torch . cuda . mem_get_info ( dev )
mem_total_torch = mem_reserved
mem_total = mem_total_cuda
if torch_total_too :
return ( mem_total , mem_total_torch )
else :
return mem_total
2025-03-13 10:05:15 -04:00
def mac_version ( ) :
try :
return tuple ( int ( n ) for n in platform . mac_ver ( ) [ 0 ] . split ( " . " ) )
except :
return None
2023-06-02 15:05:25 -04:00
total_vram = get_total_memory ( get_torch_device ( ) ) / ( 1024 * 1024 )
2026-08-27 18:57:31 +04:00
total_ram = comfy . system_memory . virtual_memory_total ( ) / ( 1024 * 1024 )
2024-03-11 13:54:56 -04:00
logging . info ( " Total VRAM {:0.0f} MB, total RAM {:0.0f} MB " . format ( total_vram , total_ram ) )
2026-08-27 18:57:31 +04:00
cgroup_ram_limit = comfy . system_memory . cgroup_memory_limit ( )
if cgroup_ram_limit is not None :
logging . info ( " RAM limited by cgroup to {:0.0f} MB (host has {:0.0f} MB) " . format ( cgroup_ram_limit / ( 1024 * 1024 ) , psutil . virtual_memory ( ) . total / ( 1024 * 1024 ) ) )
2023-06-02 15:05:25 -04:00
2024-05-20 06:22:29 -04:00
try :
2024-10-09 19:43:17 -04:00
logging . info ( " pytorch version: {} " . format ( torch_version ) )
2025-03-13 10:05:15 -04:00
mac_ver = mac_version ( )
if mac_ver is not None :
2025-03-13 15:03:18 -04:00
logging . info ( " Mac Version {} " . format ( mac_ver ) )
2024-05-20 06:22:29 -04:00
except :
pass
2023-03-22 14:49:00 -04:00
try :
OOM_EXCEPTION = torch . cuda . OutOfMemoryError
except :
OOM_EXCEPTION = Exception
2026-03-11 19:04:13 +02:00
try :
ACCELERATOR_ERROR = torch . AcceleratorError
except AttributeError :
ACCELERATOR_ERROR = RuntimeError
2026-03-09 21:41:02 -07:00
def is_oom ( e ) :
if isinstance ( e , OOM_EXCEPTION ) :
return True
2026-03-11 19:04:13 +02:00
if isinstance ( e , ACCELERATOR_ERROR ) and ( getattr ( e , ' error_code ' , None ) == 2 or " out of memory " in str ( e ) . lower ( ) ) :
2026-03-09 21:41:02 -07:00
discard_cuda_async_error ( )
return True
return False
def raise_non_oom ( e ) :
if not is_oom ( e ) :
raise e
2023-04-09 01:31:47 -04:00
XFORMERS_VERSION = " "
XFORMERS_ENABLED_VAE = True
2023-04-05 23:41:23 -04:00
if args . disable_xformers :
XFORMERS_IS_AVAILABLE = False
2023-03-13 11:36:48 -04:00
else :
try :
import xformers
import xformers . ops
2023-04-05 23:41:23 -04:00
XFORMERS_IS_AVAILABLE = True
2023-11-13 12:27:44 -05:00
try :
XFORMERS_IS_AVAILABLE = xformers . _has_cpp_library
except :
pass
2023-04-09 01:31:47 -04:00
try :
XFORMERS_VERSION = xformers . version . __version__
2024-03-11 13:54:56 -04:00
logging . info ( " xformers version: {} " . format ( XFORMERS_VERSION ) )
2023-04-09 01:31:47 -04:00
if XFORMERS_VERSION . startswith ( " 0.0.18 " ) :
2024-03-10 11:37:08 -04:00
logging . warning ( " \n WARNING: This version of xformers has a major bug where you will get black images when generating high resolution images. " )
logging . warning ( " Please downgrade or upgrade xformers to a different version. \n " )
2023-04-09 01:31:47 -04:00
XFORMERS_ENABLED_VAE = False
except :
pass
2023-03-13 11:36:48 -04:00
except :
2023-04-05 23:41:23 -04:00
XFORMERS_IS_AVAILABLE = False
2023-03-13 11:36:48 -04:00
2023-06-26 12:55:07 -04:00
def is_nvidia ( ) :
global cpu_state
if cpu_state == CPUState . GPU :
if torch . version . cuda :
return True
2023-09-02 18:22:10 -07:00
return False
2023-06-26 12:55:07 -04:00
2024-12-25 04:50:34 -05:00
def is_amd ( ) :
global cpu_state
if cpu_state == CPUState . GPU :
if torch . version . hip :
return True
return False
2024-12-22 03:06:37 -05:00
2025-09-06 20:25:22 -07:00
def amd_min_version ( device = None , min_rdna_version = 0 ) :
if not is_amd ( ) :
return False
2025-09-07 18:16:29 -07:00
if is_device_cpu ( device ) :
return False
2025-09-06 20:25:22 -07:00
arch = torch . cuda . get_device_properties ( device ) . gcnArchName
if arch . startswith ( ' gfx ' ) and len ( arch ) == 7 :
try :
cmp_rdna_version = int ( arch [ 4 ] ) + 2
except :
cmp_rdna_version = 0
if cmp_rdna_version > = min_rdna_version :
return True
return False
2024-12-22 03:06:37 -05:00
MIN_WEIGHT_MEMORY_RATIO = 0.4
if is_nvidia ( ) :
2025-02-21 06:32:11 -05:00
MIN_WEIGHT_MEMORY_RATIO = 0.0
2024-12-22 03:06:37 -05:00
2023-10-11 21:29:03 -04:00
ENABLE_PYTORCH_ATTENTION = False
if args . use_pytorch_cross_attention :
ENABLE_PYTORCH_ATTENTION = True
XFORMERS_IS_AVAILABLE = False
2023-08-27 23:06:19 -04:00
try :
if is_nvidia ( ) :
2025-02-17 04:36:45 -05:00
if torch_version_numeric [ 0 ] > = 2 :
2023-10-11 21:29:03 -04:00
if ENABLE_PYTORCH_ATTENTION == False and args . use_split_cross_attention == False and args . use_quad_cross_attention == False :
2023-06-26 12:55:07 -04:00
ENABLE_PYTORCH_ATTENTION = True
2025-07-25 01:57:36 +08:00
if is_intel_xpu ( ) or is_ascend_npu ( ) or is_mlu ( ) or is_ixuca ( ) :
2023-09-17 04:09:19 -04:00
if args . use_split_cross_attention == False and args . use_quad_cross_attention == False :
ENABLE_PYTORCH_ATTENTION = True
2023-08-27 23:06:19 -04:00
except :
pass
2025-02-14 04:17:56 -05:00
2025-06-08 11:15:34 -07:00
SUPPORT_FP8_OPS = args . supports_fp8_compute
2025-10-21 16:15:23 -07:00
2026-07-20 20:36:03 -07:00
AMD_RDNA2_AND_OLDER_ARCH = [ " gfx1030 " , " gfx1031 " , " gfx1035 " , " gfx1010 " , " gfx1011 " , " gfx1012 " , " gfx906 " , " gfx900 " , " gfx803 " ]
2025-12-19 15:01:50 -07:00
AMD_ENABLE_MIOPEN_ENV = ' COMFYUI_ENABLE_MIOPEN '
2025-10-21 16:15:23 -07:00
2025-02-14 04:17:56 -05:00
try :
if is_amd ( ) :
2026-02-26 17:16:12 -08:00
arch = torch . cuda . get_device_properties ( get_torch_device ( ) ) . gcnArchName . split ( ' : ' ) [ 0 ]
2025-10-21 16:15:23 -07:00
if not ( any ( ( a in arch ) for a in AMD_RDNA2_AND_OLDER_ARCH ) ) :
2025-12-19 15:01:50 -07:00
if os . getenv ( AMD_ENABLE_MIOPEN_ENV ) != ' 1 ' :
torch . backends . cudnn . enabled = False # Seems to improve things a lot on AMD
logging . info ( " Set: torch.backends.cudnn.enabled = False for better AMD performance. " )
2025-10-21 16:15:23 -07:00
2025-05-30 12:41:02 -07:00
try :
rocm_version = tuple ( map ( int , str ( torch . version . hip ) . split ( " . " ) [ : 2 ] ) )
except :
rocm_version = ( 6 , - 1 )
2025-10-21 16:15:23 -07:00
2026-08-12 21:40:11 -04:00
def aotriton_supported ( ) :
""" Whether pytorch reports flash attention as usable on this gpu.
can_use_flash_attention() evaluates runtime eligibility for the given
parameters; on a ROCm build that includes checking the gpu arch against the
2026-09-08 00:54:46 +03:00
arches AOTriton was built for. Querying it avoids assuming where the kernel
images live inside the torch install. The probe tensor is shaped and
2026-08-12 21:40:11 -04:00
typed to pass the unrelated SDPA checks, so False means no hardware support
rather than a rejected shape.
2026-09-08 00:54:46 +03:00
It answers True on a supported arch whose kernel image was never shipped,
and that only fails at launch, without raising. So run one attention
through the flash backend and force the pending error check.
2026-08-12 21:40:11 -04:00
"""
try :
2026-09-08 00:54:46 +03:00
device = get_torch_device ( )
2026-08-12 21:40:11 -04:00
if not torch . backends . cuda . is_flash_attention_available ( ) : # not built with flash attention
return False
2026-09-08 00:54:46 +03:00
q = torch . zeros ( ( 1 , 1 , 8 , 64 ) , dtype = torch . float16 , device = device )
2026-08-12 21:40:11 -04:00
params = torch . backends . cuda . SDPAParams ( q , q , q , None , 0.0 , False , False )
2026-09-08 00:54:46 +03:00
if not torch . backends . cuda . can_use_flash_attention ( params , False ) :
return False
from torch . nn . attention import SDPBackend , sdpa_kernel
with sdpa_kernel ( SDPBackend . FLASH_ATTENTION ) :
torch . nn . functional . scaled_dot_product_attention ( q , q , q )
torch . cuda . synchronize ( )
torch . zeros ( 1 , device = device ) . add_ ( 1 ) . item ( ) # raises if the launch above failed
return True
except Exception as e :
logging . warning ( " Could not run flash attention, disabling it: {} " . format ( e ) )
2026-08-12 21:40:11 -04:00
return False
2026-01-08 14:16:58 -08:00
2025-02-14 04:17:56 -05:00
logging . info ( " AMD arch: {} " . format ( arch ) )
2025-05-30 12:41:02 -07:00
logging . info ( " ROCm version: {} " . format ( rocm_version ) )
2025-02-14 04:17:56 -05:00
if args . use_split_cross_attention == False and args . use_quad_cross_attention == False :
2026-08-12 21:40:11 -04:00
if aotriton_supported ( ) : # AMD efficient attention implementation depends on aotriton.
2025-09-06 21:29:38 -07:00
if torch_version_numeric > = ( 2 , 7 ) : # works on 2.6 but doesn't actually seem to improve much
2026-09-02 15:13:17 -07:00
if any ( ( a in arch ) for a in [ " gfx90a " , " gfx942 " , " gfx950 " , " gfx1100 " , " gfx1101 " , " gfx1150 " , " gfx1151 " , " gfx1170 " , " gfx1171 " ] ) : # TODO: more arches, TODO: gfx950
2025-09-06 21:29:38 -07:00
ENABLE_PYTORCH_ATTENTION = True
2025-10-13 18:19:03 -07:00
if rocm_version > = ( 7 , 0 ) :
2026-08-12 21:40:11 -04:00
if any ( ( a in arch ) for a in [ " gfx1200 " , " gfx1201 " ] ) :
ENABLE_PYTORCH_ATTENTION = True
2025-06-08 11:15:34 -07:00
if torch_version_numeric > = ( 2 , 7 ) and rocm_version > = ( 6 , 4 ) :
2026-09-02 15:13:17 -07:00
if any ( ( a in arch ) for a in [ " gfx1200 " , " gfx1201 " , " gfx950 " , " gfx1170 " , " gfx1171 " ] ) : # TODO: more arches, "gfx942" gives error on pytorch nightly 2.10 1013 rocm7.0
2025-06-08 11:15:34 -07:00
SUPPORT_FP8_OPS = True
2025-02-14 04:17:56 -05:00
except :
pass
2023-04-05 23:41:23 -04:00
if ENABLE_PYTORCH_ATTENTION :
2023-03-13 12:25:19 -04:00
torch . backends . cuda . enable_math_sdp ( True )
torch . backends . cuda . enable_flash_sdp ( True )
torch . backends . cuda . enable_mem_efficient_sdp ( True )
2023-03-12 15:44:16 -04:00
2025-02-23 04:45:54 -05:00
PRIORITIZE_FP16 = False # TODO: remove and replace with something that shows exactly which dtype is faster than the other
2025-02-08 17:00:56 -05:00
try :
2025-08-09 09:49:25 -07:00
if ( is_nvidia ( ) or is_amd ( ) ) and PerformanceFeature . Fp16Accumulation in args . fast :
2025-02-08 17:00:56 -05:00
torch . backends . cuda . matmul . allow_fp16_accumulation = True
2025-02-23 04:45:54 -05:00
PRIORITIZE_FP16 = True # TODO: limit to cards where it actually boosts performance
2025-02-28 02:48:20 -05:00
logging . info ( " Enabled fp16 accumulation. " )
2025-02-08 17:00:56 -05:00
except :
pass
2026-06-11 03:54:32 +10:00
def set_cudnn_benchmark ( ) :
if torch . cuda . is_available ( ) and torch . backends . cudnn . is_available ( ) :
torch . backends . cudnn . benchmark = PerformanceFeature . AutoTune in args . fast
2025-10-18 19:35:46 -07:00
2024-12-23 03:22:48 -05:00
try :
2025-06-07 07:01:15 -07:00
if torch_version_numeric > = ( 2 , 5 ) :
2024-12-23 03:22:48 -05:00
torch . backends . cuda . allow_fp16_bf16_reduction_math_sdp ( True )
except :
logging . warning ( " Warning, could not set allow_fp16_bf16_reduction_math_sdp " )
2024-12-23 00:18:32 -08:00
2023-04-05 23:41:23 -04:00
if args . lowvram :
set_vram_to = VRAMState . LOW_VRAM
2023-05-30 12:36:41 -04:00
lowvram_available = True
2023-04-05 23:41:23 -04:00
elif args . novram :
set_vram_to = VRAMState . NO_VRAM
2023-06-15 15:21:37 -04:00
elif args . highvram or args . gpu_only :
2023-04-05 23:41:23 -04:00
vram_state = VRAMState . HIGH_VRAM
2023-03-24 14:30:43 -04:00
2023-04-07 00:27:54 -04:00
FORCE_FP32 = False
if args . force_fp32 :
2024-03-11 13:54:56 -04:00
logging . info ( " Forcing FP32, if this improves things please report it. " )
2023-04-07 00:27:54 -04:00
FORCE_FP32 = True
2023-05-30 12:36:41 -04:00
if lowvram_available :
2023-12-22 14:24:04 -05:00
if set_vram_to in ( VRAMState . LOW_VRAM , VRAMState . NO_VRAM ) :
vram_state = set_vram_to
2023-02-08 14:05:31 -05:00
2023-02-08 11:37:10 -05:00
2023-06-03 11:05:37 -04:00
if cpu_state != CPUState . GPU :
vram_state = VRAMState . DISABLED
2023-03-24 14:30:43 -04:00
2023-06-03 11:05:37 -04:00
if cpu_state == CPUState . MPS :
vram_state = VRAMState . SHARED
2023-02-08 11:37:10 -05:00
2024-03-11 13:54:56 -04:00
logging . info ( f " Set vram state to: { vram_state . name } " )
2023-02-08 11:37:10 -05:00
2023-08-17 03:12:37 -04:00
DISABLE_SMART_MEMORY = args . disable_smart_memory
if DISABLE_SMART_MEMORY :
2024-03-11 13:54:56 -04:00
logging . info ( " Disabling smart memory management " )
2023-06-03 11:05:37 -04:00
2023-05-13 17:11:27 -04:00
def get_torch_device_name ( device ) :
if hasattr ( device , ' type ' ) :
2023-06-02 15:05:25 -04:00
if device . type == " cuda " :
2023-07-17 15:18:58 -04:00
try :
allocator_backend = torch . cuda . get_allocator_backend ( )
except :
allocator_backend = " "
return " {} {} : {} " . format ( device , torch . cuda . get_device_name ( device ) , allocator_backend )
2025-07-24 12:06:25 -07:00
elif device . type == " xpu " :
return " {} {} " . format ( device , torch . xpu . get_device_name ( device ) )
2023-06-02 15:05:25 -04:00
else :
return " {} " . format ( device . type )
2023-09-02 18:22:10 -07:00
elif is_intel_xpu ( ) :
2023-08-17 03:12:17 -07:00
return " {} {} " . format ( device , torch . xpu . get_device_name ( device ) )
2024-12-27 08:36:50 +08:00
elif is_ascend_npu ( ) :
return " {} {} " . format ( device , torch . npu . get_device_name ( device ) )
2025-02-27 09:45:13 +08:00
elif is_mlu ( ) :
return " {} {} " . format ( device , torch . mlu . get_device_name ( device ) )
2023-06-02 15:05:25 -04:00
else :
return " CUDA {} : {} " . format ( device , torch . cuda . get_device_name ( device ) )
2023-05-13 17:11:27 -04:00
try :
2024-03-11 13:54:56 -04:00
logging . info ( " Device: {} " . format ( get_torch_device_name ( get_torch_device ( ) ) ) )
2023-05-13 17:11:27 -04:00
except :
2024-03-10 11:37:08 -04:00
logging . warning ( " Could not pick default device. " )
2026-05-25 18:26:40 -07:00
try :
for device in get_all_torch_devices ( exclude_current = True ) :
logging . info ( " Device: {} " . format ( get_torch_device_name ( device ) ) )
except :
pass
2023-05-13 17:11:27 -04:00
2026-05-25 18:26:40 -07:00
current_loaded_models : list [ LoadedModel ] = [ ]
2023-02-08 03:17:54 -05:00
2026-05-21 10:03:58 +10:00
DIRTY_MMAPS = set ( )
PIN_PRESSURE_HYSTERESIS = 256 * 1024 * 1024
#Freeing registerables on pressure does imply a GPU sync, so go big on
#the hysteresis so each expensive sync gives us back a good chunk.
REGISTERABLE_PIN_HYSTERESIS = 2048 * 1024 * 1024
2026-07-09 16:39:01 -07:00
WINDOWS_PIN_EVICTION_SWAP_PERCENT = 5.0
WINDOWS_PIN_EVICTION_EMERGENCY_AVAILABLE = 512 * 1024 * * 2
2026-05-21 10:03:58 +10:00
2023-12-28 21:41:10 -05:00
def module_size ( module ) :
module_mem = 0
sd = module . state_dict ( )
for k in sd :
t = sd [ k ]
2026-01-05 00:48:31 -08:00
module_mem + = t . nbytes
2023-12-28 21:41:10 -05:00
return module_mem
2026-05-21 10:03:58 +10:00
def mark_mmap_dirty ( storage ) :
mmap_refs = getattr ( storage , " _comfy_tensor_mmap_refs " , None )
if mmap_refs is not None :
DIRTY_MMAPS . add ( mmap_refs [ 0 ] )
2026-07-29 07:05:57 +10:00
PIN_SUBSETS = [ " weights " , " patches " ]
LOADED_PIN_SUBSETS = [ " weights-loaded " , " patches-loaded " ]
def models_for_pin_eviction ( active , current_prompt = None ) :
for loaded_model in current_loaded_models :
model = loaded_model . model
if model is None or not model . is_dynamic ( ) :
continue
pin_state = model . model . dynamic_pins [ model . load_device ]
if ( ( active is None or pin_state [ " active " ] == active ) and
( current_prompt is None or pin_state [ " current_prompt " ] == current_prompt ) ) :
yield model
def free_model_pins ( size , subsets , current_prompt , active , registrations = False ) :
2026-05-21 10:03:58 +10:00
freed_total = 0
2026-07-29 07:05:57 +10:00
for model in models_for_pin_eviction ( active , current_prompt = current_prompt ) :
2026-05-21 10:03:58 +10:00
if size < = 0 :
return freed_total
2026-07-29 07:05:57 +10:00
if registrations :
freed = model . unregister_inactive_pins ( size , subsets = subsets )
else :
freed = model . partially_unload_ram ( size , subsets = subsets )
freed_total + = freed
size - = freed
2026-05-21 10:03:58 +10:00
return freed_total
2026-07-29 07:05:57 +10:00
def pin_eviction_tiers ( loaded , evict_active ) :
tiers = [
( PIN_SUBSETS , False , None ) ,
( LOADED_PIN_SUBSETS , False , None ) ,
( LOADED_PIN_SUBSETS , True , None ) ,
]
if not loaded :
tiers . append ( ( PIN_SUBSETS , True , False ) )
if evict_active :
tiers . append ( ( PIN_SUBSETS , True , True ) )
return tiers
2026-08-03 01:16:06 +10:00
def registration_eviction_tiers ( evict_active ) :
subsets = PIN_SUBSETS + LOADED_PIN_SUBSETS
tiers = [
( subsets , False , False ) ,
( subsets , True , False ) ,
]
if evict_active :
tiers . extend ( [
( subsets , False , True ) ,
( subsets , True , True ) ,
] )
return tiers
2026-07-29 07:05:57 +10:00
def free_pins ( size , evict_active = False , loaded = False ) :
freed = 0
for subsets , current_prompt , active in pin_eviction_tiers ( loaded , evict_active ) :
freed + = free_model_pins ( size - freed , subsets , current_prompt , active )
return freed
2026-07-09 16:39:01 -07:00
def should_free_pins_for_ram_pressure ( shortfall ) :
if shortfall < = 0 :
return False
if not WINDOWS :
return True
2026-08-27 18:57:31 +04:00
if comfy . system_memory . virtual_memory_available ( ) < WINDOWS_PIN_EVICTION_EMERGENCY_AVAILABLE :
2026-07-09 16:39:01 -07:00
return True
2026-08-01 22:12:12 -07:00
try :
return psutil . swap_memory ( ) . percent > = WINDOWS_PIN_EVICTION_SWAP_PERCENT
except RuntimeError as err :
logging . warning ( " Could not read Windows swap usage; falling back to RAM-pressure pin eviction: %s " , err )
return True
2026-07-09 16:39:01 -07:00
2026-07-29 07:05:57 +10:00
def ensure_pin_budget ( size , evict_active = False , loaded = False ) :
2026-06-13 00:53:33 +10:00
if args . high_ram :
return True
2026-05-31 05:20:04 +10:00
if args . fast_disk :
shortfall = TOTAL_PINNED_MEMORY + size - MAX_PINNED_MEMORY
else :
2026-08-27 18:57:31 +04:00
shortfall = size + max ( comfy . memory_management . RAM_CACHE_HEADROOM / 2 , 2048 * 1024 * * 2 ) - comfy . system_memory . virtual_memory_available ( )
2026-05-21 10:03:58 +10:00
if shortfall < = 0 :
return True
to_free = shortfall + PIN_PRESSURE_HYSTERESIS
2026-07-29 07:05:57 +10:00
return free_pins ( to_free , evict_active = evict_active , loaded = loaded ) > = shortfall
2026-05-21 10:03:58 +10:00
2026-08-03 01:16:06 +10:00
def free_registrations ( shortfall , evict_active = True ) :
2026-05-21 10:03:58 +10:00
if MAX_PINNED_MEMORY < = 0 :
return False
if shortfall < = 0 :
return True
shortfall + = REGISTERABLE_PIN_HYSTERESIS
2026-08-03 01:16:06 +10:00
for subsets , current_prompt , active in registration_eviction_tiers ( evict_active ) :
2026-07-29 07:05:57 +10:00
shortfall - = free_model_pins ( shortfall , subsets , current_prompt , active , registrations = True )
2026-05-21 10:03:58 +10:00
return shortfall < = REGISTERABLE_PIN_HYSTERESIS
2026-03-13 19:18:08 -07:00
2026-08-03 01:16:06 +10:00
def ensure_pin_registerable ( size , evict_active = True ) :
return free_registrations ( TOTAL_PINNED_MEMORY + size - MAX_PINNED_MEMORY , evict_active = evict_active )
2026-06-06 01:39:35 +10:00
2023-08-17 01:06:34 -04:00
class LoadedModel :
2026-05-25 18:26:40 -07:00
def __init__ ( self , model : ModelPatcher ) :
2024-12-02 14:39:34 -05:00
self . _set_model ( model )
2023-08-17 01:06:34 -04:00
self . device = model . load_device
2024-03-28 18:01:04 -04:00
self . real_model = None
2024-06-05 19:14:56 -04:00
self . currently_used = True
2024-12-02 14:39:34 -05:00
self . model_finalizer = None
self . _patcher_finalizer = None
2026-05-25 18:26:40 -07:00
def _set_model ( self , model : ModelPatcher ) :
2024-12-02 14:39:34 -05:00
self . _model = weakref . ref ( model )
if model . parent is not None :
self . _parent_model = weakref . ref ( model . parent )
self . _patcher_finalizer = weakref . finalize ( model , self . _switch_parent )
2026-03-16 14:00:42 -06:00
self . _patcher_finalizer . atexit = False
2024-12-02 14:39:34 -05:00
def _switch_parent ( self ) :
model = self . _parent_model ( )
if model is not None :
self . _set_model ( model )
2026-05-25 18:26:40 -07:00
self . device = model . load_device
2024-12-02 14:39:34 -05:00
@property
def model ( self ) :
return self . _model ( )
2023-02-08 11:37:10 -05:00
2023-08-17 01:06:34 -04:00
def model_memory ( self ) :
return self . model . model_size ( )
2023-02-08 11:37:10 -05:00
2024-12-19 16:21:56 -05:00
def model_loaded_memory ( self ) :
return self . model . loaded_size ( )
2024-08-08 03:27:37 -04:00
def model_offloaded_memory ( self ) :
return self . model . model_size ( ) - self . model . loaded_size ( )
2023-08-17 01:06:34 -04:00
def model_memory_required ( self , device ) :
2024-08-06 13:27:48 -04:00
if device == self . model . current_loaded_device ( ) :
2024-08-10 15:29:36 -04:00
return self . model_offloaded_memory ( )
2023-08-17 01:06:34 -04:00
else :
return self . model_memory ( )
2023-02-17 21:14:07 -05:00
2024-05-11 21:46:05 -04:00
def model_load ( self , lowvram_model_memory = 0 , force_patch_weights = False ) :
2023-08-17 01:06:34 -04:00
self . model . model_patches_to ( self . device )
self . model . model_patches_to ( self . model . model_dtype ( ) )
2023-02-17 21:14:07 -05:00
2024-12-02 14:39:34 -05:00
# if self.model.loaded_size() > 0:
use_more_vram = lowvram_model_memory
if use_more_vram == 0 :
use_more_vram = 1e32
2025-11-13 07:19:53 +10:00
self . model_use_more_vram ( use_more_vram , force_patch_weights = force_patch_weights )
2025-11-09 15:51:33 -08:00
2024-12-02 14:39:34 -05:00
real_model = self . model . model
2023-02-08 03:17:54 -05:00
2023-08-17 03:12:17 -07:00
2024-12-02 14:39:34 -05:00
self . real_model = weakref . ref ( real_model )
self . model_finalizer = weakref . finalize ( real_model , cleanup_models )
2026-03-16 14:00:42 -06:00
self . model_finalizer . atexit = False
2024-12-02 14:39:34 -05:00
return real_model
2023-02-08 11:37:10 -05:00
2024-05-12 06:13:45 -04:00
def should_reload_model ( self , force_patch_weights = False ) :
2024-08-09 03:36:40 -04:00
if force_patch_weights and self . model . lowvram_patch_counter ( ) > 0 :
2024-05-12 06:13:45 -04:00
return True
return False
2024-08-08 03:27:37 -04:00
def model_unload ( self , memory_to_free = None , unpatch_weights = True ) :
if memory_to_free is not None :
if memory_to_free < self . model . loaded_size ( ) :
2024-08-13 03:57:55 -04:00
freed = self . model . partially_unload ( self . model . offload_device , memory_to_free )
if freed > = memory_to_free :
return False
2024-12-02 14:39:34 -05:00
self . model . detach ( unpatch_weights )
self . model_finalizer . detach ( )
self . model_finalizer = None
2024-03-28 18:01:04 -04:00
self . real_model = None
2024-08-08 03:27:37 -04:00
return True
2024-12-02 14:39:34 -05:00
def model_use_more_vram ( self , extra_memory , force_patch_weights = False ) :
return self . model . partially_load ( self . device , extra_memory , force_patch_weights = force_patch_weights )
2023-05-30 12:36:41 -04:00
2023-08-17 01:06:34 -04:00
def __eq__ ( self , other ) :
return self . model is other . model
2023-07-15 13:24:05 -04:00
2024-12-02 14:39:34 -05:00
def __del__ ( self ) :
if self . _patcher_finalizer is not None :
self . _patcher_finalizer . detach ( )
2024-12-02 19:49:49 -05:00
def is_dead ( self ) :
return self . real_model ( ) is not None and self . model is None
2024-12-02 14:39:34 -05:00
2024-08-08 03:27:37 -04:00
def use_more_memory ( extra_memory , loaded_models , device ) :
for m in loaded_models :
if m . device == device :
extra_memory - = m . model_use_more_vram ( extra_memory )
if extra_memory < = 0 :
break
def offloaded_memory ( loaded_models , device ) :
offloaded_mem = 0
for m in loaded_models :
if m . device == device :
offloaded_mem + = m . model_offloaded_memory ( )
return offloaded_mem
2024-09-01 17:29:31 -04:00
WINDOWS = any ( platform . win32_ver ( ) )
2024-09-01 01:01:54 -04:00
EXTRA_RESERVED_VRAM = 400 * 1024 * 1024
2024-09-01 17:29:31 -04:00
if WINDOWS :
2024-09-01 01:01:54 -04:00
EXTRA_RESERVED_VRAM = 600 * 1024 * 1024 #Windows is higher because of the shared vram issue
2025-07-29 01:07:45 -07:00
if total_vram > ( 15 * 1024 ) : # more extra reserved vram on 16GB+ cards
EXTRA_RESERVED_VRAM + = 100 * 1024 * 1024
2024-08-19 17:16:18 -04:00
if args . reserve_vram is not None :
EXTRA_RESERVED_VRAM = args . reserve_vram * 1024 * 1024 * 1024
logging . debug ( " Reserving {} MB vram for other applications. " . format ( EXTRA_RESERVED_VRAM / ( 1024 * 1024 ) ) )
def extra_reserved_memory ( ) :
return EXTRA_RESERVED_VRAM
2024-09-01 01:01:54 -04:00
def minimum_inference_memory ( ) :
return ( 1024 * 1024 * 1024 ) * 0.8 + extra_reserved_memory ( )
2026-03-13 19:18:08 -07:00
def free_memory ( memory_required , device , keep_loaded = [ ] , for_dynamic = False , pins_required = 0 , ram_required = 0 ) :
2024-12-02 14:39:34 -05:00
cleanup_models_gc ( )
2026-07-29 07:31:45 +10:00
if not for_dynamic :
detail ( " Non dynamic memory free called! memory_required= %s pins_required= %s ram_required= %s " , memory_required , pins_required , ram_required )
2024-03-24 02:36:30 -04:00
unloaded_model = [ ]
can_unload = [ ]
2024-08-06 03:22:39 -04:00
unloaded_models = [ ]
2024-03-24 02:36:30 -04:00
2023-08-17 01:06:34 -04:00
for i in range ( len ( current_loaded_models ) - 1 , - 1 , - 1 ) :
shift_model = current_loaded_models [ i ]
2026-03-27 18:34:16 -07:00
if device is None or shift_model . device == device :
2024-12-02 19:49:49 -05:00
if shift_model not in keep_loaded and not shift_model . is_dead ( ) :
2024-09-04 19:47:32 -04:00
can_unload . append ( ( - shift_model . model_offloaded_memory ( ) , sys . getrefcount ( shift_model . model ) , shift_model . model_memory ( ) , i ) )
2024-06-05 19:14:56 -04:00
shift_model . currently_used = False
2024-03-24 02:36:30 -04:00
2026-03-13 19:18:08 -07:00
can_unload_sorted = sorted ( can_unload )
for x in can_unload_sorted :
2024-03-24 02:36:30 -04:00
i = x [ - 1 ]
2026-01-31 22:01:11 -08:00
memory_to_free = 1e32
2026-05-31 05:20:33 +10:00
if not DISABLE_SMART_MEMORY or device is None :
2026-03-27 18:34:16 -07:00
memory_to_free = 0 if device is None else memory_required - get_free_memory ( device )
2026-05-31 05:20:33 +10:00
if current_loaded_models [ i ] . model . is_dynamic ( ) and for_dynamic :
2026-03-01 19:18:56 -08:00
#don't actually unload dynamic models for the sake of other dynamic models
#as that works on-demand.
memory_required - = current_loaded_models [ i ] . model . loaded_size ( )
memory_to_free = 0
2026-01-31 22:01:11 -08:00
if memory_to_free > 0 and current_loaded_models [ i ] . model_unload ( memory_to_free ) :
logging . debug ( f " Unloading { current_loaded_models [ i ] . model . model . __class__ . __name__ } " )
2024-08-08 03:27:37 -04:00
unloaded_model . append ( i )
2024-03-24 02:36:30 -04:00
for i in sorted ( unloaded_model , reverse = True ) :
2024-08-06 03:22:39 -04:00
unloaded_models . append ( current_loaded_models . pop ( i ) )
2023-08-17 01:06:34 -04:00
2026-05-31 05:20:33 +10:00
if not for_dynamic and pins_required > 0 :
ensure_pin_budget ( pins_required )
ensure_pin_registerable ( pins_required )
2024-03-24 02:36:30 -04:00
if len ( unloaded_model ) > 0 :
2023-08-17 01:06:34 -04:00
soft_empty_cache ( )
2026-03-27 18:34:16 -07:00
elif device is not None :
2023-10-22 13:53:59 -04:00
if vram_state != VRAMState . HIGH_VRAM :
mem_free_total , mem_free_torch = get_free_memory ( device , torch_free_too = True )
if mem_free_torch > mem_free_total * 0.25 :
soft_empty_cache ( )
2024-08-06 03:22:39 -04:00
return unloaded_models
2023-08-17 01:06:34 -04:00
2026-02-09 13:16:08 -08:00
def load_models_gpu ( models , memory_required = 0 , force_patch_weights = False , minimum_memory_required = None , force_full_load = False ) :
2024-12-02 14:39:34 -05:00
cleanup_models_gc ( )
2023-02-17 15:45:29 -05:00
global vram_state
2023-08-17 01:06:34 -04:00
inference_memory = minimum_inference_memory ( )
2024-08-19 17:16:18 -04:00
extra_mem = max ( inference_memory , memory_required + extra_reserved_memory ( ) )
2024-08-01 16:39:59 -04:00
if minimum_memory_required is None :
minimum_memory_required = extra_mem
else :
2024-08-19 17:16:18 -04:00
minimum_memory_required = max ( inference_memory , minimum_memory_required + extra_reserved_memory ( ) )
2023-08-17 01:06:34 -04:00
2026-05-05 05:58:06 +10:00
# Order-preserving dedup. A plain set() would randomize iteration order across runs
models_temp = { }
2025-08-20 19:26:37 -07:00
for m in models :
2026-05-05 05:58:06 +10:00
models_temp [ m ] = None
2025-08-20 19:26:37 -07:00
for mm in m . model_patches_models ( ) :
2026-05-05 05:58:06 +10:00
models_temp [ mm ] = None
2025-08-20 19:26:37 -07:00
2026-05-05 05:58:06 +10:00
models = list ( models_temp )
models . reverse ( )
2024-04-06 18:38:39 -04:00
2023-08-17 01:06:34 -04:00
models_to_load = [ ]
2024-12-02 14:39:34 -05:00
2026-01-31 22:01:11 -08:00
free_for_dynamic = True
2023-08-17 01:06:34 -04:00
for x in models :
2026-01-31 22:01:11 -08:00
if not x . is_dynamic ( ) :
free_for_dynamic = False
2023-08-17 01:06:34 -04:00
loaded_model = LoadedModel ( x )
2024-05-12 06:13:45 -04:00
try :
loaded_model_index = current_loaded_models . index ( loaded_model )
except :
loaded_model_index = None
if loaded_model_index is not None :
loaded = current_loaded_models [ loaded_model_index ]
2024-12-02 14:39:34 -05:00
loaded . currently_used = True
models_to_load . append ( loaded )
else :
2023-10-11 20:35:50 -04:00
if hasattr ( x , " model " ) :
2024-03-11 13:54:56 -04:00
logging . info ( f " Requested to load { x . model . __class__ . __name__ } " )
2023-08-17 01:06:34 -04:00
models_to_load . append ( loaded_model )
2024-12-02 14:39:34 -05:00
for loaded_model in models_to_load :
to_unload = [ ]
for i in range ( len ( current_loaded_models ) ) :
if loaded_model . model . is_clone ( current_loaded_models [ i ] . model ) :
to_unload = [ i ] + to_unload
for i in to_unload :
2025-09-25 05:35:12 +03:00
model_to_unload = current_loaded_models . pop ( i )
model_to_unload . model . detach ( unpatch_all = False )
model_to_unload . model_finalizer . detach ( )
2023-04-19 09:36:19 -04:00
2023-08-17 01:06:34 -04:00
total_memory_required = { }
2026-05-31 05:20:33 +10:00
total_pins_required = { }
2023-08-17 01:06:34 -04:00
for loaded_model in models_to_load :
2026-03-13 19:18:08 -07:00
device = loaded_model . device
total_memory_required [ device ] = total_memory_required . get ( device , 0 ) + loaded_model . model_memory_required ( device )
2026-05-31 05:20:33 +10:00
if not loaded_model . model . is_dynamic ( ) :
total_pins_required [ device ] = total_pins_required . get ( device , 0 ) + loaded_model . model_memory ( )
2023-02-16 10:38:08 -05:00
2024-12-02 14:39:34 -05:00
for device in total_memory_required :
if device != torch . device ( " cpu " ) :
2026-03-13 19:18:08 -07:00
free_memory ( total_memory_required [ device ] * 1.1 + extra_mem ,
device ,
2026-05-31 05:20:33 +10:00
for_dynamic = free_for_dynamic ,
pins_required = total_pins_required . get ( device , 0 ) )
2024-03-20 01:29:26 -04:00
2024-08-10 15:29:36 -04:00
for device in total_memory_required :
if device != torch . device ( " cpu " ) :
2024-12-02 14:39:34 -05:00
free_mem = get_free_memory ( device )
if free_mem < minimum_memory_required :
2026-01-31 22:01:11 -08:00
models_l = free_memory ( minimum_memory_required , device , for_dynamic = free_for_dynamic )
2024-12-02 14:39:34 -05:00
logging . info ( " {} models unloaded. " . format ( len ( models_l ) ) )
2024-08-10 15:29:36 -04:00
2023-08-17 01:06:34 -04:00
for loaded_model in models_to_load :
model = loaded_model . model
torch_dev = model . load_device
if is_device_cpu ( torch_dev ) :
vram_set_state = VRAMState . DISABLED
else :
vram_set_state = vram_state
lowvram_model_memory = 0
2024-08-12 23:42:21 -04:00
if lowvram_available and ( vram_set_state == VRAMState . LOW_VRAM or vram_set_state == VRAMState . NORMAL_VRAM ) and not force_full_load :
2024-12-19 16:21:56 -05:00
loaded_memory = loaded_model . model_loaded_memory ( )
current_free_mem = get_free_memory ( torch_dev ) + loaded_memory
2024-12-22 03:06:37 -05:00
2025-11-27 16:03:03 +10:00
lowvram_model_memory = max ( 0 , ( current_free_mem - minimum_memory_required ) , min ( current_free_mem * MIN_WEIGHT_MEMORY_RATIO , current_free_mem - minimum_inference_memory ( ) ) )
2025-11-09 15:51:33 -08:00
lowvram_model_memory = lowvram_model_memory - loaded_memory
if lowvram_model_memory == 0 :
lowvram_model_memory = 0.1
2023-02-08 14:05:31 -05:00
2023-08-17 01:06:34 -04:00
if vram_set_state == VRAMState . NO_VRAM :
2024-12-23 01:50:11 -05:00
lowvram_model_memory = 0.1
2023-02-17 15:45:29 -05:00
2026-09-09 07:18:47 -07:00
with comfy . utils . progress_activity ( " loading " ) :
loaded_model . model_load ( lowvram_model_memory , force_patch_weights = force_patch_weights )
2026-07-29 07:31:45 +10:00
vram_used = 0 if is_device_cpu ( torch_dev ) else loaded_model . model_loaded_memory ( )
ram_used = model . loaded_ram_size ( ) if model . is_dynamic ( ) else loaded_model . model_memory ( ) - vram_used
detail ( " Model loaded: patcher= %s model= %s ram_mb= %.1f vram_mb= %.1f " , model . __class__ . __name__ , model . model . __class__ . __name__ , ram_used / ( 1024 * * 2 ) , vram_used / ( 1024 * * 2 ) )
2023-08-17 01:06:34 -04:00
current_loaded_models . insert ( 0 , loaded_model )
return
2026-02-02 16:52:07 -08:00
def load_model_gpu ( model ) :
return load_models_gpu ( [ model ] )
2024-06-05 19:14:56 -04:00
def loaded_models ( only_currently_used = False ) :
output = [ ]
for m in current_loaded_models :
if only_currently_used :
if not m . currently_used :
continue
output . append ( m . model )
return output
2024-12-02 14:39:34 -05:00
def cleanup_models_gc ( ) :
do_gc = False
2026-01-31 22:01:11 -08:00
2024-12-02 14:39:34 -05:00
for i in range ( len ( current_loaded_models ) ) :
cur = current_loaded_models [ i ]
2024-12-02 19:49:49 -05:00
if cur . is_dead ( ) :
2024-12-02 14:39:34 -05:00
logging . info ( " Potential memory leak detected with model {} , doing a full garbage collect, for maximum performance avoid circular references in the model code. " . format ( cur . real_model ( ) . __class__ . __name__ ) )
do_gc = True
break
if do_gc :
gc . collect ( )
soft_empty_cache ( )
for i in range ( len ( current_loaded_models ) ) :
cur = current_loaded_models [ i ]
2024-12-02 19:49:49 -05:00
if cur . is_dead ( ) :
2024-12-02 14:39:34 -05:00
logging . warning ( " WARNING, memory leak with model {} . Please make sure it is not being referenced from somewhere. " . format ( cur . real_model ( ) . __class__ . __name__ ) )
2026-01-31 22:01:11 -08:00
def archive_model_dtypes ( model ) :
for name , module in model . named_modules ( ) :
for param_name , param in module . named_parameters ( recurse = False ) :
setattr ( module , f " { param_name } _comfy_model_dtype " , param . dtype )
2026-03-03 18:19:40 -08:00
for buf_name , buf in module . named_buffers ( recurse = False ) :
setattr ( module , f " { buf_name } _comfy_model_dtype " , buf . dtype )
2026-01-31 22:01:11 -08:00
2024-12-02 14:39:34 -05:00
def cleanup_models ( ) :
2023-08-17 01:06:34 -04:00
to_delete = [ ]
for i in range ( len ( current_loaded_models ) ) :
2024-12-02 14:39:34 -05:00
if current_loaded_models [ i ] . real_model ( ) is None :
to_delete = [ i ] + to_delete
2023-08-17 01:06:34 -04:00
for i in to_delete :
x = current_loaded_models . pop ( i )
del x
2023-02-17 15:45:29 -05:00
2023-08-24 17:20:54 -04:00
def dtype_size ( dtype ) :
dtype_size = 4
if dtype == torch . float16 or dtype == torch . bfloat16 :
dtype_size = 2
2023-12-04 11:52:06 -05:00
elif dtype == torch . float32 :
dtype_size = 4
else :
try :
dtype_size = dtype . itemsize
except : #Old pytorch doesn't have .itemsize
pass
2023-08-24 17:20:54 -04:00
return dtype_size
2023-07-01 13:22:51 -04:00
def unet_offload_device ( ) :
2023-07-03 00:08:30 -04:00
if vram_state == VRAMState . HIGH_VRAM :
2023-07-01 13:22:51 -04:00
return get_torch_device ( )
else :
return torch . device ( " cpu " )
2023-08-17 01:06:34 -04:00
def unet_inital_load_device ( parameters , dtype ) :
2026-03-04 13:33:14 -08:00
cpu_dev = torch . device ( " cpu " )
if comfy . memory_management . aimdo_enabled :
return cpu_dev
2023-08-17 01:06:34 -04:00
torch_dev = get_torch_device ( )
2024-12-12 06:00:31 -05:00
if vram_state == VRAMState . HIGH_VRAM or vram_state == VRAMState . SHARED :
2023-08-17 01:06:34 -04:00
return torch_dev
2025-05-26 13:39:27 -07:00
if DISABLE_SMART_MEMORY or vram_state == VRAMState . NO_VRAM :
2023-08-20 04:00:53 -04:00
return cpu_dev
2023-08-24 17:20:54 -04:00
model_size = dtype_size ( dtype ) * parameters
2023-08-17 01:06:34 -04:00
mem_dev = get_free_memory ( torch_dev )
mem_cpu = get_free_memory ( cpu_dev )
2026-03-04 13:33:14 -08:00
if mem_dev > mem_cpu and model_size < mem_dev :
2023-08-17 01:06:34 -04:00
return torch_dev
else :
return cpu_dev
2024-08-03 13:45:19 -04:00
def maximum_vram_for_weights ( device = None ) :
2024-08-05 16:24:04 -04:00
return ( get_total_memory ( device ) * 0.88 - minimum_inference_memory ( ) )
2024-08-03 13:45:19 -04:00
2025-02-27 16:39:57 -05:00
def unet_dtype ( device = None , model_params = 0 , supported_dtypes = [ torch . float16 , torch . bfloat16 , torch . float32 ] , weight_dtype = None ) :
2024-09-21 04:50:12 -04:00
if model_params < 0 :
model_params = 1000000000000000000000
2024-11-25 05:00:23 -05:00
if args . fp32_unet :
return torch . float32
if args . fp64_unet :
return torch . float64
2023-10-13 14:51:10 -04:00
if args . bf16_unet :
return torch . bfloat16
2023-12-11 18:36:29 -05:00
if args . fp16_unet :
return torch . float16
2023-12-04 11:10:00 -05:00
if args . fp8_e4m3fn_unet :
return torch . float8_e4m3fn
if args . fp8_e5m2_unet :
return torch . float8_e5m2
2025-04-22 03:17:38 -07:00
if args . fp8_e8m0fnu_unet :
return torch . float8_e8m0fnu
2024-08-03 13:45:19 -04:00
fp8_dtype = None
2025-03-25 05:23:49 -04:00
if weight_dtype in FLOAT8_TYPES :
fp8_dtype = weight_dtype
2024-08-03 13:45:19 -04:00
if fp8_dtype is not None :
2024-10-20 00:54:47 -04:00
if supports_fp8_compute ( device ) : #if fp8 compute is supported the casting is most likely not expensive
return fp8_dtype
2024-08-03 13:45:19 -04:00
free_model_memory = maximum_vram_for_weights ( device )
if model_params * 2 > free_model_memory :
return fp8_dtype
2025-02-27 16:39:57 -05:00
if PRIORITIZE_FP16 or weight_dtype == torch . float16 :
2025-02-23 04:45:54 -05:00
if torch . float16 in supported_dtypes and should_use_fp16 ( device = device , model_params = model_params ) :
return torch . float16
2024-08-07 15:00:06 -04:00
for dt in supported_dtypes :
if dt == torch . float16 and should_use_fp16 ( device = device , model_params = model_params ) :
if torch . float16 in supported_dtypes :
return torch . float16
if dt == torch . bfloat16 and should_use_bf16 ( device , model_params = model_params ) :
if torch . bfloat16 in supported_dtypes :
return torch . bfloat16
for dt in supported_dtypes :
if dt == torch . float16 and should_use_fp16 ( device = device , model_params = model_params , manual_cast = True ) :
if torch . float16 in supported_dtypes :
return torch . float16
if dt == torch . bfloat16 and should_use_bf16 ( device , model_params = model_params , manual_cast = True ) :
if torch . bfloat16 in supported_dtypes :
return torch . bfloat16
2023-10-13 14:35:21 -04:00
return torch . float32
2023-12-11 18:24:44 -05:00
# None means no manual cast
2024-02-16 10:55:08 -05:00
def unet_manual_cast ( weight_dtype , inference_device , supported_dtypes = [ torch . float16 , torch . bfloat16 , torch . float32 ] ) :
2024-11-25 05:00:23 -05:00
if weight_dtype == torch . float32 or weight_dtype == torch . float64 :
2023-12-11 18:24:44 -05:00
return None
2024-02-16 10:55:08 -05:00
fp16_supported = should_use_fp16 ( inference_device , prioritize_performance = False )
2023-12-11 18:24:44 -05:00
if fp16_supported and weight_dtype == torch . float16 :
return None
2024-02-16 10:55:08 -05:00
bf16_supported = should_use_bf16 ( inference_device )
if bf16_supported and weight_dtype == torch . bfloat16 :
return None
2024-08-21 16:38:26 -04:00
fp16_supported = should_use_fp16 ( inference_device , prioritize_performance = True )
2025-02-28 02:17:50 -05:00
if PRIORITIZE_FP16 and fp16_supported and torch . float16 in supported_dtypes :
return torch . float16
2024-08-07 15:00:06 -04:00
for dt in supported_dtypes :
if dt == torch . float16 and fp16_supported :
return torch . float16
if dt == torch . bfloat16 and bf16_supported :
return torch . bfloat16
2024-02-16 10:55:08 -05:00
2024-08-07 15:00:06 -04:00
return torch . float32
2023-12-11 18:24:44 -05:00
2023-07-01 12:37:23 -04:00
def text_encoder_offload_device ( ) :
2023-07-03 00:08:30 -04:00
if args . gpu_only :
2023-06-15 15:21:37 -04:00
return get_torch_device ( )
else :
return torch . device ( " cpu " )
2023-07-01 12:37:23 -04:00
def text_encoder_device ( ) :
2023-07-03 00:08:30 -04:00
if args . gpu_only :
2023-07-01 12:37:23 -04:00
return get_torch_device ( )
2026-03-19 12:27:55 -07:00
elif vram_state in ( VRAMState . HIGH_VRAM , VRAMState . NORMAL_VRAM ) or comfy . memory_management . aimdo_enabled :
2023-08-23 21:45:00 -04:00
if should_use_fp16 ( prioritize_performance = False ) :
2023-07-01 14:38:51 -04:00
return get_torch_device ( )
else :
return torch . device ( " cpu " )
2023-07-01 12:37:23 -04:00
else :
return torch . device ( " cpu " )
2024-08-11 23:50:01 -04:00
def text_encoder_initial_device ( load_device , offload_device , model_size = 0 ) :
2026-03-04 13:33:14 -08:00
if comfy . memory_management . aimdo_enabled :
return offload_device
2024-08-11 23:50:01 -04:00
if load_device == offload_device or model_size < = 1024 * 1024 * 1024 :
return offload_device
2024-08-12 00:23:29 -04:00
if is_device_mps ( load_device ) :
2024-12-12 06:00:31 -05:00
return load_device
2024-08-12 00:23:29 -04:00
2024-08-11 23:50:01 -04:00
mem_l = get_free_memory ( load_device )
mem_o = get_free_memory ( offload_device )
if mem_l > ( mem_o * 0.5 ) and model_size * 1.2 < mem_l :
return load_device
else :
return offload_device
2023-11-17 02:56:59 -05:00
def text_encoder_dtype ( device = None ) :
if args . fp8_e4m3fn_text_enc :
return torch . float8_e4m3fn
elif args . fp8_e5m2_text_enc :
return torch . float8_e5m2
elif args . fp16_text_enc :
return torch . float16
2025-04-01 23:18:53 +05:30
elif args . bf16_text_enc :
return torch . bfloat16
2023-11-17 02:56:59 -05:00
elif args . fp32_text_enc :
return torch . float32
2023-12-10 23:00:54 -05:00
if is_device_cpu ( device ) :
return torch . float16
2024-02-02 10:02:49 -05:00
return torch . float16
2023-11-17 02:56:59 -05:00
2023-12-08 02:35:45 -05:00
def intermediate_device ( ) :
if args . gpu_only :
return get_torch_device ( )
else :
return torch . device ( " cpu " )
2026-03-14 16:18:19 -07:00
def intermediate_dtype ( ) :
if args . fp16_intermediates :
return torch . float16
else :
return torch . float32
2023-07-01 15:22:40 -04:00
def vae_device ( ) :
2023-12-30 05:38:21 -05:00
if args . cpu_vae :
return torch . device ( " cpu " )
2023-07-01 15:22:40 -04:00
return get_torch_device ( )
def vae_offload_device ( ) :
2023-07-03 00:08:30 -04:00
if args . gpu_only :
2023-07-01 15:22:40 -04:00
return get_torch_device ( )
else :
return torch . device ( " cpu " )
2024-06-16 13:12:54 -04:00
def vae_dtype ( device = None , allowed_dtypes = [ ] ) :
if args . fp16_vae :
return torch . float16
elif args . bf16_vae :
return torch . bfloat16
elif args . fp32_vae :
return torch . float32
for d in allowed_dtypes :
2024-12-25 04:50:34 -05:00
if d == torch . float16 and should_use_fp16 ( device ) :
2024-06-16 13:12:54 -04:00
return d
2024-12-25 04:50:34 -05:00
2025-10-11 21:28:01 -07:00
if d == torch . bfloat16 and should_use_bf16 ( device ) :
2024-06-16 13:12:54 -04:00
return d
2024-12-25 04:50:34 -05:00
return torch . float32
2023-07-06 18:04:28 -04:00
2023-03-06 10:50:50 -05:00
def get_autocast_device ( dev ) :
if hasattr ( dev , ' type ' ) :
return dev . type
return " cuda "
2023-02-17 15:45:29 -05:00
2023-12-04 11:10:00 -05:00
def supports_dtype ( device , dtype ) : #TODO
if dtype == torch . float32 :
return True
2023-12-11 18:24:44 -05:00
if is_device_cpu ( device ) :
2023-12-04 11:10:00 -05:00
return False
if dtype == torch . float16 :
return True
if dtype == torch . bfloat16 :
return True
return False
2024-06-11 17:03:26 -04:00
def supports_cast ( device , dtype ) : #TODO
if dtype == torch . float32 :
return True
if dtype == torch . float16 :
return True
if directml_enabled : #TODO: test this
return False
if dtype == torch . bfloat16 :
return True
2024-08-01 09:42:17 -04:00
if is_device_mps ( device ) :
return False
2024-06-11 17:03:26 -04:00
if dtype == torch . float8_e4m3fn :
return True
if dtype == torch . float8_e5m2 :
return True
return False
2024-08-01 11:05:56 -04:00
def pick_weight_dtype ( dtype , fallback_dtype , device = None ) :
if dtype is None :
dtype = fallback_dtype
elif dtype_size ( dtype ) > dtype_size ( fallback_dtype ) :
dtype = fallback_dtype
if not supports_cast ( device , dtype ) :
dtype = fallback_dtype
return dtype
2023-12-22 14:24:04 -05:00
def device_supports_non_blocking ( device ) :
2025-08-13 16:13:35 -07:00
if args . force_non_blocking :
return True
2023-12-22 14:24:04 -05:00
if is_device_mps ( device ) :
return False #pytorch bug? mps doesn't support non blocking
2025-08-13 16:13:35 -07:00
if is_intel_xpu ( ) : #xpu does support non blocking but it is slower on iGPUs for some reason so disable by default until situation changes
return False
2024-05-30 11:07:38 -04:00
if args . deterministic : #TODO: figure out why deterministic breaks non blocking from gpu to cpu (previews)
return False
if directml_enabled :
return False
2024-05-22 13:56:28 -04:00
return True
2024-06-15 01:08:12 -04:00
def force_channels_last ( ) :
if args . force_channels_last :
return True
#TODO
return False
2023-12-22 14:24:04 -05:00
2025-04-26 13:11:21 -07:00
STREAMS = { }
2025-11-27 16:03:03 +10:00
NUM_STREAMS = 0
2025-11-27 14:46:12 -08:00
if args . async_offload is not None :
NUM_STREAMS = args . async_offload
else :
2025-12-27 15:54:15 -08:00
# Enable by default on Nvidia and AMD
if is_nvidia ( ) or is_amd ( ) :
2025-11-27 14:46:12 -08:00
NUM_STREAMS = 2
if args . disable_async_offload :
NUM_STREAMS = 0
if NUM_STREAMS > 0 :
2025-04-26 13:11:21 -07:00
logging . info ( " Using async weight offloading with {} streams " . format ( NUM_STREAMS ) )
2025-10-30 07:17:46 +10:00
def current_stream ( device ) :
if device is None :
return None
if is_device_cuda ( device ) :
return torch . cuda . current_stream ( )
elif is_device_xpu ( device ) :
return torch . xpu . current_stream ( )
else :
return None
2025-04-29 17:28:52 -07:00
stream_counters = { }
2026-01-31 22:01:11 -08:00
STREAM_CAST_BUFFERS = { }
LARGEST_CASTED_WEIGHT = ( None , 0 )
2026-05-03 09:23:24 +10:00
STREAM_AIMDO_CAST_BUFFERS = { }
LARGEST_AIMDO_CASTED_WEIGHT = ( None , 0 )
2026-08-14 02:10:08 +10:00
CROSS_STEP_STATE = weakref . WeakSet ( )
2026-05-03 09:23:24 +10:00
DEFAULT_AIMDO_CAST_BUFFER_RESERVATION_SIZE = 16 * 1024 * * 3
2026-01-31 22:01:11 -08:00
2026-08-14 02:10:08 +10:00
# NOTE: devs/agents: this is temporary and will be removed in a future comfy. Not supported for custom node use.
def _register_cross_step ( module ) :
CROSS_STEP_STATE . add ( module )
2026-01-31 22:01:11 -08:00
def get_cast_buffer ( offload_stream , device , size , ref ) :
global LARGEST_CASTED_WEIGHT
if offload_stream is not None :
wf_context = offload_stream
if hasattr ( wf_context , " as_context " ) :
wf_context = wf_context . as_context ( offload_stream )
else :
wf_context = nullcontext ( )
cast_buffer = STREAM_CAST_BUFFERS . get ( offload_stream , None )
if cast_buffer is None or cast_buffer . numel ( ) < size :
if ref is LARGEST_CASTED_WEIGHT [ 0 ] :
#If there is one giant weight we do not want both streams to
#allocate a buffer for it. It's up to the caster to get the other
#offload stream in this corner case
return None
if cast_buffer is not None and cast_buffer . numel ( ) > 50 * ( 1024 * * 2 ) :
#I want my wrongly sized 50MB+ of VRAM back from the caching allocator right now
2026-02-02 14:34:46 -08:00
synchronize ( )
2026-01-31 22:01:11 -08:00
del STREAM_CAST_BUFFERS [ offload_stream ]
del cast_buffer
2026-02-02 14:34:46 -08:00
soft_empty_cache ( )
2026-01-31 22:01:11 -08:00
with wf_context :
cast_buffer = torch . empty ( ( size ) , dtype = torch . int8 , device = device )
STREAM_CAST_BUFFERS [ offload_stream ] = cast_buffer
if size > LARGEST_CASTED_WEIGHT [ 1 ] :
LARGEST_CASTED_WEIGHT = ( ref , size )
return cast_buffer
2026-05-03 09:23:24 +10:00
def get_aimdo_cast_buffer ( offload_stream , device ) :
cast_buffer = STREAM_AIMDO_CAST_BUFFERS . get ( offload_stream , None )
if cast_buffer is None :
cast_buffer = comfy_aimdo . vram_buffer . VRAMBuffer ( DEFAULT_AIMDO_CAST_BUFFER_RESERVATION_SIZE , device . index )
STREAM_AIMDO_CAST_BUFFERS [ offload_stream ] = cast_buffer
return cast_buffer
2026-05-21 10:03:58 +10:00
2026-01-31 22:01:11 -08:00
def reset_cast_buffers ( ) :
global LARGEST_CASTED_WEIGHT
2026-05-03 09:23:24 +10:00
global LARGEST_AIMDO_CASTED_WEIGHT
2026-01-31 22:01:11 -08:00
LARGEST_CASTED_WEIGHT = ( None , 0 )
2026-05-03 09:23:24 +10:00
LARGEST_AIMDO_CASTED_WEIGHT = ( None , 0 )
2026-05-31 05:20:04 +10:00
for offload_stream in set ( STREAM_CAST_BUFFERS ) | set ( STREAM_AIMDO_CAST_BUFFERS ) :
2026-05-03 09:23:24 +10:00
if offload_stream is not None :
offload_stream . synchronize ( )
2026-03-07 09:38:08 -08:00
synchronize ( )
2026-05-03 09:23:24 +10:00
2026-05-21 10:03:58 +10:00
for mmap_obj in DIRTY_MMAPS :
mmap_obj . bounce ( )
DIRTY_MMAPS . clear ( )
2026-08-14 02:10:08 +10:00
for module in CROSS_STEP_STATE :
del module . _comfy_cross_step_state
CROSS_STEP_STATE . clear ( )
2026-05-21 10:03:58 +10:00
for loaded_model in current_loaded_models :
model = loaded_model . model
if model is not None and model . is_dynamic ( ) :
2026-05-31 05:20:04 +10:00
pin_state = model . model . dynamic_pins [ model . load_device ]
if pin_state [ " active " ] :
2026-07-29 07:05:57 +10:00
for subset in ( " weights " , " weights-loaded " ) :
* _ , buckets = pin_state [ subset ]
for size , bucket in list ( buckets . items ( ) ) :
bucket [ : ] = [ entry for entry in bucket if entry [ - 1 ] is not None ]
if not bucket :
del buckets [ size ]
2026-05-31 05:20:04 +10:00
pin_state [ " active " ] = False
2026-07-29 07:05:57 +10:00
model . partially_unload_ram ( 1e30 , subsets = [ " patches " , " patches-loaded " ] )
for subset in ( " patches " , " patches-loaded " ) :
pin_state [ subset ] = ( comfy_aimdo . host_buffer . HostBuffer ( 0 , 8 * 1024 * 1024 , pinned_hostbuf_size ( model . model_size ( ) ) ) , [ ] , [ - 1 ] , [ 0 ] , [ 0 ] , { } )
2026-05-21 10:03:58 +10:00
2026-01-31 22:01:11 -08:00
STREAM_CAST_BUFFERS . clear ( )
2026-05-03 09:23:24 +10:00
STREAM_AIMDO_CAST_BUFFERS . clear ( )
2026-02-02 14:34:46 -08:00
soft_empty_cache ( )
2026-01-31 22:01:11 -08:00
2025-04-26 13:11:21 -07:00
def get_offload_stream ( device ) :
2025-04-29 17:28:52 -07:00
stream_counter = stream_counters . get ( device , 0 )
2025-11-27 16:03:03 +10:00
if NUM_STREAMS == 0 :
2025-04-26 13:11:21 -07:00
return None
2025-11-28 13:16:46 -08:00
if torch . compiler . is_compiling ( ) :
return None
2025-04-26 13:11:21 -07:00
if device in STREAMS :
ss = STREAMS [ device ]
2025-10-30 07:17:46 +10:00
#Sync the oldest stream in the queue with the current
ss [ stream_counter ] . wait_stream ( current_stream ( device ) )
2025-04-26 13:11:21 -07:00
stream_counter = ( stream_counter + 1 ) % len ( ss )
2025-04-29 17:28:52 -07:00
stream_counters [ device ] = stream_counter
2025-10-30 07:17:46 +10:00
return ss [ stream_counter ]
2025-04-26 13:11:21 -07:00
elif is_device_cuda ( device ) :
ss = [ ]
for k in range ( NUM_STREAMS ) :
2025-11-29 07:38:12 +10:00
s1 = torch . cuda . Stream ( device = device , priority = 0 )
s1 . as_context = torch . cuda . stream
ss . append ( s1 )
2025-04-26 13:11:21 -07:00
STREAMS [ device ] = ss
s = ss [ stream_counter ]
2025-04-29 17:28:52 -07:00
stream_counters [ device ] = stream_counter
2025-04-26 13:11:21 -07:00
return s
2025-07-22 12:20:09 -07:00
elif is_device_xpu ( device ) :
ss = [ ]
for k in range ( NUM_STREAMS ) :
2025-11-29 07:38:12 +10:00
s1 = torch . xpu . Stream ( device = device , priority = 0 )
s1 . as_context = torch . xpu . stream
ss . append ( s1 )
2025-07-22 12:20:09 -07:00
STREAMS [ device ] = ss
s = ss [ stream_counter ]
stream_counters [ device ] = stream_counter
return s
2025-04-26 13:11:21 -07:00
return None
def sync_stream ( device , stream ) :
2025-10-30 07:17:46 +10:00
if stream is None or current_stream ( device ) is None :
2025-04-26 13:11:21 -07:00
return
2025-10-30 07:17:46 +10:00
current_stream ( device ) . wait_stream ( stream )
2025-04-26 13:11:21 -07:00
2026-01-31 22:01:11 -08:00
2026-05-21 10:03:58 +10:00
def cast_to_gathered ( tensors , r , non_blocking = False , stream = None , r2 = None ) :
2026-01-31 22:01:11 -08:00
wf_context = nullcontext ( )
if stream is not None :
wf_context = stream
if hasattr ( wf_context , " as_context " ) :
wf_context = wf_context . as_context ( stream )
2026-05-31 05:20:04 +10:00
dest_views = comfy . memory_management . interpret_gathered_like ( tensors , r ) if r is not None else [ None ] * len ( tensors )
2026-05-21 10:03:58 +10:00
dest2_views = comfy . memory_management . interpret_gathered_like ( tensors , r2 ) if r2 is not None else None
2026-01-31 22:01:11 -08:00
with wf_context :
for tensor in tensors :
dest_view = dest_views . pop ( 0 )
2026-05-21 10:03:58 +10:00
dest2_view = dest2_views . pop ( 0 ) if dest2_views is not None else None
2026-01-31 22:01:11 -08:00
if tensor is None :
continue
2026-05-21 10:03:58 +10:00
if comfy . memory_management . read_tensor_file_slice_into ( tensor , dest_view , stream = stream , destination2 = dest2_view ) :
2026-03-13 19:18:08 -07:00
continue
storage = tensor . _qdata . untyped_storage ( ) if isinstance ( tensor , comfy . quant_ops . QuantizedTensor ) else tensor . untyped_storage ( )
2026-05-21 10:03:58 +10:00
mark_mmap_dirty ( storage )
2026-05-31 05:20:04 +10:00
if dest_view is not None :
dest_view . copy_ ( tensor , non_blocking = non_blocking )
2026-05-21 10:03:58 +10:00
if dest2_view is not None :
2026-05-31 05:20:04 +10:00
dest2_view . copy_ ( tensor if dest_view is None else dest_view , non_blocking = non_blocking )
2026-01-31 22:01:11 -08:00
def cast_to ( weight , dtype = None , device = None , non_blocking = False , copy = False , stream = None , r = None ) :
2024-10-17 17:25:56 -04:00
if device is None or weight . device == device :
if not copy :
if dtype is None or weight . dtype == dtype :
return weight
2025-04-27 02:38:11 -07:00
if stream is not None :
2025-11-29 07:38:12 +10:00
wf_context = stream
if hasattr ( wf_context , " as_context " ) :
wf_context = wf_context . as_context ( stream )
with wf_context :
2025-04-27 02:38:11 -07:00
return weight . to ( dtype = dtype , copy = copy )
2024-10-17 17:25:56 -04:00
return weight . to ( dtype = dtype , copy = copy )
2023-09-20 17:52:41 -04:00
2025-11-29 07:38:12 +10:00
2025-04-26 13:11:21 -07:00
if stream is not None :
2025-11-29 07:38:12 +10:00
wf_context = stream
if hasattr ( wf_context , " as_context " ) :
wf_context = wf_context . as_context ( stream )
with wf_context :
2026-01-31 22:01:11 -08:00
if r is None :
r = torch . empty_like ( weight , dtype = dtype , device = device )
2025-04-26 13:11:21 -07:00
r . copy_ ( weight , non_blocking = non_blocking )
else :
2026-01-31 22:01:11 -08:00
if r is None :
r = torch . empty_like ( weight , dtype = dtype , device = device )
2025-04-26 13:11:21 -07:00
r . copy_ ( weight , non_blocking = non_blocking )
2024-10-17 17:25:56 -04:00
return r
def cast_to_device ( tensor , device , dtype , copy = False ) :
non_blocking = device_supports_non_blocking ( device )
return cast_to ( tensor , dtype = dtype , device = device , non_blocking = non_blocking , copy = copy )
2023-12-10 01:30:35 -05:00
2025-11-04 14:37:50 -08:00
PINNED_MEMORY = { }
TOTAL_PINNED_MEMORY = 0
2025-11-05 15:08:13 -08:00
MAX_PINNED_MEMORY = - 1
2026-08-03 13:29:47 -07:00
def get_disk_swap_total ( ) :
if not os . path . exists ( " /proc/swaps " ) :
return 0
total = 0
try :
with open ( " /proc/swaps " , encoding = " utf-8 " ) as swaps :
next ( swaps , None )
for line in swaps :
filename , _ , size , _ , _ = line . rsplit ( maxsplit = 4 )
if os . path . basename ( os . path . realpath ( filename ) ) . startswith ( " zram " ) :
continue
total + = int ( size ) * 1024
except :
logging . warning ( " Could not get amount of swap memory on system. " )
return total
2025-11-05 15:08:13 -08:00
if not args . disable_pinned_memory :
2025-11-05 16:11:15 -08:00
if is_nvidia ( ) or is_amd ( ) :
2026-05-21 10:03:58 +10:00
ram = get_total_memory ( torch . device ( " cpu " ) )
2025-11-05 15:08:13 -08:00
if WINDOWS :
2026-05-21 10:03:58 +10:00
MAX_PINNED_MEMORY = ram * 0.40 # Windows limit is apparently 50%
2025-11-05 15:08:13 -08:00
else :
2026-08-27 18:57:31 +04:00
swap = 0 if comfy . system_memory . cgroup_memory_limit ( ) is not None else get_disk_swap_total ( )
MAX_PINNED_MEMORY = max ( ram * 0.40 , min ( ram * 0.90 , ram - 4 * 1024 * * 3 , ram + swap - 16 * 1024 * * 3 ) )
2025-11-05 15:08:13 -08:00
logging . info ( " Enabled pinned memory {} " . format ( MAX_PINNED_MEMORY / / ( 1024 * 1024 ) ) )
2026-01-31 22:01:11 -08:00
PINNING_ALLOWED_TYPES = set ( [ " Tensor " , " Parameter " , " QuantizedTensor " ] )
2025-11-04 14:37:50 -08:00
2026-05-21 10:03:58 +10:00
def pinned_hostbuf_size ( size ) :
2026-06-13 00:53:33 +10:00
if args . high_ram :
return max ( 0 , int ( size * 2 ) )
2026-05-21 10:03:58 +10:00
return max ( 0 , int ( min ( size , MAX_PINNED_MEMORY ) * 2 ) )
2025-12-29 15:19:34 -08:00
def discard_cuda_async_error ( ) :
try :
a = torch . tensor ( [ 1 ] , dtype = torch . uint8 , device = get_torch_device ( ) )
b = torch . tensor ( [ 1 ] , dtype = torch . uint8 , device = get_torch_device ( ) )
_ = a + b
2026-02-02 14:34:46 -08:00
synchronize ( )
2026-03-11 19:04:13 +02:00
except RuntimeError :
2025-12-29 15:19:34 -08:00
#Dump it! We already know about it from the synchronous return
pass
2025-10-28 21:21:01 -07:00
def pin_memory ( tensor ) :
2025-11-04 14:37:50 -08:00
global TOTAL_PINNED_MEMORY
if MAX_PINNED_MEMORY < = 0 :
2025-10-28 21:21:01 -07:00
return False
2025-11-24 23:48:20 -08:00
if type ( tensor ) . __name__ not in PINNING_ALLOWED_TYPES :
2025-11-11 16:33:30 -08:00
return False
2025-10-28 21:21:01 -07:00
if not is_device_cpu ( tensor . device ) :
return False
2025-11-07 12:20:48 +10:00
if tensor . is_pinned ( ) :
#NOTE: Cuda does detect when a tensor is already pinned and would
#error below, but there are proven cases where this also queues an error
#on the GPU async. So dont trust the CUDA API and guard here
return False
2025-11-11 16:33:30 -08:00
if not tensor . is_contiguous ( ) :
return False
2026-01-05 18:48:58 -08:00
size = tensor . nbytes
2026-05-21 10:03:58 +10:00
comfy . memory_management . extra_ram_release ( comfy . memory_management . RAM_CACHE_HEADROOM )
ensure_pin_registerable ( size )
2025-11-04 14:37:50 -08:00
ptr = tensor . data_ptr ( )
2025-11-24 23:48:20 -08:00
if ptr == 0 :
return False
2025-11-04 14:37:50 -08:00
if torch . cuda . cudart ( ) . cudaHostRegister ( ptr , size , 1 ) == 0 :
PINNED_MEMORY [ ptr ] = size
TOTAL_PINNED_MEMORY + = size
2025-10-28 21:21:01 -07:00
return True
2025-12-29 15:19:34 -08:00
else :
2025-12-29 15:26:42 -08:00
logging . warning ( " Pin error. " )
2025-12-29 15:19:34 -08:00
discard_cuda_async_error ( )
2025-10-28 21:21:01 -07:00
return False
def unpin_memory ( tensor ) :
2025-11-04 14:37:50 -08:00
global TOTAL_PINNED_MEMORY
if MAX_PINNED_MEMORY < = 0 :
2025-10-28 21:21:01 -07:00
return False
if not is_device_cpu ( tensor . device ) :
return False
2025-11-07 08:15:05 -08:00
ptr = tensor . data_ptr ( )
2026-01-05 18:48:58 -08:00
size = tensor . nbytes
2025-11-07 08:15:05 -08:00
size_stored = PINNED_MEMORY . get ( ptr , None )
if size_stored is None :
logging . warning ( " Tried to unpin tensor not pinned by ComfyUI " )
return False
if size != size_stored :
logging . warning ( " Size of pinned tensor changed " )
2025-11-07 12:20:48 +10:00
return False
2025-11-04 14:37:50 -08:00
if torch . cuda . cudart ( ) . cudaHostUnregister ( ptr ) == 0 :
2026-05-21 10:03:58 +10:00
size = PINNED_MEMORY . pop ( ptr )
TOTAL_PINNED_MEMORY - = size
2025-10-28 21:21:01 -07:00
return True
2025-12-29 15:19:34 -08:00
else :
2025-12-29 15:26:42 -08:00
logging . warning ( " Unpin error. " )
2025-12-29 15:19:34 -08:00
discard_cuda_async_error ( )
2025-10-28 21:21:01 -07:00
return False
2024-12-18 01:56:10 -05:00
def sage_attention_enabled ( ) :
return args . use_sage_attention
2023-04-04 22:22:02 -04:00
2026-08-10 22:03:08 -07:00
def comfy_kitchen_attention_enabled ( ) :
return args . use_ck_attention
2025-03-14 08:22:41 +01:00
def flash_attention_enabled ( ) :
return args . use_flash_attention
2023-03-12 15:44:16 -04:00
def xformers_enabled ( ) :
2023-04-28 14:28:57 -04:00
global directml_enabled
2023-06-03 11:05:37 -04:00
global cpu_state
if cpu_state != CPUState . GPU :
2023-03-12 15:44:16 -04:00
return False
2023-09-02 18:22:10 -07:00
if is_intel_xpu ( ) :
2023-04-28 14:28:57 -04:00
return False
2024-12-27 08:36:50 +08:00
if is_ascend_npu ( ) :
return False
2025-02-27 09:45:13 +08:00
if is_mlu ( ) :
return False
2025-07-25 01:57:36 +08:00
if is_ixuca ( ) :
return False
2023-04-28 14:28:57 -04:00
if directml_enabled :
return False
2023-04-05 23:41:23 -04:00
return XFORMERS_IS_AVAILABLE
2023-03-12 15:44:16 -04:00
2023-04-04 22:22:02 -04:00
def xformers_enabled_vae ( ) :
enabled = xformers_enabled ( )
if not enabled :
return False
2023-04-09 01:31:47 -04:00
return XFORMERS_ENABLED_VAE
2023-04-04 22:22:02 -04:00
2023-03-13 12:25:19 -04:00
def pytorch_attention_enabled ( ) :
2023-05-06 19:58:54 -04:00
global ENABLE_PYTORCH_ATTENTION
2023-03-13 12:25:19 -04:00
return ENABLE_PYTORCH_ATTENTION
2025-02-14 05:42:14 -05:00
def pytorch_attention_enabled_vae ( ) :
if is_amd ( ) :
return False # enabling pytorch attention on AMD currently causes crash when doing high res
return pytorch_attention_enabled ( )
2023-05-06 19:58:54 -04:00
def pytorch_attention_flash_attention ( ) :
global ENABLE_PYTORCH_ATTENTION
if ENABLE_PYTORCH_ATTENTION :
#TODO: more reliable way of checking for flash attention?
2025-06-10 10:06:24 -07:00
if is_nvidia ( ) :
2023-05-06 19:58:54 -04:00
return True
2024-06-04 17:44:14 -04:00
if is_intel_xpu ( ) :
return True
2024-12-27 08:36:50 +08:00
if is_ascend_npu ( ) :
return True
2025-02-27 09:45:13 +08:00
if is_mlu ( ) :
return True
2025-02-13 08:32:36 -05:00
if is_amd ( ) :
return True #if you have pytorch attention enabled on AMD it probably supports at least mem efficient attention
2025-07-25 01:57:36 +08:00
if is_ixuca ( ) :
return True
2023-05-06 19:58:54 -04:00
return False
2024-12-25 05:18:50 -05:00
def force_upcast_attention_dtype ( ) :
upcast = args . force_upcast_attention
macos_version = mac_version ( )
2025-06-10 10:06:24 -07:00
if macos_version is not None and ( ( 14 , 5 ) < = macos_version ) : # black image bug on recent versions of macOS, I don't think it's ever getting fixed
2024-12-25 05:18:50 -05:00
upcast = True
2024-05-21 16:56:33 -04:00
if upcast :
2025-02-24 05:41:07 -05:00
return { torch . float16 : torch . float32 }
2024-05-21 16:56:33 -04:00
else :
return None
2023-03-03 03:27:33 -05:00
def get_free_memory ( dev = None , torch_free_too = False ) :
2023-04-28 14:28:57 -04:00
global directml_enabled
2023-03-03 03:27:33 -05:00
if dev is None :
2023-03-06 10:50:50 -05:00
dev = get_torch_device ( )
2023-03-03 03:27:33 -05:00
2023-03-24 14:04:50 +02:00
if hasattr ( dev , ' type ' ) and ( dev . type == ' cpu ' or dev . type == ' mps ' ) :
2026-08-27 18:57:31 +04:00
mem_free_total = comfy . system_memory . virtual_memory_available ( )
2023-03-03 03:27:33 -05:00
mem_free_torch = mem_free_total
else :
2023-04-28 14:28:57 -04:00
if directml_enabled :
mem_free_total = 1024 * 1024 * 1024 #TODO
mem_free_torch = mem_free_total
2023-09-02 18:22:10 -07:00
elif is_intel_xpu ( ) :
2023-08-17 03:12:17 -07:00
stats = torch . xpu . memory_stats ( dev )
mem_active = stats [ ' active_bytes.all.current ' ]
mem_reserved = stats [ ' reserved_bytes.all.current ' ]
2025-07-23 15:18:20 -07:00
mem_free_xpu = torch . xpu . get_device_properties ( dev ) . total_memory - mem_reserved
2023-08-17 03:12:17 -07:00
mem_free_torch = mem_reserved - mem_active
2024-05-12 03:36:30 -07:00
mem_free_total = mem_free_xpu + mem_free_torch
2024-12-27 08:36:50 +08:00
elif is_ascend_npu ( ) :
stats = torch . npu . memory_stats ( dev )
mem_active = stats [ ' active_bytes.all.current ' ]
mem_reserved = stats [ ' reserved_bytes.all.current ' ]
mem_free_npu , _ = torch . npu . mem_get_info ( dev )
mem_free_torch = mem_reserved - mem_active
mem_free_total = mem_free_npu + mem_free_torch
2025-02-27 09:45:13 +08:00
elif is_mlu ( ) :
stats = torch . mlu . memory_stats ( dev )
mem_active = stats [ ' active_bytes.all.current ' ]
mem_reserved = stats [ ' reserved_bytes.all.current ' ]
mem_free_mlu , _ = torch . mlu . mem_get_info ( dev )
mem_free_torch = mem_reserved - mem_active
mem_free_total = mem_free_mlu + mem_free_torch
2023-04-06 14:24:47 +08:00
else :
stats = torch . cuda . memory_stats ( dev )
mem_active = stats [ ' active_bytes.all.current ' ]
mem_reserved = stats [ ' reserved_bytes.all.current ' ]
mem_free_cuda , _ = torch . cuda . mem_get_info ( dev )
mem_free_torch = mem_reserved - mem_active
mem_free_total = mem_free_cuda + mem_free_torch
2023-03-03 03:27:33 -05:00
if torch_free_too :
return ( mem_free_total , mem_free_torch )
else :
return mem_free_total
2023-02-08 14:05:31 -05:00
2023-03-03 11:07:10 -05:00
def cpu_mode ( ) :
2023-06-03 11:05:37 -04:00
global cpu_state
return cpu_state == CPUState . CPU
2023-03-03 11:07:10 -05:00
2023-03-24 14:04:50 +02:00
def mps_mode ( ) :
2023-06-03 11:05:37 -04:00
global cpu_state
return cpu_state == CPUState . MPS
2023-03-24 14:04:50 +02:00
2024-02-15 21:10:10 -05:00
def is_device_type ( device , type ) :
2023-07-01 13:22:51 -04:00
if hasattr ( device , ' type ' ) :
2024-02-15 21:10:10 -05:00
if ( device . type == type ) :
2023-07-04 02:09:02 -04:00
return True
return False
2024-02-15 21:10:10 -05:00
def is_device_cpu ( device ) :
return is_device_type ( device , ' cpu ' )
2023-07-04 02:09:02 -04:00
def is_device_mps ( device ) :
2024-02-15 21:10:10 -05:00
return is_device_type ( device , ' mps ' )
2025-07-22 12:20:09 -07:00
def is_device_xpu ( device ) :
return is_device_type ( device , ' xpu ' )
2024-02-15 21:10:10 -05:00
def is_device_cuda ( device ) :
return is_device_type ( device , ' cuda ' )
2023-07-01 13:22:51 -04:00
2026-05-31 04:18:42 +03:00
def set_torch_device ( device ) :
""" Set the current device for the given torch device. Supports CUDA and XPU. """
if is_device_cuda ( device ) :
torch . cuda . set_device ( device )
elif is_device_xpu ( device ) :
torch . xpu . set_device ( device )
2025-02-11 14:11:32 -08:00
def is_directml_enabled ( ) :
global directml_enabled
if directml_enabled :
return True
return False
2024-02-04 13:23:43 -05:00
def should_use_fp16 ( device = None , model_params = 0 , prioritize_performance = True , manual_cast = False ) :
2023-08-23 21:38:28 -04:00
if device is not None :
if is_device_cpu ( device ) :
return False
2025-02-11 08:31:46 -05:00
if args . force_fp16 :
2023-07-01 22:42:35 -04:00
return True
2023-04-07 00:27:54 -04:00
if FORCE_FP32 :
return False
2025-02-20 09:29:59 -05:00
if is_directml_enabled ( ) :
return True
2023-04-28 14:28:57 -04:00
2024-12-25 05:32:51 -05:00
if ( device is not None and is_device_mps ( device ) ) or mps_mode ( ) :
2024-02-19 12:00:48 -05:00
return True
if cpu_mode ( ) :
return False
2023-03-03 11:07:10 -05:00
2023-09-02 18:22:10 -07:00
if is_intel_xpu ( ) :
2026-05-01 14:16:41 -07:00
return torch . xpu . get_device_properties ( device ) . has_fp16
2023-08-20 14:56:47 -04:00
2024-12-27 08:36:50 +08:00
if is_ascend_npu ( ) :
return True
2025-02-27 09:45:13 +08:00
if is_mlu ( ) :
return True
2025-07-25 01:57:36 +08:00
if is_ixuca ( ) :
return True
2024-02-04 20:53:35 -05:00
if torch . version . hip :
2023-03-03 11:07:10 -05:00
return True
2024-08-20 00:31:04 -04:00
props = torch . cuda . get_device_properties ( device )
2024-02-04 20:53:35 -05:00
if props . major > = 8 :
return True
2023-07-02 09:37:31 -04:00
if props . major < 6 :
return False
2024-08-21 23:23:50 -04:00
#FP16 is confirmed working on a 1080 (GP104) and on latest pytorch actually seems faster than fp32
2024-03-02 17:16:31 -05:00
nvidia_10_series = [ " 1080 " , " 1070 " , " titan x " , " p3000 " , " p3200 " , " p4000 " , " p4200 " , " p5000 " , " p5200 " , " p6000 " , " 1060 " , " 1050 " , " p40 " , " p100 " , " p6 " , " p4 " ]
2023-07-02 09:37:31 -04:00
for x in nvidia_10_series :
if x in props . name . lower ( ) :
2024-09-01 17:29:31 -04:00
if WINDOWS or manual_cast :
return True
else :
return False #weird linux behavior where fp32 is faster
2023-07-02 09:37:31 -04:00
2024-08-21 23:23:50 -04:00
if manual_cast :
2024-08-03 13:45:19 -04:00
free_model_memory = maximum_vram_for_weights ( device )
2023-08-23 21:45:00 -04:00
if ( not prioritize_performance ) or model_params * 4 > free_model_memory :
2023-07-02 09:37:31 -04:00
return True
2023-03-03 11:07:10 -05:00
if props . major < 7 :
return False
2023-07-02 09:37:31 -04:00
#FP16 is just broken on these cards
2023-10-16 16:46:41 -04:00
nvidia_16_series = [ " 1660 " , " 1650 " , " 1630 " , " T500 " , " T550 " , " T600 " , " MX550 " , " MX450 " , " CMP 30HX " , " T2000 " , " T1000 " , " T1200 " ]
2023-03-03 11:07:10 -05:00
for x in nvidia_16_series :
if x in props . name :
return False
return True
2024-02-17 08:13:17 -05:00
def should_use_bf16 ( device = None , model_params = 0 , prioritize_performance = True , manual_cast = False ) :
if device is not None :
if is_device_cpu ( device ) : #TODO ? bf16 works on CPU but is extremely slow
return False
2024-02-16 23:01:54 -05:00
if FORCE_FP32 :
return False
2024-02-17 08:13:17 -05:00
if directml_enabled :
return False
2024-12-25 05:32:51 -05:00
if ( device is not None and is_device_mps ( device ) ) or mps_mode ( ) :
2024-12-25 05:18:50 -05:00
if mac_version ( ) < ( 14 , ) :
return False
2024-08-01 16:18:14 -04:00
return True
if cpu_mode ( ) :
2024-02-17 08:13:17 -05:00
return False
2024-02-16 10:55:08 -05:00
if is_intel_xpu ( ) :
2026-05-01 14:16:41 -07:00
return torch . xpu . is_bf16_supported ( )
2025-02-12 06:49:16 -05:00
2025-02-12 19:48:11 +08:00
if is_ascend_npu ( ) :
return True
2024-02-16 10:55:08 -05:00
2025-07-25 01:57:36 +08:00
if is_ixuca ( ) :
return True
2025-02-16 05:45:08 -05:00
if is_amd ( ) :
arch = torch . cuda . get_device_properties ( device ) . gcnArchName
2025-10-21 16:15:23 -07:00
if any ( ( a in arch ) for a in AMD_RDNA2_AND_OLDER_ARCH ) : # RDNA2 and older don't support bf16
2025-02-17 04:42:40 -05:00
if manual_cast :
return True
2025-02-16 05:45:08 -05:00
return False
2024-08-20 00:31:04 -04:00
props = torch . cuda . get_device_properties ( device )
2025-02-27 09:45:13 +08:00
if is_mlu ( ) :
if props . major > 3 :
return True
2024-02-16 10:55:08 -05:00
if props . major > = 8 :
return True
2024-02-17 08:13:17 -05:00
bf16_works = torch . cuda . is_bf16_supported ( )
2025-02-18 07:28:33 -05:00
if bf16_works and manual_cast :
2024-08-03 13:45:19 -04:00
free_model_memory = maximum_vram_for_weights ( device )
2024-02-17 08:13:17 -05:00
if ( not prioritize_performance ) or model_params * 4 > free_model_memory :
return True
2024-02-16 10:55:08 -05:00
return False
2024-08-20 11:49:33 -04:00
def supports_fp8_compute ( device = None ) :
2025-06-08 11:15:34 -07:00
if SUPPORT_FP8_OPS :
2025-05-23 14:43:50 -07:00
return True
2024-10-09 19:43:17 -04:00
if not is_nvidia ( ) :
return False
2024-08-20 11:49:33 -04:00
props = torch . cuda . get_device_properties ( device )
if props . major > = 9 :
return True
if props . major < 8 :
return False
if props . minor < 9 :
return False
2024-10-09 19:43:17 -04:00
2025-06-07 07:01:15 -07:00
if torch_version_numeric < ( 2 , 3 ) :
2024-10-09 19:43:17 -04:00
return False
if WINDOWS :
2025-06-07 07:01:15 -07:00
if torch_version_numeric < ( 2 , 4 ) :
2024-10-09 19:43:17 -04:00
return False
2024-08-20 11:49:33 -04:00
return True
2026-01-06 15:07:26 -08:00
def supports_nvfp4_compute ( device = None ) :
if not is_nvidia ( ) :
return False
props = torch . cuda . get_device_properties ( device )
if props . major < 10 :
return False
return True
2026-03-15 00:36:29 +02:00
def supports_mxfp8_compute ( device = None ) :
if not is_nvidia ( ) :
return False
if torch_version_numeric < ( 2 , 10 ) :
return False
props = torch . cuda . get_device_properties ( device )
if props . major < 10 :
return False
return True
2026-04-11 18:06:36 -07:00
def supports_fp64 ( device = None ) :
2026-09-05 18:00:11 -07:00
if ( device is not None and is_device_mps ( device ) ) or mps_mode ( ) :
2026-04-11 18:06:36 -07:00
return False
if is_intel_xpu ( ) :
return False
if is_directml_enabled ( ) :
return False
if is_ixuca ( ) :
return False
return True
2026-09-08 04:41:32 +08:00
def supports_int8_compute ( device = None ) :
# The eager comfy_kitchen backend implements int8 weight-only quantized
# matmul via torch._int_mm, which PyTorch does not implement for MPS.
# https://github.com/pytorch/pytorch/issues/141287
if ( device is not None and is_device_mps ( device ) ) or mps_mode ( ) :
return False
if is_intel_xpu ( ) :
return False
if is_directml_enabled ( ) :
return False
if is_ixuca ( ) :
return False
return True
2025-06-26 00:39:09 -07:00
def extended_fp16_support ( ) :
# TODO: check why some models work with fp16 on newer torch versions but not on older
if torch_version_numeric < ( 2 , 7 ) :
return False
return True
2025-12-06 15:36:20 -08:00
LORA_COMPUTE_DTYPES = { }
def lora_compute_dtype ( device ) :
dtype = LORA_COMPUTE_DTYPES . get ( device , None )
if dtype is not None :
return dtype
if should_use_fp16 ( device ) :
dtype = torch . float16
else :
dtype = torch . float32
LORA_COMPUTE_DTYPES [ device ] = dtype
return dtype
2026-02-02 14:34:46 -08:00
def synchronize ( ) :
2026-03-04 23:39:51 -08:00
if cpu_mode ( ) :
return
2026-02-02 14:34:46 -08:00
if is_intel_xpu ( ) :
torch . xpu . synchronize ( )
elif torch . cuda . is_available ( ) :
torch . cuda . synchronize ( )
2023-09-04 00:58:18 -04:00
def soft_empty_cache ( force = False ) :
2026-03-04 23:39:51 -08:00
if cpu_mode ( ) :
return
2023-06-03 11:05:37 -04:00
global cpu_state
if cpu_state == CPUState . MPS :
2023-06-01 03:52:51 -04:00
torch . mps . empty_cache ( )
2023-09-02 18:22:10 -07:00
elif is_intel_xpu ( ) :
2026-05-01 14:16:41 -07:00
torch . xpu . synchronize ( )
2023-04-15 11:19:07 -04:00
torch . xpu . empty_cache ( )
2024-12-27 08:36:50 +08:00
elif is_ascend_npu ( ) :
torch . npu . empty_cache ( )
2025-04-03 07:24:04 +08:00
elif is_mlu ( ) :
torch . mlu . empty_cache ( )
2023-04-15 11:19:07 -04:00
elif torch . cuda . is_available ( ) :
2026-02-03 18:39:19 -08:00
torch . cuda . synchronize ( )
torch . cuda . empty_cache ( )
torch . cuda . ipc_collect ( )
2023-04-15 11:19:07 -04:00
2023-12-23 04:25:06 -05:00
def unload_all_models ( ) :
2026-05-25 18:26:40 -07:00
for device in get_all_torch_devices ( ) :
free_memory ( 1e30 , device )
def unload_model_and_clones ( model : ModelPatcher , unload_additional_models = True , all_devices = False ) :
' Unload only model and its clones - primarily for multigpu cloning purposes. '
initial_keep_loaded : list [ LoadedModel ] = current_loaded_models . copy ( )
additional_models = [ ]
if unload_additional_models :
additional_models = model . get_nested_additional_models ( )
keep_loaded = [ ]
for loaded_model in initial_keep_loaded :
if loaded_model . model is not None :
if model . clone_base_uuid == loaded_model . model . clone_base_uuid :
continue
# check additional models if they are a match
skip = False
for add_model in additional_models :
if add_model . clone_base_uuid == loaded_model . model . clone_base_uuid :
skip = True
break
if skip :
continue
keep_loaded . append ( loaded_model )
if not all_devices :
free_memory ( 1e30 , get_torch_device ( ) , keep_loaded )
else :
for device in get_all_torch_devices ( ) :
free_memory ( 1e30 , device , keep_loaded )
2023-12-23 04:25:06 -05:00
2026-01-03 19:28:38 -08:00
def debug_memory_summary ( ) :
if is_amd ( ) or is_nvidia ( ) :
return torch . cuda . memory . memory_summary ( )
return " "
2023-12-23 04:25:06 -05:00
2026-04-22 16:08:19 -06:00
class InterruptProcessingException ( BaseException ) :
2023-03-02 14:42:03 -05:00
pass
interrupt_processing_mutex = threading . RLock ( )
interrupt_processing = False
def interrupt_current_processing ( value = True ) :
global interrupt_processing
global interrupt_processing_mutex
with interrupt_processing_mutex :
interrupt_processing = value
def processing_interrupted ( ) :
global interrupt_processing
global interrupt_processing_mutex
with interrupt_processing_mutex :
return interrupt_processing
def throw_exception_if_processing_interrupted ( ) :
global interrupt_processing
global interrupt_processing_mutex
with interrupt_processing_mutex :
if interrupt_processing :
interrupt_processing = False
raise InterruptProcessingException ( )