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)