Spaces:
Sleeping
Sleeping
| """ | |
| 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: | |
| <start_of_turn>user\\n{content}<end_of_turn>\\n | |
| <start_of_turn>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 = [ | |
| "<bos>", "<eos>", "<unk>", "<pad>", | |
| "<start_of_turn>", "<end_of_turn>", | |
| "<think>", "</think>", | |
| ] | |
| 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("<pad>", self._special_to_id.get("<eos>")) | |
| self.eos_token_id = self._special_to_id.get("<eos>", 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"<start_of_turn>{role}\n{content}<end_of_turn>\n" | |
| if add_generation_prompt: | |
| formatted_text += "<start_of_turn>model\n" | |
| if add_thinking: | |
| formatted_text += "<think>\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) | |
| def vocab_size(self): | |
| return self._tok.get_vocab_size() | |