使用 Hugging Face Optimum 将 Transformers 转换为 ONNX
使用 Hugging Face Optimum 进行高级转换
Optimum 库通过 ORTModelForXxx 类提供了最友好的 ONNX 转换方法。通过在 from_pretrained() 方法中设置 from_transformers=True 标志,Optimum 会自动加载原始的 Transformers 模型,并在内部使用 transformers.onnx 包将其转换为 ONNX。
Implementation Example:
from optimum.onnxruntime import ORTModelForSequenceClassification
model = ORTModelForSequenceClassification.from_pretrained("distilbert-base-uncased-finetuned-sst-2-english", from_transformers=True)
通过 Optimum 转换的模型可以立即用于预测,或直接集成到 Hugging Face pipelines 中。
使用 transformers.onnx 进行中级转换
transformers.onnx 模块通过使用配置对象简化了转换过程,免去了用户手动定义诸如 dynamic_axes 等复杂参数的需求。
Implementation Example:
from pathlib import Path
import transformers
from transformers.onnx import FeaturesManager
from transformers import AutoConfig, AutoTokenizer, AutoModelForSequenceClassification
model_id = "distilbert-base-uncased-sst-2-english"
feature = "sequence-classification"
model = AutoModelForSequenceClassification.from_pretrained(model_id)
tokenizer = AutoTokenizer.from_pretrained(model_id)
model_kind, model_onnx_config = FeaturesManager.check_supported_model_or_raise(model, feature=feature)
onnx_inputs,
preprocessor=tokenizer,
model=model,
config=onnx_config,
opset=13,
output=Path("trfs-model.onnx")
)
使用 torch.onnx 进行低级转换
torch.onnx API 提供了最细粒度的控制,但需要手动指定多个参数,包括 input_names、output_names 和 dynamic_axes。
Implementation Example:
import torch
from transformers import AutoModelForSequenceClassification, AutoTokenizer
model_id = "distilbert-base-uncased-finetuned-sst-2-english"
model = AutoModelForSequenceClassification.from_pretrained(model_id)
tokenizer = AutoTokenizer.from_pretrained(model_id)
dummy_model_input = tokenizer("This is a sample", return_tensors="pt")
torch.onnx.export(
model,
tuple(dummy_model_input.values()),
f="torch-model.onnx",
input_names=['input_ids', 'attention_mask'],
output_names=['logits'],
dynamic_axes={'input_ids': {0: 'batch_size', 1: 'sequence'},
'attention_mask': {0: 'batch_size', 1: 'sequence'},
'logits': {0: 'batch_size', 1: 'sequence'}},
do_constant_folding=True,
opset_version=13,
)