From 38cd1fe2fda982a1281596a6627068591ec3f813 Mon Sep 17 00:00:00 2001 From: yuvraj108c Date: Mon, 8 Jun 2026 05:04:45 +0000 Subject: [PATCH] refactor + auto engine building + remove cuda --- __init__.py | 204 +---------------------- nodes/load_rife_tensorrt.py | 91 ++++++++++ nodes/rife_tensorrt.py | 56 +++++++ requirements.txt | 14 +- export_onnx.py => scripts/export_onnx.py | 0 export_trt.py => scripts/export_trt.py | 0 trt_utilities.py | 154 +++++------------ utilities.py | 41 ++++- 8 files changed, 237 insertions(+), 323 deletions(-) create mode 100644 nodes/load_rife_tensorrt.py create mode 100644 nodes/rife_tensorrt.py rename export_onnx.py => scripts/export_onnx.py (100%) rename export_trt.py => scripts/export_trt.py (100%) diff --git a/__init__.py b/__init__.py index 0d04dc9..0ead840 100644 --- a/__init__.py +++ b/__init__.py @@ -1,205 +1,5 @@ -import torch -import os -from comfy.model_management import get_torch_device -from .vfi_utilities import preprocess_frames, postprocess_frames, generate_frames_rife, logger -from .trt_utilities import Engine -from .utilities import download_file, ColoredLogger -import folder_paths -import time -from polygraphy import cuda -import comfy.model_management as mm -import tensorrt -import json - -ENGINE_DIR = os.path.join(folder_paths.models_dir, "tensorrt", "rife") - -# Image dimensions for TensorRT engine building -IMAGE_DIM_MIN = 256 -IMAGE_DIM_OPT = 512 -IMAGE_DIM_MAX = 3840 - -# Logger for this module -rife_logger = ColoredLogger("ComfyUI-Rife-Tensorrt") - -# Function to load configuration -def load_node_config(config_filename="load_rife_config.json"): - """Loads node configuration from a JSON file.""" - current_dir = os.path.dirname(__file__) - config_path = os.path.join(current_dir, config_filename) - - default_config = { - "model": { - "options": ["rife49_ensemble_True_scale_1_sim"], - "default": "rife49_ensemble_True_scale_1_sim", - "tooltip": "Default model (fallback from code)" - }, - "precision": { - "options": ["fp16", "fp32"], - "default": "fp16", - "tooltip": "Default precision (fallback from code)" - } - } - - try: - with open(config_path, 'r') as f: - config = json.load(f) - rife_logger.info(f"Successfully loaded configuration from {config_filename}") - return config - except FileNotFoundError: - rife_logger.warning(f"Configuration file '{config_path}' not found. Using default fallback configuration.") - return default_config - except json.JSONDecodeError: - rife_logger.error(f"Error decoding JSON from '{config_path}'. Using default fallback configuration.") - return default_config - except Exception as e: - rife_logger.error(f"An unexpected error occurred while loading '{config_path}': {e}. Using default fallback.") - return default_config - -# Load the configuration once when the module is imported -LOAD_RIFE_NODE_CONFIG = load_node_config() - -class LoadRifeTensorrtModel: - @classmethod - def INPUT_TYPES(cls): - # Use the pre-loaded configuration - model_config = LOAD_RIFE_NODE_CONFIG.get("model", {}) - precision_config = LOAD_RIFE_NODE_CONFIG.get("precision", {}) - - # Provide sensible defaults if keys are missing in the config - model_options = model_config.get("options", ["rife49_ensemble_True_scale_1_sim"]) - model_default = model_config.get("default", "rife49_ensemble_True_scale_1_sim") - model_tooltip = model_config.get("tooltip", "Select a RIFE model.") - - precision_options = precision_config.get("options", ["fp16", "fp32"]) - precision_default = precision_config.get("default", "fp16") - precision_tooltip = precision_config.get("tooltip", "Select precision.") - - return { - "required": { - "model": (model_options, {"default": model_default, "tooltip": model_tooltip}), - "precision": (precision_options, {"default": precision_default, "tooltip": precision_tooltip}), - } - } - - RETURN_NAMES = ("rife_trt_model",) - RETURN_TYPES = ("RIFE_TRT_MODEL",) - CATEGORY = "tensorrt" - DESCRIPTION = "Load RIFE tensorrt models, they will be built automatically if not found." - FUNCTION = "load_rife_tensorrt_model" - - def load_rife_tensorrt_model(self, model, precision): - tensorrt_models_dir = os.path.join(folder_paths.models_dir, "tensorrt", "rife") - onnx_models_dir = os.path.join(folder_paths.models_dir, "onnx") - - os.makedirs(tensorrt_models_dir, exist_ok=True) - os.makedirs(onnx_models_dir, exist_ok=True) - - onnx_model_path = os.path.join(onnx_models_dir, f"{model}.onnx") - - # Build tensorrt model path with detailed naming - engine_channel = 3 - engine_min_batch, engine_opt_batch, engine_max_batch = 1, 1, 1 - engine_min_h, engine_opt_h, engine_max_h = IMAGE_DIM_MIN, IMAGE_DIM_OPT, IMAGE_DIM_MAX - engine_min_w, engine_opt_w, engine_max_w = IMAGE_DIM_MIN, IMAGE_DIM_OPT, IMAGE_DIM_MAX - tensorrt_model_path = os.path.join(tensorrt_models_dir, f"{model}_{precision}_{engine_min_batch}x{engine_channel}x{engine_min_h}x{engine_min_w}_{engine_opt_batch}x{engine_channel}x{engine_opt_h}x{engine_opt_w}_{engine_max_batch}x{engine_channel}x{engine_max_h}x{engine_max_w}_{tensorrt.__version__}.trt") - - if not os.path.exists(tensorrt_model_path): - if not os.path.exists(onnx_model_path): - onnx_model_download_url = f"https://huggingface.co/yuvraj108c/rife-onnx/resolve/main/{model}.onnx" - rife_logger.info(f"Downloading {onnx_model_download_url}") - download_file(url=onnx_model_download_url, save_path=onnx_model_path) - else: - rife_logger.info(f"ONNX model found at: {onnx_model_path}") - - rife_logger.info(f"Building TensorRT engine for {onnx_model_path}: {tensorrt_model_path}") - mm.soft_empty_cache() - s = time.time() - engine = Engine(tensorrt_model_path) - engine.build( - onnx_path=onnx_model_path, - fp16=True if precision == "fp16" else False, - input_profile=[ - { - "img0": [(engine_min_batch, engine_channel, engine_min_h, engine_min_w), (engine_opt_batch, engine_channel, engine_opt_h, engine_opt_w), (engine_max_batch, engine_channel, engine_max_h, engine_max_w)], - "img1": [(engine_min_batch, engine_channel, engine_min_h, engine_min_w), (engine_opt_batch, engine_channel, engine_opt_h, engine_opt_w), (engine_max_batch, engine_channel, engine_max_h, engine_max_w)], - } - ], - ) - e = time.time() - rife_logger.info(f"Time taken to build: {(e-s)} seconds") - - rife_logger.info(f"Loading TensorRT engine: {tensorrt_model_path}") - mm.soft_empty_cache() - engine = Engine(tensorrt_model_path) - engine.load() - - return (engine,) - -class RifeTensorrt: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "frames": ("IMAGE", {"tooltip": "Input frames for video frame interpolation"}), - "rife_trt_model": ("RIFE_TRT_MODEL", {"tooltip": "Tensorrt model built and loaded"}), - "clear_cache_after_n_frames": ("INT", {"default": 100, "min": 1, "max": 1000, "tooltip": "Clear CUDA cache after processing this many frames"}), - "multiplier": ("INT", {"default": 2, "min": 1, "tooltip": "Frame interpolation multiplier"}), - "use_cuda_graph": ("BOOLEAN", {"default": True, "tooltip": "Use CUDA graph for better performance. Disable if experiencing high RAM usage or errors with variable input resolutions."}), - "keep_model_loaded": ("BOOLEAN", {"default": False, "tooltip": "Keep model loaded in memory after processing"}), - }, - } - - RETURN_TYPES = ("IMAGE", ) - FUNCTION = "vfi" - CATEGORY = "tensorrt" - OUTPUT_NODE=True - - def vfi( - self, - frames, - rife_trt_model, - clear_cache_after_n_frames=100, - multiplier=2, - use_cuda_graph=True, - keep_model_loaded=False, - ): - B, H, W, C = frames.shape - shape_dict = { - "img0": {"shape": (1, 3, H, W)}, - "img1": {"shape": (1, 3, H, W)}, - "output": {"shape": (1, 3, H, W)}, - } - - cudaStream = cuda.Stream() - - # Use the provided model directly - engine = rife_trt_model - logger(f"Using loaded TensorRT engine") - - # Activate and allocate buffers for the engine - engine.activate() - engine.allocate_buffers(shape_dict=shape_dict) - - frames = preprocess_frames(frames) - - def return_middle_frame(frame_0, frame_1, timestep): - timestep_t = torch.tensor([timestep], dtype=torch.float32).to(get_torch_device()) - # s = time.time() - output = engine.infer({"img0": frame_0, "img1": frame_1, "timestep": timestep_t}, cudaStream, use_cuda_graph) - # e = time.time() - # print(f"Time taken to infer: {(e-s)*1000} ms") - - result = output['output'] - return result - - result = generate_frames_rife(frames, clear_cache_after_n_frames, multiplier, return_middle_frame) - out = postprocess_frames(result) - - if not keep_model_loaded: - engine.reset() - - return (out,) - +from .nodes.load_rife_tensorrt import LoadRifeTensorrtModel +from .nodes.rife_tensorrt import RifeTensorrt NODE_CLASS_MAPPINGS = { "RifeTensorrt": RifeTensorrt, diff --git a/nodes/load_rife_tensorrt.py b/nodes/load_rife_tensorrt.py new file mode 100644 index 0000000..3be2d8b --- /dev/null +++ b/nodes/load_rife_tensorrt.py @@ -0,0 +1,91 @@ +from ..trt_utilities import Engine +from ..utilities import download_file, load_node_config, rife_logger +import folder_paths +import time +import comfy.model_management as mm +import tensorrt +import os + +# Image dimensions for TensorRT engine building +IMAGE_DIM_MIN = 256 +IMAGE_DIM_OPT = 512 +IMAGE_DIM_MAX = 3840 + +LOAD_RIFE_NODE_CONFIG = load_node_config() + +class LoadRifeTensorrtModel: + @classmethod + def INPUT_TYPES(cls): + # Use the pre-loaded configuration + model_config = LOAD_RIFE_NODE_CONFIG.get("model", {}) + precision_config = LOAD_RIFE_NODE_CONFIG.get("precision", {}) + + # Provide sensible defaults if keys are missing in the config + model_options = model_config.get("options", ["rife49_ensemble_True_scale_1_sim"]) + model_default = model_config.get("default", "rife49_ensemble_True_scale_1_sim") + model_tooltip = model_config.get("tooltip", "Select a RIFE model.") + + precision_options = precision_config.get("options", ["fp16", "fp32"]) + precision_default = precision_config.get("default", "fp16") + precision_tooltip = precision_config.get("tooltip", "Select precision.") + + return { + "required": { + "model": (model_options, {"default": model_default, "tooltip": model_tooltip}), + "precision": (precision_options, {"default": precision_default, "tooltip": precision_tooltip}), + } + } + + RETURN_NAMES = ("rife_trt_model",) + RETURN_TYPES = ("RIFE_TRT_MODEL",) + CATEGORY = "tensorrt" + DESCRIPTION = "Load RIFE tensorrt models, they will be built automatically if not found." + FUNCTION = "load_rife_tensorrt_model" + + def load_rife_tensorrt_model(self, model, precision): + tensorrt_models_dir = os.path.join(folder_paths.models_dir, "tensorrt", "rife") + onnx_models_dir = os.path.join(folder_paths.models_dir, "onnx") + + os.makedirs(tensorrt_models_dir, exist_ok=True) + os.makedirs(onnx_models_dir, exist_ok=True) + + onnx_model_path = os.path.join(onnx_models_dir, f"{model}.onnx") + + # Build tensorrt model path with detailed naming + engine_channel = 3 + engine_min_batch, engine_opt_batch, engine_max_batch = 1, 1, 1 + engine_min_h, engine_opt_h, engine_max_h = IMAGE_DIM_MIN, IMAGE_DIM_OPT, IMAGE_DIM_MAX + engine_min_w, engine_opt_w, engine_max_w = IMAGE_DIM_MIN, IMAGE_DIM_OPT, IMAGE_DIM_MAX + tensorrt_model_path = os.path.join(tensorrt_models_dir, f"{model}_{precision}_{engine_min_batch}x{engine_channel}x{engine_min_h}x{engine_min_w}_{engine_opt_batch}x{engine_channel}x{engine_opt_h}x{engine_opt_w}_{engine_max_batch}x{engine_channel}x{engine_max_h}x{engine_max_w}_{tensorrt.__version__}.trt") + + if not os.path.exists(tensorrt_model_path): + if not os.path.exists(onnx_model_path): + onnx_model_download_url = f"https://huggingface.co/yuvraj108c/rife-onnx/resolve/main/{model}.onnx" + rife_logger.info(f"Downloading {onnx_model_download_url}") + download_file(url=onnx_model_download_url, save_path=onnx_model_path) + else: + rife_logger.info(f"ONNX model found at: {onnx_model_path}") + + rife_logger.info(f"Building TensorRT engine for {onnx_model_path}: {tensorrt_model_path}") + mm.soft_empty_cache() + s = time.time() + engine = Engine(tensorrt_model_path) + engine.build( + onnx_path=onnx_model_path, + fp16=True if precision == "fp16" else False, + input_profile=[ + { + "img0": [(engine_min_batch, engine_channel, engine_min_h, engine_min_w), (engine_opt_batch, engine_channel, engine_opt_h, engine_opt_w), (engine_max_batch, engine_channel, engine_max_h, engine_max_w)], + "img1": [(engine_min_batch, engine_channel, engine_min_h, engine_min_w), (engine_opt_batch, engine_channel, engine_opt_h, engine_opt_w), (engine_max_batch, engine_channel, engine_max_h, engine_max_w)], + } + ], + ) + e = time.time() + rife_logger.info(f"Time taken to build: {(e-s)} seconds") + + rife_logger.info(f"Loading TensorRT engine: {tensorrt_model_path}") + mm.soft_empty_cache() + engine = Engine(tensorrt_model_path) + engine.load() + + return (engine,) diff --git a/nodes/rife_tensorrt.py b/nodes/rife_tensorrt.py new file mode 100644 index 0000000..81f2397 --- /dev/null +++ b/nodes/rife_tensorrt.py @@ -0,0 +1,56 @@ +import torch +import os +from comfy.model_management import get_torch_device +from ..vfi_utilities import preprocess_frames, postprocess_frames, generate_frames_rife +from ..trt_utilities import Engine +import folder_paths +import time +import comfy.model_management as mm + +class RifeTensorrt: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "frames": ("IMAGE", {"tooltip": "Input frames for video frame interpolation"}), + "rife_trt_model": ("RIFE_TRT_MODEL", {"tooltip": "Tensorrt model built and loaded"}), + "clear_cache_after_n_frames": ("INT", {"default": 100, "min": 1, "max": 1000, "tooltip": "Clear CUDA cache after processing this many frames"}), + "multiplier": ("INT", {"default": 2, "min": 1, "tooltip": "Frame interpolation multiplier"}), + }, + } + + RETURN_TYPES = ("IMAGE", ) + FUNCTION = "vfi" + CATEGORY = "tensorrt" + + def vfi( + self, + frames, + rife_trt_model, + clear_cache_after_n_frames=100, + multiplier=2, + ): + B, H, W, C = frames.shape + shape_dict = { + "img0": {"shape": (1, 3, H, W)}, + "img1": {"shape": (1, 3, H, W)}, + "output": {"shape": (1, 3, H, W)}, + } + + cudaStream = torch.cuda.current_stream().cuda_stream + engine = rife_trt_model + engine.activate() + engine.allocate_buffers(shape_dict=shape_dict) + + frames = preprocess_frames(frames) + + def return_middle_frame(frame_0, frame_1, timestep): + timestep_t = torch.tensor([timestep], dtype=torch.float32).to(get_torch_device()) + output = engine.infer({"img0": frame_0, "img1": frame_1, "timestep": timestep_t}, cudaStream, use_cuda_graph) + result = output['output'] + return result + + result = generate_frames_rife(frames, clear_cache_after_n_frames, multiplier, return_middle_frame) + out = postprocess_frames(result) + + return (out,) \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index a8ac367..c222890 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,10 +1,4 @@ -einops>=0.8.0 -colored>=1.1.0 -polygraphy>=0.49.0 -tensorrt>=10.12.0 -cuda-python>=12.0.0 -requests>=2.31.0 -tqdm>=4.66.0 -onnx>=1.20.0 -onnxsim>=0.5.0 -torch>=2.9.0 \ No newline at end of file +einops +colored +polygraphy +tensorrt \ No newline at end of file diff --git a/export_onnx.py b/scripts/export_onnx.py similarity index 100% rename from export_onnx.py rename to scripts/export_onnx.py diff --git a/export_trt.py b/scripts/export_trt.py similarity index 100% rename from export_trt.py rename to scripts/export_trt.py diff --git a/trt_utilities.py b/trt_utilities.py index 23a66fe..3ea5567 100644 --- a/trt_utilities.py +++ b/trt_utilities.py @@ -1,3 +1,20 @@ +# +# Copyright 2022 The HuggingFace Inc. team. +# SPDX-FileCopyrightText: Copyright (c) 1993-2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# import torch from torch.cuda import nvtx from collections import OrderedDict @@ -16,7 +33,6 @@ import tensorrt as trt from logging import error, warning from tqdm import tqdm import copy -import cuda.bindings.runtime as cudart TRT_LOGGER = trt.Logger(trt.Logger.ERROR) G_LOGGER.module_severity = G_LOGGER.ERROR @@ -44,17 +60,6 @@ torch_to_numpy_dtype_dict = { value: key for (key, value) in numpy_to_torch_dtype_dict.items() } -# https://github.com/Jeff-LiangF/streamv2v/blob/18c1a3bd56ff348d54a3300605936980bb13b03c/src/streamv2v/acceleration/tensorrt/utilities.py -def CUASSERT(cuda_ret): - err = cuda_ret[0] - if err != cudart.cudaError_t.cudaSuccess: - raise RuntimeError( - f"CUDA ERROR: {err}, error code reference: https://nvidia.github.io/cuda-python/module/cudart.html#cuda.cudart.cudaError_t" - ) - if len(cuda_ret) > 1: - return cuda_ret[1] - return None - class TQDMProgressMonitor(trt.IProgressMonitor): def __init__(self): trt.IProgressMonitor.__init__(self) @@ -125,6 +130,7 @@ class TQDMProgressMonitor(trt.IProgressMonitor): # There is no need to propagate this exception to TensorRT. We can simply cancel the build. return False + class Engine: def __init__( self, @@ -136,72 +142,24 @@ class Engine: self.buffers = OrderedDict() self.tensors = OrderedDict() self.cuda_graph_instance = None # cuda graph - self.graph = None def __del__(self): - # Clean up CUDA graph resources - if hasattr(self, 'cuda_graph_instance') and self.cuda_graph_instance is not None: - try: - cudart.cudaGraphDestroy(self.cuda_graph_instance) - except Exception: - pass - if hasattr(self, 'graph') and self.graph is not None: - try: - cudart.cudaGraphDestroy(self.graph) - except Exception: - pass - - if hasattr(self, 'engine'): - del self.engine - if hasattr(self, 'context'): - del self.context - if hasattr(self, 'tensors'): - for key in list(self.tensors.keys()): - del self.tensors[key] - del self.tensors - if hasattr(self, 'buffers'): - del self.buffers - if hasattr(self, 'inputs'): - del self.inputs - if hasattr(self, 'outputs'): - del self.outputs + del self.engine + del self.context + del self.buffers + del self.tensors def reset(self, engine_path=None): - # Clean up CUDA graph resources first - if hasattr(self, 'cuda_graph_instance') and self.cuda_graph_instance is not None: - try: - cudart.cudaGraphDestroy(self.cuda_graph_instance) - except Exception: - pass - self.cuda_graph_instance = None - if hasattr(self, 'graph') and self.graph is not None: - try: - cudart.cudaGraphDestroy(self.graph) - except Exception: - pass - self.graph = None - - if hasattr(self, 'engine') and self.engine is not None: - del self.engine - if hasattr(self, 'context') and self.context is not None: - del self.context - if hasattr(self, 'tensors'): - for key in list(self.tensors.keys()): - del self.tensors[key] - del self.tensors - if hasattr(self, 'buffers'): - del self.buffers - - self.engine = None - self.context = None - self.engine_path = engine_path if engine_path else self.engine_path + del self.engine + del self.context + del self.buffers + del self.tensors + self.engine_path = engine_path self.buffers = OrderedDict() self.tensors = OrderedDict() self.inputs = {} self.outputs = {} - self.cuda_graph_instance = None - self.graph = None def build( self, @@ -214,7 +172,7 @@ class Engine: timing_cache=None, update_output_names=None, ): - print(f"Building TensorRT engine for {onnx_path}: {self.engine_path}") + # print(f"Building TensorRT engine for {onnx_path}: {self.engine_path}") p = [Profile()] if input_profile: p = [Profile() for i in range(len(input_profile))] @@ -265,13 +223,10 @@ class Engine: return 0 def load(self): + # print(f"Loading TensorRT engine: {self.engine_path}") self.engine = engine_from_bytes(bytes_from_path(self.engine_path)) def activate(self, reuse_device_memory=None): - # If engine was reset, reload it - if self.engine is None: - self.load() - if reuse_device_memory: self.context = self.engine.create_execution_context_without_device_memory() # self.context.device_memory = reuse_device_memory @@ -279,20 +234,6 @@ class Engine: self.context = self.engine.create_execution_context() def allocate_buffers(self, shape_dict=None, device="cuda"): - # Clean up CUDA graph resources since tensors will be recreated - if hasattr(self, 'cuda_graph_instance') and self.cuda_graph_instance is not None: - try: - cudart.cudaGraphDestroy(self.cuda_graph_instance) - except Exception: - pass - self.cuda_graph_instance = None - if hasattr(self, 'graph') and self.graph is not None: - try: - cudart.cudaGraphDestroy(self.graph) - except Exception: - pass - self.graph = None - nvtx.range_push("allocate_buffers") for idx in range(self.engine.num_io_tensors): name = self.engine.get_tensor_name(idx) @@ -312,32 +253,25 @@ class Engine: nvtx.range_pop() def infer(self, feed_dict, stream, use_cuda_graph=False): + nvtx.range_push("set_tensors") for name, buf in feed_dict.items(): self.tensors[name].copy_(buf) for name, tensor in self.tensors.items(): self.context.set_tensor_address(name, tensor.data_ptr()) - - if use_cuda_graph: - if self.cuda_graph_instance is not None: - CUASSERT(cudart.cudaGraphLaunch(self.cuda_graph_instance, stream.ptr)) - CUASSERT(cudart.cudaStreamSynchronize(stream.ptr)) - else: - # do inference before CUDA graph capture - noerror = self.context.execute_async_v3(stream.ptr) - if not noerror: - raise ValueError("ERROR: inference failed.") - # capture cuda graph - CUASSERT( - cudart.cudaStreamBeginCapture(stream.ptr, cudart.cudaStreamCaptureMode.cudaStreamCaptureModeGlobal) - ) - self.context.execute_async_v3(stream.ptr) - self.graph = CUASSERT(cudart.cudaStreamEndCapture(stream.ptr)) - self.cuda_graph_instance = CUASSERT(cudart.cudaGraphInstantiate(self.graph, 0)) - else: - noerror = self.context.execute_async_v3(stream.ptr) - if not noerror: - raise ValueError("ERROR: inference failed.") - + nvtx.range_pop() + nvtx.range_push("execute") + noerror = self.context.execute_async_v3(stream) + if not noerror: + raise ValueError("ERROR: inference failed.") + nvtx.range_pop() return self.tensors + def __str__(self): + out = "" + for opt_profile in range(self.engine.num_optimization_profiles): + for binding_idx in range(self.engine.num_bindings): + name = self.engine.get_binding_name(binding_idx) + shape = self.engine.get_profile_shape(opt_profile, name) + out += f"\t{name} = {shape}\n" + return out \ No newline at end of file diff --git a/utilities.py b/utilities.py index f95e841..4d805d0 100644 --- a/utilities.py +++ b/utilities.py @@ -2,6 +2,8 @@ import requests from tqdm import tqdm import logging import sys +import json +import os class ColoredLogger: COLORS = { @@ -74,6 +76,8 @@ class ColoredLogger: def critical(self, message): self.logger.critical(f"{self.COLORS['MAGENTA']}{message}{self.COLORS['RESET']}") +rife_logger = ColoredLogger("ComfyUI-Rife-Tensorrt") + def download_file(url, save_path): """ Download a file from URL with progress bar @@ -99,4 +103,39 @@ def download_file(url, save_path): ) as progress_bar: for data in response.iter_content(chunk_size=1024): size = file.write(data) - progress_bar.update(size) \ No newline at end of file + progress_bar.update(size) + + +# Function to load configuration +def load_node_config(config_filename="load_rife_config.json"): + """Loads node configuration from a JSON file.""" + current_dir = os.path.dirname(__file__) + config_path = os.path.join(current_dir, config_filename) + + default_config = { + "model": { + "options": ["rife49_ensemble_True_scale_1_sim"], + "default": "rife49_ensemble_True_scale_1_sim", + "tooltip": "Default model (fallback from code)" + }, + "precision": { + "options": ["fp16", "fp32"], + "default": "fp16", + "tooltip": "Default precision (fallback from code)" + } + } + + try: + with open(config_path, 'r') as f: + config = json.load(f) + rife_logger.info(f"Successfully loaded configuration from {config_filename}") + return config + except FileNotFoundError: + rife_logger.warning(f"Configuration file '{config_path}' not found. Using default fallback configuration.") + return default_config + except json.JSONDecodeError: + rife_logger.error(f"Error decoding JSON from '{config_path}'. Using default fallback configuration.") + return default_config + except Exception as e: + rife_logger.error(f"An unexpected error occurred while loading '{config_path}': {e}. Using default fallback.") + return default_config \ No newline at end of file