From 547c3d5c9768a02ddb089656d2d8b120fa206612 Mon Sep 17 00:00:00 2001 From: Acly Date: Mon, 22 Jul 2024 16:53:09 +0200 Subject: [PATCH] Add text translation node & API --- __init__.py | 6 ++--- api.py | 20 +++++++++++++++ nsfw.py | 2 +- requirements.txt | 2 ++ translation.py | 64 ++++++++++++++++++++++++++++++++++++++++++++++++ 5 files changed, 90 insertions(+), 4 deletions(-) create mode 100644 requirements.txt create mode 100644 translation.py diff --git a/__init__.py b/__init__.py index f4ecd44..29cea77 100644 --- a/__init__.py +++ b/__init__.py @@ -1,4 +1,4 @@ -from . import api, nodes, tile, region, nsfw +from . import api, nodes, tile, region, nsfw, translation NODE_CLASS_MAPPINGS = { "ETN_LoadImageBase64": nodes.LoadImageBase64, @@ -16,6 +16,7 @@ NODE_CLASS_MAPPINGS = { "ETN_ListRegionMasks": region.ListRegionMasks, "ETN_AttentionMask": region.AttentionMask, "ETN_NSFWFilter": nsfw.NSFWFilter, + "ETN_Translate": translation.Translate, } NODE_DISPLAY_NAME_MAPPINGS = { "ETN_LoadImageBase64": "Load Image (Base64)", @@ -23,8 +24,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ETN_SendImageWebSocket": "Send Image (WebSocket)", "ETN_CropImage": "Crop Image", "ETN_ApplyMaskToImage": "Apply Mask to Image", - "ETN_ListAppend": "List 🢒 Append", - "ETN_ListElement": "List 🢒 Get Element", "ETN_TileLayout": "Create Tile Layout", "ETN_ExtractImageTile": "Extract Image Tile", "ETN_ExtractMaskTile": "Extract Mask Tile", @@ -35,4 +34,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ETN_ListRegionMasks": "List Region Masks", "ETN_AttentionMask": "Regions Attention Mask", "ETN_NSFWFilter": "NSFW Filter", + "ETN_Translate": "Translate Text", } diff --git a/api.py b/api.py index 527b06f..f5b947d 100644 --- a/api.py +++ b/api.py @@ -8,6 +8,8 @@ from comfy import model_detection import folder_paths import server +from .translation import available_languages, translate + input_block_name = "model.diffusion_model.input_blocks.0.0.weight" model_names = { @@ -95,3 +97,21 @@ if _server := getattr(server.PromptServer, "instance", None): @_server.routes.get("/api/etn/model_info") async def api_model_info(request): return await model_info(request) + + @_server.routes.get("/api/etn/languages") + async def languages(request): + try: + result = [dict(name=name, code=code) for code, name in available_languages()] + return web.json_response(result) + except Exception as e: + return web.json_response(dict(error=str(e)), status=500) + + @_server.routes.get("/api/etn/translate/{lang}/{text}") + async def translate_text(request): + try: + language = request.match_info.get("lang", "en") + text = request.match_info.get("text", "") + result = translate(text, language) + return web.json_response(result) + except Exception as e: + return web.json_response(dict(error=str(e)), status=500) diff --git a/nsfw.py b/nsfw.py index 4e9824f..755223d 100644 --- a/nsfw.py +++ b/nsfw.py @@ -68,7 +68,7 @@ class CLIPSafetyChecker(PreTrainedModel): for idx in images_to_filter: filtered = images[idx].unsqueeze(0) filtered = F.interpolate(filtered, size=64, mode="nearest") - filtered = box_blur(filtered, 7, separable=True) + filtered = box_blur(filtered, 11, separable=True) filtered = F.interpolate(filtered, size=orig_size, mode="bilinear") images[idx] = filtered.squeeze(0) return images diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..7103e63 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,2 @@ +# Optional, only required for Translate node: +argostranslate diff --git a/translation.py b/translation.py new file mode 100644 index 0000000..13547ce --- /dev/null +++ b/translation.py @@ -0,0 +1,64 @@ +"""Text translation using Argos Translate.""" + +from functools import cache + + +@cache +def available_languages(): + try: + from argostranslate.package import update_package_index, get_available_packages + + update_package_index() + list = get_available_packages() + return [(l.from_code, l.from_name) for l in list if l.to_code == "en"] + except ImportError: + return [("NOT INSTALLED", "NOT INSTALLED")] + + +def translate(text: str, language: str): + if text.strip() == "": + return text + + target = "en" + if language == target: + return text + + try: + from argostranslate.package import get_installed_packages, get_available_packages + from argostranslate.translate import translate + + installed = get_installed_packages() + if not any(p.from_code == language and p.to_code == target for p in installed): + available = get_available_packages() + pkg = next( + (p for p in available if p.from_code == language and p.to_code == target), None + ) + assert pkg, f"Couldn't find package for translation from {language}" + print("Downloading and installing translation package", pkg) + pkg.install() + + return translate(text, language, target) + + except ImportError: + raise ImportError( + "Argos Translate is not installed. Please install it with `pip install argostranslate`" + ) + + +class Translate: + @staticmethod + def INPUT_TYPES(): + return { + "required": { + "text": ("STRING", {"multiline": True}), + "language": ([code for code, _ in available_languages()],), + "target": (["en"],), + } + } + + CATEGORY = "external_tooling" + RETURN_TYPES = ("STRING",) + FUNCTION = "translate" + + def translate(self, text: str, language: str, target: str): + return (translate(text, language),)