使用 Hugging Face Optimum 將 Transformers 轉換為 ONNX
Hugging Face 提供三條不同的路徑將 Transformers 模型轉換為開放神經網路交換(ONNX)格式,讓使用者可以在細粒度控制與高階抽象之間做選擇。最簡化的方法是透過 Optimum 函式庫,它自動化轉換流程,同時保持與 Hugging Face pipelines 的相容性。
使用 Hugging Face Optimum 進行高階轉換
Optimum 函式庫提供最友善的 ONNX 轉換方式,透過 ORTModelForXxx 類別。只要在 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,
)