diff --git a/convert_bert_onnx.py b/convert_bert_onnx.py index e9cefdd..5e929d5 100644 --- a/convert_bert_onnx.py +++ b/convert_bert_onnx.py @@ -1,12 +1,16 @@ # usage: .venv/bin/python convert_bert_onnx.py --language JP # ref: https://github.com/tuna2134/sbv2-api/blob/main/convert/convert_deberta.py +import time from argparse import ArgumentParser from pathlib import Path import onnx 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 transformers.convert_slow_tokenizer import BertConverter @@ -15,6 +19,7 @@ from style_bert_vits2.nlp import bert_models if __name__ == "__main__": + start_time = time.time() parser = ArgumentParser() parser.add_argument("--language", default=Languages.JP) 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_optimized_model_path = Path(pretrained_model_name_or_path) / f"model.onnx" 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 用形式に変換して保存 tokenizer = bert_models.load_tokenizer(language) converter = BertConverter(tokenizer) 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 は途中で強制終了されてしまい変換できない class ONNXBert(nn.Module): @@ -52,6 +64,10 @@ if __name__ == "__main__": inputs = tokenizer("今日はいい天気ですね", return_tensors="pt") # モデルを 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( model=model, args=(inputs["input_ids"], inputs["token_type_ids"], inputs["attention_mask"]), @@ -64,12 +80,26 @@ if __name__ == "__main__": "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 モデルを最適化 + 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) simplified_onnx_model, check = simplify(onnx_model) onnx.save(simplified_onnx_model, onnx_optimized_model_path) - - # 最適化前の ONNX モデルを削除 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"))) diff --git a/convert_onnx.py b/convert_onnx.py index 9a8edb2..932364c 100644 --- a/convert_onnx.py +++ b/convert_onnx.py @@ -1,13 +1,17 @@ # 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 +import time from argparse import ArgumentParser from pathlib import Path from typing import cast import onnx 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 ( DEFAULT_ASSIST_TEXT_WEIGHT, @@ -23,6 +27,7 @@ from style_bert_vits2.tts_model import TTSModel if __name__ == "__main__": + start_time = time.time() parser = ArgumentParser() parser.add_argument("--model", required=True) args = parser.parse_args() @@ -37,6 +42,11 @@ if __name__ == "__main__": assert config_path.exists(), "Config 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" + 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 モデルを読み込む device = "cpu" @@ -104,6 +114,10 @@ if __name__ == "__main__": style_vec_tensor = torch.from_numpy(style_vector).to(device).unsqueeze(0) # モデルを 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( model=tts_model.net_g, args=( @@ -139,12 +153,26 @@ if __name__ == "__main__": ], 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 モデルを最適化 + 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) simplified_onnx_model, check = simplify(onnx_model) onnx.save(simplified_onnx_model, onnx_optimized_model_path) - - # 最適化前の ONNX モデルを削除 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"))) diff --git a/server_editor.py b/server_editor.py index 50921c1..a1e95d0 100644 --- a/server_editor.py +++ b/server_editor.py @@ -202,9 +202,7 @@ if args.preload_onnx_bert: ) onnx_bert_models.load_tokenizer(Languages.JP) -model_holder = TTSModelHolder( - model_dir, device, torch_device_to_onnx_providers(device) -) +model_holder = TTSModelHolder(model_dir, device, torch_device_to_onnx_providers(device)) if len(model_holder.model_names) == 0: logger.error(f"Models not found in {model_dir}.") sys.exit(1) diff --git a/style_bert_vits2/nlp/onnx_bert_models.py b/style_bert_vits2/nlp/onnx_bert_models.py index b14397d..442f2ad 100644 --- a/style_bert_vits2/nlp/onnx_bert_models.py +++ b/style_bert_vits2/nlp/onnx_bert_models.py @@ -97,7 +97,9 @@ def load_model( sess_options = onnxruntime.SessionOptions() # 基本的な最適化のみ有効化 # 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() も呼び出している sess_options.log_severity_level = 3 diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index 2ea10a7..bd56de7 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -132,7 +132,7 @@ class TTSModel: hps=self.hyper_parameters, ) 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 推論時 @@ -140,7 +140,9 @@ class TTSModel: sess_options = onnxruntime.SessionOptions() # 基本的な最適化のみ有効化 # 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() も呼び出している sess_options.log_severity_level = 3