在强化学习的世界里,当我们需要处理图像输入时,传统的状态向量往往力不从心。本文将带你深入 Stable-Baselines3 (SB3),手把手教你如何为自定义环境接入卷积神经网络(CNN)策略,实现从像素到决策的完整AI训练流程。

一、SB3自定义环境接入:两大核心准备

想要将SB3(Stable-Baselines3)与自定义环境无缝对接,核心在于两步走:首先,深入理解你的任务,明确观测空间(Observation Space)动作空间(Action Space)以及奖励函数(Reward Function)的设计;其次,熟悉Gym环境的配置规范,知道需要手写哪些函数、配置哪些参数。

之前我们使用状态向量作为输入,但许多场景下,直接读取图像更为直观。此时,只需给神经网络添加卷积层即可实现。让我们在SB3官方示例中寻找图像输入的参考实现。

在SB3的 Examples 目录下,我们可以找到 AtariGame 的示例代码。其中带有 CnnPolicy 字样的配置即表示采用图像输入模式。

二、理解A2C模型的三种策略模式

要对本地代码进行定制,首先需要理解A2C模型的 Policy 属性。A2C(Advantage Actor-Critic)提供了三种策略选择:

  • MlpPolicy:处理状态向量输入,是我们之前常用的模式。
  • CnnPolicy:专为图像输入设计,本次实战的核心。
  • MultiInputPolicy:适用于多模态输入,如多按键游戏。

因此,在创建模型时,我们选择第二个 CnnPolicy 作为策略基础:

model = A2C("CnnPolicy",
            env,
            verbose=0,
            tensorboard_log='./logs',
            learning_rate=1e-3
           )
model.policy

三、修改观测空间:图像通道的布局艺术

使用图像输入时,observation_space 必须按照官方文档规范进行修改。SB3支持前置通道(Channel First)后置通道(Channel Last)两种布局。

    def __init__(self, arg1, arg2, ...):
        super().__init__()
        # Define action and observation space
        # They must be gym.spaces objects
        # Example when using discrete actions:
        self.action_space = spaces.Discrete(N_DISCRETE_ACTIONS)
        # Example for using image as input (channel-first; channel-last also works):
        self.observation_space = spaces.Box(low=0, high=255,
                                            shape=(N_CHANNELS, HEIGHT, WIDTH), dtype=np.uint8)

根据我们定义的 OBSERVATION_SPACE_VALUES,采用的是后置通道格式。

Example for using image as input (channel-first; channel-last also works):
self.observation_space = spaces.Box(low=0, high=255,
shape=(N_CHANNELS, HEIGHT, WIDTH), dtype=np.uint8)

前置通道 vs 后置通道

  • 前置通道(PyTorch标准):格式为 (C, H, W),即先声明3个颜色通道,再声明256x256的矩阵尺寸。
  • 后置通道(TensorFlow/OpenCV标准):格式为 (H, W, C),先声明256x256矩阵,再声明每个像素点的3个数值。
(C, H, W)(3, 256, 256)(H, W, C)(256, 256, 3)

选择后置通道的原因在于 get_image() 函数中,使用“往 [x][y] 位置画颜色”的逻辑更易理解。若使用前置通道,则需先写颜色(通道数)再写位置,逻辑上不够直观。因此,N_CHANNELS(RGB通道数3)应放在第三参数位置。

    def get_image(self):
        img = np.zeros(self.OBSERVATION_SPACE_VALUES,dtype= np.uint8)
        img[self.food.x][self.food.y] = self.d[self.FOOD_N]
        img[self.player.x][self.player.y] = self.d[self.PLAYER_N]
        img[self.enemy.x][self.enemy.y] = self.d[self.ENEMY_N]
        return img
#self.observation_space = gym.spaces.Box(low=-SIZE+1, high=SIZE-1,
#shape=(4,), dtype=int)
self.observation_space = gym.spaces.Box(low=0, high=255,
shape=(self.SIZE, self.SIZE,self.N_CHANNELS), dtype=np.uint8)
self.OBSERVATION_SPACE_VALUES=(SIZE,SIZE,self.N_CHANNELS)
class envCube(gym.Env):
    # 设定三个部分的颜色分别是蓝、绿、红
    d = {1: (255, 0, 0),  # blue
        2: (0, 255, 0),  # green
        3: (0, 0, 255)}  # red
    PLAYER_N = 1
    FOOD_N = 2
    ENEMY_N = 3
    N_CHANNELS=3

⚠️ 注意:需将图像尺寸 SIZE 调整为40,因为默认模型架构的卷积核较大,10x10的输入会报错。

class envCube(gym.Env):
    # 设定三个部分的颜色分别是蓝、绿、红
    d = {1: (255, 0, 0),  # blue
        2: (0, 255, 0),  # green
        3: (0, 0, 255)}  # red
    PLAYER_N = 1
    FOOD_N = 2
    ENEMY_N = 3
    N_CHANNELS=3
    metadata = {"render_modes": ["human"], "render_fps": 30}
    def __init__(self,SIZE=40,
                 ACTION_SPACE_VALUES = 9,
                 RETURN_IMAGE = True,
                 MAX_STEP=200,
                 FOOD_REWARD = 25,
                 ENEMY_PENALITY = -300,
                 MOVE_PENALITY = -1):
        super(envCube,self).__init__()
        self.SIZE=SIZE
        self.OBSERVATION_SPACE_VALUES=(SIZE,SIZE,self.N_CHANNELS)
        self.ACTION_SPACE_VALUES=ACTION_SPACE_VALUES
        #self.OBSERVATION_SPACE_VALUES=(4,)
        #self.ACTION_SPACE_VALUES=ACTION_SPACE_VALUES
        self.RETURN_IMAGE = RETURN_IMAGE  # 考虑返回值是否图像
        self.MAX_STEP=MAX_STEP
        self.FOOD_REWARD = FOOD_REWARD  # agent获得食物的奖励
        self.ENEMY_PENALITY = ENEMY_PENALITY  # 遇上对手的惩罚
        self.MOVE_PENALITY = MOVE_PENALITY  # 每移动一步的惩罚
        self.action_space = gym.spaces.Discrete(self.ACTION_SPACE_VALUES)
        self.observation_space = gym.spaces.Box(low=0, high=255,
                                            shape=(self.SIZE, self.SIZE,self.N_CHANNELS), dtype=np.uint8)
        #self.observation_space = gym.spaces.Box(low=-SIZE+1, high=SIZE-1,
        #                                    shape=(4,), dtype=int)
    # 环境重置
    def reset(self,seed=None, options=None):
        self.player = Cube(self.SIZE)
        self.food = Cube(self.SIZE)
        self.enemy = Cube(self.SIZE)
        # 如果玩家和食物初始位置相同,重置食物的位置,直到位置不同
        while self.player == self.food:
            self.food = Cube(self.SIZE)
        # 如果敌人和玩家或食物的初始位置相同,重置敌人的位置,直到位置不同
        while self.player == self.enemy or self.food == self.enemy:
            self.enemy = Cube(self.SIZE)
        # 判断观测是图像和数字
        if self.RETURN_IMAGE:
            observation = self.get_image()
        else:
            observation = (self.player - self.food)+(self.player - self.enemy)
            observation=np.array(observation)
        self.episode_step = 0
        info={}
        return observation,info
    def step(self,action):
        self.episode_step+=1
        self.player.action(action)
        self.food.move()
        self.enemy.move()
        # 分类讨论输出new_obs
        if self.RETURN_IMAGE:
            new_observation = self.get_image()
        else:
            new_observation = (self.player - self.food)+(self.player - self.enemy)
            new_observation=np.array(new_observation)
        #获取reward值
        if self.player == self.food:
            reward=self.FOOD_REWARD
        elif self.player==self.enemy:
            reward=self.ENEMY_PENALITY
        else:
            reward=self.MOVE_PENALITY
        #检测截止情况
        terminated = False
        truncated = False
        # 4. 判断结束条件
        if self.player == self.food or self.player == self.enemy:
            terminated = True
        if self.episode_step >= self.MAX_STEP:
            truncated = True
        info={}
        return new_observation,reward,terminated,truncated,info
    def render(self,mode='human'):
        img=self.get_image()
        img = Image.fromarray(img,'RGB')
        img = img.resize((200,200))
        cv2.imshow('Predator',np.array(img))
        if self.player == self.food or self.player == self.enemy or self.episode_step >= self.MAX_STEP:
            cv2.waitKey(1500)
        else:
            cv2.waitKey(1)
    def get_image(self):
        img = np.zeros(self.OBSERVATION_SPACE_VALUES,dtype= np.uint8)
        img[self.food.x][self.food.y] = self.d[self.FOOD_N]
        img[self.player.x][self.player.y] = self.d[self.PLAYER_N]
        img[self.enemy.x][self.enemy.y] = self.d[self.ENEMY_N]
        return img

四、环境验证与A2C模型创建

完成环境配置后,需验证Gym环境的完整性。若输出形状为 (10, 10, 3),则说明环境配置正确。

env=envCube()
print(env.OBSERVATION_SPACE_VALUES)
check_env(env)

接下来,导入主循环依赖库并创建A2C模型:

import gymnasium as gym
from stable_baselines3 import DQN,A2C
from stable_baselines3.common.evaluation import evaluate_policy
import os
import stable_baselines3 as sb3
import torch as th
print(sb3.__version__)
os.environ['KMP_DUPLICATE_LIB_OK']='True'
model = A2C("CnnPolicy",
            env,
            verbose=0,
            tensorboard_log='./logs',
            learning_rate=1e-3
           )
model.policy

通过 model.policy 可以查看默认的模型架构:

该架构首先创建了一个 Sequential 容器,Conv2d 首层输入3个特征,输出32个特征,卷积核大小为8x8,步长为4。后续两层 Conv2d 隐藏层用于提取更高维度的特征。随后进入 Flatten 展平层,将特征块展开为一维向量,最后通过全连接层输出。

值得注意的是,CNN网络与Value网络架构相同,但不共享参数,采用多头输出机制。最终的全连接层输出维度分别对应动作空间(9)和价值函数。

obs
   /            \
 <128>          <128>
  |              |
 <128>          <128>
  |              |
action         value

开始训练,本次保存为 'Custom_CNN_A2C_netDefault_1M',共训练100万步。训练完成后保存权重参数,并加载模型进行测试。

# Train the agent and display a progress bar
model.learn(total_timesteps=int(1e6), progress_bar=True,tb_log_name='Custom_CNN_A2C_netDefault_1M')
model.save("Custom_CNN_A2C_netDefault_1M")
del model #清除model对象
model=A2C.load('Custom_CNN_A2C_netDefault_1M',env=env)
# Evaluate the agent
# NOTE: If you use wrappers with your environment that modify rewards,
#       this will be reflected here. To evaluate with original rewards,
#       wrap environment in a "Monitor" wrapper before other wrappers.
mean_reward, std_reward = evaluate_policy(model, model.get_env(), render=False,n_eval_episodes=10)
print(mean_reward,std_reward)

然而,训练效果并不理想:

40x40的图像对当前模型而言过于吃力,因此考虑回归10x10的图像。但问题随之而来:默认的图像输入要求大于36x36,否则需要自定义特征提取器。

五、模型优化:自定义特征提取器

整个模型有三个可优化方向:观测空间(离散/连续)、全连接层(通过 policy_kwargsnet_arch 项修改)、以及特征提取器

在官方文档中找到自定义特征提取器的示例,我们可参考并修改:

import torch as th
import torch.nn as nn
from gymnasium import spaces
from stable_baselines3 import PPO
from stable_baselines3.common.torch_layers import BaseFeaturesExtractor
class CustomCNN(BaseFeaturesExtractor):
    """
    :param observation_space: (gym.Space)
    :param features_dim: (int) Number of features extracted.
        This corresponds to the number of unit for the last layer.
    """
    def __init__(self, observation_space: spaces.Box, features_dim: int = 256):
        super().__init__(observation_space, features_dim)
        # We assume CxHxW images (channels first)
        # Re-ordering will be done by pre-preprocessing or wrapper
        n_input_channels = observation_space.shape[0]
        self.cnn = nn.Sequential(
            nn.Conv2d(n_input_channels, 32, kernel_size=8, stride=4, padding=0),
            nn.ReLU(),
            nn.Conv2d(32, 64, kernel_size=4, stride=2, padding=0),
            nn.ReLU(),
            nn.Flatten(),
        )
        # Compute shape by doing one forward pass
        with th.no_grad():
            n_flatten = self.cnn(
                th.as_tensor(observation_space.sample()[None]).float()
            ).shape[1]
        self.linear = nn.Sequential(nn.Linear(n_flatten, features_dim), nn.ReLU())
    def forward(self, observations: th.Tensor) -> th.Tensor:
        return self.linear(self.cnn(observations))
policy_kwargs = dict(
    features_extractor_class=CustomCNN,
    features_extractor_kwargs=dict(features_dim=128),
)
model = PPO("CnnPolicy", "BreakoutNoFrameskip-v4", policy_kwargs=policy_kwargs, verbose=1)
model.learn(1000)

1. 观测空间的维度调整

这里要求将通道置于首参,但若不想改变参数顺序,可使用 numpymoveaxis 函数进行轴移动。

在应用中,将待移动对象置于首参:

        # 判断观测是图像和数字
        if self.RETURN_IMAGE:
            observation = self.get_image()
            observation=np.array(observation)
            observation=np.moveaxis(observation,-1,0)
        else:
            observation = (self.player - self.food)+(self.player - self.enemy)
            observation=np.array(observation)
            observation=np.moveaxis(observation,-1,0)

同时,new_observation 也需要同样处理:

        # 分类讨论输出new_obs
        if self.RETURN_IMAGE:
            new_observation = self.get_image()
            new_observation=np.array(new_observation)
            new_observation=np.moveaxis(new_observation,-1,0)
        else:
            new_observation = (self.player - self.food)+(self.player - self.enemy)
            new_observation=np.array(new_observation)
            new_observation=np.moveaxis(new_observation,-1,0)

验证修改是否成功:

env=envCube()
print(env.observation_space.shape)

输出 (3, 10, 10),符合预期。

2. 卷积层参数调整

默认的8x8卷积核对于10x10的输入来说过于庞大,且步长4无填充会直接越界。我们改用3x3卷积核,步长为1。需注意 Conv2d 的首参为输入特征数,次参为输出特征数,层与层之间需保持维度匹配。

self.cnn = nn.Sequential(
            nn.Conv2d(n_input_channels, 32, kernel_size=8, stride=4, padding=0),
            nn.ReLU(),
            nn.Conv2d(32, 64, kernel_size=4, stride=2, padding=0),
            nn.ReLU(),
            nn.Flatten(),
        )
        n_input_channels = observation_space.shape[0]
        self.cnn = nn.Sequential(
            nn.Conv2d(n_input_channels, 32, kernel_size=(3,3), stride=(1,1), padding=0),
            nn.ReLU(),
            nn.Conv2d(32, 64, kernel_size=(3,3), stride=(1,1), padding=0),
            nn.ReLU(),
            nn.Conv2d(64, 64, kernel_size=(3,3), stride=(1,1), padding=0),
            nn.ReLU(),
            nn.Flatten(),
        )

3. 全连接层与激活函数

参考之前学习的 net_arch 修改方法,直接在 policy_kwargs 中调整:

policy_kwargs = dict(
    features_extractor_class=CustomCNN,
    features_extractor_kwargs=dict(features_dim=128),
    net_arch=dict(pi=[32, 32], vf=[64, 64])
)

同时,可以添加激活函数增强非线性表达能力:

policy_kwargs = dict(
    features_extractor_class=CustomCNN,
    features_extractor_kwargs=dict(features_dim=128),
    net_arch=dict(pi=[32, 32], vf=[64, 64]),
    activation_fn=th.nn.ReLU,
)

最终,完整的特征提取继承类代码如下:

import torch as th
import torch.nn as nn
from gymnasium import spaces
from stable_baselines3 import PPO
from stable_baselines3.common.torch_layers import BaseFeaturesExtractor
class CustomCNN(BaseFeaturesExtractor):
    """
    :param observation_space: (gym.Space)
    :param features_dim: (int) Number of features extracted.
        This corresponds to the number of unit for the last layer.
    """
    def __init__(self, observation_space: spaces.Box, features_dim: int = 256):
        super().__init__(observation_space, features_dim)
        # We assume CxHxW images (channels first)
        # Re-ordering will be done by pre-preprocessing or wrapper
        n_input_channels = observation_space.shape[0]
        self.cnn = nn.Sequential(
            nn.Conv2d(n_input_channels, 32, kernel_size=(3,3), stride=(1,1), padding=0),
            nn.ReLU(),
            nn.Conv2d(32, 64, kernel_size=(3,3), stride=(1,1), padding=0),
            nn.ReLU(),
            nn.Conv2d(64, 64, kernel_size=(3,3), stride=(1,1), padding=0),
            nn.ReLU(),
            nn.Flatten(),
        )
        # Compute shape by doing one forward pass
        with th.no_grad():
            n_flatten = self.cnn(
                th.as_tensor(observation_space.sample()[None]).float()
            ).shape[1]
        self.linear = nn.Sequential(nn.Linear(n_flatten, features_dim), nn.ReLU())
    def forward(self, observations: th.Tensor) -> th.Tensor:
        return self.linear(self.cnn(observations))
policy_kwargs = dict(
    features_extractor_class=CustomCNN,
    features_extractor_kwargs=dict(features_dim=128),
    net_arch=dict(pi=[32, 32], vf=[64, 64]),
    activation_fn=th.nn.ReLU,
)

最后两行展示了如何在训练代码中挂载自定义特征提取器,直接将 policy_kwargs 传入即可:

model = PPO("CnnPolicy", "BreakoutNoFrameskip-v4", policy_kwargs=policy_kwargs, verbose=1)
model = A2C("CnnPolicy",
            env,
            verbose=1,
            tensorboard_log='./logs',
            learning_rate=1e-3,
            policy_kwargs=policy_kwargs
           )

通过 model.policy 即可查看构建好的网络结构:

六、Callback机制:自动保存最优模型

由于训练过程存在波动,最终模型不一定是最优的。SB3提供了 Callback 机制,可每K步评估一次并保存最优模型。

参考官方文档,首先导入依赖库:

from stable_baselines3.common.callbacks import EvalCallback

设定 callbacks 参数。其中 eval_freq 决定评估频率,我们将最优参数和日志保存在 "./logs/bestModel/" 目录下:

# Use deterministic actions for evaluation
eval_callback = EvalCallback(env, best_model_save_path="./logs/bestModel/",
                             log_path="./logs/bestModel/", eval_freq=500,
                             deterministic=True, render=False)

启动训练!

# Train the agent and display a progress bar
model.learn(total_timesteps=int(2e6), progress_bar=True,tb_log_name='Custom_CNN_A2C_net128X48_2M')

训练完成后,保存模型并加载最优模型:

model.save("Custom_CNN_A2C_net128X48_2M")
del model
best_model=A2C.load('./logs/bestModel/best_model.zip',env=env)

七、完整代码实战

以下是整个流程的完整代码,涵盖环境定义、模型训练与评估:

1. 导入依赖库

#将学习环境所需要的依赖库导入
import gymnasium as gym
import numpy as np
from gymnasium import spaces
import cv2
from PIL import Image
import time
import pickle #对象到文件的处理
import matplotlib.pyplot as plt
from matplotlib import style
from stable_baselines3.common.env_checker import check_env
from stable_baselines3.common.callbacks import EvalCallback
style.use('ggplot')

2. 环境参数设定

#环境参数设定
#SIZE=10 #3个对象在10*10的范围内
EPISODES=30000 #智能体玩的游戏轮数
SHOW_EVERY=3000 #每3000局展示一次游玩过程
#奖励与惩罚
#FOOD_REWARD=25#吃食物的奖励
#ENEMY_PENALITY=300 #被敌人抓住的惩罚
#MOVE_PENALITY=1 #移动惩罚
#环境计算参数
epsilon=0.6 #在强化学习时抽取随机动作的概率 40%使用最大价值期望动作
EPS_DECAY=0.9998 #每玩一局游戏就让随机动作概率乘以这个数,到最后基本就定性了
DISCOUNT=0.95 #折扣回报 未来奖励的折扣
LEARNING_RATE=0.1 #学习率 步长
#q_table_file = "q_table_save.pkl"
#q_table = None
q_table = "qtable_1775227149.pickle"
#d = {1:(255,0,0), #蓝色——玩家
#     2:(0,255,0), #绿色——食物
#     3:(0,0,255)} #红色——敌人
#PLAYER_N=1
#FOOD_N=2
#ENEMY_N=3

3. 创建Cube对象类

#为三个对象创建类
class Cube:
    def __init__(self,size):#初始位置
        self.size=size
        self.x=np.random.randint(0,self.size-1)
        self.y=np.random.randint(0,self.size-1)
    def __str__(self): #打印当前位置
        return f'{self.x},{self.y}'
    def __sub__(self,other):#这个类的另一个实体
        return (self.x-other.x,self.y-other.y)
    def __eq__(self,other):
        return self.x == other.x and self.y == other.y
    def action(self,choise):
        if choise == 0:
            self.move(x=1,y=1)
        elif choise == 1:
            self.move(x=-1,y=1)
        elif choise == 2:
            self.move(x=1,y=-1)
        elif choise == 3:
            self.move(x=-1,y=-1)
        elif choise == 4:
            self.move(x=0,y=1)
        elif choise == 5:
            self.move(x=0,y=-1)
        elif choise == 6:
            self.move(x=1,y=0)
        elif choise == 7:
            self.move(x=-1,y=0)
        elif choise == 8:
            self.move(x=0,y=0)
    def move(self,x=False,y=False):
        if not x:#如果x没有给值
            self.x += np.random.randint(-1,2)
        else:
            self.x += x
        if not y:#如果y没有给值
            self.y += np.random.randint(-1,2)
        else:
            self.y += y
        #考虑边界
        if self.x<0:
            self.x=0
        elif self.x>=self.size:
            self.x=self.size-1
        if self.y<0:
            self.y=0
        elif self.y>=self.size:
            self.y=self.size-1

4. 创建envCube环境类

class envCube(gym.Env):
    # 设定三个部分的颜色分别是蓝、绿、红
    d = {1: (255, 0, 0),  # blue
        2: (0, 255, 0),  # green
        3: (0, 0, 255)}  # red
    PLAYER_N = 1
    FOOD_N = 2
    ENEMY_N = 3
    N_CHANNELS=3
    metadata = {"render_modes": ["human"], "render_fps": 30}
    def __init__(self,SIZE=10,
                 ACTION_SPACE_VALUES = 9,
                 RETURN_IMAGE = True,
                 MAX_STEP=200,
                 FOOD_REWARD = 25,
                 ENEMY_PENALITY = -300,
                 MOVE_PENALITY = -1):
        super(envCube,self).__init__()
        self.SIZE=SIZE
        self.OBSERVATION_SPACE_VALUES=(SIZE,SIZE,self.N_CHANNELS)
        self.ACTION_SPACE_VALUES=ACTION_SPACE_VALUES
        #self.OBSERVATION_SPACE_VALUES=(4,)
        #self.ACTION_SPACE_VALUES=ACTION_SPACE_VALUES
        self.RETURN_IMAGE = RETURN_IMAGE  # 考虑返回值是否图像
        self.MAX_STEP=MAX_STEP
        self.FOOD_REWARD = FOOD_REWARD  # agent获得食物的奖励
        self.ENEMY_PENALITY = ENEMY_PENALITY  # 遇上对手的惩罚
        self.MOVE_PENALITY = MOVE_PENALITY  # 每移动一步的惩罚
        self.action_space = gym.spaces.Discrete(self.ACTION_SPACE_VALUES)
        self.observation_space = gym.spaces.Box(low=0, high=255,
                                            shape=(self.N_CHANNELS,self.SIZE, self.SIZE), dtype=np.uint8)
        #self.observation_space = gym.spaces.Box(low=-SIZE+1, high=SIZE-1,
        #                                    shape=(4,), dtype=int)
    # 环境重置
    def reset(self,seed=None, options=None):
        self.player = Cube(self.SIZE)
        self.food = Cube(self.SIZE)
        self.enemy = Cube(self.SIZE)
        # 如果玩家和食物初始位置相同,重置食物的位置,直到位置不同
        while self.player == self.food:
            self.food = Cube(self.SIZE)
        # 如果敌人和玩家或食物的初始位置相同,重置敌人的位置,直到位置不同
        while self.player == self.enemy or self.food == self.enemy:
            self.enemy = Cube(self.SIZE)
        # 判断观测是图像和数字
        if self.RETURN_IMAGE:
            observation = self.get_image()
            observation=np.array(observation)
            observation=np.moveaxis(observation,-1,0)
        else:
            observation = (self.player - self.food)+(self.player - self.enemy)
            observation=np.array(observation)
            observation=np.moveaxis(observation,-1,0)
        self.episode_step = 0
        info={}
        return observation,info
    def step(self,action):
        self.episode_step+=1
        self.player.action(action)
        self.food.move()
        self.enemy.move()
        # 分类讨论输出new_obs
        if self.RETURN_IMAGE:
            new_observation = self.get_image()
            new_observation=np.array(new_observation)
            new_observation=np.moveaxis(new_observation,-1,0)
        else:
            new_observation = (self.player - self.food)+(self.player - self.enemy)
            new_observation=np.array(new_observation)
            new_observation=np.moveaxis(new_observation,-1,0)
        #获取reward值
        if self.player == self.food:
            reward=self.FOOD_REWARD
        elif self.player==self.enemy:
            reward=self.ENEMY_PENALITY
        else:
            reward=self.MOVE_PENALITY
        #检测截止情况
        terminated = False
        truncated = False
        # 4. 判断结束条件
        if self.player == self.food or self.player == self.enemy:
            terminated = True
        if self.episode_step >= self.MAX_STEP:
            truncated = True
        info={}
        return new_observation,reward,terminated,truncated,info
    def render(self,mode='human'):
        img=self.get_image()
        img = Image.fromarray(img,'RGB')
        img = img.resize((200,200))
        cv2.imshow('Predator',np.array(img))
        if self.player == self.food or self.player == self.enemy or self.episode_step >= self.MAX_STEP:
            cv2.waitKey(1500)
        else:
            cv2.waitKey(1)
    def get_image(self):
        img = np.zeros(self.OBSERVATION_SPACE_VALUES,dtype= np.uint8)
        img[self.food.x][self.food.y] = self.d[self.FOOD_N]
        img[self.player.x][self.player.y] = self.d[self.PLAYER_N]
        img[self.enemy.x][self.enemy.y] = self.d[self.ENEMY_N]
        return img

5. 实例化环境并检测完整性

env=envCube()
print(env.observation_space.shape)
check_env(env)

6. 导入主函数依赖库与自定义特征提取器

import gymnasium as gym
from stable_baselines3 import DQN,A2C
from stable_baselines3.common.evaluation import evaluate_policy
import os
import stable_baselines3 as sb3
import torch as th
print(sb3.__version__)
os.environ['KMP_DUPLICATE_LIB_OK']='True'
import torch as th
import torch.nn as nn
from gymnasium import spaces
from stable_baselines3 import PPO
from stable_baselines3.common.torch_layers import BaseFeaturesExtractor
class CustomCNN(BaseFeaturesExtractor):
    """
    :param observation_space: (gym.Space)
    :param features_dim: (int) Number of features extracted.
        This corresponds to the number of unit for the last layer.
    """
    def __init__(self, observation_space: spaces.Box, features_dim: int = 256):
        super().__init__(observation_space, features_dim)
        # We assume CxHxW images (channels first)
        # Re-ordering will be done by pre-preprocessing or wrapper
        n_input_channels = observation_space.shape[0]
        self.cnn = nn.Sequential(
            nn.Conv2d(n_input_channels, 32, kernel_size=(3,3), stride=(1,1), padding=0),
            nn.ReLU(),
            nn.Conv2d(32, 64, kernel_size=(3,3), stride=(1,1), padding=0),
            nn.ReLU(),
            nn.Conv2d(64, 64, kernel_size=(3,3), stride=(1,1), padding=0),
            nn.ReLU(),
            nn.Flatten(),
        )
        # Compute shape by doing one forward pass
        with th.no_grad():
            n_flatten = self.cnn(
                th.as_tensor(observation_space.sample()[None]).float()
            ).shape[1]
        self.linear = nn.Sequential(nn.Linear(n_flatten, features_dim), nn.ReLU())
    def forward(self, observations: th.Tensor) -> th.Tensor:
        return self.linear(self.cnn(observations))
policy_kwargs = dict(
    features_extractor_class=CustomCNN,
    features_extractor_kwargs=dict(features_dim=128),
    net_arch=dict(pi=[32, 32], vf=[64, 64]),
    activation_fn=th.nn.ReLU,
)

7. 创建模型实例并使用Callbacks保存最佳权重

model = A2C("CnnPolicy",
            env,
            verbose=1,
            tensorboard_log='./logs',
            learning_rate=1e-3,
            policy_kwargs=policy_kwargs
           )
model.policy
# Use deterministic actions for evaluation
eval_callback = EvalCallback(env, best_model_save_path="./logs/bestModel/",
                             log_path="./logs/bestModel/", eval_freq=500,
                             deterministic=True, render=False)

8. 训练、保存与加载模型

# Train the agent and display a progress bar
model.learn(total_timesteps=int(2e6), progress_bar=True,tb_log_name='Custom_CNN_A2C_net128X48_2M')
model.save("Custom_CNN_A2C_net128X48_2M")
del model
model=A2C.load('Custom_CNN_A2C_net128X48_2M',env=env)
best_model=A2C.load('./logs/bestModel/best_model.zip',env=env)
# Evaluate the agent
# NOTE: If you use wrappers with your environment that modify rewards,
#       this will be reflected here. To evaluate with original rewards,
#       wrap environment in a "Monitor" wrapper before other wrappers.
mean_reward, std_reward = evaluate_policy(model, model.get_env(), render=False,n_eval_episodes=10)
print(mean_reward,std_reward)

总结

自定义特征提取器是处理图像输入的关键,核心要点如下:

  • ✅ 确认图像输入是否为 Channel First,若非则使用 np.moveaxis 调整。
  • ✅ 卷积层修改在 __init__() 中完成。
  • ✅ 全连接层修改通过 policy_kwargs 中的 net_arch 实现。
  • ✅ 创建模型时传入 policy_kwargs=policy_kwargs 即可生效。

掌握这些技巧,你就能灵活运用SB3处理各类图像输入任务,让AI在像素世界中游刃有余!

[AFFILIATE_SLOT_1]

如果你觉得本文对你有帮助,欢迎关注更多强化学习实战教程!

[AFFILIATE_SLOT_2]