Keras2.0中拼接含mask_zero=True与无mask的Embedding层报错求解
你遇到的这个错误,根源在于带mask的张量和不带mask的张量在拼接时,Keras自动处理mask信息时出现了维度冲突。咱们一步步拆解问题,再给你几个可行的解决方案:
问题本质
当你给emb_a设置mask_zero=True时,这个Embedding层的输出会附带一个2维的mask张量(形状为(batch_size, 30)),用来标记输入序列中哪些位置是padding(输入值为0的位置)。而emb_b没有设置mask_zero,它的输出没有显式的mask信息,Keras会默认它的mask是1维的(形状为(batch_size,),表示整个样本是否有效)。当你用Concatenate(axis=1)拼接这两个张量时,Keras会尝试合并它们的mask,但2维mask和1维mask无法在axis=1上拼接,于是触发了维度不匹配的错误。
解决方案
根据你的实际需求,你可以选择以下任意一种方案:
方案1:给第二个Embedding也添加mask_zero(如果输入b有padding)
如果你的输入b中确实用0作为padding值,直接给emb_b也设置mask_zero=True,这样两个输出的mask都是2维的(batch_size,30),拼接后mask会自动合并为(batch_size,60),完美解决问题:
a = Input(shape=[30]) b = Input(shape=[30]) emb_a = Embedding(10, 5, mask_zero=True)(a) # 给emb_b也设置mask_zero=True emb_b = Embedding(20, 5, mask_zero=True)(b) cat = Concatenate(axis=1)([emb_a, emb_b]) model = Model(inputs=[a, b], outputs=[cat])
方案2:剥离第一个Embedding的mask(如果不需要保留mask)
如果你的后续层不需要用到mask信息,可以手动剥离emb_a的mask,这样Concatenate就不会尝试处理mask了。用Lambda层是最规范的方式:
from tensorflow.keras.layers import Lambda a = Input(shape=[30]) b = Input(shape=[30]) emb_a = Embedding(10, 5, mask_zero=True)(a) # 用Lambda层剥离mask,只保留张量本身 emb_a_no_mask = Lambda(lambda x: x)(emb_a) emb_b = Embedding(20, 5, mask_zero=False)(b) cat = Concatenate(axis=1)([emb_a_no_mask, emb_b]) model = Model(inputs=[a, b], outputs=[cat])
方案3:给第二个Embedding手动添加全有效mask(如果需要保留mask)
如果你需要保留emb_a的mask信息,同时emb_b的所有序列位置都是有效的,可以给emb_b手动创建一个全True的2维mask,让Keras能正确合并:
import tensorflow as tf from tensorflow.keras.layers import Lambda a = Input(shape=[30]) b = Input(shape=[30]) emb_a = Embedding(10, 5, mask_zero=True)(a) emb_b = Embedding(20, 5, mask_zero=False)(b) # 给emb_b添加全True的mask,形状和序列长度一致 emb_b_with_mask = Lambda( lambda x: x, mask=lambda inputs: tf.ones(tf.shape(inputs)[:-1], dtype=tf.bool) )(emb_b) cat = Concatenate(axis=1)([emb_a, emb_b_with_mask]) model = Model(inputs=[a, b], outputs=[cat])
内容的提问来源于stack exchange,提问作者Geralt Xu

