Keras添加自定义Attention层出现[32,2]与[1200,2]形状不兼容报错如何解决?
错误原因分析
- 第二层双向LSTM输出维度不匹配Attention输入要求:你第二层
Bidirectional(LSTM(units=lstm_out, go_backwards=True))没有设置return_sequences=True,默认返回最后一个时间步的输出,维度为(batch_size, 2*lstm_out)(二维张量),而你的Attention层默认需要接收三维的序列输入(batch_size, seq_len, feature_dim),这一步已经触发维度异常。 - Attention输出维度和标签维度不匹配:即使修复上一个问题,你设置Attention的
return_sequences=True,输出会是三维张量(batch_size, seq_len, feature_dim),后面直接接Dense层输出分类结果的话,得到的输出维度是(batch_size, seq_len, 2),而你的标签维度是(batch_size, 2),两个张量计算损失时形状对不上,就是你报错的[32,2] vs [1200,2]的核心原因(1200为批量大小乘序列长度的结果)。 - 自定义Attention存在通用性缺陷:你的Attention层里偏置项的形状写死为
input_shape[1],依赖固定的序列长度,后续如果输入序列长度变化会直接报错。
修复方案
- 修改第二层双向LSTM,添加
return_sequences=True参数,保证输出三维序列张量给Attention层:
model.add(Bidirectional(LSTM(units=lstm_out, return_sequences=True, go_backwards=True)))
- 调整Attention层的
return_sequences参数为False,让Attention层自动对序列维度求和压缩,输出二维张量(batch_size, feature_dim),匹配后续Dense层和标签的维度要求:
model.add(Attention(return_sequences=False))
- (可选)优化Attention层的偏置项定义,去掉对固定序列长度的依赖,修改build方法如下:
def build(self, input_shape): self.W = self.add_weight(name="att_weight", shape=(input_shape[-1], 1), initializer="normal") # 偏置只作用在特征维度,适配任意序列长度 self.b = self.add_weight(name="att_bias", shape=(1,), initializer="zeros") super(Attention, self).build(input_shape)
内容的提问来源于stack exchange,提问作者Lee93
相关产品推荐
相关产品推荐

