From d22feacb052028ea074a3d753903ed4d1b291a1b Mon Sep 17 00:00:00 2001 From: hben35096 <139383150+hben35096@users.noreply.github.com> Date: Fri, 9 May 2025 02:26:43 +0800 Subject: [PATCH] Added support for Moore Threads GPUs --- comfy/model_management.py | 37 +++++++++++++++++++++++++++++++++++++ 1 file changed, 37 insertions(+) diff --git a/comfy/model_management.py b/comfy/model_management.py index 44aff3762..f92820375 100644 --- a/comfy/model_management.py +++ b/comfy/model_management.py @@ -128,6 +128,15 @@ try: except: mlu_available = False +try: + import torch_musa + _ = torch_musa.device_count() + musa_available = torch_musa.is_available() + if musa_available: + logging.info("MUSA device detected: {}".format(torch_musa.get_device_name(0))) +except: + musa_available = False + if args.cpu: cpu_state = CPUState.CPU @@ -151,6 +160,12 @@ def is_mlu(): return True return False +def is_musa(): + global musa_available + if musa_available: + return True + return False + def get_torch_device(): global directml_enabled global cpu_state @@ -168,6 +183,8 @@ def get_torch_device(): return torch.device("npu", torch.npu.current_device()) elif is_mlu(): return torch.device("mlu", torch.mlu.current_device()) + elif is_musa(): + return torch.device('musa', torch.musa.current_device()) else: return torch.device(torch.cuda.current_device()) @@ -200,6 +217,12 @@ def get_total_memory(dev=None, torch_total_too=False): _, mem_total_mlu = torch.mlu.mem_get_info(dev) mem_total_torch = mem_reserved mem_total = mem_total_mlu + elif is_musa(): + stats = torch.musa.memory_stats(dev) + mem_reserved = stats['reserved_bytes.all.current'] + _, mem_total = torch.musa.mem_get_info(dev) + mem_total_torch = mem_reserved + else: stats = torch.cuda.memory_stats(dev) mem_reserved = stats['reserved_bytes.all.current'] @@ -1099,6 +1122,14 @@ def get_free_memory(dev=None, torch_free_too=False): 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 + elif is_musa(): + stats = torch.musa.memory_stats(dev) + mem_active = stats['active_bytes.all.current'] + mem_reserved = stats['reserved_bytes.all.current'] + mem_free_musa, _ = torch.musa.mem_get_info(dev) + mem_free_torch = mem_reserved - mem_active + mem_free_total = mem_free_musa + mem_free_torch + else: stats = torch.cuda.memory_stats(dev) mem_active = stats['active_bytes.all.current'] @@ -1171,6 +1202,9 @@ def should_use_fp16(device=None, model_params=0, prioritize_performance=True, ma if is_mlu(): return True + if is_musa(): + return True + if torch.version.hip: return True @@ -1231,6 +1265,9 @@ def should_use_bf16(device=None, model_params=0, prioritize_performance=True, ma if is_ascend_npu(): return True + if is_musa(): + return True + if is_amd(): arch = torch.cuda.get_device_properties(device).gcnArchName if any((a in arch) for a in ["gfx1030", "gfx1031", "gfx1010", "gfx1011", "gfx1012", "gfx906", "gfx900", "gfx803"]): # RDNA2 and older don't support bf16