From 260d8ebbba36659c7605ea273c6130ce9fbb2dbf Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=D0=95=D0=B2=D0=B3=D0=B5=D0=BD=D0=B8=D0=B9=20=D0=93=D1=83?= =?UTF-8?q?=D1=80=D1=8C=D0=B5=D0=B2=20=7C=20Eugene=20Gourieff=20=7C=20?= =?UTF-8?q?=E5=8F=A4=E4=BB=81?= Date: Sat, 18 Jan 2025 13:31:35 +0700 Subject: [PATCH] UPD: Processing state --- nodes.py | 11 +++++++++-- scripts/reactor_sfw.py | 8 ++++---- 2 files changed, 13 insertions(+), 6 deletions(-) diff --git a/nodes.py b/nodes.py index 9529341..baea838 100644 --- a/nodes.py +++ b/nodes.py @@ -361,11 +361,18 @@ class reactor: pil_images = batch_tensor_to_pil(input_image) # NSFW checker + logger.status("Checking for any unsafe content") pil_images_sfw = [] + tmp_img = "reactor_tmp.png" for img in pil_images: - img.save("tmp.png") - if not sfw.nsfw_image("tmp.png",NSFWDET_MODEL_PATH): + if state.interrupted or model_management.processing_interrupted(): + logger.status("Interrupted by User") + break + img.save(tmp_img) + if not sfw.nsfw_image(tmp_img, NSFWDET_MODEL_PATH): pil_images_sfw.append(img) + if os.path.exists(tmp_img): + os.remove(tmp_img) pil_images = pil_images_sfw # # # diff --git a/scripts/reactor_sfw.py b/scripts/reactor_sfw.py index 029ea55..474188d 100644 --- a/scripts/reactor_sfw.py +++ b/scripts/reactor_sfw.py @@ -7,7 +7,7 @@ SCORE = 0.85 logging.getLogger('transformers').setLevel(logging.ERROR) def nsfw_image(img_path: str, model_path: str): - img = Image.open(img_path) - predict = pipeline("image-classification", model=model_path) - result = predict(img) - return True if result[0]["score"] > SCORE else False + with Image.open(img_path) as img: + predict = pipeline("image-classification", model=model_path) + result = predict(img) + return True if result[0]["score"] > SCORE else False