Improve: .safetensors models under a directory specified with --model can be automatically converted to ONNX
This commit is contained in:
@@ -21,7 +21,7 @@ from style_bert_vits2.nlp import bert_models
|
|||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
parser = ArgumentParser()
|
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()
|
args = parser.parse_args()
|
||||||
|
|
||||||
# モデルの入出力先ファイルパスを取得
|
# モデルの入出力先ファイルパスを取得
|
||||||
@@ -70,11 +70,19 @@ if __name__ == "__main__":
|
|||||||
export_start_time = time.time()
|
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"],
|
||||||
|
),
|
||||||
f=str(onnx_temp_model_path),
|
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"],
|
output_names=["output"],
|
||||||
verbose=True,
|
|
||||||
dynamic_axes={
|
dynamic_axes={
|
||||||
"input_ids": {1: "batch_size"},
|
"input_ids": {1: "batch_size"},
|
||||||
"attention_mask": {1: "batch_size"},
|
"attention_mask": {1: "batch_size"},
|
||||||
@@ -92,13 +100,15 @@ if __name__ == "__main__":
|
|||||||
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_temp_model_path.unlink()
|
|
||||||
print(
|
print(
|
||||||
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 / "
|
||||||
|
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(Rule(characters="=", style=Style(color="blue")))
|
||||||
print("[bold cyan]Optimized model info:[/bold cyan]")
|
print("[bold cyan]Optimized model info:[/bold cyan]")
|
||||||
model_info.print_simplifying_info(onnx_model, simplified_onnx_model)
|
model_info.print_simplifying_info(onnx_model, simplified_onnx_model)
|
||||||
|
|||||||
@@ -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/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
|
# ref: https://github.com/tuna2134/sbv2-api/blob/main/convert/convert_model.py
|
||||||
|
|
||||||
import time
|
import time
|
||||||
@@ -29,15 +30,25 @@ from style_bert_vits2.tts_model import TTSModel
|
|||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
parser = ArgumentParser()
|
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()
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
# --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))
|
||||||
|
|
||||||
|
for model_path in model_paths:
|
||||||
|
|
||||||
# モデルの入出力先ファイルパスを取得
|
# モデルの入出力先ファイルパスを取得
|
||||||
model_path = Path(args.model)
|
onnx_temp_model_path = model_path.parent / f"{model_path.stem}_temp.onnx"
|
||||||
onnx_temp_model_path = Path(args.model).parent / f"{model_path.stem}_temp.onnx"
|
onnx_optimized_model_path = model_path.parent / f"{model_path.stem}.onnx"
|
||||||
onnx_optimized_model_path = Path(args.model).parent / f"{model_path.stem}.onnx"
|
config_path = model_path.parent / "config.json"
|
||||||
config_path = Path(args.model).parent / "config.json"
|
style_vec_path = model_path.parent / "style_vectors.npy"
|
||||||
style_vec_path = Path(args.model).parent / "style_vectors.npy"
|
|
||||||
assert model_path.exists(), "Model file does not exist"
|
assert model_path.exists(), "Model file does not exist"
|
||||||
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"
|
||||||
@@ -48,6 +59,13 @@ if __name__ == "__main__":
|
|||||||
print(f"[bold cyan]Style vector file:[/bold cyan] {style_vec_path}")
|
print(f"[bold cyan]Style vector file:[/bold cyan] {style_vec_path}")
|
||||||
print(Rule(characters="=", style=Style(color="blue")))
|
print(Rule(characters="=", style=Style(color="blue")))
|
||||||
|
|
||||||
|
# すでに 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
|
||||||
|
|
||||||
# PyTorch モデルを読み込む
|
# PyTorch モデルを読み込む
|
||||||
device = "cpu"
|
device = "cpu"
|
||||||
tts_model = TTSModel(
|
tts_model = TTSModel(
|
||||||
@@ -150,15 +168,6 @@ if __name__ == "__main__":
|
|||||||
),
|
),
|
||||||
f=str(onnx_temp_model_path),
|
f=str(onnx_temp_model_path),
|
||||||
verbose=False,
|
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=[
|
input_names=[
|
||||||
"x_tst",
|
"x_tst",
|
||||||
"x_tst_lengths",
|
"x_tst_lengths",
|
||||||
@@ -173,6 +182,15 @@ if __name__ == "__main__":
|
|||||||
"noise_scale_w",
|
"noise_scale_w",
|
||||||
],
|
],
|
||||||
output_names=["output"],
|
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(
|
print(
|
||||||
f"[bold green]ONNX model exported to {onnx_temp_model_path} ({time.time() - export_start_time:.2f}s)[/bold green]"
|
f"[bold green]ONNX model exported to {onnx_temp_model_path} ({time.time() - export_start_time:.2f}s)[/bold green]"
|
||||||
@@ -191,13 +209,15 @@ if __name__ == "__main__":
|
|||||||
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_temp_model_path.unlink()
|
|
||||||
print(
|
print(
|
||||||
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 / 1000 / 1000:.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(Rule(characters="=", style=Style(color="blue")))
|
||||||
print("[bold cyan]Optimized model info:[/bold cyan]")
|
print("[bold cyan]Optimized model info:[/bold cyan]")
|
||||||
model_info.print_simplifying_info(onnx_model, simplified_onnx_model)
|
model_info.print_simplifying_info(onnx_model, simplified_onnx_model)
|
||||||
|
|||||||
Reference in New Issue
Block a user