基于MLP模型的Streamlit高血压预测APP异常排查与修正
高血压多分类MLP模型Streamlit部署异常排查与修复方案
可能的异常原因
- 输入数据预处理不匹配:训练时对特征做了标准化/归一化,但Streamlit应用直接使用原始输入数据,导致数据分布与训练集偏差极大,模型默认输出训练样本占比最高的类别(比如高血压3期)。
- 模型加载失效:加载的不是训练完成的最终模型,而是未收敛的中间 checkpoint;或是Keras模型保存/加载时结构不匹配(比如丢失自定义层、未正确保存权重)。
- 预测逻辑错误:获取模型输出时硬编码了类别索引(比如固定取最后一个类别),或是对softmax输出的概率排序处理错误,导致始终返回同一结果。
- 类别映射错误:训练时的标签编码顺序与Streamlit中的类别映射不匹配,或是直接将结果硬编码为“高血压3期”。
修复方案与调整后的Streamlit代码
核心修复步骤
- 复用训练时的预处理逻辑:保存训练阶段使用的标准化/归一化器(如
StandardScaler),在应用中加载后对输入特征做相同转换。 - 正确加载模型:使用Keras官方的
load_model()方法加载完整模型,若有自定义层需指定custom_objects参数。 - 修正预测逻辑:通过
np.argmax()获取概率最高的类别索引,再映射到对应的类别名称。
调整后的完整代码示例
import streamlit as st import numpy as np from tensorflow import keras from sklearn.preprocessing import StandardScaler import joblib # 加载训练好的模型和预处理scaler model = keras.models.load_model('hypertension_mlp_model.h5') scaler = joblib.load('scaler.pkl') # 训练时保存的标准化器 # 类别映射必须与训练时的标签编码顺序完全一致 class_mapping = { 0: "正常", 1: "高血压1期", 2: "高血压2期", 3: "高血压3期" } st.title("高血压风险多分类预测") # 根据你的模型特征列表调整输入控件 age = st.number_input("年龄", min_value=18, max_value=100, value=50) systolic_bp = st.number_input("收缩压", min_value=80, max_value=220, value=120) diastolic_bp = st.number_input("舒张压", min_value=50, max_value=140, value=80) bmi = st.number_input("BMI指数", min_value=15.0, max_value=40.0, value=22.0) fasting_glucose = st.number_input("空腹血糖", min_value=3.0, max_value=20.0, value=5.0) # 整理输入为模型可接受的格式 input_features = np.array([[age, systolic_bp, diastolic_bp, bmi, fasting_glucose]]) # 执行与训练时一致的预处理 scaled_input = scaler.transform(input_features) # 模型预测与结果解析 predictions = model.predict(scaled_input, verbose=0) predicted_idx = np.argmax(predictions, axis=1)[0] predicted_class = class_mapping[predicted_idx] predicted_prob = predictions[0][predicted_idx] st.subheader("预测结果") st.write(f"**风险等级**: {predicted_class}") st.write(f"**预测概率**: {predicted_prob:.2%}")
关键注意事项
- 训练时需用
joblib.dump(scaler, 'scaler.pkl')保存标准化器,不能在应用中重新拟合scaler。 - 验证模型保存的完整性:训练完成后使用
model.save('hypertension_mlp_model.h5')保存完整模型结构与权重。 - 若训练时使用了
LabelEncoder编码类别,需同步保存编码器(joblib.dump(label_encoder, 'label_encoder.pkl')),加载后用label_encoder.inverse_transform()解析类别。 - 可添加调试代码(如
st.write(scaled_input)、st.write(predictions)),查看输入特征的缩放结果和模型输出的概率分布,快速定位问题。
内容的提问来源于stack exchange,提问作者junesoka
相关产品推荐
相关产品推荐

