From 7491b548d1c1f0b98bdc0d470148892549f186cd Mon Sep 17 00:00:00 2001 From: yuvraj108c Date: Fri, 4 Oct 2024 13:30:05 +0400 Subject: [PATCH] cuda grapsh --- __init__.py | 36 ++++++++++++++++++++++++++++++++++ requirements.txt | 3 ++- trt_utilities.py | 50 ++++++++++++++++++++++++++++++++---------------- 3 files changed, 72 insertions(+), 17 deletions(-) diff --git a/__init__.py b/__init__.py index 48d4e96..b9f444a 100644 --- a/__init__.py +++ b/__init__.py @@ -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) diff --git a/requirements.txt b/requirements.txt index a0e8f28..b479d99 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,5 @@ einops colored polygraphy -tensorrt==10.4.0 \ No newline at end of file +tensorrt==10.4.0 +cuda-python \ No newline at end of file diff --git a/trt_utilities.py b/trt_utilities.py index ebe5649..fcc1f01 100644 --- a/trt_utilities.py +++ b/trt_utilities.py @@ -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