Add text translation node & API

This commit is contained in:
Acly
2024-07-22 16:53:09 +02:00
parent 73babbd00e
commit 547c3d5c97
5 changed files with 90 additions and 4 deletions
+3 -3
View File
@@ -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",
}
+20
View File
@@ -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)
+1 -1
View File
@@ -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
+2
View File
@@ -0,0 +1,2 @@
# Optional, only required for Translate node:
argostranslate
+64
View File
@@ -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),)