| """ |
| 霜云(Shimokumo) - 分词器模块 |
| |
| 基于sentencepiece实现的中英文分词器,支持特殊token定义。 |
| """ |
|
|
| import os |
| import re |
| from typing import Dict, List, Optional, Tuple |
|
|
|
|
| class ShimokumoTokenizer: |
| """霜云专用分词器 |
| |
| 支持中英文分词,内置特殊token定义。 |
| 可使用sentencepiece后端或纯Python回退方案。 |
| |
| 特殊Token定义: |
| <BOS>: 句子开始 |
| <EOS>: 句子结束 |
| <PAD>: 填充 |
| <UNK>: 未知 |
| <user>: 用户角色标记 |
| <assistant>: 助手角色标记 |
| <narration>: 旁白标记 |
| <action>: 动作描述标记 |
| <emotion>: 情感标记 |
| """ |
|
|
| |
| SPECIAL_TOKENS = { |
| "<BOS>": 0, |
| "<EOS>": 1, |
| "<PAD>": 2, |
| "<UNK>": 3, |
| "<user>": 4, |
| "<assistant>": 5, |
| "<narration>": 6, |
| "<action>": 7, |
| "<emotion>": 8, |
| "<think_start>": 9, |
| "<think_end>": 10, |
| } |
|
|
| |
| SPECIAL_TOKEN_IDS = {v: k for k, v in SPECIAL_TOKENS.items()} |
|
|
| def __init__(self, model_path: Optional[str] = None, vocab_size: int = 32000): |
| """ |
| 初始化分词器。 |
| |
| Args: |
| model_path: sentencepiece模型文件路径,为None则使用纯Python回退 |
| vocab_size: 词表大小 |
| """ |
| self.vocab_size = vocab_size |
| self.model_path = model_path |
| self._sp_model = None |
| self._use_sentencepiece = False |
|
|
| |
| self._bpe_vocab: Dict[str, int] = {} |
| self._bpe_vocab_inv: Dict[int, str] = {} |
|
|
| if model_path and os.path.exists(model_path): |
| self._load_sentencepiece(model_path) |
| else: |
| self._init_fallback_tokenizer() |
|
|
| def _load_sentencepiece(self, model_path: str) -> None: |
| """加载sentencepiece模型""" |
| try: |
| import sentencepiece as spm |
| self._sp_model = spm.SentencePieceProcessor() |
| self._sp_model.load(model_path) |
| self._use_sentencepiece = True |
| |
| self.vocab_size = self._sp_model.get_piece_size() |
| except ImportError: |
| self._init_fallback_tokenizer() |
| except Exception: |
| self._init_fallback_tokenizer() |
|
|
| def _init_fallback_tokenizer(self) -> None: |
| """初始化纯Python回退分词器(基于字符和字节对编码)""" |
| self._use_sentencepiece = False |
|
|
| |
| base_chars = list("abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789") |
| base_chars += list(" \t\n\r",) |
| base_chars += list("!\"#$%&'()*+,-./:;<=>?@[\\]^_`{|}~") |
|
|
| |
| token_id = len(self.SPECIAL_TOKENS) |
| for char in base_chars: |
| if char not in self._bpe_vocab: |
| self._bpe_vocab[char] = token_id |
| self._bpe_vocab_inv[token_id] = char |
| token_id += 1 |
|
|
| |
| common_tokens = [ |
| "的", "了", "是", "我", "你", "他", "她", "它", "们", "这", |
| "那", "有", "不", "在", "和", "人", "都", "一", "上", "下", |
| "大", "小", "中", "来", "去", "到", "说", "会", "能", "好", |
| "就", "对", "被", "把", "让", "给", "从", "用", "过", "也", |
| "很", "最", "已", "还", "为", "与", "而", "但", "或", "如", |
| "霜", "云", "巫", "女", "网", "络", "运", "营", "商", "模", |
| "嗯", "啊", "呢", "吧", "呀", "呜", "诶", "嘻", "哈", "哦", |
| ] |
| for token in common_tokens: |
| if token not in self._bpe_vocab: |
| self._bpe_vocab[token] = token_id |
| self._bpe_vocab_inv[token_id] = token |
| token_id += 1 |
|
|
| |
| for byte_val in range(256): |
| token_str = f"<0x{byte_val:02X}>" |
| if token_str not in self._bpe_vocab and token_id < self.vocab_size: |
| self._bpe_vocab[token_str] = token_id |
| self._bpe_vocab_inv[token_id] = token_str |
| token_id += 1 |
|
|
| self.vocab_size = max(token_id, self.vocab_size) |
|
|
| def encode( |
| self, |
| text: str, |
| add_bos: bool = True, |
| add_eos: bool = False, |
| ) -> List[int]: |
| """ |
| 将文本编码为token ID序列。 |
| |
| Args: |
| text: 输入文本 |
| add_bos: 是否添加句子开始标记 |
| add_eos: 是否添加句子结束标记 |
| |
| Returns: |
| token ID列表 |
| """ |
| if not text: |
| tokens = [] |
| elif self._use_sentencepiece and self._sp_model: |
| tokens = self._sp_model.encode(text) |
| else: |
| tokens = self._fallback_encode(text) |
|
|
| |
| result: List[int] = [] |
| if add_bos: |
| result.append(self.SPECIAL_TOKENS["<BOS>"]) |
| result.extend(tokens) |
| if add_eos: |
| result.append(self.SPECIAL_TOKENS["<EOS>"]) |
|
|
| return result |
|
|
| def decode( |
| self, |
| token_ids: List[int], |
| skip_special: bool = True, |
| ) -> str: |
| """ |
| 将token ID序列解码为文本。 |
| |
| Args: |
| token_ids: token ID列表 |
| skip_special: 是否跳过特殊token |
| |
| Returns: |
| 解码后的文本 |
| """ |
| if self._use_sentencepiece and self._sp_model: |
| return self._sp_model.decode(token_ids) |
| else: |
| return self._fallback_decode(token_ids, skip_special) |
|
|
| def _fallback_encode(self, text: str) -> List[int]: |
| """ |
| 回退编码方案:字符级 + 最大匹配分词。 |
| |
| Args: |
| text: 输入文本 |
| |
| Returns: |
| token ID列表 |
| """ |
| tokens: List[int] = [] |
| i = 0 |
|
|
| while i < len(text): |
| |
| matched = False |
| for length in range(min(4, len(text) - i), 0, -1): |
| substr = text[i : i + length] |
| if substr in self._bpe_vocab: |
| tokens.append(self._bpe_vocab[substr]) |
| i += length |
| matched = True |
| break |
|
|
| if not matched: |
| |
| byte_val = ord(text[i]) |
| byte_token = f"<0x{byte_val:02X}>" |
| if byte_token in self._bpe_vocab: |
| tokens.append(self._bpe_vocab[byte_token]) |
| else: |
| tokens.append(self.SPECIAL_TOKENS["<UNK>"]) |
| i += 1 |
|
|
| return tokens |
|
|
| def _fallback_decode( |
| self, |
| token_ids: List[int], |
| skip_special: bool = True, |
| ) -> str: |
| """ |
| 回退解码方案。 |
| |
| Args: |
| token_ids: token ID列表 |
| skip_special: 是否跳过特殊token |
| |
| Returns: |
| 解码后的文本 |
| """ |
| text_parts: List[str] = [] |
| for token_id in token_ids: |
| |
| if skip_special and token_id in self.SPECIAL_TOKEN_IDS: |
| if token_id in self.SPECIAL_TOKEN_IDS: |
| continue |
| text_parts.append(self.SPECIAL_TOKEN_IDS[token_id]) |
| continue |
|
|
| |
| if token_id in self._bpe_vocab_inv: |
| token_str = self._bpe_vocab_inv[token_id] |
| |
| if token_str.startswith("<0x") and token_str.endswith(">"): |
| try: |
| byte_val = int(token_str[3:-1], 16) |
| text_parts.append(chr(byte_val)) |
| except ValueError: |
| text_parts.append(token_str) |
| else: |
| text_parts.append(token_str) |
| else: |
| text_parts.append(f"<UNK:{token_id}>") |
|
|
| return "".join(text_parts) |
|
|
| def encode_chat( |
| self, |
| messages: List[Dict[str, str]], |
| add_bos: bool = True, |
| add_eos: bool = True, |
| ) -> List[int]: |
| """ |
| 将对话消息列表编码为token ID序列。 |
| |
| Args: |
| messages: 对话消息列表,每个元素为 {"role": "user/assistant/system", "content": "..."} |
| add_bos: 是否添加BOS |
| add_eos: 是否添加EOS |
| |
| Returns: |
| 编码后的token ID序列 |
| """ |
| all_tokens: List[int] = [] |
|
|
| if add_bos: |
| all_tokens.append(self.SPECIAL_TOKENS["<BOS>"]) |
|
|
| for msg in messages: |
| role = msg.get("role", "user").lower() |
| content = msg.get("content", "") |
|
|
| |
| if role == "user": |
| all_tokens.append(self.SPECIAL_TOKENS["<user>"]) |
| elif role == "assistant": |
| all_tokens.append(self.SPECIAL_TOKENS["<assistant>"]) |
| elif role == "narration": |
| all_tokens.append(self.SPECIAL_TOKENS["<narration>"]) |
|
|
| |
| content_tokens = self.encode(content, add_bos=False, add_eos=False) |
| all_tokens.extend(content_tokens) |
|
|
| if add_eos: |
| all_tokens.append(self.SPECIAL_TOKENS["<EOS>"]) |
|
|
| return all_tokens |
|
|
| def tokenize(self, text: str) -> List[str]: |
| """ |
| 将文本分词为token字符串列表(不转换为ID)。 |
| |
| Args: |
| text: 输入文本 |
| |
| Returns: |
| token字符串列表 |
| """ |
| token_ids = self.encode(text, add_bos=False, add_eos=False) |
| return [self.decode([tid], skip_special=False) for tid in token_ids] |
|
|
| def __len__(self) -> int: |
| """返回词表大小""" |
| return self.vocab_size |
|
|
| def __repr__(self) -> str: |
| backend = "sentencepiece" if self._use_sentencepiece else "fallback" |
| return f"ShimokumoTokenizer(backend={backend}, vocab_size={self.vocab_size})" |
|
|