如何将Hugging Face Transformers中的Tokenizer保存为ONNX格式?
如何将Hugging Face Tokenizer导出为ONNX格式
Tokenizer本质是文本预处理工具,核心为字符串处理逻辑,不属于可微分的PyTorch模型范畴,无法直接用torch.onnx.export导出。要将其转为ONNX格式,需把分词的核心操作封装成可被PyTorch追踪的张量模块,让ONNX能识别并导出这些操作。
实现步骤
1. 安装依赖
除原有依赖外,需安装onnxruntime用于验证导出结果:
pip install transformers torch onnx onnxruntime
2. 封装Tokenizer为PyTorch模块
创建自定义PyTorch模块,将Tokenizer的文本转input_ids、attention_mask的逻辑包装为可追踪的张量操作:
from transformers import AutoTokenizer import torch import torch.nn as nn class TokenizerONNXWrapper(nn.Module): def __init__(self, tokenizer): super().__init__() self.tokenizer = tokenizer def forward(self, text_list): # 执行分词并返回张量格式的结果 inputs = self.tokenizer( text_list, return_tensors="pt", padding=True, truncation=True ) return inputs["input_ids"], inputs["attention_mask"] # 加载目标Tokenizer tokenizer = AutoTokenizer.from_pretrained("huawei-noah/TinyBERT_General_4L_312D") # 初始化包装模块 tokenizer_wrapper = TokenizerONNXWrapper(tokenizer) # 准备示例批量文本输入 sample_texts = [ "Hugging Face is creating a tool that democratizes AI.", "ONNX simplifies model cross-platform deployment." ] # 导出为ONNX格式 torch.onnx.export( tokenizer_wrapper, (sample_texts,), "tinybert_tokenizer.onnx", export_params=True, opset_version=13, # 高opset版本对字符串输入支持更完善 input_names=["text_list"], output_names=["input_ids", "attention_mask"], dynamic_axes={ "input_ids": {0: "batch_size", 1: "sequence_length"}, "attention_mask": {0: "batch_size", 1: "sequence_length"} } ) print("Tokenizer包装模块已成功导出为ONNX")
3. 验证导出的ONNX模型
用onnxruntime加载模型并验证结果是否与原Tokenizer一致:
import onnxruntime as ort # 加载ONNX模型 session = ort.InferenceSession("tinybert_tokenizer.onnx") # 测试输入 test_texts = ["Hello, ONNX Tokenizer!"] inputs = {"text_list": test_texts} # 执行推理 output_ids, output_mask = session.run(None, inputs) print("ONNX模型输出:") print("input_ids:", output_ids) print("attention_mask:", output_mask) # 与原Tokenizer结果对比 original_inputs = tokenizer(test_texts, return_tensors="np") print("\n原Tokenizer输出对比:") print("input_ids:", original_inputs["input_ids"]) print("attention_mask:", original_inputs["attention_mask"])
关键注意事项
- ONNX对字符串输入的支持依赖opset版本,建议使用opset 13及以上版本,避免兼容性问题。
- 若要完全脱离Python环境部署,需将文本转为字符张量输入,或使用ONNX原生字符串算子,实现复杂度会更高。
- 复杂分词逻辑(如多语言子词拆分)可能无法被PyTorch完全追踪,此时可使用Hugging Face的
optimum库简化端到端导出流程。
内容的提问来源于stack exchange,提问作者Franck Dernoncourt
相关产品推荐
相关产品推荐

