TF-Agents自定义PyEnvironment报time_step与time_step_spec不匹配
自定义PyEnvironment校验报time_step与spec不匹配问题修复
问题表现
搭建井字棋TF-Agents自定义PyEnvironment时,运行utils.validate_py_environment(environment, episodes=1)抛出如下报错:
ValueError: Given `time_step` does not match expected `time_step_spec`
已手动指定observation_spec和棋盘数组dtype为np.int32,仍未定位到规格不匹配位置。从报错详情可直接看到核心差异:
- 实际返回的observation为三维数组
shape=(1,3,3) - 预期observation_spec声明的形状不匹配,且代码中存在其他隐藏运行时bug
根因分析
- Observation形状与spec定义完全不符:代码中
_observation_spec声明shape为(1,),但实际存储状态的self._board是3×3二维数组,返回时用np.array([self._state])额外包装了一层维度,最终输出形状为(1,3,3),和spec定义完全不匹配。 - Action spec定义错误:当前action_spec设置
maximum=1,仅支持0、1两个动作,井字棋共9个落子位,合法动作范围应为0-8。 - 存在未定义变量bug:
_step方法中直接使用turn、moves变量,未加self.前缀引用实例属性,校验通过后运行会直接抛变量未定义错误。 - 落子索引逻辑错误:直接用整数action索引二维数组
self._board[action]只会取到整行数据,无法定位到具体落子格子。 - 重置逻辑遗漏:
_reset方法没有重置棋盘数组、回合数等状态,多次重置会复用旧棋盘数据。
修复代码
修正初始化方法的spec定义
def __init__(self): # 9个合法落子动作,取值0-8 self._action_spec = array_spec.BoundedArraySpec( shape=(), dtype=np.int32, minimum=0, maximum=8, name='action') # 观测为展平后的9个棋盘格子,取值0(空)/1(玩家1)/2(玩家2) self._observation_spec = array_spec.BoundedArraySpec( shape=(9,), dtype=np.int32, minimum=0, maximum=2, name='observation') self._board = np.zeros((3,3), dtype=np.int32) self._episode_ended = False self._turn = 0 self._winner = 0
修正重置方法
def _reset(self): self._board = np.zeros((3,3), dtype=np.int32) self._episode_ended = False self._turn = 0 self._winner = 0 # 返回展平后的棋盘,形状匹配(9,)的spec定义 return ts.restart(self._board.flatten())
修正步进入方法
def _step(self, action): # 回合已结束时先触发重置 if self._episode_ended: return self.reset() # 计算当前落子玩家 player = 1 if self._turn % 2 == 0 else 2 # 整数action转棋盘坐标 row, col = action // 3, action % 3 # 非法落子(位置已被占)直接给负奖励终止回合 if self._board[row, col] != 0: self._episode_ended = True return ts.termination(self._board.flatten(), reward=-1.0) # 执行落子 self._board[row, col] = player self._turn += 1 # 胜负判定 winner = 0 # 对角线判定 if self._board[0,0] == self._board[1,1] == self._board[2,2] != 0: winner = self._board[1,1] elif self._board[0,2] == self._board[1,1] == self._board[2,0] != 0: winner = self._board[1,1] # 行列判定 for i in range(3): if self._board[i,0] == self._board[i,1] == self._board[i,2] != 0: winner = self._board[i,0] break if self._board[0,i] == self._board[1,i] == self._board[2,i] != 0: winner = self._board[0,i] break # 返回对应时间步 if winner != 0: self._episode_ended = True reward = 1.0 if winner == player else -1.0 return ts.termination(self._board.flatten(), reward) elif np.all(self._board != 0): # 平局 self._episode_ended = True return ts.termination(self._board.flatten(), reward=0.0) else: return ts.transition(self._board.flatten(), reward=0.0, discount=1.0)
修复后运行环境校验可正常通过,observation形状、dtype完全匹配spec定义,所有隐藏运行时bug均已修复。
内容的提问来源于stack exchange,提问作者flowerboy
相关产品推荐
相关产品推荐

