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

Keras-RL训练井字棋智能体报错:期望dense输入2维却得到(1,1,3,3)数组

错误根因
  • reset() 方法返回值维度不匹配:你在reset中返回的是形状为(3,3)的二维数组,但你在__init__中已经明确将观测压平为长度9的一维向量,模型输入也期望一维结构。同时Keras-RL会自动给输入加上batch维度和窗口长度维度,你返回的3x3数组堆叠后就会形成报错里的(1, 1, 3, 3)结构,和模型期望的维度完全不匹配。
  • observation_space 定义不符合Gym规范:不能直接用NumPy数组套Discrete类,正确做法是用Box空间定义压平后的9维观测,取值范围对应井字棋的空、X、O三种状态。
  • 模型输入形状适配问题:Keras-RL的SequentialMemory设置了window_length=1,会自动给输入增加窗口维度,需要显式调整模型输入形状适配该结构。
修复代码

1. 自定义环境修复

from gym import Env, spaces
import numpy as np

class TTTEnv(Env):
    def __init__(self):
        self.action_space = spaces.Discrete(9)
        # 修复观测空间定义:9维一维向量,每个值取值为0/1/2
        self.observation_space = spaces.Box(low=0, high=2, shape=(9,), dtype=np.int32)
        self.game = Game()
        self.state = self.game.gameArray.flatten()
    
    def step(self, action):
        reward = 0
        done = False
        self.game.printGame()
        position = self.game.inputs[action]
        if self.game.gameArray[position[0],position[1]] != 0:
            reward -= 20
            done = True
        else:
            self.game.gameArray[position[0],position[1]] = 1
            gameOver, winner = self.game.checkWinGYM()
            if winner == "win":
                reward += 50
                done = gameOver
            elif winner == "draw":
                reward += 10
            elif winner == "ingame":
                self.game.handleBotTurn()
                gameOver, winner = self.game.checkWinGYM()
                if winner == "loss":
                    done = gameOver
                    reward -= 50
                elif winner == "draw":
                    done = gameOver
                    reward += 10
        info = {}
        return self.game.gameArray.flatten(), reward, done, info
    
    def render(self):
        pass

    def reset(self):
        self.game.resetGameArray()
        # 修复:返回压平后的一维数组,而非3x3二维数组
        self.state = self.game.gameArray.flatten()
        return self.state

2. 模型与智能体构建修复

env = TTTEnv()
actions = env.action_space.n

def build_model(actions):
    model = Sequential()
    # 调整输入形状适配Keras-RL的窗口维度
    model.add(Dense(24, activation="relu", input_shape=(1,9)))
    model.add(Dense(24, activation="relu"))
    model.add(Dense(actions, activation="linear"))
    return model

def build_agent(model, actions):
    policy = BoltzmannQPolicy()
    memory = SequentialMemory(limit=50000, window_length=1)
    dqn = DQNAgent(model=model, memory=memory, policy=policy,
                    nb_actions=actions, nb_steps_warmup=10, 
                    target_model_update=1e-2)
    return dqn

model = build_model(actions)
dqn = build_agent(model, actions)
dqn.compile(Adam(lr=1e-3), metrics=["mae"])
dqn.fit(env, nb_steps=50000, visualize=False, verbose=1)

内容的提问来源于stack exchange,提问作者Sultan Al-Rashed

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 03:15:03