Improve: Graphical display of ONNX conversion script logs

This commit is contained in:
tsukumi
2024-09-18 06:09:49 +09:00
parent 449927616a
commit 9b27ba9777
5 changed files with 74 additions and 14 deletions

View File

@@ -1,12 +1,16 @@
# usage: .venv/bin/python convert_bert_onnx.py --language JP # usage: .venv/bin/python convert_bert_onnx.py --language JP
# ref: https://github.com/tuna2134/sbv2-api/blob/main/convert/convert_deberta.py # ref: https://github.com/tuna2134/sbv2-api/blob/main/convert/convert_deberta.py
import time
from argparse import ArgumentParser from argparse import ArgumentParser
from pathlib import Path from pathlib import Path
import onnx import onnx
import torch import torch
from onnxsim import simplify from onnxsim import model_info, simplify
from rich import print
from rich.rule import Rule
from rich.style import Style
from torch import nn from torch import nn
from transformers.convert_slow_tokenizer import BertConverter from transformers.convert_slow_tokenizer import BertConverter
@@ -15,6 +19,7 @@ from style_bert_vits2.nlp import bert_models
if __name__ == "__main__": if __name__ == "__main__":
start_time = time.time()
parser = ArgumentParser() parser = ArgumentParser()
parser.add_argument("--language", default=Languages.JP) parser.add_argument("--language", default=Languages.JP)
args = parser.parse_args() args = parser.parse_args()
@@ -25,11 +30,18 @@ if __name__ == "__main__":
onnx_temp_model_path = Path(pretrained_model_name_or_path) / f"model_temp.onnx" onnx_temp_model_path = Path(pretrained_model_name_or_path) / f"model_temp.onnx"
onnx_optimized_model_path = Path(pretrained_model_name_or_path) / f"model.onnx" onnx_optimized_model_path = Path(pretrained_model_name_or_path) / f"model.onnx"
tokenizer_json_path = Path(pretrained_model_name_or_path) / "tokenizer.json" tokenizer_json_path = Path(pretrained_model_name_or_path) / "tokenizer.json"
print(Rule(characters="=", style=Style(color="blue")))
print(f"[bold cyan]Language:[/bold cyan] {language.name}")
print(f"[bold cyan]Pretrained model:[/bold cyan] {pretrained_model_name_or_path}")
print(Rule(characters="=", style=Style(color="blue")))
# トークナイザーを Fast Tokenizer 用形式に変換して保存 # トークナイザーを Fast Tokenizer 用形式に変換して保存
tokenizer = bert_models.load_tokenizer(language) tokenizer = bert_models.load_tokenizer(language)
converter = BertConverter(tokenizer) converter = BertConverter(tokenizer)
converter.converted().save(str(tokenizer_json_path)) converter.converted().save(str(tokenizer_json_path))
print(Rule(characters="=", style=Style(color="blue")))
print(f"[bold green]Tokenizer JSON saved to {tokenizer_json_path}[/bold green]")
print(Rule(characters="=", style=Style(color="blue")))
# TODO: JP, ZH は変換できるが、EN は途中で強制終了されてしまい変換できない # TODO: JP, ZH は変換できるが、EN は途中で強制終了されてしまい変換できない
class ONNXBert(nn.Module): class ONNXBert(nn.Module):
@@ -52,6 +64,10 @@ if __name__ == "__main__":
inputs = tokenizer("今日はいい天気ですね", return_tensors="pt") inputs = tokenizer("今日はいい天気ですね", return_tensors="pt")
# モデルを ONNX に変換 # モデルを ONNX に変換
print(Rule(characters="=", style=Style(color="blue")))
print(f"[bold cyan]Exporting ONNX model...[/bold cyan]")
print(Rule(characters="=", style=Style(color="blue")))
export_start_time = time.time()
torch.onnx.export( torch.onnx.export(
model=model, model=model,
args=(inputs["input_ids"], inputs["token_type_ids"], inputs["attention_mask"]), args=(inputs["input_ids"], inputs["token_type_ids"], inputs["attention_mask"]),
@@ -64,12 +80,26 @@ if __name__ == "__main__":
"attention_mask": {1: "batch_size"}, "attention_mask": {1: "batch_size"},
}, },
) )
print(
f"[bold green]ONNX model exported to {onnx_temp_model_path} ({time.time() - export_start_time:.2f}s)[/bold green]"
)
# ONNX モデルを最適化 # ONNX モデルを最適化
print(Rule(characters="=", style=Style(color="blue")))
print(f"[bold cyan]Optimizing ONNX model...[/bold cyan]")
print(Rule(characters="=", style=Style(color="blue")))
optimize_start_time = time.time()
onnx_model = onnx.load(onnx_temp_model_path) onnx_model = onnx.load(onnx_temp_model_path)
simplified_onnx_model, check = simplify(onnx_model) simplified_onnx_model, check = simplify(onnx_model)
onnx.save(simplified_onnx_model, onnx_optimized_model_path) onnx.save(simplified_onnx_model, onnx_optimized_model_path)
# 最適化前の ONNX モデルを削除
onnx_temp_model_path.unlink() onnx_temp_model_path.unlink()
print(f"ONNX model optimized and saved to {onnx_optimized_model_path}") print(
f"[bold green]ONNX model optimized and saved to {onnx_optimized_model_path} ({time.time() - optimize_start_time:.2f}s)[/bold green]"
)
print(
f"[bold]Total Time: {time.time() - start_time:.2f}s / Size: {onnx_optimized_model_path.stat().st_size / 1024 / 1024:.2f}MB[/bold]"
)
print(Rule(characters="=", style=Style(color="blue")))
print("[bold cyan]Optimized model info:[/bold cyan]")
model_info.print_simplifying_info(onnx_model, simplified_onnx_model)
print(Rule(characters="=", style=Style(color="blue")))

View File

@@ -1,13 +1,17 @@
# usage: .venv/bin/python convert_onnx.py --model model_assets/amitaro/amitaro.safetensors # usage: .venv/bin/python convert_onnx.py --model model_assets/amitaro/amitaro.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
from argparse import ArgumentParser from argparse import ArgumentParser
from pathlib import Path from pathlib import Path
from typing import cast from typing import cast
import onnx import onnx
import torch import torch
from onnxsim import simplify from onnxsim import model_info, simplify
from rich import print
from rich.rule import Rule
from rich.style import Style
from style_bert_vits2.constants import ( from style_bert_vits2.constants import (
DEFAULT_ASSIST_TEXT_WEIGHT, DEFAULT_ASSIST_TEXT_WEIGHT,
@@ -23,6 +27,7 @@ from style_bert_vits2.tts_model import TTSModel
if __name__ == "__main__": if __name__ == "__main__":
start_time = time.time()
parser = ArgumentParser() parser = ArgumentParser()
parser.add_argument("--model", required=True) parser.add_argument("--model", required=True)
args = parser.parse_args() args = parser.parse_args()
@@ -37,6 +42,11 @@ if __name__ == "__main__":
assert config_path.exists(), "Config file does not exist" assert config_path.exists(), "Config file does not exist"
assert style_vec_path.exists(), "Style vector file does not exist" assert style_vec_path.exists(), "Style vector file does not exist"
assert model_path.suffix != ".onnx", "Model file is already ONNX" assert model_path.suffix != ".onnx", "Model file is already ONNX"
print(Rule(characters="=", style=Style(color="blue")))
print(f"[bold cyan]Model file:[/bold cyan] {model_path}")
print(f"[bold cyan]Config file:[/bold cyan] {config_path}")
print(f"[bold cyan]Style vector file:[/bold cyan] {style_vec_path}")
print(Rule(characters="=", style=Style(color="blue")))
# PyTorch モデルを読み込む # PyTorch モデルを読み込む
device = "cpu" device = "cpu"
@@ -104,6 +114,10 @@ if __name__ == "__main__":
style_vec_tensor = torch.from_numpy(style_vector).to(device).unsqueeze(0) style_vec_tensor = torch.from_numpy(style_vector).to(device).unsqueeze(0)
# モデルを ONNX に変換 # モデルを ONNX に変換
print(Rule(characters="=", style=Style(color="blue")))
print(f"[bold cyan]Exporting ONNX model...[/bold cyan]")
print(Rule(characters="=", style=Style(color="blue")))
export_start_time = time.time()
torch.onnx.export( torch.onnx.export(
model=tts_model.net_g, model=tts_model.net_g,
args=( args=(
@@ -139,12 +153,26 @@ if __name__ == "__main__":
], ],
output_names=["output"], output_names=["output"],
) )
print(
f"[bold green]ONNX model exported to {onnx_temp_model_path} ({time.time() - export_start_time:.2f}s)[/bold green]"
)
# ONNX モデルを最適化 # ONNX モデルを最適化
print(Rule(characters="=", style=Style(color="blue")))
print(f"[bold cyan]Optimizing ONNX model...[/bold cyan]")
print(Rule(characters="=", style=Style(color="blue")))
optimize_start_time = time.time()
onnx_model = onnx.load(onnx_temp_model_path) onnx_model = onnx.load(onnx_temp_model_path)
simplified_onnx_model, check = simplify(onnx_model) simplified_onnx_model, check = simplify(onnx_model)
onnx.save(simplified_onnx_model, onnx_optimized_model_path) onnx.save(simplified_onnx_model, onnx_optimized_model_path)
# 最適化前の ONNX モデルを削除
onnx_temp_model_path.unlink() onnx_temp_model_path.unlink()
print(f"ONNX model optimized and saved to {onnx_optimized_model_path}") print(
f"[bold green]ONNX model optimized and saved to {onnx_optimized_model_path} ({time.time() - optimize_start_time:.2f}s)[/bold green]"
)
print(
f"[bold]Total Time: {time.time() - start_time:.2f}s / Size: {onnx_optimized_model_path.stat().st_size / 1024 / 1024:.2f}MB[/bold]"
)
print(Rule(characters="=", style=Style(color="blue")))
print("[bold cyan]Optimized model info:[/bold cyan]")
model_info.print_simplifying_info(onnx_model, simplified_onnx_model)
print(Rule(characters="=", style=Style(color="blue")))

View File

@@ -202,9 +202,7 @@ if args.preload_onnx_bert:
) )
onnx_bert_models.load_tokenizer(Languages.JP) onnx_bert_models.load_tokenizer(Languages.JP)
model_holder = TTSModelHolder( model_holder = TTSModelHolder(model_dir, device, torch_device_to_onnx_providers(device))
model_dir, device, torch_device_to_onnx_providers(device)
)
if len(model_holder.model_names) == 0: if len(model_holder.model_names) == 0:
logger.error(f"Models not found in {model_dir}.") logger.error(f"Models not found in {model_dir}.")
sys.exit(1) sys.exit(1)

View File

@@ -97,7 +97,9 @@ def load_model(
sess_options = onnxruntime.SessionOptions() sess_options = onnxruntime.SessionOptions()
# 基本的な最適化のみ有効化 # 基本的な最適化のみ有効化
# ONNX モデルの作成時にすでに onnxsim により最適化されているため、ここでは基本的な最適化のみ有効化する # ONNX モデルの作成時にすでに onnxsim により最適化されているため、ここでは基本的な最適化のみ有効化する
sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_ENABLE_BASIC sess_options.graph_optimization_level = (
onnxruntime.GraphOptimizationLevel.ORT_ENABLE_BASIC
)
# エラー以外のログを出力しない # エラー以外のログを出力しない
# 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している # 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している
sess_options.log_severity_level = 3 sess_options.log_severity_level = 3

View File

@@ -132,7 +132,7 @@ class TTSModel:
hps=self.hyper_parameters, hps=self.hyper_parameters,
) )
logger.info( logger.info(
f"Model loaded successfully from {self.model_path} to \"{self.device}\" device ({time.time() - start_time:.2f}s)" f'Model loaded successfully from {self.model_path} to "{self.device}" device ({time.time() - start_time:.2f}s)'
) )
# ONNX 推論時 # ONNX 推論時
@@ -140,7 +140,9 @@ class TTSModel:
sess_options = onnxruntime.SessionOptions() sess_options = onnxruntime.SessionOptions()
# 基本的な最適化のみ有効化 # 基本的な最適化のみ有効化
# ONNX モデルの作成時にすでに onnxsim により最適化されているため、ここでは基本的な最適化のみ有効化する # ONNX モデルの作成時にすでに onnxsim により最適化されているため、ここでは基本的な最適化のみ有効化する
sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_ENABLE_BASIC sess_options.graph_optimization_level = (
onnxruntime.GraphOptimizationLevel.ORT_ENABLE_BASIC
)
# エラー以外のログを出力しない # エラー以外のログを出力しない
# 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している # 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している
sess_options.log_severity_level = 3 sess_options.log_severity_level = 3