30 Commits
Author SHA1 Message Date
yuvraj108c b1293bf089 model resetting 2026-06-08 05:50:42 +00:00
Yuvraj Seegolam 5f97abca46 Merge pull request #22 from yuvraj108c/pr-14
Auto engine building + refactor
2026-06-08 09:28:32 +04:00
Yuvraj Seegolam 9307964f57 update readme 2026-06-08 09:26:34 +04:00
Yuvraj Seegolam 11db2d2f7f add node image 2026-06-08 09:22:42 +04:00
yuvraj108c ad6d4f8fe3 update readme + logging 2026-06-08 05:19:17 +00:00
yuvraj108c 38cd1fe2fd refactor + auto engine building + remove cuda 2026-06-08 05:04:45 +00:00
Yuvraj Seegolam bef2ac3121 Merge pull request #2 from ComfyNodePRs/publish
Add Github Action for Publishing to Comfy Registry
2026-06-08 07:49:15 +04:00
Yuvraj Seegolam a96f043aaf Merge pull request #1 from ComfyNodePRs/pyproject
Add pyproject.toml for Custom Node Registry
2026-06-08 07:48:59 +04:00
reaperhammerandClaude Opus 4.7 154310e3e9 Relax requirements to use >= pinning and add missing deps
- Remove dead #tensorrt==10.13.3.9 comment
- Add onnx and onnxsim for export scripts
- Pin minimum versions to ensure compatibility

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-05-09 20:09:54 +12:00
reaperhammerandClaude Opus 4.7 eedc0938f9 Fix bare exceptions, silent HTTP failures, and resolution display
- Replace bare except: with except Exception: in __del__ and reset() to avoid
  catching KeyboardInterrupt/SystemExit
- Add response.raise_for_status() to fail clearly on HTTP errors
- Fix resolution log output to show HxW instead of CHW shape

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-05-09 19:04:41 +12:00
reaperhammerandClaude Opus 4.7 7cf8989cd3 Improve cuda graph tooltip with RAM/resolution caveats
Advises users to disable cuda graph if experiencing high RAM usage
or errors with variable input resolutions.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-05-09 18:39:17 +12:00
reaperhammerandClaude Opus 4.7 203afb5f2a Improve CUDA graph memory cleanup in Engine class
- __del__: Add hasattr checks before del, iterate tensors dict to delete
  each tensor individually, also clean up inputs/outputs
- reset: Same improvements plus reset cuda_graph_instance and graph
  to None to prevent stale state on reuse

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-05-09 15:40:57 +12:00
reaperhammerandClaude Opus 4.7 3c3bac0dea Fix output frame count message to show actual total
Previously reported only new interpolated frames using a formula, not the
actual count. Now shows out_len which reflects true total frames output.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-05-09 15:17:31 +12:00
reaperhammer 65abe2691a minimax trial 1 2026-05-09 14:59:03 +12:00
reaperhammerandClaude Opus 4.7 03b2a8e9d6 Fix torch.load to load directly to GPU device
Without map_location, torch.load defaults to CPU then the model is moved
to GPU with .to(TORCH_DEVICE). This is inefficient and can cause issues on
systems where CPU tensors aren't properly set up.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-05-09 14:59:03 +12:00
reaperhammer f0a75cb1b1 Update requirements.txt for dependency management
Reordered and uncommented dependencies in requirements.txt
2025-12-22 02:34:43 +13:00
reaperhammer 6cd119d9e5 fix memory issues with graphs when keep_model_loaded = true 2025-09-25 08:16:02 +12:00
reaperhammer 3e898b65f7 fix engine reload on batch jobs when keep_model_loaded = false 2025-09-25 07:54:13 +12:00
reaperhammer 16994ac3cb Update readme.md
more readme updates
2025-09-25 00:42:28 +12:00
reaperhammer 8a2a1ae405 Update readme.md
Update the Python and Cuda version images as it works with these later versions that I'm using
2025-09-25 00:36:16 +12:00
reaperhammer 25d6c9f1d2 Make the node auto download/build models and update readme.md accordingly 2025-09-25 00:08:26 +12:00
reaperhammer 90b7af76d0 Merge remote-tracking branch 'refs/remotes/origin/master'
merge local branch with remote
2025-09-25 00:04:16 +12:00
reaperhammer e0d17b55e7 Merge pull request #13 from ritlo/master
fix: add export_trt.py to sys.path to fix it for use with comfyui's python_embeded/python.exe
2025-09-24 23:55:35 +12:00
Yuvraj Seegolam 41686c5c64 Merge pull request #13 from ritlo/master
fix: add export_trt.py to sys.path to fix it for use with comfyui's python_embeded/python.exe
2025-09-23 11:06:13 +04:00
Ritvick Lohani 261efd6e5b fix: add export_trt to sys.path to fix it for use with comfyui's python_embedded/python.exe 2025-09-23 02:24:44 +05:30
Yuvraj Seegolam 06b847dfa7 Merge pull request #12 from RioShiina47/master
fix cudart
2025-09-18 18:47:25 +04:00
RioShiina47 e971cfb71a fix cudart 2025-09-16 20:44:51 +08:00
snomiao 63ec06216f chore(publish): Add Github Action for Publishing to Comfy Registry 2024-10-05 08:59:46 +00:00
snomiao 970158b80f chore(pyproject): Add pyproject.toml for Custom Node Registry 2024-10-05 07:59:15 +00:00
Yuvraj Seegolam addc5be46d add node img 2024-10-04 14:23:26 +04:00
14 changed files with 483 additions and 314 deletions
+24
View File
@@ -0,0 +1,24 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
- master
paths:
- "pyproject.toml"
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
# if this is a forked repository. Skipping the workflow.
if: github.event.repository.fork == false
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+2 -1
View File
@@ -1,3 +1,4 @@
models
__pycache__
.vscode
.vscode
CLAUDE.md
+4 -83
View File
@@ -1,93 +1,14 @@
import torch
import os
from comfy.model_management import get_torch_device
from comfy.utils import ProgressBar
from .vfi_utilities import preprocess_frames, postprocess_frames, generate_frames_rife, logger
from .trt_utilities import Engine, MultiStreamEngine
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}),
"cuda_streams": ("INT", {"default": 1, "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,
cuda_streams=1,
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 = MultiStreamEngine(engine_path)
logger(f"Loading TensorRT engine: {engine_path}")
self.engine.load()
self.engine_label = engine
else:
logger(f"Using cached TensorRT engine: {engine_path}")
self.engine.set_num_streams(cuda_streams)
self.engine.activate()
logger(f"Cuda streams: {cuda_streams}")
self.engine.allocate_buffers(shape_dict=shape_dict)
logger("allocation done")
frames = preprocess_frames(frames)
logger("preprocessing done")
def return_middle_frame(data_batch):
# s = time.time()
results = self.engine.infer(data_batch)
# e = time.time()
# print(f"Time taken to infer: {(e-s)*1000} ms")
return results
result = generate_frames_rife(frames, clear_cache_after_n_frames, multiplier, return_middle_frame, cuda_streams)
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']
+16
View File
@@ -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."
}
}
+92
View File
@@ -0,0 +1,92 @@
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()
engine.model_name = model
return (engine,)
+58
View File
@@ -0,0 +1,58 @@
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)
engine.reset()
return (out,)
+15
View File
@@ -0,0 +1,15 @@
[project]
name = "comfyui-rife-tensorrt"
description = "This project provides a TensorRT implementation of [a/RIFE](https://github.com/hzwer/ECCV2022-RIFE) for ultra fast frame interpolation inside ComfyUI"
version = "1.0.0"
license = {file = "LICENSE"}
dependencies = ["einops", "colored", "polygraphy", "tensorrt==10.4.0", "cuda-python"]
[project.urls]
Repository = "https://github.com/yuvraj108c/ComfyUI-Rife-Tensorrt"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = ""
DisplayName = "ComfyUI-Rife-Tensorrt"
Icon = ""
+47 -20
View File
@@ -2,24 +2,42 @@
# ComfyUI Rife TensorRT ⚡
[![python](https://img.shields.io/badge/python-3.10.12-green)](https://www.python.org/downloads/release/python-31012/)
[![cuda](https://img.shields.io/badge/cuda-12.4-green)](https://developer.nvidia.com/cuda-downloads)
[![trt](https://img.shields.io/badge/TRT-10.4.0-green)](https://developer.nvidia.com/tensorrt)
[![python](https://img.shields.io/badge/python-3.12.3-green)](https://www.python.org/downloads/release/python-3123//)
[![cuda](https://img.shields.io/badge/cuda-13.0-green)](https://developer.nvidia.com/cuda-downloads)
[![trt](https://img.shields.io/badge/TRT-10.14.1.48-green)](https://developer.nvidia.com/tensorrt)
[![by-nc-sa/4.0](https://img.shields.io/badge/license-CC--BY--NC--SA--4.0-lightgrey)](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!
[![ComfyUI-Depth-Anything-Tensorrt](https://img.shields.io/badge/ComfyUI--Depth--Anything--Tensorrt-blue?style=flat-square)](https://github.com/yuvraj108c/ComfyUI-Depth-Anything-Tensorrt)
[![ComfyUI-Upscaler-Tensorrt](https://img.shields.io/badge/ComfyUI--Upscaler--Tensorrt-blue?style=flat-square)](https://github.com/yuvraj108c/ComfyUI-Upscaler-Tensorrt)
[![ComfyUI-Dwpose-Tensorrt](https://img.shields.io/badge/ComfyUI--Dwpose--Tensorrt-blue?style=flat-square)](https://github.com/yuvraj108c/ComfyUI-Dwpose-Tensorrt)
[![ComfyUI-Rife-Tensorrt](https://img.shields.io/badge/ComfyUI--Rife--Tensorrt-blue?style=flat-square)](https://github.com/yuvraj108c/ComfyUI-Rife-Tensorrt)
[![ComfyUI-Whisper](https://img.shields.io/badge/ComfyUI--Whisper-gray?style=flat-square)](https://github.com/yuvraj108c/ComfyUI-Whisper)
[![ComfyUI_InvSR](https://img.shields.io/badge/ComfyUI__InvSR-gray?style=flat-square)](https://github.com/yuvraj108c/ComfyUI_InvSR)
[![ComfyUI-Thera](https://img.shields.io/badge/ComfyUI--Thera-gray?style=flat-square)](https://github.com/yuvraj108c/ComfyUI-Thera)
[![ComfyUI-Video-Depth-Anything](https://img.shields.io/badge/ComfyUI--Video--Depth--Anything-gray?style=flat-square)](https://github.com/yuvraj108c/ComfyUI-Video-Depth-Anything)
[![ComfyUI-PiperTTS](https://img.shields.io/badge/ComfyUI--PiperTTS-gray?style=flat-square)](https://github.com/yuvraj108c/ComfyUI-PiperTTS)
[![buy-me-coffees](https://i.imgur.com/3MDbAtw.png)](https://www.buymeacoffee.com/yuvraj108cZ)
[![paypal-donation](https://i.imgur.com/w5jjubk.png)](https://paypal.me/yuvraj108c)
---
## ⏱️ Performance
_Note: The following results were benchmarked on FP16 engines inside ComfyUI, using 1000 frames consisting of 2 alternating similar frames, averaged 2-3 times_
_Note: The following results were benchmarked on FP16 engines inside ComfyUI, using 2000 frames consisting of 2 alternating similar frames, averaged 2-3 times_
| Device | Rife Engine | Resolution| Multiplier | FPS |
| :----: | :-: | :-: | :-: | :-: |
@@ -37,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
View File
@@ -1,5 +1,4 @@
einops
colored
polygraphy
tensorrt==10.4.0
cuda-python
einops
colored
polygraphy
tensorrt
+1 -1
View File
@@ -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
+3
View File
@@ -1,3 +1,6 @@
import sys
import os
sys.path.append(os.path.dirname(os.path.abspath(__file__)))
import torch
import time
from trt_utilities import Engine
+38 -130
View File
@@ -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,10 +33,6 @@ import tensorrt as trt
from logging import error, warning
from tqdm import tqdm
import copy
from collections import OrderedDict
from typing import List, Dict
from cuda import cudart
from polygraphy import cuda
TRT_LOGGER = trt.Logger(trt.Logger.ERROR)
G_LOGGER.module_severity = G_LOGGER.ERROR
@@ -47,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)
@@ -128,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,
@@ -147,12 +150,13 @@ class Engine:
del self.tensors
def reset(self, engine_path=None):
del self.engine
# del self.engine
del self.context
del self.buffers
del self.tensors
self.engine_path = engine_path
# self.engine_path = engine_path
self.context = None
self.buffers = OrderedDict()
self.tensors = OrderedDict()
self.inputs = {}
@@ -169,7 +173,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))]
@@ -220,6 +224,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):
@@ -249,122 +254,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())
class MultiStreamEngine(Engine):
def __init__(self, engine_path, num_streams=2):
super().__init__(engine_path)
self.num_streams = num_streams
self.streams = []
self.contexts = []
self.stream_tensors = []
def set_num_streams(self, value):
self.num_streams = value
def activate(self):
super().activate()
for _ in range(self.num_streams):
# stream = torch.cuda.Stream()
stream = cuda.Stream()
context = self.engine.create_execution_context()
self.streams.append(stream)
self.contexts.append(context)
self.stream_tensors.append(OrderedDict())
def get_torch_dtype(self, trt_dtype):
"""Convert TensorRT dtype to PyTorch dtype."""
dtype_map = {
trt.int8: torch.int8,
trt.int32: torch.int32,
trt.float16: torch.float16,
trt.float32: torch.float32,
trt.bool: torch.bool,
}
return dtype_map.get(trt_dtype, torch.float32) # Default to float32 if not found
def allocate_buffers(self, shape_dict=None, device="cuda"):
nvtx.range_push("allocate_buffers")
for stream_idx in range(self.num_streams):
tensors = self.stream_tensors[stream_idx]
context = self.contexts[stream_idx]
for idx in range(self.engine.num_io_tensors):
name = self.engine.get_tensor_name(idx)
binding = self.engine[idx]
if shape_dict and binding in shape_dict:
shape = shape_dict[binding]["shape"]
else:
shape = context.get_tensor_shape(name)
dtype = trt.nptype(self.engine.get_tensor_dtype(name))
if self.engine.get_tensor_mode(name) == trt.TensorIOMode.INPUT:
context.set_input_shape(name, shape)
torch_dtype = numpy_to_torch_dtype_dict[dtype]
tensor = torch.empty(tuple(shape), dtype=torch_dtype, device=device)
tensors[binding] = tensor
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 infer(self, feed_dicts: List[Dict[str, torch.Tensor]], use_cuda_graph=True):
results = []
num_batches = len(feed_dicts)
for i in range(0, num_batches, self.num_streams):
batch_results = []
for j in range(self.num_streams):
if i + j >= num_batches:
break
feed_dict = feed_dicts[i + j]
stream = self.streams[j]
context = self.contexts[0]
tensors = self.stream_tensors[0]
for name, buf in feed_dict.items():
tensors[name].copy_(buf)
if (i == 0):
for name, tensor in tensors.items():
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 = 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)
)
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.")
# success = context.execute_async_v3(stream.cuda_stream)
# if not success:
# raise RuntimeError(f"Inference failed for batch {i + j}")
batch_results.append({name: tensor for name, tensor in tensors.items()
if self.engine.get_tensor_mode(name) == trt.TensorIOMode.OUTPUT})
# Synchronize all used streams
for j in range(min(self.num_streams, num_batches - i)):
self.streams[j].synchronize()
results.extend(batch_results)
return results
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
View File
@@ -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
+38 -74
View File
@@ -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,110 +27,72 @@ 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")
def postprocess_frames(frames):
return einops.rearrange(frames, "n c h w -> n h w c")[..., :3].cpu()
def get_frame_0_idx(frame_itr):
frame_0_start_idx = frame_itr
frame_0_end_idx = frame_itr+1
return frame_0_start_idx, frame_0_end_idx
def get_frame_1_idx(frame_itr):
frame_1_start_idx = frame_itr+1
frame_1_end_idx = frame_itr+2
return frame_1_start_idx, frame_1_end_idx
def generate_frames_rife(
frames,
clear_cache_after_n_frames,
multiplier,
return_middle_frame_function,
cuda_streams
):
return_middle_frame_function
):
output_frames = torch.zeros(multiplier*frames.shape[0], *frames.shape[1:], device="cpu")
out_len = 0
number_of_frames_processed_since_last_cleared_cuda_cache = 0
pbar = ProgressBar(len(frames))
pbar = ProgressBar(len(frames)-1)
rife_inference_datas = []
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
# frame_0 = frames[frame_itr:frame_itr+1]
# # frame_1 = frames[frame_itr+1:frame_itr+2]
# output_frames[out_len] = frame_0 # Start with first frame
# out_len += 1
frame_0 = frames[frame_itr:frame_itr+1]
frame_1 = frames[frame_itr+1:frame_itr+2]
output_frames[out_len] = frame_0 # Start with first frame
out_len += 1
for middle_i in range(1, multiplier):
timestep = middle_i/multiplier
# rife_inference_datas.append({
# "frame_0": frame_0,
# "frame_1": frame_1,
# "timestep": timestep,
# })
middle_frame = return_middle_frame_function(frame_0, frame_1, timestep).detach().cpu()
rife_inference_datas.append({
"frame_0_idx":frame_itr,
"frame_1_idx":frame_itr + 1,
"timestep": timestep,
})
logger("data generated")
batch_size = cuda_streams
for data_idx in range(0, len(rife_inference_datas), batch_size):
current_batch_size = min(batch_size, len(rife_inference_datas) - data_idx)
data_batch = []
for batch_no in range(current_batch_size):
current_data_idx = data_idx + batch_no
rife_inference_data = rife_inference_datas[current_data_idx]
timestep = rife_inference_data["timestep"]
timestep_t = torch.tensor([timestep], dtype=torch.float32).to(get_torch_device())
frame_0_idx = rife_inference_data["frame_0_idx"]
frame_1_idx = rife_inference_data["frame_1_idx"]
frame_0 = frames[frame_0_idx].unsqueeze(0)#.to(DEVICE)
frame_1 = frames[frame_1_idx].unsqueeze(0)#to(DEVICE)
data_batch.append({
"img0": frame_0,
"img1": frame_1,
"timestep": timestep_t,
})
middle_frames = return_middle_frame_function(data_batch)
for middle_frame in middle_frames:
output_frames[out_len] = middle_frame['output'].detach().cpu()
# Copy middle frames to output
output_frames[out_len] = middle_frame
out_len +=1
# number_of_frames_processed_since_last_cleared_cuda_cache += 1
# # Try to avoid a memory overflow by clearing cuda cache regularly
# 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...")
pbar.update(current_batch_size)
# Try to avoid a memory overflow by clearing cuda cache regularly
number_of_frames_processed_since_last_cleared_cuda_cache += 1
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
# rife_logger.info("Clearing cache...") # spamming console + conflict with tqdm progress
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