Refine scoped ChatterBox torch load wrapper
This commit is contained in:
+30
-20
@@ -232,29 +232,39 @@ def load_vc_model(device: str) -> ChatterboxVC:
|
||||
# torch.load, and forces map_location onto ComfyUI core and every other pack
|
||||
# that never opted in. Scoping keeps the device-defaulting only for this
|
||||
# pack's own Chatterbox model loads.
|
||||
original_torch_load = torch.load
|
||||
def patched_torch_load(*args, **kwargs):
|
||||
if 'map_location' not in kwargs:
|
||||
# Determine the appropriate device (MPS for Mac, else CPU)
|
||||
if torch.backends.mps.is_available():
|
||||
device = "mps"
|
||||
elif torch.cuda.is_available():
|
||||
original_torch_load = torch.load
|
||||
|
||||
|
||||
def _torch_load_with_default_map_location(load_func, *args, **kwargs):
|
||||
if 'map_location' not in kwargs:
|
||||
# Determine the appropriate device (MPS for Mac, else CPU)
|
||||
if torch.backends.mps.is_available():
|
||||
device = "mps"
|
||||
elif torch.cuda.is_available():
|
||||
device = "cuda"
|
||||
else:
|
||||
device = "cpu"
|
||||
kwargs['map_location'] = torch.device(device)
|
||||
return original_torch_load(*args, **kwargs)
|
||||
else:
|
||||
device = "cpu"
|
||||
kwargs['map_location'] = torch.device(device)
|
||||
return load_func(*args, **kwargs)
|
||||
|
||||
|
||||
def patched_torch_load(*args, **kwargs):
|
||||
return _torch_load_with_default_map_location(original_torch_load, *args, **kwargs)
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def default_map_location():
|
||||
"""Temporarily install patched_torch_load for this pack's model loads."""
|
||||
previous = torch.load
|
||||
torch.load = patched_torch_load
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
torch.load = previous
|
||||
def default_map_location():
|
||||
"""Temporarily install patched_torch_load for this pack's model loads."""
|
||||
previous = torch.load
|
||||
|
||||
def scoped_torch_load(*args, **kwargs):
|
||||
return _torch_load_with_default_map_location(previous, *args, **kwargs)
|
||||
|
||||
torch.load = scoped_torch_load
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
torch.load = previous
|
||||
|
||||
|
||||
class AudioNodeBase:
|
||||
@@ -1073,4 +1083,4 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FL_ChatterboxTurboTTS": "FL Chatterbox Turbo TTS",
|
||||
"FL_ChatterboxMultilingualTTS": "FL Chatterbox Multilingual TTS",
|
||||
"FL_ChatterboxVC": "FL Chatterbox VC",
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user