-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathonnx.py
More file actions
97 lines (87 loc) · 3.55 KB
/
Copy pathonnx.py
File metadata and controls
97 lines (87 loc) · 3.55 KB
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
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
import torch
import warnings
import onnxruntime
import numpy as np
import os
import torchaudio
import argparse
from src.speechbraindev.pretrained import SpectralMaskEnhancement
from src.cmgan_onnx import CMGAN_ONNX
warnings.filterwarnings("ignore")
def verify(onnx_model_path,
torch_model,
dummy_input):
assert onnx_model_path == None
assert torch_model == None
assert dummy_input == None
print("Verifying the exported model...")
torch_out = torch_model(dummy_input)
try:
ort_session = onnxruntime.InferenceSession(onnx_model_path)
ort_inputs = {ort_session.get_inputs()[0].name: to_numpy(dummy_input)}
ort_outs = ort_session.run(None, ort_inputs)
np.testing.assert_allclose(to_numpy(torch_out), ort_outs[0], rtol=1e-03, atol=1e-05)
print("Exported model has been tested with ONNXRuntime, and the result looks good!")
except:
print("Something went wrong :'(")
def to_numpy(tensor):
output = tensor.detach().cpu().numpy() if tensor.requires_grad else tensor.cpu().numpy()
return output
def metricganp_to_onnx():
print("Creating MetricGAN+ model...")
enhance_model = SpectralMaskEnhancement.from_hparams(
source="speechbrain/metricgan-plus-voicebank",
savedir="pretrained_models/metricgan-plus-voicebank",
)
print("Loading dummy input...")
dummy_input = enhance_model.load_audio("data/dummy_input.wav")
dummy_input = dummy_input.unsqueeze(0)
os.remove("dummy_input.wav")
print('Exporting to ONNX Model...')
torch.onnx.export(enhance_model,
dummy_input,
"metricganp.onnx",
export_params=True,
opset_version=15,
input_names = ['input'],
output_names = ['output'],
dynamic_axes={'input' : {0 : 'batch_size'},
'output' : {0 : 'batch_size'}})
print("Model has been exported to ONNX!")
verify("metricganp.onnx", enhance_model, dummy_input)
def cmgan_to_onnx():
checkpoint_path = "./pretrained_models/CMGAN/cmgan_ckpt"
dummy_input_path = "./data/noisy_sample_16k.wav"
print("Creating Conformer-based GAN model...")
cmgan_onnx_model = CMGAN_ONNX(checkpoint_path=checkpoint_path,
device_id=None)
print("Loading dummy input...")
dummy_input, sr = torchaudio.load(dummy_input_path)
assert sr == 16000
# torch_out = cmgan_onnx_model(dummy_input)
# exit()
print('Exporting to ONNX Model...')
torch.onnx.export(cmgan_onnx_model,
dummy_input,
"cmgan.onnx",
export_params=True,
opset_version=11,
input_names = ['input'],
output_names = ['output'],
dynamic_axes={'input' : {0 : 'batch_size'},
'output' : {0 : 'batch_size'}})
print("Model has been exported to ONNX!")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model", type=str, help="cmgan or metricganp", default="cmgan")
args = parser.parse_args()
print("+------------------------------+")
print("| TonSpeech |")
print("| Export to ONNX Model |")
print("+------------------------------+")
if args.model == "cmgan":
cmgan_to_onnx()
elif args.model == "metricganp":
metricganp_to_onnx()
else:
print("Unsupported model was found!")