pytorch/tensordict

TensorDict is a pytorch dedicated tensor container.

해결하는 문제

TensorDict는 단일 단위로 취급해야 하는 복잡하고 중첩된 텐서 컬렉션을 관리하는 문제를 해결합니다. 표준 PyTorch에서는 텐서 딕셔너리를 관리할 때 슬라이싱, 디바이스 전송 (.to(device)), 산술 연산과 같은 작업을 모든 개별 텐서에 대해 수동으로 반복해야 합니다. TensorDict는 이러한 작업을 전체 구조에 한 번에 적용할 수 있는 통합 컨테이너를 제공하여 모든 텐서(leaf)가 일관된 배치 크기를 공유하도록 보장합니다.

작동 방식

PyTorch 텐서처럼 동작하는 배치 처리된 중첩 dict[str, Tensor]로 작동합니다. TensorDict에서 작업이 수행되면 모든 하위 텐서로 자동 디스패치됩니다. 예를 들어, TensorDict를 슬라이싱하면 동일한 구조를 가지면서 모든 leaf가 슬라이싱된 새로운 TensorDict가 반환됩니다. 고처리량 산술 연산을 위해 PyTorch foreach 커널을 사용하고, 효율적인 디바이스 이동을 위해 비차단(non-blocking) 전송을 사용합니다.

대상 사용자

데이터 단위가 단일 텐서가 아닌 구조화된 배치인 고성능 PyTorch 워크로드를 구축하는 개발자를 위해 설계되었습니다. 이는 강화 학습(RL) 궤적, LLM 사후 학습 샘플, 로보틱스 궤적 및 과학적 ML 파이프라인에서 특히 일반적입니다.

주요 특징

  • 텐서와 유사한 작업: 전체 중첩 구조에 걸친 인덱싱, 슬라이싱, 리셰이프, 스택 및 결합을 지원합니다.
  • 캐스팅 및 디바이스 관리: 모든 leaf에 대해 단일 호출로 디바이스 및 dtype 캐스팅이 가능합니다.
  • 고성능 메모리: 대규모 데이터 세트를 위한 메모리 매핑 (memmap), 레이지 스택 및 메모리 오버헤드를 줄이기 위한 사전 할당을 포함합니다.
  • 함수형 프로그래밍: 구조화된 텐서 객체를 위한 @tensorclass 및 모델 파라미터 관리를 위한 torch.vmapto_module과의 호환성을 지원합니다.
  • 컴파일 인식: 훈련 루프의 핫 패스에서 효율성을 보장하기 위해 torch.compile에 대한 전용 커버리지를 제공합니다.

관련

  • 프로젝트
  • 프로젝트
  • 프로젝트
  • 프로젝트
  • 프로젝트