""" Gemma4 Tokenizer. Reuses the same GemmaTokenizer (262k BPE vocab) as Gemma3. The only differences are additional multimodal special tokens (image, audio, video) which are not needed for text-only inference but are registered here for completeness. Chat template follows the same turn-based format as Gemma3: user\\n{content}\\n model\\n """ import re import os from typing import List, Dict, Union from pathlib import Path from tokenizers import Tokenizer from huggingface_hub import hf_hub_download from src.utils.model_store import get_model_dir class Gemma4Tokenizer: _SPECIALS = [ "", "", "", "", "", "", "", "", ] def __init__(self, model_path_or_repo: str = "google/gemma-4-E2B"): if os.path.exists(model_path_or_repo): if (Path(model_path_or_repo) / "tokenizer.json").exists(): file_path = Path(model_path_or_repo) / "tokenizer.json" else: file_path = Path(model_path_or_repo) / "tokenizer.model" else: repo_name = Path(model_path_or_repo).parts[-1] central_path = get_model_dir(repo_name) / "tokenizer.json" if central_path.exists(): file_path = central_path else: print(f"Downloading tokenizer from {model_path_or_repo}...") file_path = hf_hub_download( repo_id=model_path_or_repo, filename="tokenizer.json", local_dir=str(get_model_dir(repo_name)), ) self._tok = Tokenizer.from_file(str(file_path)) self._special_to_id = {} self._id_to_special = {} for t in self._SPECIALS: tid = self._tok.token_to_id(t) if tid is not None: self._special_to_id[t] = tid self._id_to_special[tid] = t pattern = "|".join(map(re.escape, self._SPECIALS)) self._split_re = re.compile(f"({pattern})") self.pad_token_id = self._special_to_id.get("", self._special_to_id.get("")) self.eos_token_id = self._special_to_id.get("", None) def apply_chat_template( self, messages: List[Dict[str, str]], add_generation_prompt: bool = False, add_thinking: bool = False, ) -> str: formatted_text = "" for msg in messages: role = msg["role"] content = msg["content"] formatted_text += f"{role}\n{content}\n" if add_generation_prompt: formatted_text += "model\n" if add_thinking: formatted_text += "\n" return formatted_text def encode( self, text_or_messages: Union[str, List[Dict]], add_generation_prompt: bool = False, add_thinking: bool = False, ): if isinstance(text_or_messages, list): text = self.apply_chat_template( text_or_messages, add_generation_prompt=add_generation_prompt, add_thinking=add_thinking, ) else: text = text_or_messages ids = [] parts = self._split_re.split(text) for part in parts: if not part: continue if part in self._special_to_id: ids.append(self._special_to_id[part]) else: ids.extend(self._tok.encode(part, add_special_tokens=False).ids) return ids def decode(self, ids, skip_special_tokens=False): return self._tok.decode(ids, skip_special_tokens=skip_special_tokens) @property def vocab_size(self): return self._tok.get_vocab_size()