Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8a4ee8faef | ||
|
|
1b629c50c1 | ||
|
|
11b6a8fde8 | ||
|
|
272d16ab6b |
+16
-12
@@ -1,8 +1,9 @@
|
||||
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
|
||||
from .trt_utilities import Engine, MultiStreamEngine
|
||||
import folder_paths
|
||||
import time
|
||||
from polygraphy import cuda
|
||||
@@ -18,6 +19,7 @@ class RifeTensorrt:
|
||||
"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}),
|
||||
},
|
||||
@@ -34,6 +36,7 @@ class RifeTensorrt:
|
||||
engine,
|
||||
clear_cache_after_n_frames=100,
|
||||
multiplier=2,
|
||||
cuda_streams=1,
|
||||
use_cuda_graph=True,
|
||||
keep_model_loaded=False,
|
||||
):
|
||||
@@ -47,31 +50,32 @@ class RifeTensorrt:
|
||||
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)
|
||||
self.engine = MultiStreamEngine(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.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)
|
||||
|
||||
def return_middle_frame(frame_0, frame_1, timestep):
|
||||
timestep_t = torch.tensor([timestep], dtype=torch.float32).to(get_torch_device())
|
||||
logger("preprocessing done")
|
||||
def return_middle_frame(data_batch):
|
||||
# s = time.time()
|
||||
output = self.engine.infer({"img0": frame_0, "img1": frame_1, "timestep": timestep_t}, cudaStream, use_cuda_graph)
|
||||
results = self.engine.infer(data_batch)
|
||||
# e = time.time()
|
||||
# print(f"Time taken to infer: {(e-s)*1000} ms")
|
||||
|
||||
result = output['output']
|
||||
return result
|
||||
return results
|
||||
|
||||
result = generate_frames_rife(frames, clear_cache_after_n_frames, multiplier, return_middle_frame)
|
||||
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
|
||||
|
||||
|
||||
+115
-21
@@ -16,7 +16,10 @@ 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
|
||||
@@ -252,25 +255,116 @@ class Engine:
|
||||
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.")
|
||||
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 = []
|
||||
|
||||
return self.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()
|
||||
|
||||
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
|
||||
|
||||
+67
-18
@@ -34,12 +34,23 @@ def preprocess_frames(frames):
|
||||
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
|
||||
):
|
||||
return_middle_frame_function,
|
||||
cuda_streams
|
||||
):
|
||||
|
||||
output_frames = torch.zeros(multiplier*frames.shape[0], *frames.shape[1:], device="cpu")
|
||||
out_len = 0
|
||||
@@ -47,30 +58,68 @@ def generate_frames_rife(
|
||||
number_of_frames_processed_since_last_cleared_cuda_cache = 0
|
||||
pbar = ProgressBar(len(frames))
|
||||
|
||||
rife_inference_datas = []
|
||||
|
||||
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
|
||||
middle_frame = return_middle_frame_function(frame_0, frame_1, timestep).detach().cpu()
|
||||
# rife_inference_datas.append({
|
||||
# "frame_0": frame_0,
|
||||
# "frame_1": frame_1,
|
||||
# "timestep": timestep,
|
||||
# })
|
||||
|
||||
# Copy middle frames to output
|
||||
output_frames[out_len] = middle_frame
|
||||
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()
|
||||
out_len +=1
|
||||
# number_of_frames_processed_since_last_cleared_cuda_cache += 1
|
||||
|
||||
# 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
|
||||
logger("Clearing cache...")
|
||||
# # 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)
|
||||
|
||||
pbar.update(1)
|
||||
|
||||
|
||||
# Append final frame
|
||||
output_frames[out_len] = frames[-1:]
|
||||
@@ -80,6 +129,6 @@ def generate_frames_rife(
|
||||
# clear cache for courtesy
|
||||
soft_empty_cache()
|
||||
logger("Final clearing cache done ...")
|
||||
#
|
||||
|
||||
res = output_frames[:out_len]
|
||||
return res
|
||||
Reference in New Issue
Block a user