pytorch/tensordict

TensorDict is a pytorch dedicated tensor container.

解決する課題

TensorDictは、単一のユニットとして扱う必要がある、複雑でネストされたテンソルコレクションの管理問題を解決します。標準的なPyTorchでは、テンソルの辞書を管理する場合、スライス、デバイス転送 (.to(device))、算術演算などの操作を個々のテンソルに対して手動で繰り返す必要があります。TensorDictは、これらの操作を構造全体に一度に適用できる統一されたコンテナを提供し、すべてのテンソル(リーフ)が一貫したバッチサイズを共有するようにします。

仕組み

PyTorchテンソルのように動作する、バッチ化されたネストされた dict[str, Tensor] として機能します。TensorDictに対して操作が行われると、その操作は自動的にすべての基底テンソルにディスパッチされます。例えば、TensorDictをスライスすると、同じ構造を持ちながらすべてのリーフがスライスされた新しいTensorDictが返されます。高スループットの算術演算にはPyTorchのforeachカーネルを、効率的なデバイス移動には非ブロッキング転送を使用します。

対象ユーザー

データの単位が単一のテンソルではなく、構造化されたバッチであるような、高性能なPyTorchワークロードを構築する開発者向けに設計されています。これは、強化学習 (RL) の軌跡、LLMのポストトレーニングサンプル、ロボティクスの軌跡、および科学的MLパイプラインで特に一般的です。

ハイライト

  • テンソルライクな操作: ネストされた構造全体にわたるインデックス作成、スライス、リシェイプ、スタック、結合をサポート。
  • キャストとデバイス管理: すべてのリーフに対して、単一の呼び出しでデバイスおよびdtypeのキャストが可能。
  • 高性能メモリ: 大規模データセット用のメモリマッピング (memmap)、レイジースタック、およびメモリオーバーヘッドを削減するための事前割り当てを含む。
  • 関数型プログラミング: 構造化されたテンソルオブジェクトのための @tensorclass、およびモデルパラメータを管理するための torch.vmap および to_module との互換性をサポート。
  • コンパイル対応: トレーニングループのホットパスにおける効率を確保するため、torch.compile に特化したカバレッジを提供。

関連

  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト