Keras将输入从ndarray转为tf.data时验证阶段报错排查
Siamese文本分类模型tf.data输入验证错误问题解决
问题背景
原本使用NumPy数组训练Siamese文本分类模型时一切正常,改用tf.data.Dataset加速训练后,训练阶段看似正常,但验证阶段抛出维度不兼容错误:
Input 0 of layer "lstm" is incompatible with the layer: expected ndim=3, found ndim=2. Full shape received: (124, 124)
Call arguments received by layer "model" " f"(type Functional):
• inputs=('tf.Tensor(shape=(124,), dtype=int32)', 'tf.Tensor(shape=(124,), dtype=int32)')
• training=False
• mask=None
原NumPy训练代码
epochs = 10 batch_size = 128 model.fit( x = [train_asset_text_seq, train_bug_text_seq], y = y_train.values.reshape(-1,1), epochs = epochs, batch_size=batch_size, validation_data=([val_asset_text_seq, val_bug_text_seq], y_val.values.reshape(-1,1)) )
调整后的tf.data训练代码
X_train_ds = tf.data.Dataset.from_tensor_slices((train_text_1, train_text_2)) y_train_ds = tf.data.Dataset.from_tensor_slices(y_train.values.reshape(-1,1)) X_val_ds = tf.data.Dataset.from_tensor_slices((val_text_1, val_text_2)) y_val_ds = tf.data.Dataset.from_tensor_slices(y_val.values.reshape(-1,1)) model.fit( tf.data.Dataset.zip((X_train_ds, y_train_ds)).batch(batch_size).repeat(), validation_data=tf.data.Dataset.zip((X_val_ds, y_val_ds)), epochs = epochs, steps_per_epoch=30 )
Siamese模型定义
input_1 = Input(shape=(train_asset_text_seq.shape[1],)) input_2 = Input(shape=(train_bug_text_seq.shape[1],)) common_embed = Embedding( name="synopsis_embedd", input_dim =len(t.word_index)+1, output_dim=EMBEDDING_DIM, input_length=train_asset_text_seq.shape[1], mask_zero=True ) lstm_1 = common_embed(input_1) lstm_2 = common_embed(input_2) common_lstm = LSTM(32, return_sequences=True, activation="relu") vector_1 = common_lstm(lstm_1) vector_1 = Dropout(0.5)(vector_1) vector_1 = Flatten()(vector_1) vector_2 = common_lstm(lstm_2) vector_2 = Dropout(0.5)(vector_2) vector_2 = Flatten()(vector_2) x5 = Lambda(cosine_distance, output_shape=cos_dist_output_shape)([vector_1, vector_2]) conc = Concatenate(axis=-1)([x5, vector_1, vector_2]) x = Dense(100, activation="relu", name='conc_layer')(conc) x = Dropout(0.1)(x) out = Dense(1, activation="sigmoid", name = 'out')(x) model = Model([input_1, input_2], out)
问题原因
不需要调整模型结构,错误根源在于验证集的tf.data.Dataset格式不符合Keras的期望:
- 训练集通过
tf.data.Dataset.zip((X_train_ds, y_train_ds))生成了((文本序列1, 文本序列2), 标签)的结构,Keras能正确识别输入部分和标签部分,因此训练正常。 - 验证集直接传入
tf.data.Dataset.zip((X_val_ds, y_val_ds))时,Keras错误地将整个数据集(包含输入和标签)当成了模型的输入,导致模型接收到的输入是((文本序列1, 文本序列2), 标签),而非预期的(文本序列1, 文本序列2),进而引发LSTM层的维度不兼容错误。
解决方案
1. 统一验证集的数据集结构
构建验证集时,保持和训练集一致的格式,并确保传入validation_data时,数据集的结构是(输入张量元组, 标签张量):
# 构建验证集:和训练集结构一致 val_ds = tf.data.Dataset.zip((X_val_ds, y_val_ds)).batch(batch_size) # 训练代码调整 model.fit( tf.data.Dataset.zip((X_train_ds, y_train_ds)).batch(batch_size).repeat(), validation_data=val_ds, epochs=epochs, steps_per_epoch=30, validation_steps=30 )
2. 额外注意事项
- 验证集不需要添加
repeat():validation_steps已经控制了验证的步数,重复验证集没有实际意义。 - 确保验证集序列长度和训练集一致:由于Embedding层指定了
input_length=train_asset_text_seq.shape[1],如果验证集的文本序列长度和训练集不同,也会引发维度错误,需提前对齐两者的序列长度。
内容的提问来源于stack exchange,提问作者fsulser
相关产品推荐
相关产品推荐

