diff --git a/README.md b/README.md index 4803dd8..d862931 100644 --- a/README.md +++ b/README.md @@ -88,6 +88,7 @@ Some newer ComfyUI builds capture the RAM release callback when a prompt starts. To avoid that error, this node detects that executor behavior: +- Detection covers both `execute_async()` and the newer `_execute_async()` implementation. If source inspection is unavailable, mode switching remains restricted for safety. - Newer ComfyUI builds usually start in RAM_PRESSURE/RAM cache mode. The node keeps that state, updates thresholds, and preserves per-node automatic cache release - If the executor starts in CLASSIC, the node enables the RAMPressureCache object, migrates cache data, and performs an active purge, but keeps the executor mode unchanged for the current prompt - If a workflow requests CLASSIC, the node keeps RAM_PRESSURE active on newer prompt-local callback builds to avoid making later prompts start from CLASSIC diff --git a/README_zh.md b/README_zh.md index 7e59429..b37f069 100644 --- a/README_zh.md +++ b/README_zh.md @@ -82,6 +82,7 @@ 为避免这个错误,节点会检测这种执行器实现: +- 同时检查 `execute_async()` 和新版的 `_execute_async()`;无法读取实现时,仍限制执行中切换模式,避免调用任务开始时为空的回调。 - 新版ComfyUI通常已经以RAM_PRESSURE/RAM cache模式开始,节点会保持这个状态,更新阈值,并保留每个节点执行后的自动清理行为 - 如果执行器以CLASSIC开始,节点会启用RAMPressureCache对象、迁移缓存并执行一次主动清理,但不会在当前prompt内改执行器模式 - 如果工作流请求CLASSIC,节点会在新版prompt-local回调实现中保留RAM_PRESSURE,避免把下一次prompt带入CLASSIC状态 diff --git a/nodes.py b/nodes.py index 8f1e0e8..6c28b0a 100644 --- a/nodes.py +++ b/nodes.py @@ -179,15 +179,22 @@ class DynamicRAMCacheControl: if isinstance(cache_args, dict) and 'ram_inactive' in cache_args: return True - PromptExecutor = getattr(execution, 'PromptExecutor', None) - execute_async = getattr(PromptExecutor, 'execute_async', None) - if execute_async is None: - return False + return any(source is not None and 'ram_inactive' in source + for source in self._executor_async_sources()) - try: - return 'ram_inactive' in inspect.getsource(execute_async) - except (OSError, TypeError): - return False + def _executor_async_sources(self): + PromptExecutor = getattr(execution, 'PromptExecutor', None) + sources = [] + # Newer executors keep cleanup in execute_async and execution in _execute_async. + for name in ('execute_async', '_execute_async'): + method = getattr(PromptExecutor, name, None) + if method is None: + continue + try: + sources.append(inspect.getsource(inspect.unwrap(method))) + except (OSError, TypeError, ValueError): + sources.append(None) + return sources def _can_set_executor_ram_type(self, executor): if self._is_executor_ram_type(executor): @@ -206,17 +213,12 @@ class DynamicRAMCacheControl: return ram_type is not None and getattr(executor, 'cache_type', None) == ram_type def _uses_prompt_local_ram_release_callback(self): - PromptExecutor = getattr(execution, 'PromptExecutor', None) - execute_async = getattr(PromptExecutor, 'execute_async', None) - if execute_async is None: - return True - - try: - source = inspect.getsource(execute_async) - except (OSError, TypeError): - return True - - return 'ram_release_callback' in source and 'self.cache_type == CacheType.RAM_PRESSURE' in source + sources = self._executor_async_sources() + # Unknown implementations must not enable a callback captured as None. + return not sources or any( + source is None or 'ram_release_callback' in source + for source in sources + ) def _set_executor_cache_type(self, executor, target_mode_ram): CacheType = getattr(execution, 'CacheType', None) diff --git a/tests/test_callback_compatibility.py b/tests/test_callback_compatibility.py new file mode 100644 index 0000000..b61ffd1 --- /dev/null +++ b/tests/test_callback_compatibility.py @@ -0,0 +1,179 @@ +import asyncio +import importlib.util +from pathlib import Path +import sys +import types +import unittest +from unittest.mock import patch + + +class CacheType: + CLASSIC = 0 + RAM_PRESSURE = 3 + + +class HierarchicalCache: + def __init__(self, key_class=None, **kwargs): + self.key_class = key_class + self.cache = {"retained": object()} + + +class RAMPressureCache(HierarchicalCache): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.releases = [] + + def ram_release(self, headroom, **kwargs): + self.releases.append((headroom, kwargs)) + return 0 + + +class ExecutorBase: + def __init__(self, mode): + self.cache_type = mode + self.cache_args = {"ram": 2, "ram_inactive": 1} + cache = RAMPressureCache() if mode == CacheType.RAM_PRESSURE else HierarchicalCache() + self.caches = types.SimpleNamespace(outputs=cache, all=[cache]) + + +class SplitExecutor(ExecutorBase): + async def execute_async(self, action): + await self._execute_async(action) + + async def _execute_async(self, action): + ram_inactive_headroom = self.cache_args["ram_inactive"] + ram_release_callback = self.caches.outputs.ram_release if self.cache_type == CacheType.RAM_PRESSURE else None + action() + if self.cache_type == CacheType.RAM_PRESSURE: + ram_release_callback(ram_inactive_headroom) + + +class InlineExecutor(ExecutorBase): + async def execute_async(self, action): + ram_inactive_headroom = self.cache_args["ram_inactive"] + ram_release_callback = self.caches.outputs.ram_release if self.cache_type == CacheType.RAM_PRESSURE else None + action() + if self.cache_type == CacheType.RAM_PRESSURE: + ram_release_callback(ram_inactive_headroom) + + +class LegacyExecutor(ExecutorBase): + async def execute_async(self, action): + action() + + +class CompatibilityTests(unittest.TestCase): + def setUp(self): + self.execution = types.ModuleType("execution") + self.execution.CacheType = CacheType + self.execution.PromptExecutor = SplitExecutor + cache_module = types.ModuleType("comfy_execution.caching") + cache_module.RAMPressureCache = RAMPressureCache + cache_module.HierarchicalCache = HierarchicalCache + cache_module.CacheKeySetInputSignature = object + package = types.ModuleType("comfy_execution") + package.caching = cache_module + filename = Path(__file__).resolve().parents[1] / "nodes.py" + spec = importlib.util.spec_from_file_location("dynamic_ramcache_test_nodes", filename) + self.module = importlib.util.module_from_spec(spec) + with patch.dict(sys.modules, { + "execution": self.execution, + "comfy_execution": package, + "comfy_execution.caching": cache_module, + }): + spec.loader.exec_module(self.module) + + def controller(self, executor, extreme=False): + self.execution.PromptExecutor = type(executor) + cls = self.module.RAMCacheExtremeCleanup if extreme else self.module.DynamicRAMCacheControl + node = cls() + node._find_executor = lambda: executor + return node + + def test_prompt_local_modes_across_two_prompts(self): + for executor_class in (SplitExecutor, InlineExecutor): + for start_mode in (CacheType.CLASSIC, CacheType.RAM_PRESSURE): + with self.subTest(executor=executor_class.__name__, mode=start_mode): + executor = executor_class(start_mode) + node = self.controller(executor) + original_cache = executor.caches.outputs + retained = original_cache.cache["retained"] + passthrough = object() + + def action(): + for mode in ("RAM_PRESSURE (Auto Purge)", "CLASSIC (No Eviction)", "RAM_PRESSURE (Auto Purge)"): + result = node.manage_cache(mode, 2, 1, passthrough) + self.assertIs(result[0], passthrough) + self.assertEqual(executor.cache_type, start_mode) + self.assertIs(executor.caches.outputs.cache["retained"], retained) + self.assertIs(executor.caches.all[0], executor.caches.outputs) + + for _ in range(2): + asyncio.run(executor.execute_async(action)) + self.assertTrue(executor.caches.outputs.releases) + if start_mode == CacheType.RAM_PRESSURE: + self.assertIs(executor.caches.outputs, original_cache) + + def test_legacy_executor_still_switches_both_modes(self): + executor = LegacyExecutor(CacheType.CLASSIC) + node = self.controller(executor) + node.manage_cache("RAM_PRESSURE (Auto Purge)", 2, any_input=1) + self.assertEqual(executor.cache_type, CacheType.RAM_PRESSURE) + node.manage_cache("CLASSIC (No Eviction)", 2, any_input=1) + self.assertEqual(executor.cache_type, CacheType.CLASSIC) + self.assertIsInstance(executor.caches.outputs, HierarchicalCache) + self.assertNotIsInstance(executor.caches.outputs, RAMPressureCache) + + def test_classic_prompt_can_request_ram_without_none_callback(self): + executor = SplitExecutor(CacheType.CLASSIC) + node = self.controller(executor) + asyncio.run(executor.execute_async( + lambda: node.manage_cache("RAM_PRESSURE (Auto Purge)", 2, any_input=1) + )) + self.assertEqual(executor.cache_type, CacheType.CLASSIC) + + def test_split_method_detects_inactive_arg_without_existing_key(self): + executor = SplitExecutor(CacheType.CLASSIC) + executor.cache_args.pop("ram_inactive") + node = self.controller(executor) + self.assertTrue(node._supports_inactive_cache_arg(executor)) + node.manage_cache("RAM_PRESSURE (Auto Purge)", 2, 3, any_input=1) + self.assertEqual(executor.cache_args["ram_inactive"], 3) + + def test_unknown_source_keeps_mode_unchanged(self): + executor = SplitExecutor(CacheType.CLASSIC) + node = self.controller(executor) + for error in (OSError, TypeError, ValueError): + with self.subTest(error=error), patch.object(self.module.inspect, "getsource", side_effect=error): + self.assertTrue(node._uses_prompt_local_ram_release_callback()) + node.manage_cache("RAM_PRESSURE (Auto Purge)", 2, any_input=1) + self.assertEqual(executor.cache_type, CacheType.CLASSIC) + + def test_absent_async_methods_are_conservative(self): + self.execution.PromptExecutor = ExecutorBase + node = self.module.DynamicRAMCacheControl() + self.assertTrue(node._uses_prompt_local_ram_release_callback()) + + def test_unreadable_inner_method_remains_conservative(self): + executor = SplitExecutor(CacheType.CLASSIC) + node = self.controller(executor) + with patch.object(self.module.inspect, "getsource", side_effect=[ + "async def execute_async(self): await self._execute_async()", + OSError("source unavailable"), + ]): + self.assertTrue(node._uses_prompt_local_ram_release_callback()) + + def test_extreme_cleanup_restores_thresholds_without_invalid_callback(self): + for mode in (CacheType.CLASSIC, CacheType.RAM_PRESSURE): + with self.subTest(mode=mode): + executor = SplitExecutor(mode) + original_args = dict(executor.cache_args) + node = self.controller(executor, extreme=True) + asyncio.run(executor.execute_async(lambda: node.extreme_cleanup(256, any_input=1))) + self.assertEqual(executor.cache_type, mode) + self.assertEqual(executor.cache_args["ram"], original_args["ram"]) + self.assertEqual(executor.cache_args["ram_inactive"], original_args["ram_inactive"]) + + +if __name__ == "__main__": + unittest.main()