shimokumo / src /model /tokenizer.py
Shimokumo's picture
Upload folder using huggingface_hub
94fd0b0 verified
Raw
History Blame Contribute Delete
10.2 kB
"""
霜云(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>: 情感标记
"""
# 特殊token定义
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,
}
# 特殊token的id到token的映射
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
# BPE词表(回退方案使用)
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
# 更新vocab_size为实际值
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
# 基础字符集(ASCII + 常用中文字符)
base_chars = list("abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789")
base_chars += list(" \t\n\r",)
base_chars += list("!\"#$%&'()*+,-./:;<=>?@[\\]^_`{|}~")
# 构建基础词表
token_id = len(self.SPECIAL_TOKENS) # 从特殊token之后开始
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)
# 添加特殊token
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:
# 检查是否是特殊token
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})"