pytorch/tensordict
TensorDict is a pytorch dedicated tensor container.
解決的問題
TensorDict 解決了需要將複雜的嵌套張量集合視為單一單元進行管理的難題。在標準的 PyTorch 中,管理張量字典需要手動對每個個別的張量重複執行切片、裝置轉移 (.to(device)) 和算術運算等操作。TensorDict 提供了一個統一的容器,允許一次性對整個結構應用這些操作,確保所有張量(葉子節點)共享一致的批次大小。
工作原理
它作為一個批次化、嵌套的 dict[str, Tensor] 運作,其行為類似於 PyTorch 張量。當對 TensorDict 執行操作時,該操作會自動分發到所有底層張量。例如,對 TensorDict 進行切片會回傳一個新的 TensorDict,它具有相同的結構,但所有葉子節點都經過了切片。它使用 PyTorch foreach 核心進行高吞吐量算術運算,並使用非阻塞傳輸實現高效的裝置移動。
適用對象
專為構建高性能 PyTorch 工作負載的開發人員設計,在這些工作負載中,數據單元是結構化批次而非單個張量。這在強化學習 (RL) 軌跡、LLM 後訓練樣本、機器人軌跡和科學機器學習流水線中尤為常見。
亮點
- 類張量操作:支援整個嵌套結構的索引、切片、重塑、堆疊和拼接。
- 類型轉換與裝置管理:透過單次呼叫即可對所有葉子節點進行裝置和資料類型 (dtype) 轉換。
- 高性能記憶體:包括用於大型數據集的記憶體映射 (
memmap)、延遲堆疊和用於減少記憶體開銷的預分配。 - 函數式編程:支援用於結構化張量物件的
@tensorclass,以及用於管理模型參數的torch.vmap和to_module的相容性。 - 編譯感知:針對
torch.compile提供專門的覆蓋,以確保訓練迴圈熱路徑中的效率。
相關
- 專案
- 專案
- 專案
- 專案
- 專案