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__":
|
||||
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)
|
||||
|
||||
334
convert_onnx.py
334
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")))
|
||||
|
||||
Reference in New Issue
Block a user