snap-demo / export_bert.py
softguy777's picture
Upload folder using huggingface_hub
55e66ae verified
Raw
History Blame Contribute Delete
2.74 kB
"""
SNAP BERT ONNX Export Utility
=============================
Converts a HuggingFace BERT model to ONNX for use with SNAP.
SNAP explicitly requires the 10th layer outputs (hidden_states[-3]) to maintain
compatibility with the original TTS architectures it was trained on.
Standard optimum-cli export will NOT work correctly.
Usage:
pip install transformers torch onnx
python scripts/export_bert.py --model kykim/bert-kor-base --output snap/weights/ko/bert.onnx
"""
import os
import argparse
import torch
import torch.nn as nn
from transformers import AutoTokenizer, AutoModelForMaskedLM
class BERTFeatureExtractor(nn.Module):
"""Wraps BERT to specifically output the 10th layer hidden state."""
def __init__(self, bert_model):
super().__init__()
self.bert = bert_model
def forward(self, input_ids, attention_mask, token_type_ids):
outputs = self.bert(
input_ids=input_ids,
attention_mask=attention_mask,
token_type_ids=token_type_ids,
output_hidden_states=True,
)
return outputs.hidden_states[-3]
def export_bert(model_id, output_path):
print(f"Loading {model_id}...")
tokenizer = AutoTokenizer.from_pretrained(model_id)
model = AutoModelForMaskedLM.from_pretrained(model_id)
wrapper = BERTFeatureExtractor(model).eval()
# Dummy inputs
text = "Hello, world!"
inputs = tokenizer(text, return_tensors="pt")
input_ids = inputs["input_ids"]
attention_mask = inputs["attention_mask"]
token_type_ids = inputs.get("token_type_ids", torch.zeros_like(input_ids))
os.makedirs(os.path.dirname(os.path.abspath(output_path)), exist_ok=True)
print(f"Exporting to {output_path}...")
torch.onnx.export(
wrapper,
(input_ids, attention_mask, token_type_ids),
output_path,
input_names=["input_ids", "attention_mask", "token_type_ids"],
output_names=["hidden_states"],
dynamic_axes={
"input_ids": {0: "batch", 1: "seq_len"},
"attention_mask": {0: "batch", 1: "seq_len"},
"token_type_ids": {0: "batch", 1: "seq_len"},
"hidden_states": {0: "batch", 1: "seq_len"},
},
opset_version=17,
do_constant_folding=True,
)
tok_dir = os.path.dirname(os.path.abspath(output_path))
tokenizer.save_pretrained(tok_dir)
print(f"Done! Tokenizer saved to {tok_dir}")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model", required=True, help="HuggingFace model ID")
parser.add_argument("--output", required=True, help="Output ONNX path")
args = parser.parse_args()
export_bert(args.model, args.output)