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提供专门的覆盖,以确保训练循环热路径中的效率。
相关
- 项目
- 项目
- 项目
- 项目
- 项目