financial-rag / src /tokenization /gemma4_tokenizer.py
tolivert's picture
deploy: financial_rag streamlit app
4e316d6
Raw
History Blame Contribute Delete
3.87 kB
"""
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)
@property
def vocab_size(self):
return self._tok.get_vocab_size()