TT Lab
はじめる
学ぶ 学習パス コース

AIダイエット失敗事件

32件を入れたら拒否された — 死んだバッチ軸を生き返らせる

TT Labで続きを見る

目標

同じ重みで、バッチ軸がシンボルの/root/onnxq-shape/dyn.onnxと、1で固定された/root/onnxq-shape/fixed.onnxを作り、軸を読み取って実際に入れてみるツール/root/onnxq-shape/axes.pyを作ります。最後に、定数のシェイプが埋め込まれてバッチ軸が死んだモデルを見つけ出し、生き返らせます。

なぜ重要なのか

ONNXの軸は、整数で固定されているか、シンボル名で開いているか、そもそも指定されていないかのどれかです。3つはまったく別の意味ですが、ファイルをざっと見ただけでは区別がつきません。バッチを大きくしてスループットを測る実験が最初の1行で止まる原因は、たいていここにあります。 拒否がセッションを開くときではなく、値を入れるときに起きるという点も重要です。セッションが開いたからモデルの問題ではないと結論づけると、見当違いの場所を掘ることになります。 同じシンボル名が2つの入力に使われている場合、それは「両方とも動的」ではなく、「2つの行数が互いに同じでなければならない」という制約です。破ると拒否されますが、ランタイムはシンボル名を指摘せず、最適化を経たあとのノード名を挙げます。原因と症状の間に1層挟まっています。 最も静かな事故は、Reshapeに定数のシェイプが埋め込まれている場合です。入力はシンボルで開いているのに、その箇所からバッチ軸が整数に変わります。エラーも警告もなく、シェイプ推論はその定数をそのまま信じて、辻褄の合った答えを出します。 採点ツールは、皆さんが書き出した文言を信じません。一時ディレクトリに採点ツールが自分で作ったファイルを用意し、皆さんのツールを実際に実行して、同じファイルを採点ツールが読んで動かして得た答えと照合します。軸名とシェイプ、埋め込む行数は、実行のたびに変わります。

ステップ

  1. /root/onnxq-shape/build_shapes.pyを作成して実行し、2つのファイルを作ってください(出力先: /root/onnxq-shape/dyn.onnx、/root/onnxq-shape/fixed.onnx)。
  2. /root/onnxq-shape/axes.pyにdimsを作り、入出力の軸が整数かシンボルかを読み取らせてください。
  3. feedを追加し、指定された行数で値を作って実際に入れてみて、結果や拒否を書かせてください。
  4. inferを追加し、onnx.shape_inferenceが埋めたものと埋められなかったものを分けさせてください。
  5. symbolsを追加し、同じシンボル名がどこに使われているかを集めさせてください。
  6. scanを追加し、定数のシェイプが埋め込まれたReshapeを見つけさせてください。
  7. repairを追加し、埋め込まれたシェイプのバッチの位置を元に戻して、出力軸にシンボル名を付け直させてください。
  8. レポートを作成してください(出力先: /root/onnxq-shape/shape_report.json、/root/onnxq-shape/shape_report.md)。

参考

軸の宣言だけが違う2つのバージョンを作る

/root/onnxq-shape/build_shapes.pyを作成して実行し、2つのファイルを作ってください(出力先: /root/onnxq-shape/dyn.onnx、/root/onnxq-shape/fixed.onnx)。重みは2つのファイルで完全に同じで、軸0の宣言だけが違います。

make_tensor_value_infoのシェイプのリストで、位置0に文字列を入れるとシンボル、整数を入れるとその値で固定されます。重みは一度だけ作り、2つのグラフに同じオブジェクトを入れてください。値が違うと、あとで2つのモデルを見比べられません。

軸が開いているか固定されているかを読み取る

/root/onnxq-shape/axes.pyにdims <모델>(プレースホルダーはモデルです)を作り、ランタイムの入力と出力の軸を読み取らせてください。整数なら整数で、シンボルなら文字列で、なければnullで書きます。

軸1つは、dim_param(シンボル)かdim_value(整数)のどちらかを持つか、両方とも持たないかです。3つを区別して書いてください。そして、initializerで埋められる名前はランタイムの入力ではないので、inputsから除いてください。

実際に入れてみる

feed <모델> <행수...>(プレースホルダーはモデルと行数です)を追加し、ランタイムの入力ごとにその行数で値を作って入れてみて、{"status", "error_type", "message", "shapes"}を出力させてください。拒否されても、終了コードは0です。

後ろの軸は宣言された整数をそのまま使い、軸0だけを指定された行数に変えます。拒否は失敗ではなく答えです。例外を捕まえてクラス名と最初の1行を書いておけば、あとでどの軸で何を期待していたのかを、その文言が教えてくれます。セッションを開く段階と、値を入れる段階を、別々に囲んでください。2つは、別の箇所で失敗します。

推論が埋めたものと埋められなかったもの

infer <모델>(プレースホルダーはモデルです)を追加し、onnx.shape_inferenceが埋めた中間テンソルと、埋められなかった名前を分けさせてください。応答は{"value_info", "unknown"}です。

infer_shapesは新しいModelProtoを返し、埋めた結果はgraph.value_infoに入ります。unknownは、ノードが作り出すもののグラフの出力ではない名前のうち、埋められなかったものです。中間テンソルがそもそもないグラフでは、value_infoは空のまま返ってきます。それは失敗ではなく、埋めるものがなかったという意味です。

同じ名前は同じ値である

symbols <모델>(プレースホルダーはモデルです)を追加し、シンボル名ごとに、それが使われている位置を"텐서이름:축번호"の形式で集めさせてください(プレースホルダーは、テンソル名と軸番号です)。そして、同じシンボルを使う2つの入力に異なる行数を入れて、feedで拒否を確認してください。

2つの入力の軸0に同じ名前が書かれていれば、ランタイムにとって2つの行数が同じでなければならないという意味です。破ると拒否されますが、ランタイムはシンボル名を指摘せず、最適化を経たあとのノード名を挙げます。そのため、この表を事前に作っておけば、そのエラーを解釈できます。

定数のシェイプが埋め込まれたReshapeを見つける

scan <모델>(プレースホルダーはモデルです)を追加し、ターゲットシェイプがinitializerで、その位置0が正の数であるReshapeを見つけさせてください。応答は{"frozen_reshape", "shapes"}です。

位置0が-1なら、残りの軸から計算されるので、軸は生きています。0は、入力のその軸をそのまま使うという意味なので、やはり生きています。正の数だけを選ぶ必要があります。ターゲットシェイプがinitializerではなく、ほかのノードの出力なら、実行時に決まるものなので、この検査の対象ではありません。

死んだバッチ軸を生き返らせる

repair <모델> <출력>(プレースホルダーはモデルと出力です)を追加し、埋め込まれた定数のシェイプの位置0を-1に書き換え、グラフの出力の軸0を、最初のランタイム入力のシンボル名で宣言し直して、新しいファイルに保存させてください。

initializerをその場で書き換えるには、新しいテンソルを作ってCopyFromで上書きすればよいです。出力軸は、dim_valueを消してdim_paramに名前を入れます。重みとノードには手を触れないでください。直したファイルが、元々入っていた行数で同じ値を出してこそ、直したと言えます。

2つのバージョンを並べて報告する

/root/onnxq-shape/shape_report.jsonにdyn_input・fixed_input・trialsを書き、## 어떤 축이 열려 있나、## 굳은 축에 무엇을 넣었나、## 추론이 못 채운 곳、## 보내는 쪽에 요청할 것の4つのセクションでMarkdownのレポートを書いてください(出力先: /root/onnxq-shape/shape_report.md。韓国語の見出しは、順に「どの軸が開いているか」「固定された軸に何を入れたか」「推論が埋められなかった箇所」「送る側に依頼すること」という意味です)。

trialsは、どのモデルに何行を入れて何が出たかの記録です。採点ツールが同じ組み合わせを自分で入れて照合するので、実際に測って書いてください。レポートの本文には、バッチ軸のシンボル名を書く必要があります。そうすれば、受け取る側がその名前で話せます。