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)