Keras中不同形状张量拼接问题:注意力机制实现报错排查
Keras注意力机制中张量拼接的维度不匹配问题解决办法
嘿,我看你在实现注意力机制时遇到了拼接维度不匹配的问题,咱们来一步步搞定它。
首先看你遇到的报错:
ValueError: A
Concatenatelayer requires inputs with matching shapes except for the concat axis. Got inputs shapes: [(None, 1, 1024), (None, 38, 1024)]
问题根源
你想拼接的两个张量维度不对齐:
context_vector实际是 (None, 1, 1024)(虽然你说它是(?,1024),大概率是代码里不小心多扩展了一次维度,或者后续操作悄悄修改了它的形状)decoder_embedding是 (None, 38, 1024)
拼接时除了指定的concat轴(你选的是axis=-1,也就是最后一维),其他维度必须完全匹配。这里第二个维度(时间步维度)一个是1,一个是38,自然会触发报错。
解决思路
我们需要把context_vector的时间步维度扩展到和decoder_embedding一致(也就是38个时间步),这样解码器的每个时间步都能带上注意力输出的全局上下文信息。具体来说,先把context_vector扩展成(?,1,1024),再把这个单时间步的张量重复38次,变成(?,38,1024),这样就能和decoder_embedding在最后一维顺利拼接了。
修改后的代码
我给你调整了拼接前的关键部分,还补全了函数参数的小疏漏:
def B_Attention_layer(state_h, state_c, encoder_outputs, decoder_embedding): # 新增decoder_embedding参数,不然函数内无法调用 d0 = tf.keras.layers.Dense(1024,name='dense_layer_1') d1 = tf.keras.layers.Dense(1024,name='dense_layer_2') d2 = tf.keras.layers.Dense(1024,name='dense_layer_3') # 处理LSTM隐藏状态,扩展时间维度 hidden_with_time_axis_1 = tf.keras.backend.expand_dims(state_h, 1) hidden_with_time_axis_2 = tf.keras.backend.expand_dims(state_c, 1) # 计算注意力分数与权重 score = d0(tf.keras.activations.tanh(encoder_outputs) + d1(hidden_with_time_axis_1) + d2(hidden_with_time_axis_2)) attention_weights = tf.keras.activations.softmax(score, axis=1) # 计算上下文向量 context_vector = attention_weights * encoder_outputs context_vector = tf.keras.backend.sum(context_vector, axis=1) # 此时shape=(?,1024) # 关键修改:扩展并重复上下文向量,匹配解码器嵌入的时间步 context_vector_expanded = tf.expand_dims(context_vector, axis=1) # shape变为(?,1,1024) time_steps = tf.shape(decoder_embedding)[1] # 动态获取时间步数量,避免硬编码38 context_vector_repeated = tf.tile(context_vector_expanded, [1, time_steps, 1]) # shape变为(?,38,1024) # 现在两个张量维度完全匹配,可在最后一维拼接 input_to_decoder = tf.keras.layers.Concatenate(axis=-1)([context_vector_repeated, decoder_embedding]) return input_to_decoder , attention_weights
额外说明
- 原来的函数参数里没有
decoder_embedding,我给加上了——不然函数内部根本没法调用这个变量,这也是个容易忽略的小细节。 - 用
tf.shape(decoder_embedding)[1]动态获取时间步数量,比硬写38更灵活,以后调整解码器的时间步长度时,代码不用跟着改。
这样修改后,两个张量的维度就完全对齐了,拼接操作就能正常运行啦!
内容的提问来源于stack exchange,提问作者ML_Guy
相关产品推荐
相关产品推荐

