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))的作用拆解
这行代码其实完成了两个关键操作:
- 定义输入样本的形状:
input_shape=(1,) + obs_space指定了单个输入样本的结构。比如如果环境返回的观测形状是(4,),这里就把输入样本设置为(1,4)——可以理解为把单个4维观测包裹在一个长度为1的额外维度里(有些环境会默认返回带这类额外维度的观测,或者你需要强制模型接受这种格式的输入)。 - 压平多维输入:Flatten层的作用是把多维输入直接压平成一维。对于(1,4)的输入,Flatten后会变成(4,),这样后面的全连接层(Dense)就能直接处理这个一维特征向量(Dense层默认对最后一个维度做计算)。
补充说明:Keras里的input_shape参数不包含batch维度,所以模型摘要里的None代表batch维度(可以是任意大小的批次)。你之前看到的[(None,4)]只是Functional API输入层的显示格式,本质和Sequential的(None,4)是同类结构,但因为输入形状本身不匹配才导致运行问题,修改后就能完全统一。
内容的提问来源于stack exchange,提问作者GGuemez
相关产品推荐
相关产品推荐

