You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.26 08:54:22