model resetting
This commit is contained in:
@@ -87,5 +87,6 @@ class LoadRifeTensorrtModel:
|
||||
mm.soft_empty_cache()
|
||||
engine = Engine(tensorrt_model_path)
|
||||
engine.load()
|
||||
engine.model_name = model
|
||||
|
||||
return (engine,)
|
||||
|
||||
@@ -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,)
|
||||
+3
-2
@@ -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 = {}
|
||||
|
||||
+1
-1
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user