如何在Keras模型中为输入张量添加序列编号并完成拼接?
为Keras输入张量每行添加序列编号的解决方法
问题说明
要构建Keras模型,给输入张量的每一行添加序列编号,但原代码因张量形状不匹配报错。
原尝试代码:
input_layer = Input(shape=(3, 3)) seq = tf.range(3) seq = tf.reshape(seq, (3, 1)) concatenated = Concatenate(axis=-1)([input_layer, seq]) additional_layer = Dense(4, activation="relu")(concatenated) ...
输入层形状为(None, 3, 3)(None代表批量维度),而seq形状是(3,1);即使调整seq为(1,3,1),拼接时仍因批量维度不匹配报错。
测试代码及报错信息
测试代码:
a = np.array([ [[10.0,20.0,30.0], [20.0,30.0,40.0], [40.0,50.0,50.0]], [[10.0,20.0,30.0], [20.0,30.0,40.0], [40.0,50.0,50.0]], [[10.0,20.0,30.0], [20.0,30.0,40.0], [40.0,50.0,50.0]] ]) x_tf = tf.convert_to_tensor(a) input_layer = tf.keras.layers.Input(shape=(3, 3)) seq = tf.range(3, dtype=tf.float32) seq = tf.reshape(seq, (1, 3, 1)) concatenated = tf.keras.layers.Lambda(lambda x:tf.concat([x, seq], axis=-1))(input_layer) model = Model(inputs=input_layer, outputs=concatenated) print(model(x_tf))
报错:
InvalidArgumentError: Exception encountered when calling layer 'lambda_5' (type Lambda). {{function_node __wrapped__ConcatV2_N_2_device_/job:localhost/replica:0/task:0/device:CPU:0}} ConcatOp : Dimension 0 in both shapes must be equal: shape[0] = [3,3,3] vs. shape[1] = [1,3,1] [Op:ConcatV2] name: concat Call arguments received by layer 'lambda_5' (type Lambda): • inputs=tf.Tensor(shape=(3, 3, 3), dtype=float32) • mask=None • training=None
解决方案
核心是让序列编号张量的批量维度和输入张量保持一致,通过tf.tile动态复制序列编号的批量维度,匹配输入的样本数量。
修正后的代码:
import tensorflow as tf from tensorflow.keras.layers import Input, Lambda from tensorflow.keras.models import Model import numpy as np # 定义输入层 input_layer = Input(shape=(3, 3)) # 创建基础序列编号,形状为(1, 3, 1) seq = tf.range(3, dtype=tf.float32) seq = tf.reshape(seq, (1, 3, 1)) # 在Lambda层中动态匹配批量维度并拼接 concatenated = Lambda(lambda x: tf.concat([ x, # 将seq的批量维度复制为和输入x相同的数量 tf.tile(seq, [tf.shape(x)[0], 1, 1]) ], axis=-1))(input_layer) # 构建模型 model = Model(inputs=input_layer, outputs=concatenated) # 测试验证 a = np.array([ [[10.0,20.0,30.0], [20.0,30.0,40.0], [40.0,50.0,50.0]], [[10.0,20.0,30.0], [20.0,30.0,40.0], [40.0,50.0,50.0]], [[10.0,20.0,30.0], [20.0,30.0,40.0], [40.0,50.0,50.0]] ]) x_tf = tf.convert_to_tensor(a) print(model(x_tf))
原理说明
tf.shape(x)[0]获取输入张量的批量大小(测试数据中为3);tf.tile(seq, [tf.shape(x)[0], 1, 1])将seq的第一个维度(批量)复制3次,使seq形状变为(3,3,1),和输入张量的(3,3,3)在除拼接轴外的维度完全匹配;- 此时在
axis=-1(最后一维)拼接,就能得到形状为(3,3,4)的输出,每个样本的每行都添加了对应的序列编号。
内容的提问来源于stack exchange,提问作者Marko Zadravec
相关产品推荐
相关产品推荐

