Stable Baselines3
Visão Geral
Stable Baselines3 (SB3) é uma biblioteca baseada em PyTorch que fornece implementações confiáveis de algoritmos de aprendizado por reforço. Esta skill oferece orientação abrangente para treinar agentes de RL, criar ambientes customizados, implementar callbacks e otimizar workflows de treinamento usando a API unificada do SB3.
Capacidades Principais
1. Treinamento de Agentes de RL
Padrão de Treinamento Básico:
import gymnasium as gym
from stable_baselines3 import PPO
# Create environment
env = gym.make("CartPole-v1")
# Initialize agent
model = PPO("MlpPolicy", env, verbose=1)
# Train the agent
model.learn(total_timesteps=10000)
# Save the model
model.save("ppo_cartpole")
# Load the model (without prior instantiation)
model = PPO.load("ppo_cartpole", env=env)
Observações Importantes:
total_timestepsé um limite inferior; o treinamento real pode exceder isso devido à coleta de batches- Use
model.load()como um método estático, não em uma instância existente - O buffer de replay NÃO é salvo com o modelo para economizar espaço
Seleção de Algoritmo:
Use references/algorithms.md para orientação detalhada sobre características de algoritmos e seleção. Referência rápida:
- PPO/A2C: Uso geral, suporta todos os tipos de espaço de ação, bom para multiprocessamento
- SAC/TD3: Controle contínuo, off-policy, eficiente em amostragem
- DQN: Ações discretas, off-policy
- HER: Tarefas com objetivos condicionados
Veja scripts/train_rl_agent.py para um template completo de treinamento com melhores práticas.
2. Ambientes Customizados
Requisitos:
Ambientes customizados devem herdar de gymnasium.Env e implementar:
__init__(): Definir action_space e observation_spacereset(seed, options): Retornar observação inicial e dicionário infostep(action): Retornar observation, reward, terminated, truncated, inforender(): Visualização (opcional)close(): Limpeza de recursos
Restrições-Chave:
- Observações de imagem devem ser
np.uint8no intervalo [0, 255] - Use formato channel-first quando possível (channels, height, width)
- SB3 normaliza imagens automaticamente dividindo por 255
- Defina
normalize_images=Falseem policy_kwargs se pré-normalizado - SB3 NÃO suporta espaços
DiscreteouMultiDiscretecomstart!=0
Validação:
from stable_baselines3.common.env_checker import check_env
check_env(env, warn=True)
Veja scripts/custom_env_template.py para um template completo de ambiente customizado e references/custom_environments.md para orientação abrangente.
3. Ambientes Vetorizados
Propósito: Ambientes vetorizados executam múltiplas instâncias do ambiente em paralelo, acelerando o treinamento e habilitando certos wrappers (frame-stacking, normalização).
Tipos:
- DummyVecEnv: Execução sequencial no processo atual (para ambientes leves)
- SubprocVecEnv: Execução paralela entre processos (para ambientes computacionalmente pesados)
Configuração Rápida:
from stable_baselines3.common.env_util import make_vec_env
# Create 4 parallel environments
env = make_vec_env("CartPole-v1", n_envs=4, vec_env_cls=SubprocVecEnv)
model = PPO("MlpPolicy", env, verbose=1)
model.learn(total_timesteps=25000)
Otimização Off-Policy:
Ao usar múltiplos ambientes com algoritmos off-policy (SAC, TD3, DQN), defina gradient_steps=-1 para executar uma atualização de gradiente por passo de ambiente, equilibrando tempo de relógio e eficiência de amostragem.
Diferenças de API:
reset()retorna apenas observações (info disponível emvec_env.reset_infos)step()retorna tupla de 4 elementos:(obs, rewards, dones, infos)não 5-tupla- Ambientes auto-reiniciam após episódios
- Observações terminais disponíveis via
infos[env_idx]["terminal_observation"]
Veja references/vectorized_envs.md para informações detalhadas sobre wrappers e uso avançado.
4. Callbacks para Monitoramento e Controle
Propósito: Callbacks habilitam monitoramento de métricas, salvamento de checkpoints, implementação de parada antecipada e lógica de treinamento customizada sem modificar algoritmos principais.
Callbacks Comuns:
- EvalCallback: Avaliar periodicamente e salvar melhor modelo
- CheckpointCallback: Salvar checkpoints de modelo em intervalos
- StopTrainingOnRewardThreshold: Parar quando recompensa alvo é alcançada
- ProgressBarCallback: Exibir progresso de treinamento com temporizações
Estrutura de Callback Customizado:
from stable_baselines3.common.callbacks import BaseCallback
class CustomCallback(BaseCallback):
def _on_training_start(self):
# Called before first rollout
pass
def _on_step(self):
# Called after each environment step
# Return False to stop training
return True
def _on_rollout_end(self):
# Called at end of rollout
pass
Atributos Disponíveis:
self.model: A instância do algoritmo de RLself.num_timesteps: Total de passos do ambienteself.training_env: O ambiente de treinamento
Encadeamento de Callbacks:
from stable_baselines3.common.callbacks import CallbackList
callback = CallbackList([eval_callback, checkpoint_callback, custom_callback])
model.learn(total_timesteps=10000, callback=callback)
Veja references/callbacks.md para documentação abrangente de callbacks.
5. Persistência e Inspeção de Modelos
Salvamento e Carregamento:
# Save model
model.save("model_name")
# Save normalization statistics (if using VecNormalize)
vec_env.save("vec_normalize.pkl")
# Load model
model = PPO.load("model_name", env=env)
# Load normalization statistics
vec_env = VecNormalize.load("vec_normalize.pkl", vec_env)
Acesso a Parâmetros:
# Get parameters
params = model.get_parameters()
# Set parameters
model.set_parameters(params)
# Access PyTorch state dict
state_dict = model.policy.state_dict()
6. Avaliação e Gravação
Avaliação:
from stable_baselines3.common.evaluation import evaluate_policy
mean_reward, std_reward = evaluate_policy(
model,
env,
n_eval_episodes=10,
deterministic=True
)
Gravação de Vídeo:
from stable_baselines3.common.vec_env import VecVideoRecorder
# Wrap environment with video recorder
env = VecVideoRecorder(
env,
"videos/",
record_video_trigger=lambda x: x % 2000 == 0,
video_length=200
)
Veja scripts/evaluate_agent.py para um template completo de avaliação e gravação.
7. Recursos Avançados
Agendamentos de Taxa de Aprendizado:
def linear_schedule(initial_value):
def func(progress_remaining):
# progress_remaining goes from 1 to 0
return progress_remaining * initial_value
return func
model = PPO("MlpPolicy", env, learning_rate=linear_schedule(0.001))
Políticas Multi-Input (Observações Dict):
model = PPO("MultiInputPolicy", env, verbose=1)
Use quando observações são dicionários (p.ex., combinando imagens com dados de sensores).
Hindsight Experience Replay:
from stable_baselines3 import SAC, HerReplayBuffer
model = SAC(
"MultiInputPolicy",
env,
replay_buffer_class=HerReplayBuffer,
replay_buffer_kwargs=dict(
n_sampled_goal=4,
goal_selection_strategy="future",
),
)
Integração com TensorBoard:
model = PPO("MlpPolicy", env, tensorboard_log="./tensorboard/")
model.learn(total_timesteps=10000)
Orientação de Workflow
Iniciando um Novo Projeto de RL:
- Defina o problema: Identifique espaço de observação, espaço de ação e estrutura de recompensa
- Escolha algoritmo: Use
references/algorithms.mdpara orientação de seleção - Crie/adapte ambiente: Use
scripts/custom_env_template.pyse necessário - Valide ambiente: Sempre execute
check_env()antes de treinar - Configure treinamento: Use
scripts/train_rl_agent.pycomo template inicial - Adicione monitoramento: Implemente callbacks para avaliação e checkpointing
- Otimize desempenho: Considere ambientes vetorizados para velocidade
- Avalie e itere: Use
scripts/evaluate_agent.pypara avaliação
Problemas Comuns:
- Erros de memória: Reduza
buffer_sizepara algoritmos off-policy ou use menos ambientes em paralelo - Treinamento lento: Considere SubprocVecEnv para ambientes em paralelo
- Treinamento instável: Tente algoritmos diferentes, ajuste hiperparâmetros ou verifique escalonamento de recompensa
- Erros de importação: Garanta que
stable_baselines3está instalado:uv pip install stable-baselines3[extra]
Recursos
scripts/
train_rl_agent.py: Template de script de treinamento completo com melhores práticasevaluate_agent.py: Template de avaliação de agente e gravação de vídeocustom_env_template.py: Template de ambiente Gym customizado
references/
algorithms.md: Comparação detalhada de algoritmos e guia de seleçãocustom_environments.md: Guia abrangente para criação de ambiente customizadocallbacks.md: Referência completa do sistema de callbacksvectorized_envs.md: Uso de ambientes vetorizados e wrappers
Instalação
# Basic installation
uv pip install stable-baselines3
# With extra dependencies (Tensorboard, etc.)
uv pip install stable-baselines3[extra]