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

能否为Keras模型保存特征名并加载,以处理独热编码列缺失问题?

解决Keras中独热编码特征列匹配问题的方案

完全可以实现将独热编码的特征名列表与Keras模型绑定保存的需求,以下是具体步骤和代码示例:

1. 训练阶段:保存特征名到模型

训练时先完成独热编码,提取完整的特征名列表,再将其作为自定义属性附加到Keras模型上,最后用SavedModel格式保存模型(该格式支持保留自定义属性,HDF5格式可能不兼容)。

import pandas as pd
import tensorflow as tf
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Dense

# 1. 准备训练数据并做独热编码
train_df = pd.DataFrame({
    'cat_feature1': ['A', 'B', 'C', 'A', 'B'],
    'cat_feature2': ['X', 'Y', 'X', 'Y', 'X'],
    'num_feature': [10, 20, 30, 40, 50],
    'target': [1, 0, 1, 0, 1]
})

# 对类别特征做独热编码
train_encoded = pd.get_dummies(train_df, columns=['cat_feature1', 'cat_feature2'])
# 获取完整的特征名列表(包含所有可能的独热编码列)
full_feature_names = train_encoded.drop('target', axis=1).columns.tolist()

# 2. 构建并训练Keras模型
model = Sequential([
    Dense(32, activation='relu', input_shape=(len(full_feature_names),)),
    Dense(16, activation='relu'),
    Dense(1, activation='sigmoid')
])
model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'])
model.fit(train_encoded.drop('target', axis=1), train_encoded['target'], epochs=20)

# 3. 将特征名列表附加到模型,并用SavedModel格式保存
model.full_feature_names = full_feature_names
model.save('keras_model_with_features', save_format='tf')

2. 预测阶段:恢复特征名并补全缺失列

加载模型后提取保存的特征名列表,对预测数据做独热编码后,补全缺失的列(设为0)并调整列顺序与训练时一致,确保输入形状和特征匹配。

# 1. 加载模型并恢复特征名列表
loaded_model = tf.keras.models.load_model('keras_model_with_features')
saved_feature_names = loaded_model.full_feature_names

# 2. 准备预测数据(仅包含部分类别)
test_df = pd.DataFrame({
    'cat_feature1': ['A', 'A'],
    'cat_feature2': ['X', 'X'],
    'num_feature': [60, 70]
})

# 3. 对预测数据做独热编码
test_encoded = pd.get_dummies(test_df, columns=['cat_feature1', 'cat_feature2'])

# 4. 补全缺失的特征列并调整顺序
# 补全缺失列,值设为0
for col in saved_feature_names:
    if col not in test_encoded.columns:
        test_encoded[col] = 0
# 强制列顺序与训练时一致
test_encoded = test_encoded[saved_feature_names]

# 5. 执行预测
predictions = loaded_model.predict(test_encoded)

关键注意事项

  • 确保训练和预测时pd.get_dummies的参数完全一致(比如prefix、drop_first、dtype等),避免特征名或编码规则不匹配。
  • 必须使用SavedModel格式保存模型(即save_format='tf'),HDF5格式无法保留自定义属性。
  • 补全列后一定要调整列顺序,否则即使列数正确,特征顺序错位也会导致预测结果错误。

内容的提问来源于stack exchange,提问作者Tim Romero

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 16:35:55