Hugging FaceのTensorFlow哲学
Hugging FaceはTensorFlowに対して「Keras優先」の哲学を採用し、Kerasを回避すべき障壁ではなく、主要なハイレベルAPIとして扱います。このアプローチにより、トランスフォーマーモデルはfit()、compile()、predict()を含む標準的なKerasワークフローと完全に互換性を持ち、XLAを活用してJAXやPyTorchと同等のパフォーマンスを実現します。
Kerasとの深い統合
transformersライブラリのすべてのTensorFlowモデルとレイヤーはKeras ModelおよびLayerオブジェクトとして実装されています。この設計により、ユーザーは低レベルのトレーニングループを書くことなく、標準的なKerasメソッドを用いてトレーニングや推論を行うことができます。
モデル構成と柔軟性
Kerasのサブクラス化によりハイブリッドモデルの作成が可能になります。ユーザーは複数の事前学習済みモデル(例:言語モデルとビジョントランスフォーマーの統合)を単一のKerasモデルに結合できます。これにより、ハイレベルAPIの利点を保ちつつ、複雑なアーキテクチャの開発が可能になります。
自動ロス関数
トレーニングプロセスを簡素化するため、Hugging Faceはベースモデルと出力タイプに合わせたデフォルトのロス関数を提供します。ユーザーがcompile()をロス引数なしで呼び出すと、ライブラリはパディングとマスキングを正しく処理するロス関数(例:BERTのマスク付き言語モデリングロス)を自動的に提供します。ユーザーはcompile()でカスタムロスを指定するか、サブクラス化したモデルで独自のtrain_step()を実装することでこれを上書きできます。
標準化されたラベル処理
ラベルは、入力辞書に含めるのではなく、標準的なKerasの慣例(別個の引数として、または(inputs, labels)タプルの一部として)で渡されるようになりました。この変更により、標準的なKerasメトリクスとの互換性が確保され、ユーザーの混乱が減少します。
データパイプラインの最適化
トークナイズされたデータセット全体をRAMにロードする際のメモリオーバーヘッドを回避するため、Hugging Faceはdatasetsライブラリとtf.dataを統合しています。
prepare_tf_dataset()による効率的なストリーミング
小規模なデータセットはNumPy配列に変換できますが、大規模なデータセットはprepare_tf_dataset()メソッドの恩恵を受けます。このメソッドはデータセットをtf.data.Datasetオブジェクトでラップし、以下を可能にします:
- オンザフライロード:データはメモリにロードされるのではなく、ディスクからストリーミングされます。
- 動的パディング:パディングはデータセット全体ではなくバッチ単位で適用され、パディングトークン数が減少し、トレーニング速度が向上します。
- 自動フィルタリング:モデルは特定のアーキテクチャに対して有効な入力名でないデータセット列を自動的に除外します。
パフォーマンスとデプロイメント
XLAによるアクセラレーション
Hugging FaceはXLA(Accelerated Linear Algebra)を利用しています。XLAはTensorFlowとJAXが共有するJITコンパイラで、線形代数コードを最適化し、実行速度の向上とメモリ使用量の削減を実現します。
主なパフォーマンス向上点は以下の通りです:
- 生成速度:XLAを使用した更新された
generate()コードにより、テキスト生成速度がPyTorchより速く、JAXに匹敵する速度となりました。 - トレーニング速度:TFモデルは言語モデルのトレーニングなどのタスクでJAXに近い速度を達成しています。
XLAの制限の一つは静的な入力形状が必要なことです。シーケンス長が可変の場合、再コンパイルが頻繁に発生し、パフォーマンス向上が相殺される可能性があります。
エンドツーエンドデプロイメント
TF ServingやTFXを通じたデプロイを簡素化するため、Hugging Faceはトークナイゼーションをモデルアーティファクトに直接埋め込む作業を進めています。これにより、推論時に外部トークナイザーライブラリへの依存がなくなります。BERTなどの一般的なモデルでは、トークナイザーとモデルを単一のKeras ModelにラップしてEndToEndModelを作成でき、モデルが生の文字列を入力として受け取れるようになります。
コミュニティとモデル共有
push_to_hub()を使用してモデルをHugging Face Hubにアップロードすると、モデルページと自動生成されたモデルカードが作成されます。これにより、ファインチューニングされたモデルも基盤モデルと同じAPIで扱えるようになり、共有アーティファクトと実践のオープンなエコシステムが促進されます。