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

如何在Java中调用Python开发的ML Classifier Model?

在Java环境中调用Python机器学习分类模型的可行方案

方案1:将模型导出为ONNX格式,用Java ONNX Runtime加载

  • 步骤1:Python端导出模型为ONNX
    以Scikit-learn模型为例,使用skl2onnx完成转换:
    from skl2onnx import convert_sklearn
    from skl2onnx.common.data_types import FloatTensorType
    import joblib
    import numpy as np
    
    # 加载训练好的分类器模型
    model = joblib.load("your_classifier.pkl")
    
    # 根据模型输入特征维度定义初始类型,示例为4维特征
    initial_type = [('float_input', FloatTensorType([None, 4]))]
    onnx_model = convert_sklearn(model, initial_types=initial_type)
    
    # 保存ONNX格式模型
    with open("classifier.onnx", "wb") as f:
        f.write(onnx_model.SerializeToString())
    
  • 步骤2:Java端加载模型并推理
    先通过Maven引入ONNX Runtime依赖:
    <dependency>
        <groupId>com.microsoft.onnxruntime</groupId>
        <artifactId>onnxruntime</artifactId>
        <version>1.16.3</version> <!-- 使用最新稳定版本 -->
    </dependency>
    
    推理代码示例:
    import ai.onnxruntime.*;
    import java.util.Map;
    
    public class ONNXInference {
        public static void main(String[] args) throws OrtException {
            try (OrtEnvironment env = OrtEnvironment.getEnvironment();
                 OrtSession session = env.createSession("classifier.onnx", new OrtSession.SessionOptions())) {
    
                // 构造输入特征数据,匹配模型输入维度
                float[][] inputData = {{1.2f, 3.4f, 5.6f, 7.8f}};
                OnnxTensor inputTensor = OnnxTensor.createTensor(env, inputData);
    
                // 执行推理并获取结果
                Map<String, OnnxTensor> inputs = Map.of("float_input", inputTensor);
                try (OrtSession.Result results = session.run(inputs)) {
                    float[][] output = (float[][]) results.get(0).getValue();
                    System.out.println("预测结果:" + output[0][0]);
                }
            }
        }
    }
    
  • 注意事项:部分复杂自定义模型可能无法直接转换为ONNX,需提前验证兼容性。

方案2:搭建Python API服务,Java通过HTTP请求调用

  • 步骤1:用FastAPI编写模型接口
    from fastapi import FastAPI
    import joblib
    import numpy as np
    
    app = FastAPI()
    # 加载本地模型文件
    model = joblib.load("your_classifier.pkl")
    
    @app.post("/predict")
    def predict(features: list[float]):
        input_data = np.array(features).reshape(1, -1)
        prediction = model.predict(input_data)[0]
        return {"prediction": int(prediction)}
    
    启动服务:uvicorn main:app --host 0.0.0.0 --port 8000
  • 步骤2:Java发送HTTP请求调用接口
    使用Java原生HttpClient实现:
    import java.net.URI;
    import java.net.http.HttpClient;
    import java.net.http.HttpRequest;
    import java.net.http.HttpResponse;
    import com.google.gson.JsonObject;
    import com.google.gson.JsonParser;
    
    public class APIPredictor {
        public static void main(String[] args) throws Exception {
            HttpClient client = HttpClient.newHttpClient();
            JsonObject requestBody = new JsonObject();
            requestBody.add("features", JsonParser.parseString("[1.2,3.4,5.6,7.8]"));
    
            HttpRequest request = HttpRequest.newBuilder()
                    .uri(URI.create("http://localhost:8000/predict"))
                    .header("Content-Type", "application/json")
                    .POST(HttpRequest.BodyPublishers.ofString(requestBody.toString()))
                    .build();
    
            HttpResponse<String> response = client.send(request, HttpResponse.BodyHandlers.ofString());
            JsonObject responseJson = JsonParser.parseString(response.body()).getAsJsonObject();
            System.out.println("预测结果:" + responseJson.get("prediction").getAsInt());
        }
    }
    
  • 注意事项:生产环境需添加接口认证、请求限流等机制保障稳定性。

方案3:使用Py4J实现Java与Python进程通信

  • 步骤1:Python端启动Py4J网关
    先安装Py4J:pip install py4j
    编写网关服务代码:
    from py4j.java_gateway import GatewayServer
    import joblib
    import numpy as np
    
    model = joblib.load("your_classifier.pkl")
    
    class ModelService:
        def predict(self, features):
            input_data = np.array(features).reshape(1, -1)
            return int(model.predict(input_data)[0])
    
    if __name__ == "__main__":
        # 启动网关,默认端口25333
        gateway = GatewayServer(ModelService())
        gateway.start()
        print("Py4J Gateway started successfully")
    
  • 步骤2:Java端连接网关并调用模型
    通过Maven引入Py4J依赖:
    <dependency>
        <groupId>net.sf.py4j</groupId>
        <artifactId>py4j</artifactId>
        <version>0.10.9.7</version> <!-- 使用最新稳定版本 -->
    </dependency>
    
    调用代码示例:
    import py4j.java_gateway.JavaGateway;
    
    public class Py4JPredictor {
        public static void main(String[] args) {
            // 连接本地Py4J网关
            JavaGateway gateway = new JavaGateway();
            ModelService service = gateway.getEntryPoint();
    
            double[] features = {1.2, 3.4, 5.6, 7.8};
            int prediction = service.predict(features);
            System.out.println("预测结果:" + prediction);
    
            // 关闭网关连接
            gateway.shutdown();
        }
    }
    
  • 注意事项:确保Java与Python端的Py4J版本一致,避免兼容性问题。

方案4:将Python脚本打包为可执行文件,Java通过Process调用

  • 步骤1:用PyInstaller打包推理脚本
    安装PyInstaller:pip install pyinstaller
    编写推理脚本predict.py:
    import joblib
    import numpy as np
    import sys
    
    if __name__ == "__main__":
        model = joblib.load("your_classifier.pkl")
        # 从命令行参数获取特征数据
        features = list(map(float, sys.argv[1:]))
        input_data = np.array(features).reshape(1, -1)
        prediction = model.predict(input_data)[0]
        # 输出预测结果供Java读取
        print(int(prediction))
    
    打包为单文件可执行程序:pyinstaller --onefile predict.py
  • 步骤2:Java调用可执行文件
    import java.io.BufferedReader;
    import java.io.InputStreamReader;
    
    public class ProcessPredictor {
        public static void main(String[] args) throws Exception {
            // 替换为实际的可执行文件路径和特征参数
            String[] cmd = {"/path/to/predict", "1.2", "3.4", "5.6", "7.8"};
            Process process = Runtime.getRuntime().exec(cmd);
    
            // 读取程序输出的预测结果
            BufferedReader reader = new BufferedReader(new InputStreamReader(process.getInputStream()));
            String prediction = reader.readLine();
            System.out.println("预测结果:" + prediction);
    
            process.waitFor();
        }
    }
    
  • 注意事项:需将模型文件放在可执行文件可访问的路径,跨平台部署时需针对不同系统打包对应版本的可执行文件。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 08:15:11