diff --git a/py/imagefunc.py b/py/imagefunc.py index 7b060c7..8449645 100644 --- a/py/imagefunc.py +++ b/py/imagefunc.py @@ -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 diff --git a/py/mask_edge_ultra_detail_v2.py b/py/mask_edge_ultra_detail_v2.py index f78bb38..b1a2cf2 100644 --- a/py/mask_edge_ultra_detail_v2.py +++ b/py/mask_edge_ultra_detail_v2.py @@ -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",), diff --git a/py/mask_edge_ultra_detail_v3.py b/py/mask_edge_ultra_detail_v3.py index 82747ab..e1b2a8b 100644 --- a/py/mask_edge_ultra_detail_v3.py +++ b/py/mask_edge_ultra_detail_v3.py @@ -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",), diff --git a/py/rmbg_ultra_v2.py b/py/rmbg_ultra_v2.py index 2b41f02..433e8ba 100644 --- a/py/rmbg_ultra_v2.py +++ b/py/rmbg_ultra_v2.py @@ -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",), diff --git a/py/segformer_ultra.py b/py/segformer_ultra.py index 2740225..cfee5d4 100644 --- a/py/segformer_ultra.py +++ b/py/segformer_ultra.py @@ -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,), diff --git a/py/vqa_prompt.py b/py/vqa_prompt.py index 6d8cce1..bf00eb3 100644 --- a/py/vqa_prompt.py +++ b/py/vqa_prompt.py @@ -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):