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

DQN网络Sequential与Functional API实现的输入形状差异及Flatten层作用问询

解决Sequential与Functional API输入形状匹配问题及Flatten层作用解析

一、用Functional API复刻Sequential模型的输入结构

你遇到的核心问题是两个模型的输入样本形状不一致:Sequential版本的输入样本形状是(1,) + obs_space(比如obs_space为(4,)时就是(1,4)),经过Flatten层处理后转为一维的(4,);而你最初写的Functional API直接用obs_space作为输入形状,相当于跳过了原始的(1,4)输入和Flatten步骤,才导致运行异常。

要完全对齐Sequential模型的结构,只需在Functional API里还原Flatten层和对应的输入形状,修改后的代码如下:

obs_space = env.observation_space.shape
# 匹配Sequential的输入样本形状:(1,) + obs_space
i = Input(shape=(1,) + obs_space)
# 添加Flatten层,和Sequential中的作用完全一致
x = Flatten()(i)
x = Dense(16, activation='relu')(x)
x = Dense(32, activation='relu')(x)
o = Dense(env.action_space.n, activation='linear')(x)
model = Model(i, o)
model.summary()

运行后你会看到,输入层的输出形状为(None, 1, 4)(假设obs_space是(4,)),经过Flatten后转为(None, 4),和Sequential模型的结构、参数总数(690个)完全对齐,就能正常运行了。

二、model.add(Flatten(input_shape=(1,) + obs_space))的作用拆解

这行代码其实完成了两个关键操作:

  1. 定义输入样本的形状:input_shape=(1,) + obs_space指定了单个输入样本的结构。比如如果环境返回的观测形状是(4,),这里就把输入样本设置为(1,4)——可以理解为把单个4维观测包裹在一个长度为1的额外维度里(有些环境会默认返回带这类额外维度的观测,或者你需要强制模型接受这种格式的输入)。
  2. 压平多维输入:Flatten层的作用是把多维输入直接压平成一维。对于(1,4)的输入,Flatten后会变成(4,),这样后面的全连接层(Dense)就能直接处理这个一维特征向量(Dense层默认对最后一个维度做计算)。

补充说明:Keras里的input_shape参数不包含batch维度,所以模型摘要里的None代表batch维度(可以是任意大小的批次)。你之前看到的[(None,4)]只是Functional API输入层的显示格式,本质和Sequential的(None,4)是同类结构,但因为输入形状本身不匹配才导致运行问题,修改后就能完全统一。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.30 09:38:13