Hugging Face에서 Stable-Baselines3 사용하기

Hugging Face에서 Stable-Baselines3 사용하기

stable-baselines3는 PyTorch로 구현된 강화학습 알고리즘의 신뢰할 수 있는 구현 모음이에요.

출처: 문서

본문

Hub에서 Stable-Baselines3 탐색하기

모델 페이지 왼쪽에서 필터링해 Stable-Baselines3 모델을 찾을 수 있어요.

Hub의 모든 모델에는 유용한 기능이 있어요:

  1. 설명, 학습 구성 등이 담긴 자동 생성 모델 카드.
  2. 발견 가능성을 돕는 메타데이터 태그.
  3. 다른 모델과 비교할 평가 결과.
  4. 에이전트의 성능을 보여주는 비디오 위젯.

라이브러리 설치

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: 다운로드할 파일.

모델 공유하기

두 함수로 모델을 쉽게 업로드할 수 있어요:

  1. 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에 푸시할 파일.
  1. 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.

더 알아보기 (Learn more)

  • Hugging Face Stable-Baselines3 문서를 참고하세요.
  • Stable-Baselines3 문서를 참고하세요.