4 Commits
Author SHA1 Message Date
yuvraj108c 8a4ee8faef multi stream cuda graph 2024-10-04 14:12:56 +04:00
yuvraj108c 1b629c50c1 - 2024-10-04 13:52:07 +04:00
yuvraj108c 11b6a8fde8 cuda graphs multi stream 2024-10-04 13:50:17 +04:00
yuvraj108c 272d16ab6b multi stream 2024-10-04 11:49:16 +04:00
3 changed files with 198 additions and 51 deletions
+16 -12
View File
@@ -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
View File
@@ -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
View File
@@ -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