グラフ最適化が変えるもの — ノードを数えて確かめる
一言でいうと
onnxruntimeは、セッションを開くときにグラフを書き直します。何がどこまで変わるかは最適化レベルが決め、その結果はファイルに取り出してノードを数えることで確認できます。
なぜ必要なのか
量子化の前後で速度を比べる実験で、数字がどうしても合いません。同じモデル、同じマシンなのに、測るたびに違い、ときには量子化したほうが遅くなります。
原因が最適化レベルであることが多くあります。片方の測定はデフォルト(すべてオン)で動き、もう片方は何らかの理由で最適化がオフのまま動いたのなら、その比較は量子化を測ったのではなく、最適化を測ったことになります。もっとよくある事故は、最適化されたファイルを保存してデプロイすることです。そのファイルは、作った環境の外では開けないことがあります。
4つのレベルは何をするのか
Graph optimizationsのドキュメントは、レベルを4つに分けています。
- ORT_DISABLE_ALL: 何もしません。グラフをそのまま実行します。
- ORT_ENABLE_BASIC: 意味を変えない整理です。定数だけで計算できる部分を先に畳み込み(定数畳み込み)、なくてもよいノードを削除し、よく出てくる形を標準の演算子にまとめます。
- ORT_ENABLE_EXTENDED: 特定の実行プロバイダーに合わせた融合が加わります。ここで、標準ではないドメインのノードが登場します。
- ORT_ENABLE_ALL: ここに、レイアウト変換まで加わります。デフォルトはこのレベルです。
この説明はドキュメントの言葉であり、実際にどこまで行われるかは、モデルとマシンが決めます。そのため、読むよりも取り出して数えるほうが速くて正確です。
取り出して数える方法
SessionOptions.optimized_model_filepathにパスを渡すと、セッションを開きながら書き直したグラフが、そのパスに保存されます。そのファイルをonnx.loadで開いてノードを数えれば、何が起きたかがそのまま見えます。
원본 Add(상수,상수) MatMul Add Relu Identity Mul 노드 6개
ORT_DISABLE_ALL Add MatMul Add Relu Identity Mul 노드 6개
ORT_ENABLE_BASIC Gemm Relu Mul 노드 3개
ORT_ENABLE_EXTENDED FusedGemm Mul 노드 2개
ここで、3つのことが一度に見えます。定数どうしを足していたノードは消え、その結果がinitializerとして収まります。Identityはなくてもよいので削除されます。MatMulのあとのAddはGemm1つにまとめられ、さらに次のレベルでReluまで飲み込んでFusedGemmになります。
ノードが減ったからといって、答えが変わるわけではありません。 同じ入力を4つのレベルに入れて結果を比べれば確認できます。ただし、float32は有効桁が7桁しかないため、融合されたカーネルが乗算の順序を変えると、最後の桁が揺れることがあります。そのため、「正確に同じか」ではなく、許容誤差を決めておいて、その範囲内かを問います。
現場での姿
1つ目は、最適化されたファイルをデプロイすることです。取り出したファイルは小さくてノードも少ないので、「これを送ればよさそうだ」という気になります。しかし、そのファイルには、com.microsoftのような標準ではないドメインのノードが入っています。標準の演算子しか知らないほかのランタイムは、そのファイルを開けません。ORT自身も、保存するときに同じ環境でだけ使うようにという警告を出します。
2つ目は、onnx.checkerがそのファイルを通すことです。チェッカーは、知らないドメインを誰かの拡張と見なして通り過ぎるからです。そのため、「チェッカー通過」を移植可能の根拠にすると、この事故を捕まえられません。根拠は、ノードのドメインのリストです。
3つ目は、比較のベースラインを揃えないことです。量子化の効果を測るには、2つの測定が同じ最適化レベルでなければなりません。ベースラインをORT_DISABLE_ALLで取っておけば、揺れません。そして、実際の運用に出す数字は、運用と同じレベルで別に測る必要があります。
4つ目は、ノード数を性能として読むことです。ノードが減ったからといって、必ず速くなるわけではありません。このMacのようにエミュレーションが入ると、時間は2倍まで揺れます。そのため、このラボは時間を測りません。構造がどう変わったかと答えが同じかだけを判定します。時間は、運用と同じマシンで別に測ります。
実務で本当に大切なこと
- 取り出して数えます。 ドキュメントから推測せず、最適化されたグラフをファイルに取り出してノードを数えます。
- 答えが同じかも一緒に測ります。 許容誤差を決めて、その根拠を書いておきます。
- 移植可能かどうかは、ドメインで判断します。 チェッカーの通過は根拠ではありません。
- デプロイには元のファイルを送ります。 最適化は、受け取る側が自分の環境で行うようにします。
次のラボですること
定数畳み込み・なくてもよいノード・融合される3つのノードを1つのグラフに集めて自分で組み立て、ツールoptlevel.pyをステップごとに育てます。4つのレベルをすべて実行してノードを数え、並べて置き、レベルごとに何が消えて何ができたかを書き、採点ツールが決めたシードで同じ入力を4つのレベルに入れて、答えが同じかどうかを測ります。最後に、取り出したファイルのドメインのリストで、移植可能かどうかを判定します。採点ツールは、毎回異なるシェイプ・重み・シードで自分のグラフを作り、皆さんのツールを実際に動かして答えを照合します。