feat: expose timing on all interpolation nodes
This commit is contained in:
@@ -1,6 +1,7 @@
|
|||||||
import math
|
import math
|
||||||
import os
|
import os
|
||||||
import glob
|
import glob
|
||||||
|
from functools import wraps
|
||||||
import logging
|
import logging
|
||||||
import re
|
import re
|
||||||
import shutil
|
import shutil
|
||||||
@@ -28,6 +29,25 @@ from .gimm_vfi_arch import clear_gimm_caches
|
|||||||
logger = logging.getLogger("Tween")
|
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():
|
def _get_torch_device():
|
||||||
"""Honor ComfyUI's selected device instead of assuming the default CUDA GPU."""
|
"""Honor ComfyUI's selected device instead of assuming the default CUDA GPU."""
|
||||||
if model_management is not None:
|
if model_management is not None:
|
||||||
@@ -356,8 +376,8 @@ class BIMVFIInterpolate:
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("IMAGE", "IMAGE")
|
RETURN_TYPES = ("IMAGE", "IMAGE", "FLOAT")
|
||||||
RETURN_NAMES = ("images", "oversampled")
|
RETURN_NAMES = ("images", "oversampled", "elapsed_seconds")
|
||||||
FUNCTION = "interpolate"
|
FUNCTION = "interpolate"
|
||||||
CATEGORY = "video/BIM-VFI"
|
CATEGORY = "video/BIM-VFI"
|
||||||
|
|
||||||
@@ -420,6 +440,7 @@ class BIMVFIInterpolate:
|
|||||||
n = 2 * n - 1
|
n = 2 * n - 1
|
||||||
return total
|
return total
|
||||||
|
|
||||||
|
@_with_elapsed_seconds()
|
||||||
def interpolate(self, images, model, multiplier, clear_cache_after_n_frames,
|
def interpolate(self, images, model, multiplier, clear_cache_after_n_frames,
|
||||||
keep_device, all_on_gpu, batch_size, chunk_size,
|
keep_device, all_on_gpu, batch_size, chunk_size,
|
||||||
source_fps=0.0, target_fps=0.0, settings=None, seed=None):
|
source_fps=0.0, target_fps=0.0, settings=None, seed=None):
|
||||||
@@ -554,11 +575,12 @@ class BIMVFISegmentInterpolate(BIMVFIInterpolate):
|
|||||||
})
|
})
|
||||||
return base
|
return base
|
||||||
|
|
||||||
RETURN_TYPES = ("IMAGE", "BIM_VFI_MODEL")
|
RETURN_TYPES = ("IMAGE", "BIM_VFI_MODEL", "FLOAT")
|
||||||
RETURN_NAMES = ("images", "model")
|
RETURN_NAMES = ("images", "model", "elapsed_seconds")
|
||||||
FUNCTION = "interpolate"
|
FUNCTION = "interpolate"
|
||||||
CATEGORY = "video/BIM-VFI"
|
CATEGORY = "video/BIM-VFI"
|
||||||
|
|
||||||
|
@_with_elapsed_seconds()
|
||||||
def interpolate(self, images, model, multiplier, clear_cache_after_n_frames,
|
def interpolate(self, images, model, multiplier, clear_cache_after_n_frames,
|
||||||
keep_device, all_on_gpu, batch_size, chunk_size,
|
keep_device, all_on_gpu, batch_size, chunk_size,
|
||||||
segment_index, segment_size,
|
segment_index, segment_size,
|
||||||
@@ -653,7 +675,7 @@ class BIMVFISegmentInterpolate(BIMVFIInterpolate):
|
|||||||
|
|
||||||
# Standard multiplier mode
|
# Standard multiplier mode
|
||||||
is_continuation = segment_index > 0
|
is_continuation = segment_index > 0
|
||||||
(result, _) = super().interpolate(
|
(result, _, _) = super().interpolate(
|
||||||
segment_images, model, multiplier, clear_cache_after_n_frames,
|
segment_images, model, multiplier, clear_cache_after_n_frames,
|
||||||
keep_device, all_on_gpu, batch_size, chunk_size,
|
keep_device, all_on_gpu, batch_size, chunk_size,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
@@ -1123,8 +1145,8 @@ class EMAVFIInterpolate:
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("IMAGE", "IMAGE")
|
RETURN_TYPES = ("IMAGE", "IMAGE", "FLOAT")
|
||||||
RETURN_NAMES = ("images", "oversampled")
|
RETURN_NAMES = ("images", "oversampled", "elapsed_seconds")
|
||||||
FUNCTION = "interpolate"
|
FUNCTION = "interpolate"
|
||||||
CATEGORY = "video/EMA-VFI"
|
CATEGORY = "video/EMA-VFI"
|
||||||
|
|
||||||
@@ -1182,6 +1204,7 @@ class EMAVFIInterpolate:
|
|||||||
n = 2 * n - 1
|
n = 2 * n - 1
|
||||||
return total
|
return total
|
||||||
|
|
||||||
|
@_with_elapsed_seconds("EMA-VFI")
|
||||||
def interpolate(self, images, model, multiplier, clear_cache_after_n_frames,
|
def interpolate(self, images, model, multiplier, clear_cache_after_n_frames,
|
||||||
keep_device, all_on_gpu, batch_size, chunk_size,
|
keep_device, all_on_gpu, batch_size, chunk_size,
|
||||||
source_fps=0.0, target_fps=0.0, settings=None):
|
source_fps=0.0, target_fps=0.0, settings=None):
|
||||||
@@ -1310,11 +1333,12 @@ class EMAVFISegmentInterpolate(EMAVFIInterpolate):
|
|||||||
})
|
})
|
||||||
return base
|
return base
|
||||||
|
|
||||||
RETURN_TYPES = ("IMAGE", "EMA_VFI_MODEL")
|
RETURN_TYPES = ("IMAGE", "EMA_VFI_MODEL", "FLOAT")
|
||||||
RETURN_NAMES = ("images", "model")
|
RETURN_NAMES = ("images", "model", "elapsed_seconds")
|
||||||
FUNCTION = "interpolate"
|
FUNCTION = "interpolate"
|
||||||
CATEGORY = "video/EMA-VFI"
|
CATEGORY = "video/EMA-VFI"
|
||||||
|
|
||||||
|
@_with_elapsed_seconds("EMA-VFI segment")
|
||||||
def interpolate(self, images, model, multiplier, clear_cache_after_n_frames,
|
def interpolate(self, images, model, multiplier, clear_cache_after_n_frames,
|
||||||
keep_device, all_on_gpu, batch_size, chunk_size,
|
keep_device, all_on_gpu, batch_size, chunk_size,
|
||||||
segment_index, segment_size,
|
segment_index, segment_size,
|
||||||
@@ -1401,7 +1425,7 @@ class EMAVFISegmentInterpolate(EMAVFIInterpolate):
|
|||||||
|
|
||||||
# Standard multiplier mode
|
# Standard multiplier mode
|
||||||
is_continuation = segment_index > 0
|
is_continuation = segment_index > 0
|
||||||
(result, _) = super().interpolate(
|
(result, _, _) = super().interpolate(
|
||||||
segment_images, model, multiplier, clear_cache_after_n_frames,
|
segment_images, model, multiplier, clear_cache_after_n_frames,
|
||||||
keep_device, all_on_gpu, batch_size, chunk_size,
|
keep_device, all_on_gpu, batch_size, chunk_size,
|
||||||
)
|
)
|
||||||
@@ -1550,8 +1574,8 @@ class SGMVFIInterpolate:
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("IMAGE", "IMAGE")
|
RETURN_TYPES = ("IMAGE", "IMAGE", "FLOAT")
|
||||||
RETURN_NAMES = ("images", "oversampled")
|
RETURN_NAMES = ("images", "oversampled", "elapsed_seconds")
|
||||||
FUNCTION = "interpolate"
|
FUNCTION = "interpolate"
|
||||||
CATEGORY = "video/SGM-VFI"
|
CATEGORY = "video/SGM-VFI"
|
||||||
|
|
||||||
@@ -1609,6 +1633,7 @@ class SGMVFIInterpolate:
|
|||||||
n = 2 * n - 1
|
n = 2 * n - 1
|
||||||
return total
|
return total
|
||||||
|
|
||||||
|
@_with_elapsed_seconds("SGM-VFI")
|
||||||
def interpolate(self, images, model, multiplier, clear_cache_after_n_frames,
|
def interpolate(self, images, model, multiplier, clear_cache_after_n_frames,
|
||||||
keep_device, all_on_gpu, batch_size, chunk_size,
|
keep_device, all_on_gpu, batch_size, chunk_size,
|
||||||
source_fps=0.0, target_fps=0.0, settings=None):
|
source_fps=0.0, target_fps=0.0, settings=None):
|
||||||
@@ -1737,11 +1762,12 @@ class SGMVFISegmentInterpolate(SGMVFIInterpolate):
|
|||||||
})
|
})
|
||||||
return base
|
return base
|
||||||
|
|
||||||
RETURN_TYPES = ("IMAGE", "SGM_VFI_MODEL")
|
RETURN_TYPES = ("IMAGE", "SGM_VFI_MODEL", "FLOAT")
|
||||||
RETURN_NAMES = ("images", "model")
|
RETURN_NAMES = ("images", "model", "elapsed_seconds")
|
||||||
FUNCTION = "interpolate"
|
FUNCTION = "interpolate"
|
||||||
CATEGORY = "video/SGM-VFI"
|
CATEGORY = "video/SGM-VFI"
|
||||||
|
|
||||||
|
@_with_elapsed_seconds("SGM-VFI segment")
|
||||||
def interpolate(self, images, model, multiplier, clear_cache_after_n_frames,
|
def interpolate(self, images, model, multiplier, clear_cache_after_n_frames,
|
||||||
keep_device, all_on_gpu, batch_size, chunk_size,
|
keep_device, all_on_gpu, batch_size, chunk_size,
|
||||||
segment_index, segment_size,
|
segment_index, segment_size,
|
||||||
@@ -1828,7 +1854,7 @@ class SGMVFISegmentInterpolate(SGMVFIInterpolate):
|
|||||||
|
|
||||||
# Standard multiplier mode
|
# Standard multiplier mode
|
||||||
is_continuation = segment_index > 0
|
is_continuation = segment_index > 0
|
||||||
(result, _) = super().interpolate(
|
(result, _, _) = super().interpolate(
|
||||||
segment_images, model, multiplier, clear_cache_after_n_frames,
|
segment_images, model, multiplier, clear_cache_after_n_frames,
|
||||||
keep_device, all_on_gpu, batch_size, chunk_size,
|
keep_device, all_on_gpu, batch_size, chunk_size,
|
||||||
)
|
)
|
||||||
@@ -1984,11 +2010,12 @@ class LDFVFIInterpolate:
|
|||||||
FUNCTION = "interpolate"
|
FUNCTION = "interpolate"
|
||||||
CATEGORY = "video/LDF-VFI"
|
CATEGORY = "video/LDF-VFI"
|
||||||
|
|
||||||
|
@_with_elapsed_seconds("LDF-VFI")
|
||||||
def interpolate(self, images, model, temporal_factor, sampling_steps,
|
def interpolate(self, images, model, temporal_factor, sampling_steps,
|
||||||
t_shift, t_cond, seed, offload_after,
|
t_shift, t_cond, seed, offload_after,
|
||||||
source_fps=0.0, target_fps=0.0):
|
source_fps=0.0, target_fps=0.0):
|
||||||
if images.shape[0] < 2:
|
if images.shape[0] < 2:
|
||||||
return (images, images, 0.0)
|
return (images, images)
|
||||||
device = _get_torch_device()
|
device = _get_torch_device()
|
||||||
if device.type != "cuda":
|
if device.type != "cuda":
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
@@ -2004,7 +2031,7 @@ class LDFVFIInterpolate:
|
|||||||
selected = _select_target_fps_frames(
|
selected = _select_target_fps_frames(
|
||||||
source, source_fps, target_fps, 1, source.shape[0]
|
source, source_fps, target_fps, 1, source.shape[0]
|
||||||
).permute(0, 2, 3, 1).cpu()
|
).permute(0, 2, 3, 1).cpu()
|
||||||
return (selected, images, 0.0)
|
return (selected, images)
|
||||||
temporal_factor = math.ceil(ratio)
|
temporal_factor = math.ceil(ratio)
|
||||||
if temporal_factor > 16:
|
if temporal_factor > 16:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -2027,7 +2054,6 @@ class LDFVFIInterpolate:
|
|||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
model.to(device)
|
model.to(device)
|
||||||
interpolation_started = time.perf_counter()
|
|
||||||
generated = model.interpolate_sequence(
|
generated = model.interpolate_sequence(
|
||||||
source,
|
source,
|
||||||
temporal_factor=temporal_factor,
|
temporal_factor=temporal_factor,
|
||||||
@@ -2037,7 +2063,6 @@ class LDFVFIInterpolate:
|
|||||||
seed=seed,
|
seed=seed,
|
||||||
progress_callback=update_progress,
|
progress_callback=update_progress,
|
||||||
)
|
)
|
||||||
elapsed_seconds = time.perf_counter() - interpolation_started
|
|
||||||
finally:
|
finally:
|
||||||
generation_failed = sys.exc_info()[0] is not None
|
generation_failed = sys.exc_info()[0] is not None
|
||||||
cleanup_error = None
|
cleanup_error = None
|
||||||
@@ -2063,11 +2088,8 @@ class LDFVFIInterpolate:
|
|||||||
temporal_factor, source.shape[0],
|
temporal_factor, source.shape[0],
|
||||||
)
|
)
|
||||||
result = generated.permute(0, 2, 3, 1).cpu()
|
result = generated.permute(0, 2, 3, 1).cpu()
|
||||||
logger.info(
|
logger.info("LDF-VFI: done, %s output frames", result.shape[0])
|
||||||
"LDF-VFI: done, %s output frames in %.2f seconds",
|
return (result, generated_sequence)
|
||||||
result.shape[0], elapsed_seconds,
|
|
||||||
)
|
|
||||||
return (result, generated_sequence, round(elapsed_seconds, 3))
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -2145,6 +2167,8 @@ class LoadSPEEDVFIModel:
|
|||||||
class SPEEDVFIInterpolate(BIMVFIInterpolate):
|
class SPEEDVFIInterpolate(BIMVFIInterpolate):
|
||||||
MODEL_LABEL = "SPEED"
|
MODEL_LABEL = "SPEED"
|
||||||
CATEGORY = "video/SPEED"
|
CATEGORY = "video/SPEED"
|
||||||
|
RETURN_TYPES = ("IMAGE", "IMAGE", "FLOAT")
|
||||||
|
RETURN_NAMES = ("images", "oversampled", "elapsed_seconds")
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
@@ -2172,10 +2196,10 @@ class SPEEDVFIInterpolate(BIMVFIInterpolate):
|
|||||||
)
|
)
|
||||||
return inputs
|
return inputs
|
||||||
|
|
||||||
|
|
||||||
class SPEEDVFISegmentInterpolate(BIMVFISegmentInterpolate):
|
class SPEEDVFISegmentInterpolate(BIMVFISegmentInterpolate):
|
||||||
MODEL_LABEL = "SPEED"
|
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"
|
CATEGORY = "video/SPEED"
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -2347,8 +2371,8 @@ class GIMMVFIInterpolate:
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("IMAGE", "IMAGE")
|
RETURN_TYPES = ("IMAGE", "IMAGE", "FLOAT")
|
||||||
RETURN_NAMES = ("images", "oversampled")
|
RETURN_NAMES = ("images", "oversampled", "elapsed_seconds")
|
||||||
FUNCTION = "interpolate"
|
FUNCTION = "interpolate"
|
||||||
CATEGORY = "video/GIMM-VFI"
|
CATEGORY = "video/GIMM-VFI"
|
||||||
|
|
||||||
@@ -2447,6 +2471,7 @@ class GIMMVFIInterpolate:
|
|||||||
n = 2 * n - 1
|
n = 2 * n - 1
|
||||||
return total
|
return total
|
||||||
|
|
||||||
|
@_with_elapsed_seconds("GIMM-VFI")
|
||||||
def interpolate(self, images, model, multiplier, single_pass,
|
def interpolate(self, images, model, multiplier, single_pass,
|
||||||
clear_cache_after_n_frames, keep_device, all_on_gpu,
|
clear_cache_after_n_frames, keep_device, all_on_gpu,
|
||||||
batch_size, chunk_size,
|
batch_size, chunk_size,
|
||||||
@@ -2594,11 +2619,12 @@ class GIMMVFISegmentInterpolate(GIMMVFIInterpolate):
|
|||||||
})
|
})
|
||||||
return base
|
return base
|
||||||
|
|
||||||
RETURN_TYPES = ("IMAGE", "GIMM_VFI_MODEL")
|
RETURN_TYPES = ("IMAGE", "GIMM_VFI_MODEL", "FLOAT")
|
||||||
RETURN_NAMES = ("images", "model")
|
RETURN_NAMES = ("images", "model", "elapsed_seconds")
|
||||||
FUNCTION = "interpolate"
|
FUNCTION = "interpolate"
|
||||||
CATEGORY = "video/GIMM-VFI"
|
CATEGORY = "video/GIMM-VFI"
|
||||||
|
|
||||||
|
@_with_elapsed_seconds("GIMM-VFI segment")
|
||||||
def interpolate(self, images, model, multiplier, single_pass,
|
def interpolate(self, images, model, multiplier, single_pass,
|
||||||
clear_cache_after_n_frames, keep_device, all_on_gpu,
|
clear_cache_after_n_frames, keep_device, all_on_gpu,
|
||||||
batch_size, chunk_size, segment_index, segment_size,
|
batch_size, chunk_size, segment_index, segment_size,
|
||||||
@@ -2695,7 +2721,7 @@ class GIMMVFISegmentInterpolate(GIMMVFIInterpolate):
|
|||||||
|
|
||||||
# Standard multiplier mode
|
# Standard multiplier mode
|
||||||
is_continuation = segment_index > 0
|
is_continuation = segment_index > 0
|
||||||
(result, _) = super().interpolate(
|
(result, _, _) = super().interpolate(
|
||||||
segment_images, model, multiplier, single_pass,
|
segment_images, model, multiplier, single_pass,
|
||||||
clear_cache_after_n_frames, keep_device, all_on_gpu,
|
clear_cache_after_n_frames, keep_device, all_on_gpu,
|
||||||
batch_size, chunk_size,
|
batch_size, chunk_size,
|
||||||
|
|||||||
@@ -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
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user