| """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: |
| spaces = None |
|
|
|
|
| |
| _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: |
| |
| 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) |
|
|