From c1fce3fec7065716f88c9f0f2b0e79289fe722d8 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Thu, 19 Dec 2024 07:16:59 +0900 Subject: [PATCH] Improve: Make it possible to convert BERT language models to FP16 --- convert_bert_onnx.py | 267 +++++++++++++++++++++++++++++++++++++++-- requirements-colab.txt | 1 + requirements-infer.txt | 1 + requirements.txt | 1 + 4 files changed, 259 insertions(+), 11 deletions(-) diff --git a/convert_bert_onnx.py b/convert_bert_onnx.py index f30935c..46f962e 100644 --- a/convert_bert_onnx.py +++ b/convert_bert_onnx.py @@ -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: {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 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"))) diff --git a/requirements-colab.txt b/requirements-colab.txt index 0664647..726a2d2 100644 --- a/requirements-colab.txt +++ b/requirements-colab.txt @@ -10,6 +10,7 @@ nltk<=3.8.1 num2words numpy<2 onnx +onnxconverter-common onnxruntime onnxruntime-gpu onnxsim diff --git a/requirements-infer.txt b/requirements-infer.txt index f856be7..6d4c7b0 100644 --- a/requirements-infer.txt +++ b/requirements-infer.txt @@ -12,6 +12,7 @@ nltk<=3.8.1 num2words numpy<2 onnx +onnxconverter-common onnxruntime onnxruntime-directml; sys_platform == 'win32' onnxruntime-gpu; sys_platform != 'darwin' diff --git a/requirements.txt b/requirements.txt index 360de3b..3e58e33 100644 --- a/requirements.txt +++ b/requirements.txt @@ -12,6 +12,7 @@ nltk<=3.8.1 num2words numpy<2 onnx +onnxconverter-common onnxruntime onnxruntime-directml; sys_platform == 'win32' onnxruntime-gpu; sys_platform != 'darwin'