如何使用Keras将ANN模型保存为h5并在其他项目中加载预测
保存Keras ANN模型为H5格式并加载预测
嘿,针对你的客户流失预测ANN模型,我来一步步教你搞定模型的保存和复用:
一、保存训练好的模型为H5格式
Keras提供了非常便捷的方法保存完整模型——包括结构、权重、编译配置,加载后无需重新编译就能直接用。
在你的代码中,等模型训练完成后(也就是classifier.fit(...)执行结束),添加这一行即可:
# 将模型保存为H5格式文件 classifier.save('churn_prediction_model.h5')
这个H5文件会打包所有必要信息:
- 模型的网络层结构配置
- 各层训练好的权重参数
- 模型的编译信息(优化器、损失函数等)
- 甚至能保留训练状态(如果后续需要继续训练的话)
二、在新项目中加载模型并预测
要在其他项目里复用这个模型,只需要加载H5文件,再做好数据预处理就能预测了:
1. 加载模型
先导入load_model函数,再加载保存好的H5文件:
# 导入必要依赖 import numpy as np import pandas as pd from keras.models import load_model from sklearn.preprocessing import StandardScaler, LabelEncoder, OneHotEncoder # 加载保存的模型 classifier = load_model('churn_prediction_model.h5')
2. 处理新数据(关键!)
新数据必须和训练数据做完全一致的预处理,否则预测结果会失真。比如你有一条新客户数据:
# 示例新数据(特征顺序和训练集一致:CreditScore, Geography, Gender, Age, Tenure, Balance, NumOfProducts, HasCrCard, IsActiveMember, EstimatedSalary) new_customer = np.array([[620, 'Germany', 'Female', 35, 5, 85000, 1, 0, 1, 75000]])
先做和训练时一样的编码:
# 注意:实际项目中建议保存训练时的编码器,这里为演示重新初始化 labelencoder_geo = LabelEncoder() new_customer[:, 1] = labelencoder_geo.fit_transform(new_customer[:, 1]) labelencoder_gender = LabelEncoder() new_customer[:, 2] = labelencoder_gender.fit_transform(new_customer[:, 2]) # OneHot编码并去掉第一列(避免虚拟变量陷阱) onehotencoder = OneHotEncoder(categorical_features=[1]) new_customer = onehotencoder.fit_transform(new_customer).toarray() new_customer = new_customer[:, 1:]
再用训练时的标准化器做特征缩放(绝对不能用新数据重新拟合scaler):
# 实际项目中建议保存训练时的scaler,比如用joblib # 保存:import joblib; joblib.dump(sc, 'scaler.pkl') # 加载:sc = joblib.load('scaler.pkl') sc = StandardScaler() new_customer_scaled = sc.transform(new_customer)
3. 执行预测
最后用加载好的模型输出结果:
# 预测流失概率 churn_prob = classifier.predict(new_customer_scaled) # 转换为二分类结果(阈值设为0.5) churn_pred = (churn_prob > 0.5) print(f"客户流失预测结果:{'会流失' if churn_pred[0][0] else '不会流失'}") print(f"流失概率:{churn_prob[0][0]:.2f}")
额外提示:保存预处理组件
为了确保新数据预处理的一致性,建议把训练时用的LabelEncoder、OneHotEncoder和StandardScaler也保存下来,比如用joblib:
import joblib # 保存预处理工具 joblib.dump(sc, 'scaler.pkl') joblib.dump(labelencoder_geo, 'labelencoder_geo.pkl') joblib.dump(labelencoder_gender, 'labelencoder_gender.pkl')
加载时直接调用:
sc = joblib.load('scaler.pkl') labelencoder_geo = joblib.load('labelencoder_geo.pkl')
这样就能彻底避免因预处理不一致导致的预测偏差啦。
内容的提问来源于stack exchange,提问作者Ind
相关产品推荐
相关产品推荐

