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

如何使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 09:08:54