diff --git a/nodes/load_rife_tensorrt.py b/nodes/load_rife_tensorrt.py index 3be2d8b..0798b31 100644 --- a/nodes/load_rife_tensorrt.py +++ b/nodes/load_rife_tensorrt.py @@ -87,5 +87,6 @@ class LoadRifeTensorrtModel: mm.soft_empty_cache() engine = Engine(tensorrt_model_path) engine.load() + engine.model_name = model return (engine,) diff --git a/nodes/rife_tensorrt.py b/nodes/rife_tensorrt.py index f85d529..b7d152d 100644 --- a/nodes/rife_tensorrt.py +++ b/nodes/rife_tensorrt.py @@ -53,4 +53,6 @@ class RifeTensorrt: result = generate_frames_rife(frames, clear_cache_after_n_frames, multiplier, return_middle_frame) out = postprocess_frames(result) + engine.reset() + return (out,) \ No newline at end of file diff --git a/trt_utilities.py b/trt_utilities.py index 3ea5567..bbac1fc 100644 --- a/trt_utilities.py +++ b/trt_utilities.py @@ -150,12 +150,13 @@ class Engine: del self.tensors def reset(self, engine_path=None): - del self.engine + # del self.engine del self.context del self.buffers del self.tensors - self.engine_path = engine_path + # self.engine_path = engine_path + self.context = None self.buffers = OrderedDict() self.tensors = OrderedDict() self.inputs = {} diff --git a/vfi_utilities.py b/vfi_utilities.py index 88d7b16..3e89cb3 100644 --- a/vfi_utilities.py +++ b/vfi_utilities.py @@ -74,7 +74,7 @@ def generate_frames_rife( 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 - rife_logger.info("Clearing cache...") + # rife_logger.info("Clearing cache...") # spamming console + conflict with tqdm progress pbar.update(1) progress_bar.update(1)