cuda grapsh

This commit is contained in:
yuvraj108c
2024-10-04 13:30:05 +04:00
parent 116ce22440
commit 7491b548d1
3 changed files with 72 additions and 17 deletions
+36
View File
@@ -4,6 +4,11 @@ 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
<<<<<<< Updated upstream
=======
import time
from polygraphy import cuda
>>>>>>> Stashed changes
ENGINE_DIR = os.path.join(folder_paths.models_dir, "tensorrt", "rife")
@@ -16,6 +21,10 @@ class RifeTensorrt:
"engine": (os.listdir(ENGINE_DIR),),
"clear_cache_after_n_frames": ("INT", {"default": 50, "min": 1, "max": 1000}),
"multiplier": ("INT", {"default": 2, "min": 1}),
<<<<<<< Updated upstream
=======
"use_cuda_graph": ("BOOLEAN", {"default": True}),
>>>>>>> Stashed changes
"keep_model_loaded": ("BOOLEAN", {"default": False}),
},
}
@@ -23,6 +32,10 @@ class RifeTensorrt:
RETURN_TYPES = ("IMAGE", )
FUNCTION = "vfi"
CATEGORY = "tensorrt"
<<<<<<< Updated upstream
=======
OUTPUT_NODE=True
>>>>>>> Stashed changes
def vfi(
self,
@@ -30,7 +43,12 @@ class RifeTensorrt:
engine,
clear_cache_after_n_frames=50,
multiplier=2,
<<<<<<< Updated upstream
keep_model_loaded=False
=======
use_cuda_graph=True,
keep_model_loaded=False,
>>>>>>> Stashed changes
):
B, H, W, C = frames.shape
shape_dict = {
@@ -39,8 +57,12 @@ class RifeTensorrt:
"output": {"shape": (1, 3, H, W)},
}
<<<<<<< Updated upstream
# cache tensorrt engine in memory
cudaStream = torch.cuda.current_stream().cuda_stream
=======
cudaStream = cuda.Stream()
>>>>>>> Stashed changes
engine_path = os.path.join(ENGINE_DIR, engine)
if (not hasattr(self, 'engine') or self.engine_label != engine):
self.engine = Engine(engine_path)
@@ -54,6 +76,7 @@ class RifeTensorrt:
self.engine.allocate_buffers(shape_dict=shape_dict)
frames = preprocess_frames(frames)
<<<<<<< Updated upstream
def return_middle_frame(frame_0, frame_1, timestep):
timestep_t = torch.tensor([timestep], dtype=torch.float32).to(get_torch_device())
@@ -62,6 +85,19 @@ class RifeTensorrt:
result = output['output']
return result
=======
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
>>>>>>> Stashed changes
result = generate_frames_rife(frames, clear_cache_after_n_frames, multiplier, return_middle_frame)
out = postprocess_frames(result)
+2 -1
View File
@@ -1,4 +1,5 @@
einops
colored
polygraphy
tensorrt==10.4.0
tensorrt==10.4.0
cuda-python
+34 -16
View File
@@ -16,7 +16,7 @@ import tensorrt as trt
from logging import error, warning
from tqdm import tqdm
import copy
# from .vfi_utilities import logger
from cuda import cudart
TRT_LOGGER = trt.Logger(trt.Logger.ERROR)
G_LOGGER.module_severity = G_LOGGER.ERROR
@@ -44,6 +44,17 @@ 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)
@@ -114,7 +125,6 @@ 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,
@@ -207,7 +217,6 @@ class Engine:
return 0
def load(self):
# logger(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):
@@ -237,22 +246,31 @@ 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())
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):
for idx in range(self.engine.num_io_tensors):
name = self.engine.get_tensor_name(idx)
shape = self.context.get_tensor_shape(name)
print(name, shape)
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.")
return self.tensors