hysts's picture
hysts HF Staff
Update
b7cb245
Raw
History Blame Contribute Delete
7.77 kB
#!/usr/bin/env python
import os
import gradio as gr
import spaces
import torch
from transformers import AutoModelForImageTextToText, AutoProcessor
MAX_NEW_TOKENS_LIMIT = int(os.getenv("MAX_NEW_TOKENS_LIMIT", "2000"))
MAX_NEW_TOKENS_DEFAULT = int(os.getenv("MAX_NEW_TOKENS_DEFAULT", "400"))
DEFAULT_SOURCE_LANG = os.getenv("DEFAULT_SOURCE_LANG", "en")
DEFAULT_TARGET_LANG = os.getenv("DEFAULT_TARGET_LANG", "fr_FR")
ALLOWED_MODEL_IDS = {
"google/translategemma-4b-it",
"google/translategemma-12b-it",
"google/translategemma-27b-it",
}
model_id = os.getenv("MODEL_ID", "google/translategemma-27b-it")
if model_id not in ALLOWED_MODEL_IDS:
msg = f"MODEL_ID={model_id!r} is not supported. Choose one of: {', '.join(sorted(ALLOWED_MODEL_IDS))}"
raise ValueError(msg)
processor = AutoProcessor.from_pretrained(model_id)
model = AutoModelForImageTextToText.from_pretrained(model_id, device_map="auto")
LANG_CODE_TO_NAME = {
"en": "English",
"ar_EG": "Arabic (Egypt)",
"ar_SA": "Arabic (Saudi Arabia)",
"bg_BG": "Bulgarian",
"bn_IN": "Bengali (India)",
"ca_ES": "Catalan (Spain)",
"cs_CZ": "Czech (Czechia)",
"da_DK": "Danish (Denmark)",
"de_DE": "German (Germany)",
"el_GR": "Greek (Greece)",
"es_MX": "Spanish (Mexico)",
"et_EE": "Estonian (Estonia)",
"fa_IR": "Persian (Iran)",
"fi_FI": "Finnish (Finland)",
"fil_PH": "Filipino (Philippines)",
"fr_CA": "French (Canada)",
"fr_FR": "French (France)",
"gu_IN": "Gujarati (India)",
"he_IL": "Hebrew (Israel)",
"hi_IN": "Hindi (India)",
"hr_HR": "Croatian (Croatia)",
"hu_HU": "Hungarian (Hungary)",
"id_ID": "Indonesian (Indonesia)",
"is_IS": "Icelandic (Iceland)",
"it_IT": "Italian (Italy)",
"ja_JP": "Japanese (Japan)",
"kn_IN": "Kannada (India)",
"ko_KR": "Korean (Korea)",
"lt_LT": "Lithuanian (Lithuania)",
"lv_LV": "Latvian (Latvia)",
"ml_IN": "Malayalam (India)",
"mr_IN": "Marathi (India)",
"nl_NL": "Dutch (Netherlands)",
"no_NO": "Norwegian (Norway)",
"pa_IN": "Punjabi (India)",
"pl_PL": "Polish (Poland)",
"pt_BR": "Portuguese (Brazil)",
"pt_PT": "Portuguese (Portugal)",
"ro_RO": "Romanian (Romania)",
"ru_RU": "Russian (Russia)",
"sk_SK": "Slovak (Slovakia)",
"sl_SI": "Slovenian (Slovenia)",
"sr_RS": "Serbian (Serbia)",
"sv_SE": "Swedish (Sweden)",
"sw_KE": "Swahili (Kenya)",
"sw_TZ": "Swahili (Tanzania)",
"ta_IN": "Tamil (India)",
"te_IN": "Telugu (India)",
"th_TH": "Thai (Thailand)",
"tr_TR": "Turkish (Turkey)",
"uk_UA": "Ukrainian (Ukraine)",
"ur_PK": "Urdu (Pakistan)",
"vi_VN": "Vietnamese (Vietnam)",
"zh_CN": "Chinese (Simplified, China)",
"zh_TW": "Chinese (Traditional, Taiwan)",
"zu_ZA": "Zulu (South Africa)",
}
LANG_CHOICES = [
("English (en)", "en"),
*sorted(
[(f"{name} ({code})", code) for code, name in LANG_CODE_TO_NAME.items() if code != "en"],
key=lambda x: x[0].lower(),
),
]
_valid_lang_codes = set(LANG_CODE_TO_NAME.keys())
if DEFAULT_SOURCE_LANG not in _valid_lang_codes:
msg = f"DEFAULT_SOURCE_LANG={DEFAULT_SOURCE_LANG!r} is not a supported language code."
raise ValueError(msg)
if DEFAULT_TARGET_LANG not in _valid_lang_codes:
msg = f"DEFAULT_TARGET_LANG={DEFAULT_TARGET_LANG!r} is not a supported language code."
raise ValueError(msg)
def _build_messages(text: str, source_lang_code: str, target_lang_code: str) -> list[dict]:
# TODO:Remove this once the chat template is fixed. # noqa: FIX002, TD002, TD003, TD007
# Temporary workaround: Model expects zh_CH instead of zh_CN due to a bug
if source_lang_code == "zh_CN":
source_lang_code = "zh_CH"
if target_lang_code == "zh_CN":
target_lang_code = "zh_CH"
return [
{
"role": "user",
"content": [
{
"type": "text",
"source_lang_code": source_lang_code,
"target_lang_code": target_lang_code,
"text": text,
}
],
}
]
def swap_languages(source: str, target: str) -> tuple[str, str]:
return target, source
def count_tokens(text: str, source_lang_code: str, target_lang_code: str) -> str:
"""Count input tokens without GPU. Returns a short info string."""
if not text:
return ""
if source_lang_code not in LANG_CODE_TO_NAME or target_lang_code not in LANG_CODE_TO_NAME:
return ""
messages = _build_messages(text, source_lang_code, target_lang_code)
inputs = processor.apply_chat_template(messages, tokenize=True, add_generation_prompt=True, return_dict=True)
return f"Input tokens: {len(inputs['input_ids'][0])}"
@spaces.GPU(duration=120)
@torch.inference_mode()
def translate(text: str, source_lang_code: str, target_lang_code: str, max_new_tokens: int) -> str:
if not text:
raise gr.Error("Please enter text to translate")
if source_lang_code not in LANG_CODE_TO_NAME:
error_message = f"Invalid source language: {source_lang_code}"
raise gr.Error(error_message)
if target_lang_code not in LANG_CODE_TO_NAME:
error_message = f"Invalid target language: {target_lang_code}"
raise gr.Error(error_message)
messages = _build_messages(text, source_lang_code, target_lang_code)
inputs = processor.apply_chat_template(
messages, tokenize=True, add_generation_prompt=True, return_dict=True, return_tensors="pt"
).to(model.device, dtype=torch.bfloat16)
input_len = len(inputs["input_ids"][0])
if input_len + max_new_tokens > MAX_NEW_TOKENS_LIMIT:
error_message = f"Input ({input_len} tokens) + max output ({max_new_tokens} tokens) exceeds the total limit of {MAX_NEW_TOKENS_LIMIT} tokens."
raise gr.Error(error_message)
generation = model.generate(
**inputs, do_sample=False, max_new_tokens=max_new_tokens, cache_implementation="dynamic"
)
generation = generation[0][input_len:]
return processor.decode(generation, skip_special_tokens=True)
with gr.Blocks() as demo:
gr.Markdown("# Translategemma-27b-it")
with gr.Row():
source_lang_code = gr.Dropdown(label="Source Language", choices=LANG_CHOICES, value=DEFAULT_SOURCE_LANG)
swap_button = gr.Button("⇆", elem_id="swap-btn")
target_lang_code = gr.Dropdown(label="Target Language", choices=LANG_CHOICES, value=DEFAULT_TARGET_LANG)
max_new_tokens = gr.Slider(
label="Max New Tokens",
info="Higher values allow longer translations but take more time",
minimum=50,
maximum=MAX_NEW_TOKENS_LIMIT,
step=10,
value=MAX_NEW_TOKENS_DEFAULT,
)
with gr.Row():
with gr.Column():
text = gr.Textbox(label="Input", lines=10, placeholder="Enter text to translate")
token_info = gr.Textbox(label="Token Count", lines=1)
translate_button = gr.Button("Translate", variant="primary")
with gr.Column():
output = gr.Textbox(label="Translation", lines=10, placeholder="Translation will appear here")
token_count_inputs = [text, source_lang_code, target_lang_code]
for component in token_count_inputs:
component.change(fn=count_tokens, inputs=token_count_inputs, outputs=token_info)
swap_button.click(
fn=swap_languages,
inputs=[source_lang_code, target_lang_code],
outputs=[source_lang_code, target_lang_code],
)
translate_button.click(
fn=translate,
inputs=[text, source_lang_code, target_lang_code, max_new_tokens],
outputs=output,
)
if __name__ == "__main__":
demo.launch(css_paths="style.css")