使用 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_namesoutput_namesdynamic_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,
)

Sources