使用BERT预训练模型构建文本分类模型时遇输入不匹配错误
问题解决:BERT模型训练时输入不匹配报错
错误原因
你的模型定义了两个输入层(input_ids 和 attention_masks),但调用 model.fit() 时只传入了 X_train(仅对应input_ids),缺少BERT必需的注意力掩码输入,导致模型收到的输入数量不匹配。
解决方案
1. 生成注意力掩码
BERT需要注意力掩码来区分真实文本token和填充的0值。对输入序列生成对应掩码:非0位置标记为1,0位置标记为0。
在数据划分后添加以下代码:
# 生成训练集和测试集的注意力掩码 attention_mask_train = np.where(X_train != 0, 1, 0) attention_mask_test = np.where(X_test != 0, 1, 0)
2. 修改model.fit的输入参数
将输入改为包含input_ids和注意力掩码的列表,替换原有的model.fit调用:
model.fit( [X_train, attention_mask_train], y_train, epochs=32, batch_size=32, verbose=1, validation_data=([X_test, attention_mask_test], y_test) # 可选,添加验证集监控训练 )
完整修改后的关键代码片段
# 数据划分后添加掩码生成 X_train, X_test, y_train, y_test = train_test_split( x, y, test_size=0.2, random_state=42) # 生成注意力掩码 attention_mask_train = np.where(X_train != 0, 1, 0) attention_mask_test = np.where(X_test != 0, 1, 0) # ... 模型定义部分保持不变 ... # 修改训练调用 model.fit( [X_train, attention_mask_train], y_train, epochs=32, batch_size=32, verbose=1 )
内容的提问来源于stack exchange,提问作者Ihsan ullah khan
相关产品推荐
相关产品推荐

