Fix broken callbacks.py

This commit is contained in:
oobabooga 2023-03-23 22:12:24 -03:00 committed by GitHub
parent 9bdb3c784d
commit d1327f99f9
No known key found for this signature in database
GPG Key ID: 4AEE18F83AFDEB23

View File

@ -4,8 +4,6 @@ from threading import Thread
import torch import torch
import transformers import transformers
from modules.text_generation import clear_torch_cache
# Copied from https://github.com/PygmalionAI/gradio-ui/ # Copied from https://github.com/PygmalionAI/gradio-ui/
class _SentinelTokenStoppingCriteria(transformers.StoppingCriteria): class _SentinelTokenStoppingCriteria(transformers.StoppingCriteria):
@ -89,3 +87,8 @@ class Iteratorize:
def __exit__(self, exc_type, exc_val, exc_tb): def __exit__(self, exc_type, exc_val, exc_tb):
self.stop_now = True self.stop_now = True
clear_torch_cache() clear_torch_cache()
def clear_torch_cache():
gc.collect()
if not shared.args.cpu:
torch.cuda.empty_cache()