使用Keras函数式API实现Transformer时look_ahead_mask维度报错如何解决
报错根因
错误由create_look_ahead_mask函数实现逻辑错误导致,你在生成全1矩阵时,传入tf.fill或tf.ones的形状参数是3维张量,但TensorFlow要求构造张量时传入的形状参数必须是1维的。
解决步骤
- 核对
create_look_ahead_mask函数的实现,你大概率在生成掩码时错误引入了多余的维度。正确的前瞻掩码生成逻辑参考如下:
def create_look_ahead_mask(dec_inputs): # 取Decoder输入的序列长度,dec_inputs形状为 (batch_size, seq_len) seq_len = tf.shape(dec_inputs)[1] # 生成下三角矩阵,形状为 (seq_len, seq_len),传入的形状参数是1维列表,符合要求 look_ahead_mask = 1 - tf.linalg.band_part(tf.ones((seq_len, seq_len)), -1, 0) # 扩展为 (1, seq_len, seq_len) 匹配注意力层的广播规则 return look_ahead_mask[tf.newaxis, ...]
- 替换你原有错误的
create_look_ahead_mask实现后,重新实例化Transformer模型即可正常构造。 - 如果你需要保留原有Lambda层的
output_shape=(1, None, None)声明,和上面的实现逻辑完全兼容,不需要额外修改。
为什么单独调用Decoder时没有报错
你单独实例化Decoder时没有动态生成掩码,直接传入了符合形状要求的固定掩码,所以不会触发形状校验错误。
内容的提问来源于stack exchange,提问作者R.tahir
相关产品推荐
相关产品推荐

