Vernacular / main /pipeline /clients.py
bhardwaj08sarthak's picture
Upload folder using huggingface_hub
478fb0c verified
Raw
History Blame
11.9 kB
"""In-process Hugging Face clients for the two translation models.
Stage 1 - `translate`: google/translategemma-12b-it via its official
TranslateGemma chat template.
Stage 2 - `adjust_tone`: google/gemma-4-12B-it via the normal Gemma 4 chat
template, rewriting a draft translation in a character's voice using their
wiki.
"""
from __future__ import annotations
import gc
import re
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Any
sys.path.insert(0, str(Path(__file__).parent.parent))
import config
try:
import spaces
except ImportError: # Local/dev installs do not need the Spaces runtime.
spaces = None
# Runtime-injected tokens like "Person1" that must survive translation verbatim.
_PLACEHOLDER_RE = re.compile(r"\b([A-Za-z]+)(\d+)\b")
TONE_SYSTEM_TEMPLATE = """\
You are a localization editor for "Riverstone", a narrative mystery mobile game \
told through phone chats. Below is the voice wiki for {name}, a game character.
{wiki}
You will receive one of {name}'s chat lines in English and a draft {language} \
translation. Rewrite the {language} draft so it reads like {name} texting in \
{language}: match the wiki's tone, register, slang level, emoji and punctuation \
habits.
Rules:
- Keep the exact meaning of the English line; never add or drop information.
- Keep placeholder tokens (e.g. Person1), proper names, emoji and special \
symbols exactly as written.
- Use the formality level the wiki implies for {name} (casual characters use \
informal address).
- If the draft already sounds right, return it unchanged.
- Output ONLY the final {language} line - no quotes, no commentary.
"""
TONE_USER_TEMPLATE = """\
{context}English line: {source}
Draft {language} translation: {draft}
Final {language} line:"""
WIKI_MAX_NEW_TOKENS = 1200
@dataclass
class _LoadedModel:
processor: Any
model: Any
_MODELS: dict[str, _LoadedModel] = {}
_MODEL_IDS = {
"translate": config.TRANSLATE_MODEL_ID,
"tone": config.TONE_MODEL_ID,
}
def _gpu(fn):
if spaces is None:
return fn
return spaces.GPU(
duration=config.ZERO_GPU_DURATION_S,
size=config.ZERO_GPU_SIZE,
)(fn)
def _torch():
try:
import torch
except ImportError as exc:
raise RuntimeError(
"Hugging Face inference requires torch. Install the Space/runtime "
"dependencies from requirements.txt."
) from exc
return torch
def _hf_classes():
try:
from transformers import AutoModelForMultimodalLM, AutoProcessor
return AutoProcessor, AutoModelForMultimodalLM
except ImportError:
try:
from transformers import AutoModelForImageTextToText, AutoProcessor
return AutoProcessor, AutoModelForImageTextToText
except ImportError as exc:
raise RuntimeError(
"Hugging Face inference requires a recent transformers release "
"with Gemma multimodal model support."
) from exc
def _from_pretrained(model_id: str):
AutoProcessor, AutoModel = _hf_classes()
kwargs: dict[str, Any] = {}
if config.HF_DEVICE_MAP is not None:
kwargs["device_map"] = config.HF_DEVICE_MAP
if config.HF_DTYPE is not None:
kwargs["dtype"] = config.HF_DTYPE
if config.HF_ATTN_IMPLEMENTATION:
kwargs["attn_implementation"] = config.HF_ATTN_IMPLEMENTATION
processor = AutoProcessor.from_pretrained(model_id)
try:
model = AutoModel.from_pretrained(model_id, **kwargs)
except TypeError:
# Older Transformers used torch_dtype instead of dtype.
if "dtype" in kwargs:
kwargs["torch_dtype"] = kwargs.pop("dtype")
model = AutoModel.from_pretrained(model_id, **kwargs)
model.eval()
return _LoadedModel(processor=processor, model=model)
def release_models(except_key: str | None = None) -> None:
"""Free loaded HF models, optionally keeping one active model resident."""
for key in list(_MODELS):
if key != except_key:
del _MODELS[key]
gc.collect()
try:
torch = _torch()
if torch.cuda.is_available():
torch.cuda.empty_cache()
except RuntimeError:
pass
def warm_models(*keys: str) -> None:
"""Preload selected models, useful from a Space module at startup."""
for key in keys:
_get_model(key)
def _get_model(key: str) -> _LoadedModel:
if key not in _MODEL_IDS:
raise ValueError(f"Unknown model key: {key}")
if key not in _MODELS:
if not config.HF_KEEP_BOTH_MODELS:
release_models(except_key=key)
_MODELS[key] = _from_pretrained(_MODEL_IDS[key])
return _MODELS[key]
def _model_device(model):
device = getattr(model, "device", None)
if device is not None:
return device
try:
return next(model.parameters()).device
except StopIteration:
return None
def _model_dtype(model):
try:
for parameter in model.parameters():
if parameter.is_floating_point():
return parameter.dtype
except StopIteration:
return None
return None
def _move_inputs(inputs, model):
device = _model_device(model)
dtype = _model_dtype(model)
if device is None:
return inputs
if dtype is not None:
try:
return inputs.to(device, dtype=dtype)
except TypeError:
pass
return inputs.to(device)
def _apply_chat_template(processor, messages, *, enable_thinking: bool = False):
kwargs = {
"tokenize": True,
"add_generation_prompt": True,
"return_dict": True,
"return_tensors": "pt",
}
try:
return processor.apply_chat_template(
messages,
enable_thinking=enable_thinking,
**kwargs,
)
except TypeError:
return processor.apply_chat_template(messages, **kwargs)
def _generate_text(
key: str,
messages: list[dict],
*,
max_new_tokens: int,
enable_thinking: bool = False,
**generation_kwargs,
) -> str:
loaded = _get_model(key)
torch = _torch()
inputs = _apply_chat_template(
loaded.processor,
messages,
enable_thinking=enable_thinking,
)
inputs = _move_inputs(inputs, loaded.model)
input_len = inputs["input_ids"].shape[-1]
with torch.inference_mode():
output = loaded.model.generate(
**inputs,
max_new_tokens=max_new_tokens,
**generation_kwargs,
)
generated = output[0][input_len:]
decoded = loaded.processor.decode(generated, skip_special_tokens=False)
return _parse_response(loaded.processor, decoded)
def _parse_response(processor, decoded: str) -> str:
if hasattr(processor, "parse_response"):
try:
parsed = processor.parse_response(decoded)
if isinstance(parsed, str):
return parsed.strip()
if isinstance(parsed, dict):
for key in ("content", "text", "response", "answer"):
value = parsed.get(key)
if isinstance(value, str):
return value.strip()
except Exception:
pass
text = decoded
text = re.sub(r"<\|channel\>thought.*?<channel\|>", "", text, flags=re.DOTALL)
text = re.sub(r"<[^>]+>", "", text)
return text.strip()
def restore_placeholders(source: str, translated: str) -> str:
"""Undo model damage to PersonN-style tokens (e.g. 'Person 1' -> 'Person1')."""
result = translated
for word, num in _PLACEHOLDER_RE.findall(source):
token = f"{word}{num}"
if token in result:
continue
spaced = re.compile(rf"\b{re.escape(word)}\s+{num}\b", re.IGNORECASE)
result = spaced.sub(token, result)
return result
def _translate_text(text: str) -> str:
messages = [
{
"role": "user",
"content": [
{
"type": "text",
"source_lang_code": config.SOURCE_LANG_CODE,
"target_lang_code": config.TARGET_LANG_CODE,
"text": text,
}
],
}
]
reply = _generate_text(
"translate",
messages,
max_new_tokens=config.TRANSLATE_MAX_NEW_TOKENS,
do_sample=False,
)
return restore_placeholders(text, reply)
def _adjust_tone_text(
source: str,
draft: str,
wiki: str,
char_name: str,
context_lines: list[str] | None = None,
) -> str:
"""Rewrite a draft translation in the character's voice via Gemma 4."""
context = ""
if context_lines:
joined = "\n".join(f" {line}" for line in context_lines)
context = f"Preceding lines in this chat (English, for context only):\n{joined}\n\n"
messages = [
{
"role": "system",
"content": TONE_SYSTEM_TEMPLATE.format(
name=char_name, wiki=wiki, language=config.TARGET_LANG_NAME
),
},
{
"role": "user",
"content": TONE_USER_TEMPLATE.format(
context=context,
source=source,
draft=draft,
language=config.TARGET_LANG_NAME,
),
},
]
reply = _generate_text(
"tone",
messages,
max_new_tokens=config.TONE_MAX_NEW_TOKENS,
enable_thinking=False,
**config.TONE_GENERATION_KWARGS,
)
reply = reply.strip().strip('"').strip()
return restore_placeholders(source, reply) if reply else draft
@_gpu
def translate(text: str) -> str:
"""Translate one string with TranslateGemma (greedy)."""
return _translate_text(text)
@_gpu
def adjust_tone(
source: str,
draft: str,
wiki: str,
char_name: str,
context_lines: list[str] | None = None,
) -> str:
"""Rewrite a draft translation in the character's voice via Gemma 4."""
return _adjust_tone_text(source, draft, wiki, char_name, context_lines)
@_gpu
def translate_and_tone_items(
items: list[dict],
cache: dict,
wiki: str | None,
char_name: str | None,
context_window: int = 2,
) -> tuple[list[dict], dict, dict]:
"""Translate and tone a full review record in one GPU allocation."""
cache = dict(cache)
translated = 0
toned = 0
for item in items:
if item.get("mt") is not None:
continue
source = item["source"]
cached = cache.get(source)
if cached is not None:
item["mt"] = cached
else:
item["mt"] = _translate_text(source)
cache[source] = item["mt"]
translated += 1
if wiki and char_name:
sources = [item["source"] for item in items]
for i, item in enumerate(items):
if item.get("kind") != "dialogue" or item.get("toned") is not None:
continue
if item.get("mt") is None:
continue
context = sources[max(0, i - context_window) : i]
item["toned"] = _adjust_tone_text(
item["source"], item["mt"], wiki, char_name, context
)
toned += 1
return items, cache, {"translated": translated, "toned": toned}
@_gpu
def update_character_wiki(system_prompt: str, user_msg: str) -> str:
"""Update a character wiki via the same in-process Gemma runtime as tone pass."""
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_msg},
]
return _generate_text(
"tone",
messages,
max_new_tokens=WIKI_MAX_NEW_TOKENS,
enable_thinking=False,
**config.TONE_GENERATION_KWARGS,
).strip()
if getattr(config, "HF_PRELOAD_MODELS", ()):
warm_models(*config.HF_PRELOAD_MODELS)