text-to-speech-train-code / export_onnx.py
Isa0's picture
add train code
e498f5b
Raw
History Blame Contribute Delete
2.24 kB
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)