File size: 2,236 Bytes
e498f5b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 | import argparse
from pathlib import Path
import numpy as np
import onnxruntime as ort
import torch
from transformers import Wav2Vec2ForCTC, Wav2Vec2Processor
DEFAULT_MODEL_DIR = "wav2vec2-ljspeech"
DEFAULT_OUTPUT = "wav2vec2-ljspeech.onnx"
SAMPLE_RATE = 16000
def export(model_dir: str, output: str, opset: int) -> None:
model_path = Path(model_dir)
processor = Wav2Vec2Processor.from_pretrained(model_path)
model = Wav2Vec2ForCTC.from_pretrained(model_path)
model.eval()
seq_len = SAMPLE_RATE
dummy = torch.zeros(1, seq_len, dtype=torch.float32)
dynamic_axes = {
"input_values": {0: "batch", 1: "time"},
"logits": {0: "batch", 1: "time"},
}
output_path = Path(output)
output_path.parent.mkdir(parents=True, exist_ok=True)
with torch.no_grad():
torch.onnx.export(
model,
(dummy,),
str(output_path),
input_names=["input_values"],
output_names=["logits"],
dynamic_axes=dynamic_axes,
opset_version=opset,
do_constant_folding=True,
)
print(f"Exported ONNX model to {output_path}")
print(f" opset: {opset}, vocab_size: {len(processor.tokenizer)}")
print(f" size: {output_path.stat().st_size / (1024 * 1024):.1f} MB")
validate(output_path, processor, seq_len)
def validate(onnx_path: Path, processor, seq_len: int) -> None:
session = ort.InferenceSession(str(onnx_path), providers=["CPUExecutionProvider"])
audio = np.zeros((1, seq_len), dtype=np.float32)
outputs = session.run(None, {"input_values": audio})
logits = outputs[0]
print(f"Validated with ONNX Runtime: output shape = {logits.shape}")
pred_ids = np.argmax(logits, axis=-1)
text = processor.tokenizer.batch_decode(pred_ids)[0]
print(f"Decoded dummy input -> {text!r}")
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Export Wav2Vec2 CTC to ONNX")
parser.add_argument("--model-dir", default=DEFAULT_MODEL_DIR)
parser.add_argument("--output", default=DEFAULT_OUTPUT)
parser.add_argument("--opset", type=int, default=17)
args = parser.parse_args()
export(args.model_dir, args.output, args.opset)
|