Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 9 additions & 5 deletions openlrc/validators.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,11 +5,15 @@
import re

from json_repair import repair_json
from langcodes import Language
from langcodes import Language as Langcode
from lingua import Language as LinguaLanguage
from lingua import LanguageDetectorBuilder

from openlrc.defaults import supported_languages_lingua
from openlrc.logger import logger

_LINGUA_LANGUAGES = [getattr(LinguaLanguage, name) for name in supported_languages_lingua]

ORIGINAL_PREFIX = "Original>"
TRANSLATION_PREFIX = "Translation>"
PROOFREAD_PREFIX = "Proofread>"
Expand All @@ -33,7 +37,7 @@ def validate(self, user_input, generated_content):

class ChunkedTranslateValidator(BaseValidator):
def __init__(self, target_lang):
self.lan_detector = LanguageDetectorBuilder.from_all_languages().build()
self.lan_detector = LanguageDetectorBuilder.from_languages(*_LINGUA_LANGUAGES).build()
self.target_lang = target_lang

def _extract_translation(self, content: str) -> list[str]:
Expand Down Expand Up @@ -64,7 +68,7 @@ def _is_translation_in_target_language(self, translation: list[str]) -> bool:
return True
translated_lang = detected_lang.name.lower()

target_lang = Language.get(self.target_lang).language_name().lower()
target_lang = Langcode.get(self.target_lang).language_name().lower()
if translated_lang != target_lang:
logger.warning(f"Translated language is {translated_lang}, not {target_lang}.")
return False
Expand Down Expand Up @@ -111,7 +115,7 @@ def validate(self, user_input, generated_content):

class AtomicTranslateValidator(BaseValidator):
def __init__(self, target_lang):
self.lan_detector = LanguageDetectorBuilder.from_all_languages().build()
self.lan_detector = LanguageDetectorBuilder.from_languages(*_LINGUA_LANGUAGES).build()
self.target_lang = target_lang

def validate(self, user_input, generated_content):
Expand All @@ -124,7 +128,7 @@ def validate(self, user_input, generated_content):
return True

translated_lang = detected_lang.name.lower()
target_lang = Language.get(self.target_lang).language_name().lower()
target_lang = Langcode.get(self.target_lang).language_name().lower()
if translated_lang != target_lang:
logger.warning(f'Translated text: "{generated_content}" is {translated_lang}, not {target_lang}.')
return False
Expand Down
Loading