使用 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_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