Merge pull request #22 from yuvraj108c/pr-14
Auto engine building + refactor
This commit is contained in:
+2
-1
@@ -1,3 +1,4 @@
|
||||
models
|
||||
__pycache__
|
||||
.vscode
|
||||
.vscode
|
||||
CLAUDE.md
|
||||
|
||||
+4
-79
@@ -1,89 +1,14 @@
|
||||
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
|
||||
import folder_paths
|
||||
import time
|
||||
from polygraphy import cuda
|
||||
|
||||
ENGINE_DIR = os.path.join(folder_paths.models_dir, "tensorrt", "rife")
|
||||
|
||||
class RifeTensorrt:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"frames": ("IMAGE", ),
|
||||
"engine": (os.listdir(ENGINE_DIR),),
|
||||
"clear_cache_after_n_frames": ("INT", {"default": 100, "min": 1, "max": 1000}),
|
||||
"multiplier": ("INT", {"default": 2, "min": 1}),
|
||||
"use_cuda_graph": ("BOOLEAN", {"default": True}),
|
||||
"keep_model_loaded": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", )
|
||||
FUNCTION = "vfi"
|
||||
CATEGORY = "tensorrt"
|
||||
OUTPUT_NODE=True
|
||||
|
||||
def vfi(
|
||||
self,
|
||||
frames,
|
||||
engine,
|
||||
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()
|
||||
engine_path = os.path.join(ENGINE_DIR, engine)
|
||||
if (not hasattr(self, 'engine') or self.engine_label != engine):
|
||||
self.engine = Engine(engine_path)
|
||||
logger(f"Loading TensorRT engine: {engine_path}")
|
||||
self.engine.load()
|
||||
self.engine.activate()
|
||||
self.engine_label = engine
|
||||
else:
|
||||
logger(f"Using cached TensorRT engine: {engine_path}")
|
||||
|
||||
self.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 = self.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:
|
||||
del self.engine, self.engine_label
|
||||
|
||||
return (out,)
|
||||
|
||||
from .nodes.load_rife_tensorrt import LoadRifeTensorrtModel
|
||||
from .nodes.rife_tensorrt import RifeTensorrt
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"RifeTensorrt": RifeTensorrt,
|
||||
"LoadRifeTensorrtModel": LoadRifeTensorrtModel,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"RifeTensorrt": "⚡ Rife Tensorrt",
|
||||
"LoadRifeTensorrtModel": "Load Rife Tensorrt Model",
|
||||
}
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
{
|
||||
"model": {
|
||||
"options": [
|
||||
"rife47_ensemble_True_scale_1_sim",
|
||||
"rife48_ensemble_True_scale_1_sim",
|
||||
"rife49_ensemble_True_scale_1_sim"
|
||||
],
|
||||
"default": "rife49_ensemble_True_scale_1_sim",
|
||||
"tooltip": "RIFE models for video frame interpolation. These models have been tested with tensorrt. Loaded from config."
|
||||
},
|
||||
"precision": {
|
||||
"options": ["fp16", "fp32"],
|
||||
"default": "fp16",
|
||||
"tooltip": "Precision to build the tensorrt engines. Loaded from config."
|
||||
}
|
||||
}
|
||||
@@ -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,)
|
||||
@@ -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)
|
||||
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,)
|
||||
@@ -2,21 +2,36 @@
|
||||
|
||||
# ComfyUI Rife TensorRT ⚡
|
||||
|
||||
[](https://www.python.org/downloads/release/python-31012/)
|
||||
[](https://developer.nvidia.com/cuda-downloads)
|
||||
[](https://developer.nvidia.com/tensorrt)
|
||||
[](https://www.python.org/downloads/release/python-3123//)
|
||||
[](https://developer.nvidia.com/cuda-downloads)
|
||||
[](https://developer.nvidia.com/tensorrt)
|
||||
[](https://creativecommons.org/licenses/by-nc-sa/4.0/deed.en)
|
||||
|
||||

|
||||
|
||||
|
||||
</div>
|
||||
|
||||
This project provides a [TensorRT](https://github.com/NVIDIA/TensorRT) implementation of [RIFE](https://github.com/hzwer/ECCV2022-RIFE) for ultra fast frame interpolation inside ComfyUI
|
||||
|
||||
This project is licensed under [CC BY-NC-SA](https://creativecommons.org/licenses/by-nc-sa/4.0/), everyone is FREE to access, use, modify and redistribute with the same license.
|
||||
**Last tested**: 08 June 2026 (ComfyUI v0.23.0 | Torch 2.12.0 | Python 3.12.3 | L40S | CUDA 13.0 | Ubuntu 24.04)
|
||||
|
||||
If you like the project, please give me a star! ⭐
|
||||
<img width="938" height="236" alt="Screenshot 2026-06-08 at 09 21 51" src="https://github.com/user-attachments/assets/8228bd5f-7683-4b66-a476-220d5f667808" />
|
||||
|
||||
</div>
|
||||
|
||||
## ⭐ Support
|
||||
If you like my projects and wish to see updates and new features, please consider supporting me. It helps a lot!
|
||||
|
||||
[](https://github.com/yuvraj108c/ComfyUI-Depth-Anything-Tensorrt)
|
||||
[](https://github.com/yuvraj108c/ComfyUI-Upscaler-Tensorrt)
|
||||
[](https://github.com/yuvraj108c/ComfyUI-Dwpose-Tensorrt)
|
||||
[](https://github.com/yuvraj108c/ComfyUI-Rife-Tensorrt)
|
||||
|
||||
[](https://github.com/yuvraj108c/ComfyUI-Whisper)
|
||||
[](https://github.com/yuvraj108c/ComfyUI_InvSR)
|
||||
[](https://github.com/yuvraj108c/ComfyUI-Thera)
|
||||
[](https://github.com/yuvraj108c/ComfyUI-Video-Depth-Anything)
|
||||
[](https://github.com/yuvraj108c/ComfyUI-PiperTTS)
|
||||
|
||||
[](https://www.buymeacoffee.com/yuvraj108cZ)
|
||||
[](https://paypal.me/yuvraj108c)
|
||||
|
||||
---
|
||||
|
||||
@@ -40,26 +55,35 @@ cd ./ComfyUI-Rife-Tensorrt
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
## 🛠️ Building Tensorrt Engine
|
||||
## 🛠️ Supported Models
|
||||
|
||||
1. Download one of the following onnx models:
|
||||
- [rife49_ensemble_True_scale_1_sim.onnx](https://huggingface.co/yuvraj108c/rife-onnx/resolve/main/rife49_ensemble_True_scale_1_sim.onnx)
|
||||
- [rife48_ensemble_True_scale_1_sim.onnx](https://huggingface.co/yuvraj108c/rife-onnx/resolve/main/rife48_ensemble_True_scale_1_sim.onnx)
|
||||
- [rife47_ensemble_True_scale_1_sim.onnx](https://huggingface.co/yuvraj108c/rife-onnx/resolve/main/rife47_ensemble_True_scale_1_sim.onnx)
|
||||
2. Edit onnx/trt paths inside [export_trt.py](./export_trt.py) and build tensorrt engine by running:
|
||||
- `python export_trt.py`
|
||||
The following RIFE models are supported and will be automatically downloaded and built:
|
||||
- **rife49_ensemble_True_scale_1_sim** (default) - Latest and most accurate
|
||||
- **rife48_ensemble_True_scale_1_sim** - Good balance of speed and quality
|
||||
- **rife47_ensemble_True_scale_1_sim** - Fastest option
|
||||
|
||||
3. Place the exported engine inside ComfyUI `/models/tensorrt/rife` directory
|
||||
Models are automatically downloaded from [HuggingFace](https://huggingface.co/yuvraj108c/rife-onnx) and TensorRT engines are built on first use.
|
||||
|
||||
## ☀️ Usage
|
||||
|
||||
- Insert node by `Right Click -> tensorrt -> Rife Tensorrt`
|
||||
- Image resolutions between `256x256` and `3840x3840` will work with the tensorrt engines
|
||||
1. **Load Model**: Insert `Right Click -> Add Node -> tensorrt -> Load Rife Tensorrt Model`
|
||||
- Choose your preferred RIFE model (rife47, rife48, or rife49)
|
||||
- Select precision (fp16 recommended for speed, fp32 for maximum accuracy)
|
||||
- The model will be automatically downloaded and TensorRT engine built on first use
|
||||
|
||||
## 🤖 Environment tested
|
||||
2. **Process Frames**: Insert `Right Click -> Add Node -> tensorrt -> Rife Tensorrt`
|
||||
- Connect the loaded model from step 1
|
||||
- Input your video frames
|
||||
- Configure interpolation settings (multiplier, etc.)
|
||||
- Image resolutions between `256x256` and `3840x3840` are supported
|
||||
|
||||
- Ubuntu 22.04 LTS, Cuda 12.4, Tensorrt 10.4.0, Python 3.10, RTX 3070 GPU
|
||||
- Windows (Not tested, but should work)
|
||||
|
||||
## 🚨 Updates
|
||||
|
||||
### 08 June 2026
|
||||
- **Automatic Model Management**: No more manual downloads! Models are automatically downloaded from HuggingFace and TensorRT engines are built on demand. [PR#14](https://github.com/yuvraj108c/ComfyUI-Rife-Tensorrt/pull/14) by [@reaperhammer](https://github.com/reaperhammer)
|
||||
- **Improved Workflow + Codebase**: New two-node system with `Load Rife Tensorrt Model` + `Rife Tensorrt` for better organization
|
||||
- **Remove cuda-python**: No more cuda installation issues on windows
|
||||
|
||||
## 👏 Credits
|
||||
|
||||
|
||||
+4
-5
@@ -1,5 +1,4 @@
|
||||
einops
|
||||
colored
|
||||
polygraphy
|
||||
tensorrt==10.4.0
|
||||
cuda-python
|
||||
einops
|
||||
colored
|
||||
polygraphy
|
||||
tensorrt
|
||||
@@ -87,7 +87,7 @@ def export_onnx(ckpt_name, ensemble, scale_factor):
|
||||
model_path = load_file_from_github_release(MODEL_TYPE, ckpt_name)
|
||||
arch_ver = CKPT_NAME_VER_DICT[ckpt_name]
|
||||
interpolation_model = IFNet(arch_ver=arch_ver)
|
||||
interpolation_model.load_state_dict(torch.load(model_path))
|
||||
interpolation_model.load_state_dict(torch.load(model_path, map_location=TORCH_DEVICE))
|
||||
interpolation_model.eval().to(TORCH_DEVICE)
|
||||
|
||||
# # dummy data
|
||||
+35
-35
@@ -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,
|
||||
@@ -166,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))]
|
||||
@@ -217,6 +223,7 @@ 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):
|
||||
@@ -246,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
|
||||
+141
@@ -0,0 +1,141 @@
|
||||
import requests
|
||||
from tqdm import tqdm
|
||||
import logging
|
||||
import sys
|
||||
import json
|
||||
import os
|
||||
|
||||
class ColoredLogger:
|
||||
COLORS = {
|
||||
'RED': '\033[91m',
|
||||
'GREEN': '\033[92m',
|
||||
'YELLOW': '\033[93m',
|
||||
'BLUE': '\033[94m',
|
||||
'MAGENTA': '\033[95m',
|
||||
'RESET': '\033[0m'
|
||||
}
|
||||
|
||||
LEVEL_COLORS = {
|
||||
'DEBUG': COLORS['BLUE'],
|
||||
'INFO': COLORS['GREEN'],
|
||||
'WARNING': COLORS['YELLOW'],
|
||||
'ERROR': COLORS['RED'],
|
||||
'CRITICAL': COLORS['MAGENTA']
|
||||
}
|
||||
|
||||
def __init__(self, name="MY-APP"):
|
||||
self.logger = logging.getLogger(name)
|
||||
self.logger.setLevel(logging.DEBUG)
|
||||
self.app_name = name
|
||||
|
||||
# Prevent message propagation to parent loggers
|
||||
self.logger.propagate = False
|
||||
|
||||
# Clear existing handlers
|
||||
self.logger.handlers = []
|
||||
|
||||
# Create console handler
|
||||
handler = logging.StreamHandler(sys.stdout)
|
||||
handler.setLevel(logging.DEBUG)
|
||||
|
||||
# Custom formatter class to handle colored components
|
||||
class ColoredFormatter(logging.Formatter):
|
||||
def format(self, record):
|
||||
# Color the level name according to severity
|
||||
level_color = ColoredLogger.LEVEL_COLORS.get(record.levelname, '')
|
||||
colored_levelname = f"{level_color}{record.levelname}{ColoredLogger.COLORS['RESET']}"
|
||||
|
||||
# Color the logger name in blue
|
||||
colored_name = f"{ColoredLogger.COLORS['BLUE']}{record.name}{ColoredLogger.COLORS['RESET']}"
|
||||
|
||||
# Set the colored components
|
||||
record.levelname = colored_levelname
|
||||
record.name = colored_name
|
||||
|
||||
return super().format(record)
|
||||
|
||||
# Create formatter with the new format
|
||||
formatter = ColoredFormatter('[%(name)s|%(levelname)s] - %(message)s')
|
||||
handler.setFormatter(formatter)
|
||||
|
||||
self.logger.addHandler(handler)
|
||||
|
||||
|
||||
def debug(self, message):
|
||||
self.logger.debug(f"{self.COLORS['BLUE']}{message}{self.COLORS['RESET']}")
|
||||
|
||||
def info(self, message):
|
||||
self.logger.info(f"{self.COLORS['GREEN']}{message}{self.COLORS['RESET']}")
|
||||
|
||||
def warning(self, message):
|
||||
self.logger.warning(f"{self.COLORS['YELLOW']}{message}{self.COLORS['RESET']}")
|
||||
|
||||
def error(self, message):
|
||||
self.logger.error(f"{self.COLORS['RED']}{message}{self.COLORS['RESET']}")
|
||||
|
||||
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
|
||||
|
||||
Args:
|
||||
url (str): URL of the file to download
|
||||
save_path (str): Path to save the file as
|
||||
"""
|
||||
GREEN = '\033[92m'
|
||||
RESET = '\033[0m'
|
||||
response = requests.get(url, stream=True)
|
||||
response.raise_for_status()
|
||||
total_size = int(response.headers.get('content-length', 0))
|
||||
|
||||
with open(save_path, 'wb') as file, tqdm(
|
||||
desc=save_path,
|
||||
total=total_size,
|
||||
unit='iB',
|
||||
unit_scale=True,
|
||||
unit_divisor=1024,
|
||||
colour='green',
|
||||
bar_format=f'{GREEN}{{l_bar}}{{bar}}{RESET}{GREEN}{{r_bar}}{RESET}'
|
||||
) as progress_bar:
|
||||
for data in response.iter_content(chunk_size=1024):
|
||||
size = file.write(data)
|
||||
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
|
||||
+22
-9
@@ -7,7 +7,9 @@ import einops
|
||||
from comfy.model_management import soft_empty_cache, get_torch_device
|
||||
import numpy as np
|
||||
from comfy.utils import ProgressBar
|
||||
from colored import Fore, Back, Style
|
||||
from colored import Fore, Back, Style
|
||||
from .utilities import rife_logger
|
||||
from tqdm import tqdm
|
||||
|
||||
DEVICE = get_torch_device()
|
||||
|
||||
@@ -25,9 +27,6 @@ def load_file_from_github_release(model_type, ckpt_name):
|
||||
error_str = '\n\n'.join(error_strs)
|
||||
raise Exception(f"Tried all GitHub base urls to download {ckpt_name} but no suceess. Below is the error log:\n\n{error_str}")
|
||||
|
||||
def logger(msg):
|
||||
print(f'{Style.reset}{Fore.cyan}⚡ [Rife Tensorrt] - {msg}{Style.reset}')
|
||||
|
||||
def preprocess_frames(frames):
|
||||
return einops.rearrange(frames[..., :3], "n h w c -> n c h w")
|
||||
|
||||
@@ -45,7 +44,15 @@ def generate_frames_rife(
|
||||
out_len = 0
|
||||
|
||||
number_of_frames_processed_since_last_cleared_cuda_cache = 0
|
||||
pbar = ProgressBar(len(frames))
|
||||
pbar = ProgressBar(len(frames)-1)
|
||||
|
||||
bar_format = "[\033[94mComfyUI-Rife-Tensorrt\033[0m|\033[92mINFO\033[0m] - \033[92m{desc}: {percentage:3.0f}%|{bar}| {n_fmt}/{total_fmt} [{elapsed}<{remaining}]"
|
||||
progress_bar = tqdm(
|
||||
total=len(frames)-1,
|
||||
desc="Interpolating",
|
||||
bar_format=bar_format,
|
||||
disable=((len(frames)-1) == 1)
|
||||
)
|
||||
|
||||
for frame_itr in range(len(frames) - 1): # Skip the final frame since there are no frames after it
|
||||
|
||||
@@ -67,19 +74,25 @@ def generate_frames_rife(
|
||||
if number_of_frames_processed_since_last_cleared_cuda_cache >= clear_cache_after_n_frames:
|
||||
soft_empty_cache()
|
||||
number_of_frames_processed_since_last_cleared_cuda_cache = 0
|
||||
logger("Clearing cache...")
|
||||
rife_logger.info("Clearing cache...")
|
||||
|
||||
pbar.update(1)
|
||||
progress_bar.update(1)
|
||||
|
||||
|
||||
progress_bar.refresh()
|
||||
progress_bar.close()
|
||||
|
||||
# Append final frame
|
||||
output_frames[out_len] = frames[-1:]
|
||||
logger(f"done! - {(len(frames) -1) * (multiplier-1)} new frames generated at resolution: {output_frames[0].shape}")
|
||||
# Get actual frame shape from first interpolated frame (CHW format)
|
||||
actual_frame = output_frames[0]
|
||||
h, w = actual_frame.shape[1], actual_frame.shape[2]
|
||||
rife_logger.info(f"done! - {out_len} total frames output at resolution: {h}x{w}")
|
||||
out_len += 1
|
||||
|
||||
# clear cache for courtesy
|
||||
soft_empty_cache()
|
||||
logger("Final clearing cache done ...")
|
||||
rife_logger.info("Final clearing cache done ...")
|
||||
#
|
||||
res = output_frames[:out_len]
|
||||
return res
|
||||
Reference in New Issue
Block a user