使用Stable Baselines时,如何避免智能体重复选择同一动作组的动作?
解决stable-baselines中禁止重复选择动作组的问题
针对你遇到的问题——用stable-baselines时无法阻止智能体重复选同一动作组,且用-inf奖励导致死循环的情况,以下是几个实用的解决方案:
方法1:重写策略的动作选择逻辑(推荐)
stable-baselines的算法支持自定义策略,你可以继承对应算法的Policy类,重写predict方法,在动作选择前主动屏蔽已选组对应的所有动作,从根源上避免智能体选无效动作。以DQN为例:
from stable_baselines3 import DQN from stable_baselines3.dqn.policies import DQNPolicy import torch class CustomDQNPolicy(DQNPolicy): def predict(self, observation, state=None, episode_start=None, deterministic=False): # 从观测中提取已选组的标记(需要你的环境把这个状态放到observation里) selected_groups = observation[-4:] # 假设最后4位是四个组的选中状态(0未选/1已选) # 获取模型输出的原始Q值 q_values = self.q_net(torch.tensor(observation).unsqueeze(0).to(self.device)) # 构建掩码:已选组对应的所有动作设为-inf mask = torch.zeros_like(q_values) # 按动作范围分组:组1(0-29)、组2(30-59)、组3(60-100059)、组4(100060-100089)、DONE(100090) if selected_groups[0] == 1: mask[0, 0:30] = -float('inf') if selected_groups[1] == 1: mask[0, 30:60] = -float('inf') if selected_groups[2] == 1: mask[0, 60:100060] = -float('inf') if selected_groups[3] == 1: mask[0, 100060:100089] = -float('inf') # 应用掩码后再选动作 masked_q_values = q_values + mask if deterministic: action = masked_q_values.argmax(dim=1).item() else: action = torch.multinomial(torch.softmax(masked_q_values, dim=1), num_samples=1).item() return action, state # 初始化模型时使用自定义策略 model = DQN(CustomDQNPolicy, env, verbose=1)
核心是要把已选动作组的状态嵌入到环境的观测中,这样策略才能动态判断哪些动作需要屏蔽。
方法2:改用分层动作空间
把单一的大动作空间拆分为“选组+选组内动作”的分层结构,用stable-baselines支持的Dict动作空间实现,从动作定义层面限制重复选组:
- 第一层:选动作组(含DONE),用
Discrete(5)(对应组1、组2、组3、组4、DONE) - 第二层:每个组内的动作,用
Dict存储各小组的动作空间
示例环境代码:
from stable_baselines3 import PPO from gymnasium import spaces import gymnasium as gym class CustomEnv(gym.Env): def __init__(self): super().__init__() # 分层动作空间 self.action_space = spaces.Dict({ "group_choice": spaces.Discrete(5), # 0-3对应四个组,4是DONE "group_action": spaces.Dict({ "group1": spaces.Discrete(30), "group2": spaces.Discrete(30), "group3": spaces.Discrete(100000), "group4": spaces.Discrete(30) }) }) self.selected_groups = [False]*4 # 标记组是否已被选中 def step(self, action): group_choice = action["group_choice"] # 处理DONE动作 if group_choice == 4: return self._get_observation(), 0.0, True, False, {} # 检查是否重复选组 if self.selected_groups[group_choice]: # 重复选组给负奖励,不推进状态 return self._get_observation(), -5.0, False, False, {} # 处理组内动作 group_action = action["group_action"][f"group{group_choice+1}"] # 这里添加你的环境逻辑:执行动作、更新状态等 self.selected_groups[group_choice] = True # 判断是否所有组都选完,结束episode done = all(self.selected_groups) reward = 10.0 # 根据你的任务设置合理奖励 return self._get_observation(), reward, done, False, {} def _get_observation(self): # 返回包含已选组状态的观测 return {"selected_groups": self.selected_groups, "other_state": ...} # 用MultiInputPolicy适配Dict动作空间 model = PPO("MultiInputPolicy", CustomEnv(), verbose=1)
这种方法不仅解决了重复选组的问题,还能大幅降低动作空间的复杂度,提升训练效率,尤其适合你有10万级动作的组3场景。
方法3:在环境step中强制终止无效动作
如果不想改动策略或动作空间,可以在环境的step方法中检测到无效动作时,直接终止当前episode并给负奖励:
def step(self, action): # 先判断当前动作属于哪个组(根据你的动作范围映射) group = self._map_action_to_group(action) # 检测是否选了已选组的动作 if group is not None and self.selected_groups[group]: # 直接结束episode,给惩罚 return self._get_observation(), -20.0, True, False, {} # 正常处理有效动作的逻辑...
不过这种方法效率较低,因为智能体仍会尝试选无效动作,只是会被强制终止,不如前两种方法从根源上解决问题。
内容的提问来源于stack exchange,提问作者Nitin Prasad
相关产品推荐
相关产品推荐

