Improve: Graphical display of ONNX conversion script logs
This commit is contained in:
@@ -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")))
|
||||||
|
|||||||
@@ -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")))
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user