pytorch/rl
A modular, primitive-first, python-first PyTorch library for Reinforcement Learning.
TorchRL – PyTorch 기반 강화학습 툴킷
무엇인가요 – TorchRL은 PyTorch 위에 구축된 라이브러리로, 강화학습(RL) 연구 및 생산을 위한 조합 가능한 구성 요소를 제공합니다. 이는 단일 알고리즘을 제공하지 않습니다. 대신, 통합된 데이터 컨테이너(TensorDict)와 교환 가능한 모듈(환경, 정책, 수집기, 리플레이 버퍼, 손실 객체, 트레이너 등)을 제공하여, 네이티브 PyTorch 프로그래밍 모델에 가깝게 자유롭게 조합할 수 있습니다.
핵심 아이디어
| 아이디어 | TorchRL의 구현 방식 |
|---|---|
| 명시적이고 이름이 지정된 데이터 | 모든 텐서는 필드 이름, 배치 차원, 장치 배치 정보를 유지하는 TensorDict 내부를 이동하며, 훈련 루프 전체에서 일관성 유지. |
| 모듈식 스택 | 환경, 정책, 수집기, 리플레이 버퍼, 손실 모듈은 독립적이고 교체 가능한 구성 요소. |
| 프로토타입에서 생산까지 확장 가능 | 동기 프로세스, 비동기 멀티프로세스 수집기, 분산 훈련, 컴파일(torch.compile) 또는 CUDA 가속 파이프라인에서도 코드 변경 없이 동일한 API 사용 가능. |
주요 구성 요소
- TensorDict – PyTorch 연산, 장치 전송, 공유 메모리 및 메모리 매핑을 완전히 지원하는 사전 유사 컨테이너.
- 환경 및 변환 – 네이티브
PendulumEnv, MuJoCo 작업, Gymnasium, DM-Control, Brax, PettingZoo 등에 대한 래퍼; 변환(관측 정규화, 액션 스케일링, 프레임 스택, HER 등)은 일급 모듈로 제공. - 수집기 – 동기, 비동기, 멀티프로세스 및 분산 수집기로, 트래잭토리를 배치 처리하고 데이터를 적절한 장치로 이동하며 실시간으로 정책 가중치 업데이트 가능.
- 리플레이 버퍼 – 모듈식 저장소(메모리 내, 지연 메모리 매핑, CUDA 인식), 우선순위 샘플링, HER, 오프라인 데이터셋 처리.
- 모듈 / 정책 – 명시적인 입력/출력 키 계약을 가진 일반
nn.Module로 래핑; 확률적 액터, 크리틱, 순환 네트워크, 분포 래퍼, 월드 모델 구성 요소 포함. - 목표 함수 – PPO, SAC, DQN, TD3, REDQ, IQL, CQL, Decision Transformers, DreamerV3, MAPPO/IPPO, QMIX/VDN, 행동 클로닝 등 손실 모듈로, 모두 이름이 지정된 키를 읽고 씁니다.
- 트레이너 및 Hydra 설정 – 환경, 수집기, 손실, 최적화기, 로깅을 재현 가능한 레시피로 연결하는 고수준 유틸리티.
최근 주목할 만한 기능 (v0.13)
- Triton 기반 GRU/LSTM 리셋을 통한 더 빠른 순환 경로.
- 새로운 MuJoCo 환경, 위성 예제, 매크로 제어 프리미티브.
- 다중 에이전트 지원 확장: MAPPO, IPPO,
MultiAgentGAE, 가치 정규화 유틸리티. - 비동기 우선순위 리플레이 버퍼 쓰기, 컴팩트한 관측 저장, 선택적 CUDA 커널 리플레이.
- 추가 변환(
ActionScaling,FlattenAction,NextObservationDelta등) 및 가치 추정기 개선.
일반적인 워크플로우 (빠른 데모)
import torch
from tensordict.nn import TensorDictModule
from torch import nn
from torchrl.envs import PendulumEnv, StepCounter, TransformedEnv
# 간단한 변환 스택이 있는 환경
env = TransformedEnv(PendulumEnv(), StepCounter(max_steps=200))
# 명시적인 TensorDict 키를 가진 일반 nn.Module로 표현된 정책
policy = TensorDictModule(
nn.Sequential(nn.LazyLinear(64), nn.Tanh(), nn.Linear(64, 1), nn.Tanh()),
in_keys=["observation"],
out_keys=["action"],
)
# 1단계 롤아웃 – 결과는 관측, 액션, 보상 등을 포함하는 TensorDict
rollout = env.rollout(max_steps=32, policy=policy)
print(rollout.batch_size) # torch.Size([32])
print(rollout["next", "reward"].shape) # torch.Size([32])
같은 TensorDict는 수집기, 리플레이 버퍼에 저장하거나 손실 모듈에 공급할 수 있으며, 형식 변환 없이 사용 가능.
누구에게 적합한가요?
- 연구자 – 새로운 RL 알고리즘 개발 (유연한 데이터 흐름, 구성 요소 교체 용이, PyTorch의 autograd/compile과의 밀접한 통합 필요).
- 엔지니어 – 많은 CPU/GPU 또는 분산 클러스터로 RL 파이프라인 확장 (비동기 수집기, CUDA 인식 리플레이 버퍼, 메모리 매핑 저장).
- 로보틱스/시뮬레이션 팀 – 네이티브 MuJoCo 또는 사용자 정의 환경과 장치 내부에서 실행되는 변환 필요.
- 다중 에이전트 또는 모델 기반 RL 프로젝트 (VMAS, PettingZoo 래퍼, DreamerV3, Decision-Transformer 구성 요소 제공).
- LLM 후처리 실험 (TorchRL은 GRPO/SFT 스타일의 피니튜닝을 위한 작은 LLM 스택 포함).
설치
# 안정 버전 (CPU 전용 리플레이 버퍼)
pip install torchrl
# CUDA 지원 리플레이 버퍼 (cu118를 사용 중인 CUDA 버전으로 교체)
pip install "torchrl==0.13.0+cu118" \
--extra-index-url https://download.pytorch.org/whl/cu118
# 선택적 확장 (추가 환경/유틸리티)
pip install "torchrl[utils]" # Hydra, 로깅 등
pip install "torchrl[gym_continuous]" # Gymnasium 연속 제어
pip install "torchrl[atari]" # Atari 지원
pip install "torchrl[marl]" # 다중 에이전트 라이브러리 (PettingZoo, VMAS, …)
문서 및 리소스
- 시작하기 및 튜토리얼: https://pytorch.org/rl/stable/index.html#getting-started
- API 참조: https://pytorch.org/rl/stable/reference/index.html
- 최고 성능 구현:
sota-implementations/폴더 (PPO, SAC, TD3 등) - 벤치마크 및 CI 상태: README의 배지가 실시간 대시보드로 연결.
- 커뮤니티: 지원을 위한 Discord 서버, Twitter, GitHub 이슈.
결론: TorchRL은 코드베이스를 깔끔하고 PyTorch 기반으로 유지하면서, 다양한 RL, 다중 에이전트, 모델 기반, 심지어 LLM 피니튜닝 워크플로우의 프로토타이핑, 확장, 실험을 가능하게 하는 기능이 풍부한 PyTorch 중심의 RL 엔지니어링 프레임워크입니다.
관련
- 프로젝트
- 프로젝트
- 프로젝트
- Dispatch
- 프로젝트