You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Stable Baselines3简单掷硬币游戏示例运行报错排查

掷硬币预测游戏报错解决

错误信息

Using cpu device
Traceback (most recent call last):
  File "/home/user/python/simplegame.py", line 40, in <module>
    model.learn(total_timesteps=10000)
  File "/home/user/python/mypython3.10/lib/python3.10/site-packages/stable_baselines3/ppo/ppo.py", line 315, in learn
    return super().learn(
  File "/home/user/python/mypython3.10/lib/python3.10/site-packages/stable_baselines3/common/on_policy_algorithm.py", line 264, in learn
    total_timesteps, callback = self._setup_learn(
  File "/home/user/python/mypython3.10/lib/python3.10/site-packages/stable_baselines3/common/base_class.py", line 423, in _setup_learn
    self._last_obs = self.env.reset()  # type: ignore[assignment]
  File "/home/user/python/mypython3.10/lib/python3.10/site-packages/stable_baselines3/common/vec_env/dummy_vec_env.py", line 77, in reset
    obs, self.reset_infos[env_idx] = self.envs[env_idx].reset(seed=self._seeds[env_idx], **maybe_options)
TypeError: CoinFlipEnv.reset() got an unexpected keyword argument 'seed'

原问题代码

import gymnasium as gym
import numpy as np
from stable_baselines3 import PPO
from stable_baselines3.common.vec_env import DummyVecEnv

class CoinFlipEnv(gym.Env):
    def __init__(self, heads_probability=0.8):
        super(CoinFlipEnv, self).__init__()
        self.action_space = gym.spaces.Discrete(2)  # 0 for heads, 1 for tails
        self.observation_space = gym.spaces.Discrete(2)  # 0 for heads, 1 for tails
        self.heads_probability = heads_probability
        self.flip_result = None

    def reset(self):
        # Reset the environment
        self.flip_result = None
        return self._get_observation()

    def step(self, action):
        # Perform the action (0 for heads, 1 for tails)
        self.flip_result = int(np.random.rand() < self.heads_probability)

        # Compute the reward (1 for correct prediction, -1 for incorrect)
        reward = 1 if self.flip_result == action else -1

        # Return the observation, reward, done, and info
        return self._get_observation(), reward, True, {}

    def _get_observation(self):
        # Return the current coin flip result
        return self.flip_result

# Create the environment with heads probability of 0.8
env = DummyVecEnv([lambda: CoinFlipEnv(heads_probability=0.8)])

# Create the PPO model
model = PPO("MlpPolicy", env, verbose=1)

# Train the model
model.learn(total_timesteps=10000)

# Save the model
model.save("coin_flip_model")

# Evaluate the model
obs = env.reset()
for _ in range(10):
    action, _states = model.predict(obs)
    obs, rewards, dones, info = env.step(action)
    print(f"Action: {action}, Observation: {obs}, Reward: {rewards}")

问题原因及修复

核心问题

  1. reset方法参数不兼容:导入的gymnasium要求reset方法必须支持seed和options参数,且返回(observation, info)格式结果,但你的reset方法未接收这些参数,也未返回合法格式。
  2. step方法返回格式错误:Gymnasium的step方法要求返回5个值:(observation, reward, terminated, truncated, info),你仅返回4个,不符合规范。
  3. 初始观测非法:reset时flip_result为None,但观测空间是Discrete(2),None不属于该空间,会触发后续报错。

修复后的完整代码

import gymnasium as gym
import numpy as np
from stable_baselines3 import PPO
from stable_baselines3.common.vec_env import DummyVecEnv

class CoinFlipEnv(gym.Env):
    def __init__(self, heads_probability=0.8):
        super(CoinFlipEnv, self).__init__()
        self.action_space = gym.spaces.Discrete(2)  # 0=正面, 1=反面
        # 掷硬币是无状态任务,观测用固定值即可,这里用Discrete(1)表示无额外状态信息
        self.observation_space = gym.spaces.Discrete(1)
        self.heads_probability = heads_probability

    def reset(self, seed=None, options=None):
        # 调用父类方法处理种子,保证随机性可复现
        super().reset(seed=seed)
        # 返回符合观测空间的初始值和空info字典
        return 0, {}

    def step(self, action):
        # 执行掷硬币动作
        flip_result = int(np.random.rand() < self.heads_probability)
        # 计算奖励:预测正确得1分,错误扣1分
        reward = 1 if flip_result == action else -1
        # 每轮任务结束,terminated=True;未被截断,truncated=False
        return 0, reward, True, False, {}

# 创建环境
env = DummyVecEnv([lambda: CoinFlipEnv(heads_probability=0.8)])

# 初始化PPO模型
model = PPO("MlpPolicy", env, verbose=1)

# 训练模型
model.learn(total_timesteps=10000)

# 保存训练好的模型
model.save("coin_flip_model")

# 评估模型效果
obs = env.reset()
for _ in range(10):
    action, _states = model.predict(obs)
    obs, rewards, dones, info = env.step(action)
    print(f"预测结果: {'正面' if action[0]==0 else '反面'}, 实际结果: {'正面' if rewards[0]==1 else '反面'}, 奖励: {rewards[0]}")

关键修改说明

  • 给reset方法添加seed和options参数,调用父类reset处理种子,同时返回(观测值, 信息字典)的标准格式。
  • 调整观测空间为Discrete(1),适配掷硬币无状态的任务特性;若需传递历史结果,可改为Discrete(2)并在reset时返回合法初始值。
  • 修改step方法返回5个值,符合Gymnasium规范。
  • 优化评估环节输出,让结果更直观。

内容的提问来源于stack exchange,提问作者Simd

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.03 01:22:09