keep_model_loaded
This commit is contained in:
+2
-1
@@ -1,2 +1,3 @@
|
||||
models
|
||||
__pycache__
|
||||
__pycache__
|
||||
.vscode
|
||||
+7
-6
@@ -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,)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user