keep_model_loaded

This commit is contained in:
yuvraj108c
2024-10-03 23:11:13 +04:00
parent 333d068dc1
commit 116ce22440
2 changed files with 9 additions and 7 deletions
+2 -1
View File
@@ -1,2 +1,3 @@
models
__pycache__
__pycache__
.vscode
+7 -6
View File
@@ -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,)