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

如何让Vertex AI自定义模型批量预测返回结果包含置信度得分

结论

Vertex AI完全支持批量预测返回scikit-learn分类模型的置信度得分,你当前只拿到分类标签的问题和批量预测接口配置无关,是默认上传的scikit-learn模型仅调用了predict()方法返回分类结果,没有调用predict_proba()输出置信度导致的。

解决方案

需要先重新封装上传你的scikit-learn模型,自定义预测逻辑同时返回标签和置信度,再发起批量预测即可。

步骤1:自定义预测逻辑并重新上传模型

在你打包上传的模型 artifacts 目录中新增predictor.py文件,定义预测处理类:

import joblib
import os

class ScikitLearnPredictor:
    def load(self, artifacts_uri: str):
        # 加载本地预训练好的scikit-learn模型文件
        model_path = os.path.join(artifacts_uri, "model.joblib")
        self.model = joblib.load(model_path)
    
    def predict(self, instances):
        # 同时获取分类标签和置信度得分
        predictions = self.model.predict(instances)
        # 二分类场景取正类置信度,多分类场景可直接返回predict_proba的完整结果
        confidences = self.model.predict_proba(instances)[:, 1].tolist()
        # 组装返回结构
        return [
            {"prediction": bool(pred), "confidence": conf}
            for pred, conf in zip(predictions, confidences)
        ]

使用Vertex AI SDK上传自定义预测逻辑的模型:

from google.cloud import aiplatform

aiplatform.init(project="你的GCP项目ID", location="部署区域")

model = aiplatform.Model.upload_scikit_learn_model_file(
    model_file_path="本地预训练模型的存储路径",
    display_name="你的模型名称",
    predictor_scheme="predictor.ScikitLearnPredictor", # 指定自定义预测类
)

步骤2:发起批量预测

你原有批量预测代码不需要做任何修改,直接使用新上传的模型发起任务即可,返回的结果会自动包含prediction(分类结果)和confidence(置信度得分)两个字段:

batch_prediction_job = model.batch_predict(
    job_display_name = job_display_name,
    gcs_source = input_path,
    instances_format = "jsonl", # 按需选择你的输入格式
    gcs_destination_prefix = output_path,
    starting_replica_count = 1,
    max_replica_count = 10,
    sync = True,
)

batch_prediction_job.wait()
注意事项
  • 打包模型时需要同步新增requirements.txt文件,指定和训练环境版本一致的scikit-learn依赖,避免模型加载失败
  • 多分类场景可以按需调整predict方法的返回结构,返回所有类别的置信度即可

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 01:54:07