Use ComfyUI's device selection instead of hardcoded CUDA checks

Several model-loading and processing paths pick their device with
`"cuda" if torch.cuda.is_available() else "cpu"`, and move tensors with
`.cuda()` guarded by the same check:

- `imagefunc.load_RMBG_model` / `RMBG`
- `imagefunc.get_masked_VITMatte_image` (vitmatte)
- `imagefunc.UformGen2QwenChat`
- `imagefunc.clear_memory`
- `blendmodes.rgb_to_hsv_via_torch` / `hsv_to_rgb_via_torch`

On any non-CUDA accelerator that ComfyUI itself supports (e.g. Ascend NPU
via torch_npu, which ComfyUI's own `comfy.model_management.get_torch_device`
returns a `torch.device("npu")` for), these paths silently fall back to the
CPU or, for `.cuda()`, move the input to a CUDA device that does not exist.

This repo already imports `comfy.model_management` (e.g. in `purge_vram.py`),
so use its device selection for consistency with the rest of ComfyUI:

- `load_RMBG_model`, `UformGen2QwenChat`, and the blendmode helpers now use
  `comfy.model_management.get_torch_device()`.
- `RMBG` moves the input to the device of the (lru-cached) model itself
  instead of probing CUDA, so it also follows device changes correctly.
- the vitmatte path uses the ComfyUI device unless the caller explicitly
  asked for CPU; a hardcoded "cuda" request without CUDA available now logs
  and uses the ComfyUI default.
- cache clearing also calls `torch.accelerator.empty_cache()` when a
  non-CUDA accelerator is present (`torch.accelerator` is the device-agnostic
  PyTorch API).

Verified on Ascend 910B (torch 2.14.0a0 + torch_npu 2.14.0) with
ComfyUI-master's `comfy.model_management` implementation:
`get_torch_device()` returns `torch.device("npu", 0)`, a conv2d forward and
an RMBG-style model/input flow run on `npu:0`, and
`torch.accelerator.empty_cache()` is callable. On CUDA machines the selected
device is unchanged (`cuda`).
This commit is contained in:
li-lizhe
2026-09-02 16:02:58 +08:00
parent 557d882e18
commit 6555d35004
2 changed files with 17 additions and 12 deletions
+5 -4
View File
@@ -7,9 +7,10 @@ import numpy as np
import torch
import torch.nn.functional as F
from colorsys import rgb_to_hsv
import comfy.model_management
from blend_modes import difference, normal, screen, soft_light, lighten_only, dodge, \
addition, darken_only, multiply, hard_light, \
grain_extract, grain_merge, divide, overlay
addition, darken_only, multiply, hard_light, \
grain_extract, grain_merge, divide, overlay
def dissolve(backdrop, source, opacity):
# Normalize the RGB and alpha values to 0-1
@@ -56,7 +57,7 @@ def rgb_to_hsv_via_torch(rgb_numpy: np.ndarray, device=None) -> torch.Tensor:
The hue (H) will be in the range [0, 1], while S and V will be in the range [0, 1].
"""
if device is None:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
device = comfy.model_management.get_torch_device()
rgb = torch.from_numpy(rgb_numpy).float().permute(2, 0, 1).to(device)
r, g, b = rgb[0], rgb[1], rgb[2]
@@ -99,7 +100,7 @@ def hsv_to_rgb_via_torch(hsv_numpy: np.ndarray, device=None) -> torch.Tensor:
The RGB values will be in the range [0, 1].
"""
if device is None:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
device = comfy.model_management.get_torch_device()
hsv = torch.from_numpy(hsv_numpy).float().permute(2, 0, 1).to(device)
h, s, v = hsv[0], hsv[1], hsv[2]
+12 -8
View File
@@ -1486,7 +1486,7 @@ def create_mask_from_color_tensor(image:Image, color:str, tolerance:int=0) -> Im
def load_RMBG_model():
from .briarmbg import BriaRMBG
current_directory = os.path.dirname(os.path.abspath(__file__))
device = "cuda" if torch.cuda.is_available() else "cpu"
device = comfy.model_management.get_torch_device()
net = BriaRMBG()
model_path = ""
try:
@@ -1510,8 +1510,8 @@ def RMBG(image:Image) -> Image:
im_tensor = torch.tensor(im_np, dtype=torch.float32).permute(2, 0, 1)
im_tensor = torch.divide(torch.unsqueeze(im_tensor, 0), 255.0)
im_tensor = TF.normalize(im_tensor, [0.5, 0.5, 0.5], [1.0, 1.0, 1.0])
if torch.cuda.is_available():
im_tensor = im_tensor.cuda()
device = rmbgmodel.get_device() if hasattr(rmbgmodel, 'get_device') else next(rmbgmodel.parameters()).device
im_tensor = im_tensor.to(device)
result = rmbgmodel(im_tensor)
result = torch.squeeze(F.interpolate(result[0][0], size=(h, w), mode='bilinear'), 0)
ma = torch.max(result)
@@ -1601,11 +1601,11 @@ def generate_VITMatte(image:Image, trimap:Image, local_files_only:bool=False, de
if device=="cpu":
device = torch.device('cpu')
else:
if torch.cuda.is_available():
if device == "cuda" and not torch.cuda.is_available():
log("vitmatte device is set to cuda, but not available, using ComfyUI default device instead.")
device = comfy.model_management.get_torch_device()
if device.type == "cpu" and torch.cuda.is_available():
device = torch.device('cuda')
else:
log("vitmatte device is set to cuda, but not available, using cpu instead.")
device = torch.device('cpu')
if method == "vitmatte-base-composition-1k":
model_name = "hustvl/vitmatte-base-composition-1k"
vit_matte_model = load_VITMatte_base_model(model_name=model_name, local_files_only=local_files_only)
@@ -1621,6 +1621,8 @@ def generate_VITMatte(image:Image, trimap:Image, local_files_only:bool=False, de
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
elif torch.accelerator.is_available():
torch.accelerator.empty_cache()
mask = tensor2pil(predictions).convert('L')
mask = mask.crop(
(0, 0, image.width, image.height)) # remove padding that the prediction appends (works in 32px tiles)
@@ -2155,6 +2157,8 @@ def clear_memory():
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
elif torch.accelerator.is_available():
torch.accelerator.empty_cache()
def tensor_info(tensor:object) -> str:
value = ''
@@ -2206,7 +2210,7 @@ class UformGen2QwenChat:
# local_files_only=False, # Set to False to allow downloading if not available locally
# local_dir_use_symlinks="auto") # or set to True/False based on your symlink preference
self.model_path = files_for_uform_gen2_qwen
self.device = "cuda" if torch.cuda.is_available() else "cpu"
self.device = comfy.model_management.get_torch_device()
self.model = AutoModel.from_pretrained(self.model_path, trust_remote_code=True).to(self.device)
self.processor = AutoProcessor.from_pretrained(self.model_path, trust_remote_code=True)