Improve: ONNX conversion logic improved, and support for noise_scale(_w) model argument
This commit is contained in:
174
convert_onnx.py
174
convert_onnx.py
@@ -1,4 +1,4 @@
|
|||||||
# usage: .venv/bin/python convert_onnx.py --model model_assets/amitaro/amitaro.safetensors
|
# usage: .venv/bin/python convert_onnx.py --model model_assets/koharune-ami/koharune-ami.safetensors
|
||||||
# ref: https://github.com/tuna2134/sbv2-api/blob/main/convert/convert_model.py
|
# ref: https://github.com/tuna2134/sbv2-api/blob/main/convert/convert_model.py
|
||||||
|
|
||||||
import time
|
import time
|
||||||
@@ -59,37 +59,10 @@ if __name__ == "__main__":
|
|||||||
tts_model.load()
|
tts_model.load()
|
||||||
style_id = tts_model.style2id[DEFAULT_STYLE]
|
style_id = tts_model.style2id[DEFAULT_STYLE]
|
||||||
assert tts_model.net_g is not None, "Model is not loaded"
|
assert tts_model.net_g is not None, "Model is not loaded"
|
||||||
assert (
|
|
||||||
tts_model.hyper_parameters.data.use_jp_extra is True
|
|
||||||
), "Normal model is not supported yet"
|
|
||||||
|
|
||||||
# SynthesizerTrnJPExtra の forward メソッドをオーバーライド
|
|
||||||
def forward(
|
|
||||||
x: torch.Tensor,
|
|
||||||
x_lengths: torch.Tensor,
|
|
||||||
sid: torch.Tensor,
|
|
||||||
tone: torch.Tensor,
|
|
||||||
language: torch.Tensor,
|
|
||||||
bert: torch.Tensor,
|
|
||||||
style_vec: torch.Tensor,
|
|
||||||
length_scale: float = 1.0,
|
|
||||||
sdp_ratio: float = 0.0,
|
|
||||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, tuple[torch.Tensor, ...]]:
|
|
||||||
return cast(SynthesizerTrnJPExtra, tts_model.net_g).infer(
|
|
||||||
x,
|
|
||||||
x_lengths,
|
|
||||||
sid,
|
|
||||||
tone,
|
|
||||||
language,
|
|
||||||
bert,
|
|
||||||
style_vec,
|
|
||||||
sdp_ratio=sdp_ratio,
|
|
||||||
length_scale=length_scale,
|
|
||||||
)
|
|
||||||
|
|
||||||
tts_model.net_g.forward = forward # type: ignore
|
|
||||||
|
|
||||||
# 音声合成に必要な BERT 特徴量・音素列・アクセント列・言語 ID を取得
|
# 音声合成に必要な BERT 特徴量・音素列・アクセント列・言語 ID を取得
|
||||||
|
# JP-Extra モデルアーキテクチャの場合、bert (中国語の BERT 特徴量) や en_bert (英語の BERT 特徴量) は
|
||||||
|
# torch.zeros() で適当に埋められており、推論には ja_bert (日本語の BERT 特徴量) のみが使用される
|
||||||
bert, ja_bert, en_bert, phones, tones, lang_ids = get_text(
|
bert, ja_bert, en_bert, phones, tones, lang_ids = get_text(
|
||||||
"今日はいい天気ですね。",
|
"今日はいい天気ですね。",
|
||||||
Languages.JP,
|
Languages.JP,
|
||||||
@@ -104,6 +77,7 @@ if __name__ == "__main__":
|
|||||||
# スタイルベクトルを取得
|
# スタイルベクトルを取得
|
||||||
style_vector = tts_model.get_style_vector(style_id, DEFAULT_STYLE_WEIGHT)
|
style_vector = tts_model.get_style_vector(style_id, DEFAULT_STYLE_WEIGHT)
|
||||||
|
|
||||||
|
# モデルの入力を作成
|
||||||
x_tst = phones.to(device).unsqueeze(0)
|
x_tst = phones.to(device).unsqueeze(0)
|
||||||
tones = tones.to(device).unsqueeze(0)
|
tones = tones.to(device).unsqueeze(0)
|
||||||
lang_ids = lang_ids.to(device).unsqueeze(0)
|
lang_ids = lang_ids.to(device).unsqueeze(0)
|
||||||
@@ -112,50 +86,102 @@ if __name__ == "__main__":
|
|||||||
en_bert = en_bert.to(device).unsqueeze(0)
|
en_bert = en_bert.to(device).unsqueeze(0)
|
||||||
x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device)
|
x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device)
|
||||||
style_vec_tensor = torch.from_numpy(style_vector).to(device).unsqueeze(0)
|
style_vec_tensor = torch.from_numpy(style_vector).to(device).unsqueeze(0)
|
||||||
|
sid = 0
|
||||||
|
sid_tensor = torch.LongTensor([sid]).to(device)
|
||||||
|
length_scale = torch.tensor(1.0)
|
||||||
|
sdp_ratio = torch.tensor(0.0)
|
||||||
|
noise_scale = torch.tensor(0.667)
|
||||||
|
noise_scale_w = torch.tensor(0.8)
|
||||||
|
|
||||||
# モデルを ONNX に変換
|
# JP-Extra モデルアーキテクチャ向けの ONNX 変換ロジック
|
||||||
print(Rule(characters="=", style=Style(color="blue")))
|
if tts_model.hyper_parameters.data.use_jp_extra is True:
|
||||||
print(f"[bold cyan]Exporting ONNX model...[/bold cyan]")
|
|
||||||
print(Rule(characters="=", style=Style(color="blue")))
|
# SynthesizerTrnJPExtra の forward メソッドをオーバーライド
|
||||||
export_start_time = time.time()
|
def forward(
|
||||||
torch.onnx.export(
|
x: torch.Tensor,
|
||||||
model=tts_model.net_g,
|
x_lengths: torch.Tensor,
|
||||||
args=(
|
sid: torch.Tensor,
|
||||||
x_tst,
|
tone: torch.Tensor,
|
||||||
x_tst_lengths,
|
language: torch.Tensor,
|
||||||
torch.LongTensor([0]).to(device),
|
bert: torch.Tensor,
|
||||||
tones,
|
style_vec: torch.Tensor,
|
||||||
lang_ids,
|
length_scale: float = 1.0,
|
||||||
bert,
|
sdp_ratio: float = 0.0,
|
||||||
style_vec_tensor,
|
noise_scale: float = 0.667,
|
||||||
torch.tensor(1.0),
|
noise_scale_w: float = 0.8,
|
||||||
torch.tensor(0.0),
|
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, tuple[torch.Tensor, ...]]:
|
||||||
),
|
return cast(SynthesizerTrnJPExtra, tts_model.net_g).infer(
|
||||||
f=str(onnx_temp_model_path),
|
x,
|
||||||
verbose=True,
|
x_lengths,
|
||||||
dynamic_axes={
|
sid,
|
||||||
"x_tst": {1: "batch_size"},
|
tone,
|
||||||
"x_tst_lengths": {0: "batch_size"},
|
language,
|
||||||
"tones": {1: "batch_size"},
|
bert,
|
||||||
"language": {1: "batch_size"},
|
style_vec,
|
||||||
"bert": {2: "batch_size"},
|
length_scale=length_scale,
|
||||||
},
|
sdp_ratio=sdp_ratio,
|
||||||
input_names=[
|
noise_scale=noise_scale,
|
||||||
"x_tst",
|
noise_scale_w=noise_scale_w,
|
||||||
"x_tst_lengths",
|
)
|
||||||
"sid",
|
|
||||||
"tones",
|
tts_model.net_g.forward = forward # type: ignore
|
||||||
"language",
|
|
||||||
"bert",
|
# モデルを ONNX に変換
|
||||||
"style_vec",
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
"length_scale",
|
print(
|
||||||
"sdp_ratio",
|
f"[bold cyan]Exporting ONNX model... (Architecture: JP-Extra)[/bold cyan]"
|
||||||
],
|
)
|
||||||
output_names=["output"],
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
)
|
export_start_time = time.time()
|
||||||
print(
|
torch.onnx.export(
|
||||||
f"[bold green]ONNX model exported to {onnx_temp_model_path} ({time.time() - export_start_time:.2f}s)[/bold green]"
|
model=tts_model.net_g,
|
||||||
)
|
args=(
|
||||||
|
x_tst,
|
||||||
|
x_tst_lengths,
|
||||||
|
sid_tensor,
|
||||||
|
tones,
|
||||||
|
lang_ids,
|
||||||
|
ja_bert,
|
||||||
|
style_vec_tensor,
|
||||||
|
length_scale,
|
||||||
|
sdp_ratio,
|
||||||
|
noise_scale,
|
||||||
|
noise_scale_w,
|
||||||
|
),
|
||||||
|
f=str(onnx_temp_model_path),
|
||||||
|
verbose=False,
|
||||||
|
dynamic_axes={
|
||||||
|
"x_tst": {0: "batch_size", 1: "x_tst_max_length"},
|
||||||
|
"x_tst_lengths": {0: "batch_size"},
|
||||||
|
"sid": {0: "batch_size"},
|
||||||
|
"tones": {0: "batch_size", 1: "x_tst_max_length"},
|
||||||
|
"language": {0: "batch_size", 1: "x_tst_max_length"},
|
||||||
|
"bert": {0: "batch_size", 2: "x_tst_max_length"},
|
||||||
|
"style_vec": {0: "batch_size"},
|
||||||
|
},
|
||||||
|
input_names=[
|
||||||
|
"x_tst",
|
||||||
|
"x_tst_lengths",
|
||||||
|
"sid",
|
||||||
|
"tones",
|
||||||
|
"language",
|
||||||
|
"bert",
|
||||||
|
"style_vec",
|
||||||
|
"length_scale",
|
||||||
|
"sdp_ratio",
|
||||||
|
"noise_scale",
|
||||||
|
"noise_scale_w",
|
||||||
|
],
|
||||||
|
output_names=["output"],
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"[bold green]ONNX model exported to {onnx_temp_model_path} ({time.time() - export_start_time:.2f}s)[/bold green]"
|
||||||
|
)
|
||||||
|
|
||||||
|
else:
|
||||||
|
raise NotImplementedError(
|
||||||
|
"non-JP-Extra model architecture is not implemented yet"
|
||||||
|
)
|
||||||
|
|
||||||
# ONNX モデルを最適化
|
# ONNX モデルを最適化
|
||||||
print(Rule(characters="=", style=Style(color="blue")))
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
@@ -170,7 +196,7 @@ if __name__ == "__main__":
|
|||||||
f"[bold green]ONNX model optimized and saved to {onnx_optimized_model_path} ({time.time() - optimize_start_time:.2f}s)[/bold green]"
|
f"[bold green]ONNX model optimized and saved to {onnx_optimized_model_path} ({time.time() - optimize_start_time:.2f}s)[/bold green]"
|
||||||
)
|
)
|
||||||
print(
|
print(
|
||||||
f"[bold]Total Time: {time.time() - start_time:.2f}s / Size: {onnx_optimized_model_path.stat().st_size / 1024 / 1024:.2f}MB[/bold]"
|
f"[bold]Total Time: {time.time() - start_time:.2f}s / Size: {onnx_optimized_model_path.stat().st_size / 1000 / 1000:.2f}MB[/bold]"
|
||||||
)
|
)
|
||||||
print(Rule(characters="=", style=Style(color="blue")))
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
print("[bold cyan]Optimized model info:[/bold cyan]")
|
print("[bold cyan]Optimized model info:[/bold cyan]")
|
||||||
|
|||||||
@@ -230,10 +230,10 @@ def infer(
|
|||||||
lang_ids,
|
lang_ids,
|
||||||
ja_bert,
|
ja_bert,
|
||||||
style_vec=style_vec_tensor,
|
style_vec=style_vec_tensor,
|
||||||
|
length_scale=length_scale,
|
||||||
sdp_ratio=sdp_ratio,
|
sdp_ratio=sdp_ratio,
|
||||||
noise_scale=noise_scale,
|
noise_scale=noise_scale,
|
||||||
noise_scale_w=noise_scale_w,
|
noise_scale_w=noise_scale_w,
|
||||||
length_scale=length_scale,
|
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
output = cast(SynthesizerTrn, net_g).infer(
|
output = cast(SynthesizerTrn, net_g).infer(
|
||||||
@@ -246,10 +246,10 @@ def infer(
|
|||||||
ja_bert,
|
ja_bert,
|
||||||
en_bert,
|
en_bert,
|
||||||
style_vec=style_vec_tensor,
|
style_vec=style_vec_tensor,
|
||||||
|
length_scale=length_scale,
|
||||||
sdp_ratio=sdp_ratio,
|
sdp_ratio=sdp_ratio,
|
||||||
noise_scale=noise_scale,
|
noise_scale=noise_scale,
|
||||||
noise_scale_w=noise_scale_w,
|
noise_scale_w=noise_scale_w,
|
||||||
length_scale=length_scale,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
audio = output[0][0, 0].data.cpu().float().numpy()
|
audio = output[0][0, 0].data.cpu().float().numpy()
|
||||||
|
|||||||
@@ -167,6 +167,8 @@ def infer_onnx(
|
|||||||
style_vec_tensor,
|
style_vec_tensor,
|
||||||
np.array([length_scale], dtype=np.float32),
|
np.array([length_scale], dtype=np.float32),
|
||||||
np.array([sdp_ratio], dtype=np.float32),
|
np.array([sdp_ratio], dtype=np.float32),
|
||||||
|
np.array([noise_scale], dtype=np.float32),
|
||||||
|
np.array([noise_scale_w], dtype=np.float32),
|
||||||
]
|
]
|
||||||
|
|
||||||
first_provider = onnx_session.get_providers()[0]
|
first_provider = onnx_session.get_providers()[0]
|
||||||
|
|||||||
Reference in New Issue
Block a user