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

Android中基于ONNX部署Scikit SVM如何获取预测概率值?

问题描述

我用Scikit-learn在Python中训练了一个开启概率输出的SVM模型,代码如下:

model = svm.SVC(gamma=params["gamma"], C=params["C"], probability=True)
model.fit(X_train, y_train)
initial_type = [('float_input', FloatTensorType([None, X_train.shape[1]]))]
onnx_model = convert_sklearn(model, initial_types=initial_type,options={"zipmap": False},  target_opset=12)

同时编写了Kotlin类用于在Android中调用该模型,目前能得到正确的预测标签,但无法获取预测概率值,该如何实现?

对应的Kotlin调用代码:

class Classifier(context: Context, modelFileName: String="") {

    private val session: OrtSession
    private val env: OrtEnvironment

    init {
        val modelBytes = context.resources.openRawResource(R.raw.svm_model2).readBytes()
        env = OrtEnvironment.getEnvironment()
        session = try {
            env.createSession(modelBytes)
        } catch (e: Exception) {
            Log.e("Classifier", "Error initializing model: ${e.message}")
            throw e
        }
    }

    suspend fun runInference(inputData: FloatArray): String {
        val inputName = session.inputNames?.iterator()?.next()
        val floatBufferInputs = FloatBuffer.wrap(inputData)
        val inputTensor = OnnxTensor.createTensor(env, floatBufferInputs, longArrayOf(1,inputData.size.toLong()))
        val result = session.run(mapOf(inputName to inputTensor))
        val re = result[0].value as Array<*>
        val outputArray = re.map { it.toString() }.toTypedArray()

        Log.v("classifier", outputArray.contentToString())

        return outputArray.contentToString()
    }
}
解决方案

要获取预测概率,需要处理ONNX模型的两个输出节点(标签和概率),具体修改如下:

1. 确认Python端导出的模型包含概率输出

当SVC设置probability=True时,sklearn-onnx转换后的ONNX模型会自动生成两个输出:

  • 第一个输出:预测的类别标签(对应你当前拿到的结果)
  • 第二个输出:每个类别的预测概率分布

你当前的导出代码无需修改,options={"zipmap": False}仅影响标签的输出格式,不会屏蔽概率输出。

2. 修改Kotlin端代码获取概率

session.run()返回的result是包含所有输出的列表,读取索引1的元素即可拿到概率值,修改后的代码示例:

class Classifier(context: Context, modelFileName: String="") {

    private val session: OrtSession
    private val env: OrtEnvironment

    init {
        val modelBytes = context.resources.openRawResource(R.raw.svm_model2).readBytes()
        env = OrtEnvironment.getEnvironment()
        session = try {
            env.createSession(modelBytes)
        } catch (e: Exception) {
            Log.e("Classifier", "Error initializing model: ${e.message}")
            throw e
        }
    }

    // 自定义数据类,同时返回标签和概率结果
    data class InferenceResult(val predictedLabel: String, val probabilities: FloatArray) {
        override fun equals(other: Any?): Boolean {
            if (this === other) return true
            if (javaClass != other?.javaClass) return false
            other as InferenceResult
            if (predictedLabel != other.predictedLabel) return false
            if (!probabilities.contentEquals(other.probabilities)) return false
            return true
        }

        override fun hashCode(): Int {
            var result = predictedLabel.hashCode()
            result = 31 * result + probabilities.contentHashCode()
            return result
        }
    }

    suspend fun runInference(inputData: FloatArray): InferenceResult {
        val inputName = session.inputNames?.iterator()?.next()
        val floatBufferInputs = FloatBuffer.wrap(inputData)
        val inputTensor = OnnxTensor.createTensor(env, floatBufferInputs, longArrayOf(1,inputData.size.toLong()))
        val result = session.run(mapOf(inputName to inputTensor))

        // 获取预测标签(第一个输出)
        val labelArray = result[0].value as Array<Float>
        val predictedLabel = labelArray[0].toInt().toString() // 根据你的标签类型调整转换逻辑

        // 获取概率分布(第二个输出)
        val probabilityArray = result[1].value as Array<Array<Float>>
        val probabilities = probabilityArray[0] // 单条输入取第一个批量维度的数组

        Log.v("classifier", "Predicted label: $predictedLabel")
        Log.v("classifier", "Probabilities: ${probabilities.contentToString()}")

        return InferenceResult(predictedLabel, probabilities)
    }
}

代码说明

  • 新增InferenceResult数据类,方便同时返回标签和概率
  • 概率输出的结构为Array<Array<Float>>,外层是批量维度,内层是每个类别的概率值
  • 根据实际标签类型调整predictedLabel的转换逻辑(如果标签是字符串,需自行适配)

注意事项

  • 若不确定输出顺序,可通过session.outputNames查看输出节点名称,概率输出通常包含probabilities关键字
  • 确保重新导出ONNX模型时,SVC的probability=True参数已正确设置

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 06:33:22