| from __future__ import annotations |
|
|
| import json |
| from pathlib import Path |
| from types import SimpleNamespace |
|
|
| from transformers import PreTrainedTokenizerFast |
|
|
|
|
| PROMPT_PREFIX = { |
| "query": "Query: ", |
| "document": "Document: ", |
| } |
|
|
|
|
| def load_text_runtime_config(model_dir: str | Path): |
| model_path = Path(model_dir) |
| with open(model_path / "config.json", encoding="utf-8") as f: |
| raw_config = json.load(f) |
|
|
| text_config = dict(raw_config.get("text_config") or raw_config) |
| text_config["model_type"] = text_config.get("model_type") or raw_config.get("model_type") or "qwen3" |
| text_config["rms_norm_eps"] = text_config.get("rms_norm_eps", 1e-6) |
| text_config["max_position_embeddings"] = text_config.get("max_position_embeddings", 32768) |
| text_config["pad_token_id"] = text_config.get("pad_token_id") |
| text_config["rope_parameters"] = text_config.get("rope_parameters", {}) |
| return SimpleNamespace(**text_config) |
|
|
|
|
| def load_tokenizer(model_dir: str | Path): |
| model_path = Path(model_dir) |
| with open(model_path / "tokenizer_config.json", encoding="utf-8") as f: |
| tokenizer_config = json.load(f) |
|
|
| init_kwargs = { |
| "tokenizer_file": str(model_path / "tokenizer.json"), |
| "bos_token": tokenizer_config.get("bos_token"), |
| "eos_token": tokenizer_config.get("eos_token"), |
| "pad_token": tokenizer_config.get("pad_token"), |
| "unk_token": tokenizer_config.get("unk_token"), |
| "mask_token": tokenizer_config.get("mask_token"), |
| "padding_side": tokenizer_config.get("padding_side", "left"), |
| } |
| tokenizer = PreTrainedTokenizerFast(**{k: v for k, v in init_kwargs.items() if v is not None}) |
|
|
| chat_template_path = model_path / "chat_template.jinja" |
| if chat_template_path.exists(): |
| tokenizer.chat_template = chat_template_path.read_text(encoding="utf-8") |
|
|
| return tokenizer |
|
|
|
|
| def build_text_inputs( |
| tokenizer, |
| *, |
| text: str, |
| prompt_name: str, |
| max_length: int, |
| ): |
| if prompt_name not in PROMPT_PREFIX: |
| raise ValueError(f"Unsupported prompt_name: {prompt_name}") |
| encoded = tokenizer( |
| [f"{PROMPT_PREFIX[prompt_name]}{text}"], |
| return_tensors="np", |
| padding=True, |
| truncation=True, |
| max_length=max_length, |
| ) |
| return encoded |
|
|