Add text translation node & API
This commit is contained in:
+3
-3
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
# Optional, only required for Translate node:
|
||||
argostranslate
|
||||
@@ -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),)
|
||||
Reference in New Issue
Block a user