@@ -5,19 +5,204 @@ import time
from argparse import ArgumentParser
from pathlib import Path
import numpy as np
import onnx
import torch
from onnxconverter_common import float16 as float16_converter
from onnxruntime import InferenceSession
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 import PreTrainedTokenizerBase
from transformers . convert_slow_tokenizer import BertConverter
from style_bert_vits2 . constants import DEFAULT_BERT_MODEL_PATHS , Languages
from style_bert_vits2 . nlp import bert_models
def validate_model_outputs (
language : Languages ,
original_model : nn . Module ,
onnx_session : InferenceSession ,
tokenizer : PreTrainedTokenizerBase ,
max_diff_threshold : float = 1e-3 ,
mean_diff_threshold : float = 1e-4 ,
) - > tuple [ bool , str ] :
""" ONNXモデルの出力を検証 """
if language == Languages . JP :
test_texts = [
" 今日はすっごく楽しかったよ!また遊ぼうね! " ,
" えー、そんなの嫌だよ。もう二度と行きたくないな… " ,
" わーい!プレゼントありがとう!大好き! " ,
" あのね、実は昨日泣いちゃったんだ。寂しくて… " ,
" もう!怒ったからね!知らないもん! " ,
" ごめんなさい…私が悪かったです。許してください。 " ,
" やったー! テストで100点取れたよ! すっごく嬉しい! " ,
" あら、素敵なお洋服ね。とてもお似合いですわ。 " ,
" うーん、それは難しい選択ですね…よく考えましょう。 " ,
" おはよう!今日も一日頑張ろうね!元気いっぱいだよ! " ,
" こんなに美味しいご飯、初めて食べました!感動です! " ,
" ちょっと待って!その話、すっごく気になる! " ,
" はぁ…疲れた。今日は本当に大変な一日だったよ。 " ,
" きゃー!虫!虫がいるよ!誰か助けて! " ,
" 私ね、将来は宇宙飛行士になりたいの!夢があるでしょ? " ,
" あれ?鍵がない…どこに置いたっけ…困ったなぁ… " ,
" お誕生日おめでとう!素敵な一年になりますように! " ,
" えっと…その…好きです!付き合ってください! " ,
" まったく、いつも心配ばかりかけて…でも、ありがとう。 " ,
" よーし!今日は徹夜で全部やっちゃうぞー! " ,
" 本日のニュースをお伝えします。 " ,
" それは、静かな冬の朝のことでした。雪が街を真っ白に染め上げ、人々はまだ深い眠りの中にいました。 " ,
" ただいまより、会議を始めさせていただきます。 " ,
" 彼女は窓際に立ち、遠く輝く星々を見上げていました。その瞳には、かすかな涙が光っていたのです。 " ,
" 只今の時刻は10時を回りました。 " ,
" 本日のハイライトをお届けいたします。 " ,
" 古びた図書館の奥で、一冊の不思議な本を見つけた少年は、そっとページを開きました。 " ,
" 本製品の特徴について、ご説明いたします。 " ,
" 春風が桜の花びらを舞わせる中、彼は十年ぶりに故郷の駅に降り立ちました。 " ,
" 続いて、天気予報です。 " ,
" お客様へのご案内を申し上げます。 " ,
" 深い森の中、こだまする鳥のさえずりと、せせらぎの音だけが時を刻んでいました。 " ,
" 月明かりに照らされた海は、まるで無数の宝石を散りばめたように輝いていました。 " ,
" このたびの新商品発表会にようこそ。 " ,
" 古城の階段を上るたび、彼女の心臓は激しく鼓動を打ちました。この先で何が待ち受けているのか… " ,
" 本日のメニューをご紹介いたします。 " ,
" 時計の針が真夜中を指す瞬間、不思議な出来事の幕が開けたのです。 " ,
" 次の停車駅は、東京駅です。 " ,
" 霧深い港町で、一通の差出人不明の手紙が彼を待っていました。 " ,
" 幼い頃に聞いた祖母の物語は、今でも鮮明に心に刻まれています。 " ,
]
elif language == Languages . EN :
test_texts = [
" Today was so much fun! Let ' s play again! " ,
" Ugh, I hate that. I never want to go there again... " ,
" Yay! Thank you for the present! I love it! " ,
" You know, I actually cried yesterday. I was feeling lonely... " ,
" That ' s it! I ' m angry! I don ' t care anymore! " ,
" I ' m sorry... It was my fault. Please forgive me. " ,
" I did it! I got 100 on the test! I ' m so happy! " ,
" My, what a lovely dress. It suits you perfectly. " ,
" Hmm, that ' s a difficult choice... Let ' s think about it carefully. " ,
" Good morning! Let ' s do our best today! I ' m full of energy! " ,
" This is the most delicious food I ' ve ever had! I ' m moved! " ,
" Wait a minute! That story sounds really interesting! " ,
" Sigh... I ' m tired. Today was really tough. " ,
" Eek! A bug! A bug! Someone help! " ,
" You know what? I want to become an astronaut! Isn ' t that a great dream? " ,
" Huh? Where are my keys... Where did I put them... This is troubling... " ,
" Happy birthday! May you have a wonderful year ahead! " ,
" Um... well... I like you! Please go out with me! " ,
" Geez, you always make me worry... but thank you. " ,
" Alright! I ' m going to pull an all-nighter and finish everything today! " ,
" Now for today ' s news. " ,
" It was a quiet winter morning. Snow had painted the town white, and people were still deep in slumber. " ,
" Let us now begin the meeting. " ,
" She stood by the window, gazing at the distant stars. Tears glistened faintly in her eyes. " ,
" The time is now just past 10 o ' clock. " ,
" Here are today ' s highlights. " ,
" In the depths of an old library, a boy found a mysterious book and gently opened its pages. " ,
" Let me explain the features of this product. " ,
" As spring winds scattered cherry blossoms, he stepped onto his hometown station platform for the first time in ten years. " ,
" And now, the weather forecast. " ,
" An announcement for our customers. " ,
" Deep in the forest, only the echoing birdsong and the sound of flowing water marked the passage of time. " ,
" The moonlit sea sparkled like countless scattered jewels. " ,
" Welcome to today ' s new product announcement. " ,
" With each step up the castle stairs, her heart beat faster, wondering what awaited ahead... " ,
" Let me introduce today ' s menu. " ,
" As the clock struck midnight, the curtain rose on a strange occurrence. " ,
" The next stop is Tokyo Station. " ,
" In the foggy port town, a letter with no sender awaited him. " ,
" The story my grandmother told me in my childhood remains vivid in my heart. " ,
]
elif language == Languages . ZH :
test_texts = [
" 今天真是太开心了!下次再一起玩吧! " ,
" 唉,我讨厌那样。我再也不想去那里了... " ,
" 耶!谢谢你的礼物!我好喜欢! " ,
" 其实呢,我昨天哭了。因为感到很寂寞... " ,
" 够了!我生气了!我不管了! " ,
" 对不起...都是我的错。请原谅我。 " ,
" 太棒了! 考试得了100分! 我好高兴! " ,
" 哎呀,多漂亮的衣服啊。真适合你。 " ,
" 嗯,这是个难决定...让我们好好考虑一下。 " ,
" 早安!今天也要加油哦!我充满干劲! " ,
" 这是我吃过最好吃的饭!太感动了! " ,
" 等一下!那个故事听起来很有趣! " ,
" 唉...好累啊。今天真是辛苦的一天。 " ,
" 啊!虫子!虫子!谁来帮帮我! " ,
" 你知道吗?我想成为宇航员!这是个好梦想吧? " ,
" 咦?钥匙呢...放在哪里了...真是麻烦... " ,
" 生日快乐!祝你度过美好的一年! " ,
" 那个...就是...我喜欢你!请和我交往! " ,
" 真是的,总是让人担心...不过,谢谢你。 " ,
" 好!今天我要熬夜把所有事情都做完! " ,
" 现在播报今日新闻。 " ,
" 那是个安静的冬日早晨。白雪覆盖了整个城镇,人们还在沉睡中。 " ,
" 现在开始会议。 " ,
" 她站在窗边,仰望着远处的星星。她的眼中闪烁着微弱的泪光。 " ,
" 现在时间刚过十点。 " ,
" 为您播报今日要闻。 " ,
" 在古老图书馆的深处,一个男孩发现了一本神秘的书,轻轻地翻开了书页。 " ,
" 让我为您介绍本产品的特点。 " ,
" 春风吹散樱花花瓣时,他时隔十年重返故乡车站。 " ,
" 接下来是天气预报。 " ,
" 现在为顾客播报通知。 " ,
" 在深邃的森林中,只有鸟鸣的回声和溪水的声音在记录着时间的流逝。 " ,
" 月光照耀下的大海,闪烁得像撒满了无数宝石。 " ,
" 欢迎参加今天的新产品发布会。 " ,
" 每上一级城堡的台阶,她的心跳就加快一分,不知前方等待着什么... " ,
" 让我为您介绍今日菜单。 " ,
" 当时钟指向午夜时分,一个奇异的事件拉开了序幕。 " ,
" 下一站是东京站。 " ,
" 在雾气弥漫的港口小镇,一封没有署名的信在等待着他。 " ,
" 童年时奶奶讲的故事,至今仍清晰地铭刻在我的心中。 " ,
]
max_diff = 0
mean_diff = 0
# セッションの入力名を取得
input_names = [ input . name for input in onnx_session . get_inputs ( ) ]
original_model . eval ( )
with torch . no_grad ( ) :
for text in test_texts :
# PyTorch
inputs = tokenizer ( text , return_tensors = " pt " )
torch_output = original_model (
inputs [ " input_ids " ] ,
inputs [ " token_type_ids " ] ,
inputs [ " attention_mask " ] ,
) . numpy ( )
# ONNX
onnx_inputs = { }
if " input_ids " in input_names :
onnx_inputs [ " input_ids " ] = inputs [ " input_ids " ] . numpy ( ) . astype ( np . int64 ) # type: ignore
if " token_type_ids " in input_names :
onnx_inputs [ " token_type_ids " ] = inputs [ " token_type_ids " ] . numpy ( ) . astype ( np . int64 ) # type: ignore
if " attention_mask " in input_names :
onnx_inputs [ " attention_mask " ] = inputs [ " attention_mask " ] . numpy ( ) . astype ( np . int64 ) # type: ignore
onnx_output = onnx_session . run ( None , onnx_inputs ) [ 0 ]
# 差分を計算
diff = np . abs ( torch_output - onnx_output )
max_diff = max ( max_diff , np . max ( diff ) )
mean_diff = max ( mean_diff , np . mean ( diff ) )
is_valid = max_diff < max_diff_threshold and mean_diff < mean_diff_threshold
message = (
f " Validation { ' passed ' if is_valid else ' failed ' } \n "
f " Max difference: { max_diff : .6f } (threshold: { max_diff_threshold } ) \n "
f " Mean difference: { mean_diff : .6f } (threshold: { mean_diff_threshold } ) "
)
return is_valid , message
if __name__ == " __main__ " :
start_time = time . time ( )
parser = ArgumentParser ( )
@@ -32,8 +217,10 @@ if __name__ == "__main__":
language = Languages ( args . language )
pretrained_model_name_or_path = DEFAULT_BERT_MODEL_PATHS [ language ]
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_fp32 _model_path = Path ( pretrained_model_name_or_path ) / f " model.onnx "
onnx_fp16_model_path = Path ( pretrained_model_name_or_path ) / f " model_fp16.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 } " )
@@ -93,7 +280,7 @@ if __name__ == "__main__":
} ,
)
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 ( { time . time ( ) - export_start_time : .2f } s)[/bold green] "
)
# ONNX モデルを最適化
@@ -103,17 +290,75 @@ if __name__ == "__main__":
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 . save ( simplified_onnx_model , onnx_fp32 _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] "
f " [bold green]ONNX model optimized ( { time . time ( ) - optimize_start_time : .2f } s)[/bold green] "
)
print (
f " [bold]Total Time: { tim e . tim e( ) - 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 = " = " , styl e= Styl e( color = " blue " ) ) )
print ( " [bold cyan]Optimized ONNX model info:[/bold cyan] " )
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 )
# FP32 モデルの検証
print ( Rule ( characters = " = " , style = Style ( color = " blue " ) ) )
print ( f " [bold cyan]Validating FP32 model...[/bold cyan] " )
session = InferenceSession (
str ( onnx_fp32_model_path ) ,
providers = [ " CPUExecutionProvider " ] ,
)
is_valid , message = validate_model_outputs ( language , model , session , tokenizer )
color = " green " if is_valid else " red "
print ( f " [bold { color } ] { message } [/bold { color } ] " )
if is_valid :
# FP16 への変換
print ( Rule ( characters = " = " , style = Style ( color = " blue " ) ) )
print ( f " [bold cyan]Converting to FP16...[/bold cyan] " )
print ( Rule ( characters = " = " , style = Style ( color = " blue " ) ) )
fp16_start_time = time . time ( )
fp16_model = float16_converter . convert_float_to_float16 (
simplified_onnx_model ,
keep_io_types = True , # 入出力は float32 のまま
disable_shape_infer = True ,
)
onnx . save ( fp16_model , onnx_fp16_model_path )
print (
f " [bold green]FP16 conversion completed ( { time . time ( ) - fp16_start_time : .2f } s)[/bold green] "
)
# FP16 モデルの検証
print ( Rule ( characters = " = " , style = Style ( color = " blue " ) ) )
print ( f " [bold cyan]Validating FP16 model...[/bold cyan] " )
session = InferenceSession (
str ( onnx_fp16_model_path ) ,
providers = [ " CPUExecutionProvider " ] ,
)
is_valid , message = validate_model_outputs (
language ,
model ,
session ,
tokenizer ,
max_diff_threshold = 1e-2 , # FP16なのでより緩い閾値を設定
mean_diff_threshold = 1e-3 ,
)
color = " green " if is_valid else " red "
print ( f " [bold { color } ] { message } [/bold { color } ] " )
# サイズ情報の表示
print ( Rule ( characters = " = " , style = Style ( color = " blue " ) ) )
print ( " [bold cyan]Model size information:[/bold cyan] " )
original_size = onnx_temp_model_path . stat ( ) . st_size / 1000 / 1000
fp32_size = onnx_fp32_model_path . stat ( ) . st_size / 1000 / 1000
fp16_size = onnx_fp16_model_path . stat ( ) . st_size / 1000 / 1000
print ( f " Original: { original_size : .2f } MB " )
print ( f " Optimized (FP32): { fp32_size : .2f } MB " )
print ( f " Optimized (FP16): { fp16_size : .2f } MB " )
print ( f " Size reduction: { ( 1 - fp16_size / original_size ) * 100 : .1f } % " )
# 一時ファイルの削除
onnx_temp_model_path . unlink ( )
print ( Rule ( characters = " = " , style = Style ( color = " blue " ) ) )
print ( f " [bold green]Total time: { time . time ( ) - start_time : .2f } s[/bold green] " )
print ( Rule ( characters = " = " , style = Style ( color = " blue " ) ) )