分类变量Embedding输入Keras网络调用fit方法报错排查
Keras Embedding训练报错InvalidArgumentError的问题排查与修复
核心错误原因及修复步骤
1. 输入数据格式不匹配
你把训练/验证数据包装成了列表(X_train_actual = []然后append数组),但Keras模型期望的是二维numpy数组(shape=(样本数量, 1)),列表格式会导致输入维度不兼容,触发InvalidArgumentError。
修复代码:
# 替换原来的X_train_actual和X_valid_actual定义 X_train_actual = np.array(X_train["data"]).reshape(-1, 1) X_valid_actual = np.array(X_valid["data"]).reshape(-1, 1)
2. OrdinalEncoder的错误使用
你在验证集和测试集上调用了fit_transform,这会重新拟合编码器,导致训练/验证/测试集的编码规则不一致,可能出现大于categories=7的类别索引,和Embedding层的输入要求冲突。
修复代码:
enc = OrdinalEncoder() # 仅在训练集fit X_train[["data"]] = enc.fit_transform(X_train[["data"]]) # 验证集和测试集用transform X_valid[["data"]] = enc.transform(X_valid[["data"]]) X_test[["data"]] = enc.transform(X_test[["data"]])
3. LabelEncoder的错误使用
同样,你对每个数据集单独使用fit_transform,会导致类别映射不一致(比如训练集里的类别A编码为0,验证集里可能编码为1),进而影响分类结果和损失计算。
修复代码:
le = LabelEncoder() # 仅在训练集fit Y_train["data_type"] = le.fit_transform(Y_train["data_type"]) # 验证集和测试集用transform Y_valid["data_type"] = le.transform(Y_valid["data_type"]) Y_test["data_type"] = le.transform(Y_test["data_type"])
4. 损失函数选择错误
你的模型输出是softmax激活的多分类结果,却使用了回归任务的mse损失函数,这不仅不符合分类任务的损失逻辑,也可能导致训练过程中的维度或数值不匹配问题。
修复代码:
# 替换compile中的损失函数 model.compile(loss='categorical_crossentropy', optimizer=optimizer, metrics=['accuracy'])
完整修正后的关键代码片段
# 数据预处理修正 X_train, X_rem, y_train, y_rem = train_test_split(X,y, train_size=0.8) X_valid, X_test, y_valid, y_test = train_test_split(X_rem,y_rem, test_size=0.5) Y_train=pd.DataFrame({'data_type':y_train}) Y_valid=pd.DataFrame({'data_type':y_valid}) Y_test=pd.DataFrame({'data_type':y_test}) # OrdinalEncoder修正 enc = OrdinalEncoder() X_train[["data"]] = enc.fit_transform(X_train[["data"]]) X_valid[["data"]] = enc.transform(X_valid[["data"]]) X_test[["data"]] = enc.transform(X_test[["data"]]) # 输入数据格式修正 X_train_actual = np.array(X_train["data"]).reshape(-1, 1) X_valid_actual = np.array(X_valid["data"]).reshape(-1, 1) # LabelEncoder修正 le = LabelEncoder() Y_train["data_type"] = le.fit_transform(Y_train["data_type"]) Y_valid["data_type"] = le.transform(Y_valid["data_type"]) Y_test["data_type"] = le.transform(Y_test["data_type"]) Y_train_actual= to_categorical(Y_train.to_numpy()) Y_valid_actual= to_categorical(Y_valid.to_numpy()) # 模型编译修正 model.compile(loss='categorical_crossentropy', optimizer=optimizer, metrics=['accuracy']) # 训练调用 history = model.fit(X_train_actual, Y_train_actual, batch_size=100, epochs=500, validation_data =(X_valid_actual, Y_valid_actual), callbacks=[checkpoint, early_stop], verbose=1)
内容的提问来源于stack exchange,提问作者Anantha
相关产品推荐
相关产品推荐

