Hugging Face에서 Stable-Baselines3 사용하기
Hugging Face에서 Stable-Baselines3 사용하기
stable-baselines3는 PyTorch로 구현된 강화학습 알고리즘의 신뢰할 수 있는 구현 모음이에요.
출처: 문서
본문
Hub에서 Stable-Baselines3 탐색하기
모델 페이지 왼쪽에서 필터링해 Stable-Baselines3 모델을 찾을 수 있어요.
Hub의 모든 모델에는 유용한 기능이 있어요:
- 설명, 학습 구성 등이 담긴 자동 생성 모델 카드.
- 발견 가능성을 돕는 메타데이터 태그.
- 다른 모델과 비교할 평가 결과.
- 에이전트의 성능을 보여주는 비디오 위젯.
라이브러리 설치
stable-baselines3를 설치하려면 두 패키지를 설치해야 해요:
stable-baselines3: Stable-Baselines3 라이브러리.huggingface-sb3: Hub에서 Stable-baselines3 모델을 불러오고 업로드하는 추가 코드.
pip install stable-baselines3
pip install huggingface-sb3
기존 모델 사용하기
load_from_hub 함수로 Hub에서 모델을 간단히 다운로드할 수 있어요.
checkpoint = load_from_hub(
repo_id="sb3/demo-hf-CartPole-v1",
filename="ppo-CartPole-v1.zip",
)
두 파라미터를 정의해야 해요:
--repo-id: 다운로드할 Hugging Face 저장소 이름.--filename: 다운로드할 파일.
모델 공유하기
두 함수로 모델을 쉽게 업로드할 수 있어요:
package_to_hub(): 모델을 저장하고, 평가하고, 모델 카드를 생성하고, 에이전트의 리플레이 비디오를 기록한 뒤 저장소 전체를 Hub에 푸시해요.
package_to_hub(model=model,
model_name="ppo-LunarLander-v2",
model_architecture="PPO",
env_id=env_id,
eval_env=eval_env,
repo_id="ThomasSimonini/ppo-LunarLander-v2",
commit_message="Test commit")
일곱 파라미터를 정의해야 해요:
--model: 학습된 모델.--model_architecture: 모델 아키텍처 이름(DQN, PPO, A2C, SAC...).--env_id: 환경 이름.--eval_env: 에이전트 평가에 쓰는 환경.--repo-id: 생성하거나 업데이트할 Hugging Face 저장소 이름.<your huggingface username>/<the repo name>형식이에요.--commit-message.--filename: Hub에 푸시할 파일.
push_to_hub(): 파일 하나를 Hub에 간단히 푸시해요.
push_to_hub(
repo_id="ThomasSimonini/ppo-LunarLander-v2",
filename="ppo-LunarLander-v2.zip",
commit_message="Added LunarLander-v2 model trained with PPO",
)
세 파라미터를 정의해야 해요:
--repo-id: 생성하거나 업데이트할 Hugging Face 저장소 이름.<your huggingface username>/<the repo name>형식이에요.--filename: Hub에 푸시할 파일.--commit-message.