cuda grapsh
This commit is contained in:
+36
@@ -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
@@ -1,4 +1,5 @@
|
||||
einops
|
||||
colored
|
||||
polygraphy
|
||||
tensorrt==10.4.0
|
||||
tensorrt==10.4.0
|
||||
cuda-python
|
||||
+34
-16
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user