From eefbcbd9a220c8a2f48821f633b56f67025fa0c1 Mon Sep 17 00:00:00 2001 From: Ethanfel Date: Sat, 15 Aug 2026 21:56:59 +0200 Subject: [PATCH] feat: expose timing on all interpolation nodes --- nodes.py | 88 +++++++++++++++++++++++------------- tests/test_timing_outputs.py | 69 ++++++++++++++++++++++++++++ 2 files changed, 126 insertions(+), 31 deletions(-) create mode 100644 tests/test_timing_outputs.py diff --git a/nodes.py b/nodes.py index b65550f..e28cd76 100644 --- a/nodes.py +++ b/nodes.py @@ -1,6 +1,7 @@ import math import os import glob +from functools import wraps import logging import re import shutil @@ -28,6 +29,25 @@ from .gimm_vfi_arch import clear_gimm_caches logger = logging.getLogger("Tween") +def _with_elapsed_seconds(label=None): + """Append wall-clock node execution time without disturbing existing outputs.""" + def decorate(method): + @wraps(method) + def timed(self, *args, **kwargs): + started = time.perf_counter() + outputs = method(self, *args, **kwargs) + elapsed_seconds = time.perf_counter() - started + output_label = label or getattr(self, "MODEL_LABEL", self.__class__.__name__) + logger.info( + "%s: node completed in %.2f seconds", output_label, elapsed_seconds + ) + return (*outputs, round(elapsed_seconds, 3)) + + return timed + + return decorate + + def _get_torch_device(): """Honor ComfyUI's selected device instead of assuming the default CUDA GPU.""" if model_management is not None: @@ -356,8 +376,8 @@ class BIMVFIInterpolate: }, } - RETURN_TYPES = ("IMAGE", "IMAGE") - RETURN_NAMES = ("images", "oversampled") + RETURN_TYPES = ("IMAGE", "IMAGE", "FLOAT") + RETURN_NAMES = ("images", "oversampled", "elapsed_seconds") FUNCTION = "interpolate" CATEGORY = "video/BIM-VFI" @@ -420,6 +440,7 @@ class BIMVFIInterpolate: n = 2 * n - 1 return total + @_with_elapsed_seconds() def interpolate(self, images, model, multiplier, clear_cache_after_n_frames, keep_device, all_on_gpu, batch_size, chunk_size, source_fps=0.0, target_fps=0.0, settings=None, seed=None): @@ -554,11 +575,12 @@ class BIMVFISegmentInterpolate(BIMVFIInterpolate): }) return base - RETURN_TYPES = ("IMAGE", "BIM_VFI_MODEL") - RETURN_NAMES = ("images", "model") + RETURN_TYPES = ("IMAGE", "BIM_VFI_MODEL", "FLOAT") + RETURN_NAMES = ("images", "model", "elapsed_seconds") FUNCTION = "interpolate" CATEGORY = "video/BIM-VFI" + @_with_elapsed_seconds() def interpolate(self, images, model, multiplier, clear_cache_after_n_frames, keep_device, all_on_gpu, batch_size, chunk_size, segment_index, segment_size, @@ -653,7 +675,7 @@ class BIMVFISegmentInterpolate(BIMVFIInterpolate): # Standard multiplier mode is_continuation = segment_index > 0 - (result, _) = super().interpolate( + (result, _, _) = super().interpolate( segment_images, model, multiplier, clear_cache_after_n_frames, keep_device, all_on_gpu, batch_size, chunk_size, seed=seed, @@ -1123,8 +1145,8 @@ class EMAVFIInterpolate: }, } - RETURN_TYPES = ("IMAGE", "IMAGE") - RETURN_NAMES = ("images", "oversampled") + RETURN_TYPES = ("IMAGE", "IMAGE", "FLOAT") + RETURN_NAMES = ("images", "oversampled", "elapsed_seconds") FUNCTION = "interpolate" CATEGORY = "video/EMA-VFI" @@ -1182,6 +1204,7 @@ class EMAVFIInterpolate: n = 2 * n - 1 return total + @_with_elapsed_seconds("EMA-VFI") def interpolate(self, images, model, multiplier, clear_cache_after_n_frames, keep_device, all_on_gpu, batch_size, chunk_size, source_fps=0.0, target_fps=0.0, settings=None): @@ -1310,11 +1333,12 @@ class EMAVFISegmentInterpolate(EMAVFIInterpolate): }) return base - RETURN_TYPES = ("IMAGE", "EMA_VFI_MODEL") - RETURN_NAMES = ("images", "model") + RETURN_TYPES = ("IMAGE", "EMA_VFI_MODEL", "FLOAT") + RETURN_NAMES = ("images", "model", "elapsed_seconds") FUNCTION = "interpolate" CATEGORY = "video/EMA-VFI" + @_with_elapsed_seconds("EMA-VFI segment") def interpolate(self, images, model, multiplier, clear_cache_after_n_frames, keep_device, all_on_gpu, batch_size, chunk_size, segment_index, segment_size, @@ -1401,7 +1425,7 @@ class EMAVFISegmentInterpolate(EMAVFIInterpolate): # Standard multiplier mode is_continuation = segment_index > 0 - (result, _) = super().interpolate( + (result, _, _) = super().interpolate( segment_images, model, multiplier, clear_cache_after_n_frames, keep_device, all_on_gpu, batch_size, chunk_size, ) @@ -1550,8 +1574,8 @@ class SGMVFIInterpolate: }, } - RETURN_TYPES = ("IMAGE", "IMAGE") - RETURN_NAMES = ("images", "oversampled") + RETURN_TYPES = ("IMAGE", "IMAGE", "FLOAT") + RETURN_NAMES = ("images", "oversampled", "elapsed_seconds") FUNCTION = "interpolate" CATEGORY = "video/SGM-VFI" @@ -1609,6 +1633,7 @@ class SGMVFIInterpolate: n = 2 * n - 1 return total + @_with_elapsed_seconds("SGM-VFI") def interpolate(self, images, model, multiplier, clear_cache_after_n_frames, keep_device, all_on_gpu, batch_size, chunk_size, source_fps=0.0, target_fps=0.0, settings=None): @@ -1737,11 +1762,12 @@ class SGMVFISegmentInterpolate(SGMVFIInterpolate): }) return base - RETURN_TYPES = ("IMAGE", "SGM_VFI_MODEL") - RETURN_NAMES = ("images", "model") + RETURN_TYPES = ("IMAGE", "SGM_VFI_MODEL", "FLOAT") + RETURN_NAMES = ("images", "model", "elapsed_seconds") FUNCTION = "interpolate" CATEGORY = "video/SGM-VFI" + @_with_elapsed_seconds("SGM-VFI segment") def interpolate(self, images, model, multiplier, clear_cache_after_n_frames, keep_device, all_on_gpu, batch_size, chunk_size, segment_index, segment_size, @@ -1828,7 +1854,7 @@ class SGMVFISegmentInterpolate(SGMVFIInterpolate): # Standard multiplier mode is_continuation = segment_index > 0 - (result, _) = super().interpolate( + (result, _, _) = super().interpolate( segment_images, model, multiplier, clear_cache_after_n_frames, keep_device, all_on_gpu, batch_size, chunk_size, ) @@ -1984,11 +2010,12 @@ class LDFVFIInterpolate: FUNCTION = "interpolate" CATEGORY = "video/LDF-VFI" + @_with_elapsed_seconds("LDF-VFI") def interpolate(self, images, model, temporal_factor, sampling_steps, t_shift, t_cond, seed, offload_after, source_fps=0.0, target_fps=0.0): if images.shape[0] < 2: - return (images, images, 0.0) + return (images, images) device = _get_torch_device() if device.type != "cuda": raise RuntimeError( @@ -2004,7 +2031,7 @@ class LDFVFIInterpolate: selected = _select_target_fps_frames( source, source_fps, target_fps, 1, source.shape[0] ).permute(0, 2, 3, 1).cpu() - return (selected, images, 0.0) + return (selected, images) temporal_factor = math.ceil(ratio) if temporal_factor > 16: raise ValueError( @@ -2027,7 +2054,6 @@ class LDFVFIInterpolate: ) try: model.to(device) - interpolation_started = time.perf_counter() generated = model.interpolate_sequence( source, temporal_factor=temporal_factor, @@ -2037,7 +2063,6 @@ class LDFVFIInterpolate: seed=seed, progress_callback=update_progress, ) - elapsed_seconds = time.perf_counter() - interpolation_started finally: generation_failed = sys.exc_info()[0] is not None cleanup_error = None @@ -2063,11 +2088,8 @@ class LDFVFIInterpolate: temporal_factor, source.shape[0], ) result = generated.permute(0, 2, 3, 1).cpu() - logger.info( - "LDF-VFI: done, %s output frames in %.2f seconds", - result.shape[0], elapsed_seconds, - ) - return (result, generated_sequence, round(elapsed_seconds, 3)) + logger.info("LDF-VFI: done, %s output frames", result.shape[0]) + return (result, generated_sequence) # --------------------------------------------------------------------------- @@ -2145,6 +2167,8 @@ class LoadSPEEDVFIModel: class SPEEDVFIInterpolate(BIMVFIInterpolate): MODEL_LABEL = "SPEED" CATEGORY = "video/SPEED" + RETURN_TYPES = ("IMAGE", "IMAGE", "FLOAT") + RETURN_NAMES = ("images", "oversampled", "elapsed_seconds") @classmethod def INPUT_TYPES(cls): @@ -2172,10 +2196,10 @@ class SPEEDVFIInterpolate(BIMVFIInterpolate): ) return inputs - class SPEEDVFISegmentInterpolate(BIMVFISegmentInterpolate): MODEL_LABEL = "SPEED" - RETURN_TYPES = ("IMAGE", "SPEED_VFI_MODEL") + RETURN_TYPES = ("IMAGE", "SPEED_VFI_MODEL", "FLOAT") + RETURN_NAMES = ("images", "model", "elapsed_seconds") CATEGORY = "video/SPEED" @classmethod @@ -2347,8 +2371,8 @@ class GIMMVFIInterpolate: }, } - RETURN_TYPES = ("IMAGE", "IMAGE") - RETURN_NAMES = ("images", "oversampled") + RETURN_TYPES = ("IMAGE", "IMAGE", "FLOAT") + RETURN_NAMES = ("images", "oversampled", "elapsed_seconds") FUNCTION = "interpolate" CATEGORY = "video/GIMM-VFI" @@ -2447,6 +2471,7 @@ class GIMMVFIInterpolate: n = 2 * n - 1 return total + @_with_elapsed_seconds("GIMM-VFI") def interpolate(self, images, model, multiplier, single_pass, clear_cache_after_n_frames, keep_device, all_on_gpu, batch_size, chunk_size, @@ -2594,11 +2619,12 @@ class GIMMVFISegmentInterpolate(GIMMVFIInterpolate): }) return base - RETURN_TYPES = ("IMAGE", "GIMM_VFI_MODEL") - RETURN_NAMES = ("images", "model") + RETURN_TYPES = ("IMAGE", "GIMM_VFI_MODEL", "FLOAT") + RETURN_NAMES = ("images", "model", "elapsed_seconds") FUNCTION = "interpolate" CATEGORY = "video/GIMM-VFI" + @_with_elapsed_seconds("GIMM-VFI segment") def interpolate(self, images, model, multiplier, single_pass, clear_cache_after_n_frames, keep_device, all_on_gpu, batch_size, chunk_size, segment_index, segment_size, @@ -2695,7 +2721,7 @@ class GIMMVFISegmentInterpolate(GIMMVFIInterpolate): # Standard multiplier mode is_continuation = segment_index > 0 - (result, _) = super().interpolate( + (result, _, _) = super().interpolate( segment_images, model, multiplier, single_pass, clear_cache_after_n_frames, keep_device, all_on_gpu, batch_size, chunk_size, diff --git a/tests/test_timing_outputs.py b/tests/test_timing_outputs.py new file mode 100644 index 0000000..376b6f5 --- /dev/null +++ b/tests/test_timing_outputs.py @@ -0,0 +1,69 @@ +import ast +from pathlib import Path + + +NODE_SOURCE = Path(__file__).resolve().parents[1] / "nodes.py" +TIMED_CLASSES = { + "BIMVFIInterpolate", + "BIMVFISegmentInterpolate", + "EMAVFIInterpolate", + "EMAVFISegmentInterpolate", + "SGMVFIInterpolate", + "SGMVFISegmentInterpolate", + "LDFVFIInterpolate", + "SPEEDVFIInterpolate", + "SPEEDVFISegmentInterpolate", + "GIMMVFIInterpolate", + "GIMMVFISegmentInterpolate", +} +DIRECTLY_DECORATED = TIMED_CLASSES - { + "SPEEDVFIInterpolate", + "SPEEDVFISegmentInterpolate", +} + + +def _class_definitions(): + tree = ast.parse(NODE_SOURCE.read_text(encoding="utf-8")) + return { + node.name: node + for node in tree.body + if isinstance(node, ast.ClassDef) and node.name in TIMED_CLASSES + } + + +def _literal_assignment(class_node, name): + for statement in class_node.body: + if ( + isinstance(statement, ast.Assign) + and len(statement.targets) == 1 + and isinstance(statement.targets[0], ast.Name) + and statement.targets[0].id == name + ): + return ast.literal_eval(statement.value) + raise AssertionError(f"{class_node.name} does not define {name}") + + +def test_all_interpolation_nodes_expose_elapsed_seconds_last(): + classes = _class_definitions() + assert classes.keys() == TIMED_CLASSES + + for class_node in classes.values(): + assert _literal_assignment(class_node, "RETURN_TYPES")[-1] == "FLOAT" + assert _literal_assignment(class_node, "RETURN_NAMES")[-1] == "elapsed_seconds" + + +def test_direct_interpolation_methods_append_timing_output(): + classes = _class_definitions() + for class_name in DIRECTLY_DECORATED: + interpolate = next( + statement + for statement in classes[class_name].body + if isinstance(statement, ast.FunctionDef) + and statement.name == "interpolate" + ) + assert any( + isinstance(decorator, ast.Call) + and isinstance(decorator.func, ast.Name) + and decorator.func.id == "_with_elapsed_seconds" + for decorator in interpolate.decorator_list + )