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

如何从y_pred获取模型预测类别?预测结果超出预期范围求助

解决模型预测类别超出预期范围的问题

核心修复:使用模型的predict方法做预测

你当前代码的关键错误是直接用loaded_model(X_test)调用模型,这不是Keras模型进行预测的正确方式,应该调用predict方法:

# 替换错误代码
predict_x = loaded_model.predict(X_test)
y_pred = np.argmax(predict_x, axis=1)

额外排查步骤(如果修复后仍有问题)

  1. 验证模型加载的正确性

    • 检查输出层的神经元数量是否对应你的5个类别(0-4),确保输出层定义为Dense(5, activation='softmax')。
    • 确认模型结构和权重的加载代码正确:
      from tensorflow.keras.models import model_from_json
      
      # 加载结构
      with open('model_structure.json', 'r') as f:
          loaded_model = model_from_json(f.read())
      # 加载权重
      loaded_model.load_weights('model_weights.h5')
      # 必须编译模型(匹配训练时的损失函数和优化器)
      loaded_model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy'])
      
  2. 确保测试数据和训练数据预处理一致

    • 如果训练时对输入做了归一化、缩放或维度转换,测试数据必须执行完全相同的操作(不能重新拟合预处理工具):
      # 示例:训练时用了StandardScaler
      from sklearn.preprocessing import StandardScaler
      scaler = StandardScaler()
      # 训练时:X_train = scaler.fit_transform(X_train)
      # 测试时:必须用transform,不能fit
      X_test = scaler.transform(X_test)
      
    • 检查输入数据的形状是否和模型输入一致,比如模型输入是(None, 28, 28, 1),测试数据需要reshape到对应形状。
  3. 确认np.argmax的axis参数正确

    • 若predict_x的形状是(样本数, 5),axis=1是正确的(取每个样本概率最大的类别索引);如果输出是其他维度(比如CNN的多维度输出),需要调整axis值。

完整验证代码示例

import numpy as np
import pandas as pd
from tensorflow.keras.models import model_from_json

# 加载模型
with open('model_structure.json', 'r') as f:
    loaded_model = model_from_json(f.read())
loaded_model.load_weights('model_weights.h5')
loaded_model.compile(loss='sparse_categorical_crossentropy', optimizer='adam', metrics=['accuracy'])

# 加载测试数据(根据你的CSV结构调整)
test_data = pd.read_csv('test_labels.csv')
X_test = test_data.drop('true_label', axis=1).values
y_true = test_data['true_label'].values

# 预处理测试数据(和训练时一致)
X_test = X_test / 255.0  # 示例:归一化到0-1
# X_test = X_test.reshape(-1, 28, 28, 1)  # 如果是CNN模型,调整维度

# 预测
predict_probs = loaded_model.predict(X_test)
y_pred = np.argmax(predict_probs, axis=1)

# 对比真实标签和预测标签
accuracy = np.mean(y_pred == y_true)
print(f"测试准确率:{accuracy:.4f}")

# 遍历查看结果
for idx, (pred, true) in enumerate(zip(y_pred, y_true)):
    print(f"样本{idx}: 预测={pred}, 真实={true}")

内容的提问来源于stack exchange,提问作者بتول

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 01:31:17