| 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) |
|
|