From 16115a366b4d5c9780ba0a07809611820da37649 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Sun, 22 Sep 2024 16:27:22 +0900 Subject: [PATCH] Improve: .safetensors models under a directory specified with --model can be automatically converted to ONNX --- convert_bert_onnx.py | 22 ++- convert_onnx.py | 334 +++++++++++++++++++++++-------------------- 2 files changed, 193 insertions(+), 163 deletions(-) diff --git a/convert_bert_onnx.py b/convert_bert_onnx.py index 5e929d5..a37a971 100644 --- a/convert_bert_onnx.py +++ b/convert_bert_onnx.py @@ -21,7 +21,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) + parser.add_argument("--language", default=Languages.JP, help="Language of the BERT model to be converted") args = parser.parse_args() # モデルの入出力先ファイルパスを取得 @@ -70,11 +70,19 @@ if __name__ == "__main__": export_start_time = time.time() torch.onnx.export( model=model, - args=(inputs["input_ids"], inputs["token_type_ids"], inputs["attention_mask"]), + args=( + inputs["input_ids"], + inputs["token_type_ids"], + inputs["attention_mask"], + ), f=str(onnx_temp_model_path), - input_names=["input_ids", "token_type_ids", "attention_mask"], + verbose=False, + input_names=[ + "input_ids", + "token_type_ids", + "attention_mask", + ], output_names=["output"], - verbose=True, dynamic_axes={ "input_ids": {1: "batch_size"}, "attention_mask": {1: "batch_size"}, @@ -92,13 +100,15 @@ if __name__ == "__main__": 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_temp_model_path.unlink() 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]" + f"[bold]Total Time: {time.time() - start_time:.2f}s / " + f"Size: {onnx_temp_model_path.stat().st_size / 1000 / 1000:.2f}MB -> " + f"{onnx_optimized_model_path.stat().st_size / 1000 / 1000:.2f}MB[/bold]" ) + onnx_temp_model_path.unlink() 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) diff --git a/convert_onnx.py b/convert_onnx.py index 4ecd863..6f34ceb 100644 --- a/convert_onnx.py +++ b/convert_onnx.py @@ -1,4 +1,5 @@ # usage: .venv/bin/python convert_onnx.py --model model_assets/koharune-ami/koharune-ami.safetensors +# usage: .venv/bin/python convert_onnx.py --model model_assets/ (All models in the directory will be converted) # ref: https://github.com/tuna2134/sbv2-api/blob/main/convert/convert_model.py import time @@ -29,176 +30,195 @@ from style_bert_vits2.tts_model import TTSModel if __name__ == "__main__": start_time = time.time() parser = ArgumentParser() - parser.add_argument("--model", required=True) + parser.add_argument("--model", required=True, help="Path to the model file or directory") + parser.add_argument("--force-convert", action="store_true", help="Already converted models will be overwritten") args = parser.parse_args() - # モデルの入出力先ファイルパスを取得 - model_path = Path(args.model) - onnx_temp_model_path = Path(args.model).parent / f"{model_path.stem}_temp.onnx" - onnx_optimized_model_path = Path(args.model).parent / f"{model_path.stem}.onnx" - config_path = Path(args.model).parent / "config.json" - style_vec_path = Path(args.model).parent / "style_vectors.npy" - assert model_path.exists(), "Model 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 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"))) + # --model に指定されたパスがディレクトリの時、配下にある全ての .safetensors ファイルを対象に変換する + model_paths: list[Path] = [] + if Path(args.model).is_dir(): + for path in Path(args.model).glob("**/*.safetensors"): + model_paths.append(path) + else: + model_paths.append(Path(args.model)) - # PyTorch モデルを読み込む - device = "cpu" - tts_model = TTSModel( - model_path=model_path, - config_path=config_path, - style_vec_path=style_vec_path, - device=device, - ) - tts_model.load() - style_id = tts_model.style2id[DEFAULT_STYLE] - assert tts_model.net_g is not None, "Model is not loaded" + for model_path in model_paths: - # 音声合成に必要な 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( - "今日はいい天気ですね。", - Languages.JP, - tts_model.hyper_parameters, - device, - assist_text=None, - assist_text_weight=DEFAULT_ASSIST_TEXT_WEIGHT, - given_phone=None, - given_tone=None, - ) + # モデルの入出力先ファイルパスを取得 + onnx_temp_model_path = model_path.parent / f"{model_path.stem}_temp.onnx" + onnx_optimized_model_path = model_path.parent / f"{model_path.stem}.onnx" + config_path = model_path.parent / "config.json" + style_vec_path = model_path.parent / "style_vectors.npy" + assert model_path.exists(), "Model 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 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"))) - # スタイルベクトルを取得 - style_vector = tts_model.get_style_vector(style_id, DEFAULT_STYLE_WEIGHT) + # すでに ONNX モデルが存在する場合、--force-convert オプションが指定されていない場合はスキップ + if onnx_optimized_model_path.exists() and not args.force_convert: + print(f"[bold yellow]ONNX model already exists: {onnx_optimized_model_path}[/bold yellow]") + print("[bold]If you want to overwrite it, use the --force-convert option.[/bold]") + print(Rule(characters="=", style=Style(color="blue"))) + continue - # モデルの入力を作成 - x_tst = phones.to(device).unsqueeze(0) - tones = tones.to(device).unsqueeze(0) - lang_ids = lang_ids.to(device).unsqueeze(0) - bert = bert.to(device).unsqueeze(0) - ja_bert = ja_bert.to(device).unsqueeze(0) - en_bert = en_bert.to(device).unsqueeze(0) - x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device) - 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) + # PyTorch モデルを読み込む + device = "cpu" + tts_model = TTSModel( + model_path=model_path, + config_path=config_path, + style_vec_path=style_vec_path, + device=device, + ) + tts_model.load() + style_id = tts_model.style2id[DEFAULT_STYLE] + assert tts_model.net_g is not None, "Model is not loaded" - # JP-Extra モデルアーキテクチャ向けの ONNX 変換ロジック - if tts_model.hyper_parameters.data.use_jp_extra is True: + # 音声合成に必要な 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( + "今日はいい天気ですね。", + Languages.JP, + tts_model.hyper_parameters, + device, + assist_text=None, + assist_text_weight=DEFAULT_ASSIST_TEXT_WEIGHT, + given_phone=None, + given_tone=None, + ) - # 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, - noise_scale: float = 0.667, - noise_scale_w: float = 0.8, - ) -> 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, - length_scale=length_scale, - sdp_ratio=sdp_ratio, - noise_scale=noise_scale, - noise_scale_w=noise_scale_w, + # スタイルベクトルを取得 + style_vector = tts_model.get_style_vector(style_id, DEFAULT_STYLE_WEIGHT) + + # モデルの入力を作成 + x_tst = phones.to(device).unsqueeze(0) + tones = tones.to(device).unsqueeze(0) + lang_ids = lang_ids.to(device).unsqueeze(0) + bert = bert.to(device).unsqueeze(0) + ja_bert = ja_bert.to(device).unsqueeze(0) + en_bert = en_bert.to(device).unsqueeze(0) + x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device) + 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) + + # JP-Extra モデルアーキテクチャ向けの ONNX 変換ロジック + if tts_model.hyper_parameters.data.use_jp_extra is True: + + # 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, + noise_scale: float = 0.667, + noise_scale_w: float = 0.8, + ) -> 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, + length_scale=length_scale, + sdp_ratio=sdp_ratio, + noise_scale=noise_scale, + noise_scale_w=noise_scale_w, + ) + + tts_model.net_g.forward = forward # type: ignore + + # モデルを ONNX に変換 + print(Rule(characters="=", style=Style(color="blue"))) + print( + f"[bold cyan]Exporting ONNX model... (Architecture: JP-Extra)[/bold cyan]" + ) + print(Rule(characters="=", style=Style(color="blue"))) + export_start_time = time.time() + torch.onnx.export( + 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, + 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"], + 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"}, + }, + ) + print( + f"[bold green]ONNX model exported to {onnx_temp_model_path} ({time.time() - export_start_time:.2f}s)[/bold green]" ) - tts_model.net_g.forward = forward # type: ignore + else: + raise NotImplementedError( + "non-JP-Extra model architecture is not implemented yet" + ) - # モデルを 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) + simplified_onnx_model, check = simplify(onnx_model) + onnx.save(simplified_onnx_model, onnx_optimized_model_path) print( - f"[bold cyan]Exporting ONNX model... (Architecture: JP-Extra)[/bold cyan]" - ) - print(Rule(characters="=", style=Style(color="blue"))) - export_start_time = time.time() - torch.onnx.export( - 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"], + 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 green]ONNX model exported to {onnx_temp_model_path} ({time.time() - export_start_time:.2f}s)[/bold green]" + f"[bold]Total Time: {time.time() - start_time:.2f}s / " + f"Size: {onnx_temp_model_path.stat().st_size / 1000 / 1000:.2f}MB -> " + f"{onnx_optimized_model_path.stat().st_size / 1000 / 1000:.2f}MB[/bold]" ) - - else: - raise NotImplementedError( - "non-JP-Extra model architecture is not implemented yet" - ) - - # 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_temp_model_path.unlink() - 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 / 1000 / 1000:.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"))) + onnx_temp_model_path.unlink() + 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")))