From 596850bc61665e9318914841b41ee4154253f020 Mon Sep 17 00:00:00 2001 From: filliptm Date: Mon, 22 Jun 2026 14:24:44 -0500 Subject: [PATCH] Refine scoped ChatterBox torch load wrapper --- chatterbox_node.py | 50 +++++++++++++++++++++++++++------------------- 1 file changed, 30 insertions(+), 20 deletions(-) diff --git a/chatterbox_node.py b/chatterbox_node.py index 0473331..bf09bae 100644 --- a/chatterbox_node.py +++ b/chatterbox_node.py @@ -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", -} \ No newline at end of file +}