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

如何在Python中用自建ANN模型预测新输入的多分类标签

自定义ANN模型的保存、加载与多分类预测实现

一、模型保存

根据你构建ANN的方式,选择对应的保存方法:

  • 若用TensorFlow/Keras实现:直接用内置方法保存完整模型(含结构、权重)
# 假设你的模型对象是model
model.save('custom_ann_churn_model.h5')
  • 若用Numpy手动实现自定义ANN类:用pickle序列化整个模型实例
import pickle

# 假设你的自定义模型实例是ann_model
with open('custom_ann_churn_model.pkl', 'wb') as f:
    pickle.dump(ann_model, f)

注意:如果训练时用到了标准化器、编码器这类预处理组件,也要一并保存,比如StandardScaler、OneHotEncoder,否则预测时新数据无法对齐训练逻辑

二、模型加载

对应保存方式加载即可:

Keras模型加载

from tensorflow.keras.models import load_model

loaded_model = load_model('custom_ann_churn_model.h5')

自定义Numpy模型加载

import pickle

with open('custom_ann_churn_model.pkl', 'rb') as f:
    loaded_model = pickle.load(f)

别忘了加载之前保存的预处理组件,比如标准化器:

with open('scaler.pkl', 'rb') as f:
    loaded_scaler = pickle.load(f)

三、新输入预测与自动填充

1. 对齐训练时的预处理逻辑

新输入的特征必须和训练数据做完全一致的处理,比如训练时对CreditScore做了标准化,新输入值也要用同一个标准化器转换:

import numpy as np

# 假设用户输入是字典格式,比如从Streamlit表单收集的
user_input = {
    'CreditScore': 620,
    'Age': 35,
    'Tenure': 4,
    # 其他特征...
}

# 转换为模型接受的输入格式(二维数组)
input_features = np.array([[
    user_input['CreditScore'],
    user_input['Age'],
    user_input['Tenure'],
    # 按训练时的特征顺序排列
]])

# 标准化处理
input_scaled = loaded_scaler.transform(input_features)

2. 多分类预测输出

针对Exited的三类(no/yes/maybe),模型输出的是各类别的概率,取概率最高的类别作为预测结果:

# Keras模型预测
probabilities = loaded_model.predict(input_scaled)[0]  # 获取单样本的概率数组
class_labels = ['no', 'yes', 'maybe']  # 必须和训练时的类别编码顺序一致
predicted_class = class_labels[np.argmax(probabilities)]

# 自定义Numpy模型预测(假设模型有predict方法返回概率)
probabilities = loaded_model.predict(input_scaled)
predicted_class = class_labels[np.argmax(probabilities)]

3. Streamlit中自动填充结果

在Streamlit里,你可以用数据编辑器直接展示带预测结果的表格,或者在表单后自动填充结果单元格:

import streamlit as st
import pandas as pd

# 收集用户输入
col1, col2 = st.columns(2)
with col1:
    credit_score = st.number_input('信用评分', min_value=300, max_value=850)
    age = st.number_input('年龄', min_value=18, max_value=100)
with col2:
    tenure = st.number_input('在网时长(年)', min_value=0, max_value=10)
    # 其他输入组件...

# 预测按钮触发逻辑
if st.button('预测客户流失类别'):
    # 预处理+预测
    input_features = np.array([[credit_score, age, tenure]])
    input_scaled = loaded_scaler.transform(input_features)
    probabilities = loaded_model.predict(input_scaled)[0]
    class_labels = ['no', 'yes', 'maybe']
    predicted_class = class_labels[np.argmax(probabilities)]
    
    # 生成带结果的表格并展示(自动填充结果列)
    result_data = {
        '信用评分': [credit_score],
        '年龄': [age],
        '在网时长': [tenure],
        # 其他特征列...
        '流失预测结果': [predicted_class]
    }
    result_df = pd.DataFrame(result_data)
    st.data_editor(result_df, disabled=True)  # 禁用编辑,只展示

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 19:15:42