Merge pull request #610 from li-lizhe/feat/device-auto-list
Add 'auto' device option to nodes with hardcoded CUDA/CPU device lists
This commit is contained in:
@@ -59,6 +59,22 @@ except ImportError as e:
|
||||
|
||||
|
||||
|
||||
'''device selection'''
|
||||
|
||||
|
||||
# Shared device options for node UI dropdowns.
|
||||
# 'auto' follows ComfyUI's default device (CUDA/NPU/XPU/MPS/CPU depending on the runtime).
|
||||
DEVICE_LIST_OPTIONS = ['auto', 'cuda', 'cpu']
|
||||
|
||||
|
||||
def get_device(device_str: str = "auto"):
|
||||
"""Resolve a user-facing device option to a torch.device.
|
||||
'auto' returns ComfyUI's default device (handles CUDA/NPU/XPU/MPS/CPU)."""
|
||||
if device_str == "cpu":
|
||||
return torch.device("cpu")
|
||||
return comfy.model_management.get_torch_device()
|
||||
|
||||
|
||||
'''warpper'''
|
||||
|
||||
# create a wrapper function that can apply a function to multiple images in a batch while passing all other arguments to the function
|
||||
|
||||
@@ -2,6 +2,7 @@ import torch
|
||||
from PIL import Image
|
||||
from .imagefunc import log, tensor2pil, pil2tensor, image2mask, expand_mask, mask_fix
|
||||
from .imagefunc import guided_filter_alpha, histogram_remap, mask_edge_detail ,RGB2RGBA, generate_VITMatte, generate_VITMatte_trimap
|
||||
from .imagefunc import DEVICE_LIST_OPTIONS
|
||||
|
||||
|
||||
|
||||
@@ -13,7 +14,7 @@ class MaskEdgeUltraDetailV2:
|
||||
def INPUT_TYPES(cls):
|
||||
|
||||
method_list = ['VITMatte', 'VITMatte(local)', 'vitmatte-base-composition-1k', 'PyMatting', 'GuidedFilter', ]
|
||||
device_list = ['cuda','cpu']
|
||||
device_list = DEVICE_LIST_OPTIONS
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
|
||||
@@ -2,6 +2,7 @@ import torch
|
||||
from PIL import Image
|
||||
from .imagefunc import log, tensor2pil, pil2tensor, image2mask, mask2image, expand_mask, mask_fix, gaussian_blur, pixel_spread
|
||||
from .imagefunc import guided_filter_alpha, histogram_remap, mask_edge_detail ,RGB2RGBA, generate_VITMatte, generate_VITMatte_trimap
|
||||
from .imagefunc import DEVICE_LIST_OPTIONS
|
||||
|
||||
|
||||
|
||||
@@ -13,7 +14,7 @@ class MaskEdgeUltraDetailV3:
|
||||
def INPUT_TYPES(cls):
|
||||
|
||||
method_list = ['VITMatte', 'VITMatte(local)', 'vitmatte-base-composition-1k', 'PyMatting', 'GuidedFilter', ]
|
||||
device_list = ['cuda','cpu']
|
||||
device_list = DEVICE_LIST_OPTIONS
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
|
||||
+2
-1
@@ -2,6 +2,7 @@ import torch
|
||||
from PIL import Image
|
||||
from .imagefunc import log, tensor2pil, pil2tensor, image2mask, mask2image, RMBG, RGB2RGBA, mask_edge_detail
|
||||
from .imagefunc import guided_filter_alpha, histogram_remap, generate_VITMatte, generate_VITMatte_trimap
|
||||
from .imagefunc import DEVICE_LIST_OPTIONS
|
||||
|
||||
|
||||
|
||||
@@ -13,7 +14,7 @@ class RmBgUltraV2:
|
||||
def INPUT_TYPES(cls):
|
||||
|
||||
method_list = ['VITMatte', 'VITMatte(local)', 'vitmatte-base-composition-1k', 'PyMatting', 'GuidedFilter', ]
|
||||
device_list = ['cuda','cpu']
|
||||
device_list = DEVICE_LIST_OPTIONS
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
|
||||
@@ -10,6 +10,7 @@ import torch.nn as nn
|
||||
import folder_paths
|
||||
from .imagefunc import log, tensor2pil, pil2tensor, mask2image, image2mask, RGB2RGBA
|
||||
from .imagefunc import guided_filter_alpha, mask_edge_detail, histogram_remap, generate_VITMatte, generate_VITMatte_trimap
|
||||
from .imagefunc import DEVICE_LIST_OPTIONS
|
||||
|
||||
|
||||
class SegformerPipeline:
|
||||
@@ -70,7 +71,7 @@ class Segformer_B2_Clothes:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
method_list = ['VITMatte', 'VITMatte(local)', 'vitmatte-base-composition-1k', 'PyMatting', 'GuidedFilter', ]
|
||||
device_list = ['cuda', 'cpu']
|
||||
device_list = DEVICE_LIST_OPTIONS
|
||||
return {"required":
|
||||
{
|
||||
"image": ("IMAGE",),
|
||||
@@ -462,7 +463,7 @@ class SegformerUltraV2:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
method_list = ['VITMatte', 'VITMatte(local)', 'vitmatte-base-composition-1k', 'PyMatting', 'GuidedFilter', ]
|
||||
device_list = ['cuda', 'cpu']
|
||||
device_list = DEVICE_LIST_OPTIONS
|
||||
return {"required":
|
||||
{
|
||||
"image": ("IMAGE",),
|
||||
@@ -801,7 +802,7 @@ class LS_LoadSegformerModel:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
model_list = ['segformer_b3_clothes', 'segformer_b2_clothes', 'segformer_b3_fashion']
|
||||
device_list = ['cuda', 'cpu']
|
||||
device_list = DEVICE_LIST_OPTIONS
|
||||
return {"required":
|
||||
{
|
||||
"model_name": (model_list,),
|
||||
|
||||
+3
-2
@@ -5,7 +5,7 @@ import re
|
||||
from transformers import pipeline
|
||||
import folder_paths
|
||||
|
||||
from .imagefunc import log, tensor2pil
|
||||
from .imagefunc import log, tensor2pil, DEVICE_LIST_OPTIONS, get_device
|
||||
|
||||
vqa_model_path = os.path.join(folder_paths.models_dir, 'VQA')
|
||||
|
||||
@@ -34,7 +34,7 @@ class LS_LoadVQAModel:
|
||||
def INPUT_TYPES(s):
|
||||
model_list = list(vqa_model_repos.keys())
|
||||
precision_list = ["fp16", "fp32"]
|
||||
device_list = ['cuda','cpu']
|
||||
device_list = DEVICE_LIST_OPTIONS
|
||||
return {
|
||||
"required": {
|
||||
"model": (model_list,),
|
||||
@@ -49,6 +49,7 @@ class LS_LoadVQAModel:
|
||||
CATEGORY = '😺dzNodes/LayerUtility'
|
||||
|
||||
def load_vqa_model(self, model, precision, device):
|
||||
device = str(get_device(device))
|
||||
|
||||
if (model == self.model_name and precision == self.precision and device == self.device
|
||||
and self.model is not None and self.processor is not None):
|
||||
|
||||
Reference in New Issue
Block a user