"""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.*?", "", 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)