From 116ce22440e9b8933b75c8c7d4496a7555da6a6c Mon Sep 17 00:00:00 2001 From: yuvraj108c Date: Thu, 3 Oct 2024 23:11:13 +0400 Subject: [PATCH] keep_model_loaded --- .gitignore | 3 ++- __init__.py | 13 +++++++------ 2 files changed, 9 insertions(+), 7 deletions(-) diff --git a/.gitignore b/.gitignore index 10a149b..236f609 100644 --- a/.gitignore +++ b/.gitignore @@ -1,2 +1,3 @@ models -__pycache__ \ No newline at end of file +__pycache__ +.vscode \ No newline at end of file diff --git a/__init__.py b/__init__.py index 27eb812..48d4e96 100644 --- a/__init__.py +++ b/__init__.py @@ -1,11 +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 import folder_paths -import time ENGINE_DIR = os.path.join(folder_paths.models_dir, "tensorrt", "rife") @@ -18,6 +16,7 @@ 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}), + "keep_model_loaded": ("BOOLEAN", {"default": False}), }, } @@ -30,7 +29,8 @@ class RifeTensorrt: frames, engine, clear_cache_after_n_frames=50, - multiplier=2 + multiplier=2, + keep_model_loaded=False ): B, H, W, C = frames.shape shape_dict = { @@ -57,16 +57,17 @@ class RifeTensorrt: 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) - # e = time.time() - # print(f"Time taken to infer: {(e-s)*1000} ms") result = output['output'] return result result = generate_frames_rife(frames, clear_cache_after_n_frames, multiplier, return_middle_frame) out = postprocess_frames(result) + + if not keep_model_loaded: + del self.engine, self.engine_label + return (out,)