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
相关产品推荐
相关产品推荐

