Java模型部署案例

wen java案例 3

本文目录导读:

Java模型部署案例

  1. 项目结构
  2. Maven依赖配置
  3. 配置类
  4. 数据模型
  5. 服务层实现
  6. 控制器层
  7. 主应用类
  8. 配置Swagger
  9. 高级功能实现
  10. Docker部署
  11. 性能优化配置
  12. 测试用例

我来为您介绍一个完整的Java模型部署案例,使用Spring Boot框架集成机器学习模型。

项目结构

ml-model-deployment/
├── src/main/java/com/example/mlapp/
│   ├── controller/
│   │   └── PredictionController.java
│   ├── model/
│   │   ├── PredictionRequest.java
│   │   └── PredictionResponse.java
│   ├── service/
│   │   ├── ModelService.java
│   │   └── ModelServiceImpl.java
│   └── MlApplication.java
├── src/main/resources/
│   ├── models/
│   │   └── model.pkl (或 .onnx)
│   └── application.yml
└── pom.xml

Maven依赖配置

<!-- pom.xml -->
<?xml version="1.0" encoding="UTF-8"?>
<project xmlns="http://maven.apache.org/POM/4.0.0">
    <modelVersion>4.0.0</modelVersion>
    <groupId>com.example</groupId>
    <artifactId>ml-model-deployment</artifactId>
    <version>1.0.0</version>
    <parent>
        <groupId>org.springframework.boot</groupId>
        <artifactId>spring-boot-starter-parent</artifactId>
        <version>2.7.14</version>
    </parent>
    <properties>
        <java.version>11</java.version>
    </properties>
    <dependencies>
        <!-- Spring Boot Web -->
        <dependency>
            <groupId>org.springframework.boot</groupId>
            <artifactId>spring-boot-starter-web</artifactId>
        </dependency>
        <!-- PMML (跨平台模型格式) -->
        <dependency>
            <groupId>org.jpmml</groupId>
            <artifactId>pmml-evaluator</artifactId>
            <version>1.6.6</version>
        </dependency>
        <dependency>
            <groupId>org.jpmml</groupId>
            <artifactId>pmml-evaluator-extension</artifactId>
            <version>1.6.6</version>
        </dependency>
        <!-- ONNX Runtime -->
        <dependency>
            <groupId>com.microsoft.onnxruntime</groupId>
            <artifactId>onnxruntime</artifactId>
            <version>1.15.1</version>
        </dependency>
        <!-- TensorFlow Java (可选) -->
        <dependency>
            <groupId>org.tensorflow</groupId>
            <artifactId>tensorflow-core-api</artifactId>
            <version>0.5.0</version>
        </dependency>
        <!-- Lombok -->
        <dependency>
            <groupId>org.projectlombok</groupId>
            <artifactId>lombok</artifactId>
            <optional>true</optional>
        </dependency>
        <!-- Validation -->
        <dependency>
            <groupId>org.springframework.boot</groupId>
            <artifactId>spring-boot-starter-validation</artifactId>
        </dependency>
        <!-- Swagger/OpenAPI -->
        <dependency>
            <groupId>io.springfox</groupId>
            <artifactId>springfox-boot-starter</artifactId>
            <version>3.0.0</version>
        </dependency>
    </dependencies>
</project>

配置类

// application.yml
server:
  port: 8080
spring:
  application:
    name: ml-model-service
model:
  path: classpath:models/model.pmml
  onnx-path: classpath:models/model.onnx
  type: pmml  # pmml, onnx, tensorflow
# 性能配置
prediction:
  cache:
    enabled: true
    size: 1000
    ttl: 3600

数据模型

// PredictionRequest.java
package com.example.mlapp.model;
import lombok.Data;
import javax.validation.constraints.NotNull;
import java.util.Map;
@Data
public class PredictionRequest {
    @NotNull(message = "Features cannot be null")
    private Map<String, Object> features;
    private String modelVersion;
    private boolean returnProbability = false;
}
// PredictionResponse.java
package com.example.mlapp.model;
import lombok.Builder;
import lombok.Data;
import java.util.Map;
@Data
@Builder
public class PredictionResponse {
    private Object prediction;
    private Map<String, Double> probabilities;
    private double confidence;
    private long processingTime;
    private String modelVersion;
}

服务层实现

// ModelService.java
package com.example.mlapp.service;
import com.example.mlapp.model.PredictionRequest;
import com.example.mlapp.model.PredictionResponse;
public interface ModelService {
    PredictionResponse predict(PredictionRequest request);
    String getModelInfo();
}
// ModelServiceImpl.java
package com.example.mlapp.service;
import org.jpmml.evaluator.*;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.stereotype.Service;
import org.springframework.core.io.Resource;
import javax.annotation.PostConstruct;
import java.io.*;
import java.util.*;
@Service
public class ModelServiceImpl implements ModelService {
    @Value("${model.path}")
    private Resource modelResource;
    private Evaluator evaluator;
    @PostConstruct
    public void init() throws Exception {
        // 加载PMML模型
        try (InputStream is = modelResource.getInputStream()) {
            PMML pmml = org.jpmml.model.PMMLUtil.unmarshal(is);
            evaluator = ModelEvaluatorFactory.newInstance().newModelEvaluator(pmml);
            evaluator.verify();
            System.out.println("Model loaded successfully: " + evaluator.getSummary());
        }
    }
    @Override
    public PredictionResponse predict(PredictionRequest request) {
        long startTime = System.currentTimeMillis();
        // 准备输入特征
        Map<String, FieldValue> arguments = new HashMap<>();
        for (Map.Entry<String, Object> entry : request.getFeatures().entrySet()) {
            FieldName fieldName = new FieldName(entry.getKey());
            InputField inputField = evaluator.getActiveField(fieldName);
            FieldValue value = inputField.prepare(entry.getValue());
            arguments.put(fieldName.getValue(), value);
        }
        // 执行预测
        Map<FieldName, ?> result = evaluator.evaluate(arguments);
        // 解析结果
        Object prediction = null;
        Map<String, Double> probabilities = new HashMap<>();
        for (Map.Entry<FieldName, ?> entry : result.entrySet()) {
            FieldName fieldName = entry.getKey();
            if (evaluator.getTargetField().equals(fieldName)) {
                prediction = entry.getValue();
                if (entry.getValue() instanceof Computable) {
                    prediction = ((Computable) entry.getValue()).getResult();
                }
            }
        }
        long processingTime = System.currentTimeMillis() - startTime;
        return PredictionResponse.builder()
                .prediction(prediction)
                .probabilities(probabilities)
                .confidence(calculateConfidence(probabilities))
                .processingTime(processingTime)
                .modelVersion("1.0.0")
                .build();
    }
    private double calculateConfidence(Map<String, Double> probabilities) {
        return probabilities.values().stream()
                .max(Double::compare)
                .orElse(0.0);
    }
    @Override
    public String getModelInfo() {
        return evaluator.getSummary();
    }
}

控制器层

// PredictionController.java
package com.example.mlapp.controller;
import com.example.mlapp.model.PredictionRequest;
import com.example.mlapp.model.PredictionResponse;
import com.example.mlapp.service.ModelService;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.http.ResponseEntity;
import org.springframework.web.bind.annotation.*;
import io.swagger.annotations.Api;
import io.swagger.annotations.ApiOperation;
import javax.validation.Valid;
@RestController
@RequestMapping("/api/v1/predict")
@Api(tags = "Machine Learning Prediction API")
public class PredictionController {
    @Autowired
    private ModelService modelService;
    @PostMapping
    @ApiOperation("Make a prediction using the deployed model")
    public ResponseEntity<PredictionResponse> predict(
            @Valid @RequestBody PredictionRequest request) {
        PredictionResponse response = modelService.predict(request);
        return ResponseEntity.ok(response);
    }
    @GetMapping("/health")
    @ApiOperation("Health check endpoint")
    public ResponseEntity<String> health() {
        return ResponseEntity.ok("Model service is running");
    }
    @GetMapping("/info")
    @ApiOperation("Get model information")
    public ResponseEntity<String> getModelInfo() {
        return ResponseEntity.ok(modelService.getModelInfo());
    }
}

主应用类

// MlApplication.java
package com.example.mlapp;
import org.springframework.boot.SpringApplication;
import org.springframework.boot.autoconfigure.SpringBootApplication;
import org.springframework.cache.annotation.EnableCaching;
import springfox.documentation.swagger2.annotations.EnableSwagger2;
@SpringBootApplication
@EnableSwagger2
@EnableCaching
public class MlApplication {
    public static void main(String[] args) {
        SpringApplication.run(MlApplication.class, args);
    }
}

配置Swagger

// SwaggerConfig.java
package com.example.mlapp.config;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import springfox.documentation.builders.ApiInfoBuilder;
import springfox.documentation.builders.PathSelectors;
import springfox.documentation.builders.RequestHandlerSelectors;
import springfox.documentation.service.ApiInfo;
import springfox.documentation.spi.DocumentationType;
import springfox.documentation.spring.web.plugins.Docket;
@Configuration
public class SwaggerConfig {
    @Bean
    public Docket api() {
        return new Docket(DocumentationType.SWAGGER_2)
                .select()
                .apis(RequestHandlerSelectors.basePackage("com.example.mlapp.controller"))
                .paths(PathSelectors.any())
                .build()
                .apiInfo(apiInfo());
    }
    private ApiInfo apiInfo() {
        return new ApiInfoBuilder()
                .title("ML Model Deployment API")
                .description("API for deploying and serving machine learning models")
                .version("1.0.0")
                .build();
    }
}

高级功能实现

批量预测

@PostMapping("/batch")
@ApiOperation("Batch prediction")
public ResponseEntity<List<PredictionResponse>> predictBatch(
        @Valid @RequestBody List<PredictionRequest> requests) {
    List<PredictionResponse> responses = requests.stream()
            .map(modelService::predict)
            .collect(Collectors.toList());
    return ResponseEntity.ok(responses);
}

模型版本管理

// ModelVersionManager.java
package com.example.mlapp.service;
import org.springframework.stereotype.Component;
import java.util.concurrent.ConcurrentHashMap;
@Component
public class ModelVersionManager {
    private final ConcurrentHashMap<String, ModelService> modelVersions = new ConcurrentHashMap<>();
    public void addModel(String version, ModelService model) {
        modelVersions.put(version, model);
    }
    public ModelService getModel(String version) {
        return modelVersions.getOrDefault(version, modelVersions.get("latest"));
    }
}

Docker部署

# Dockerfile
FROM openjdk:11-jre-slim
WORKDIR /app
COPY target/ml-model-deployment-1.0.0.jar app.jar
COPY src/main/resources/models/ models/
EXPOSE 8080
ENTRYPOINT ["java", "-jar", "app.jar"]
# docker-compose.yml
version: '3.8'
services:
  ml-service:
    build: .
    ports:
      - "8080:8080"
    environment:
      - SPRING_PROFILES_ACTIVE=prod
      - MODEL_PATH=/app/models/model.pmml
    volumes:
      - ./models:/app/models
    healthcheck:
      test: ["CMD", "curl", "-f", "http://localhost:8080/api/v1/predict/health"]
      interval: 30s
      timeout: 10s
      retries: 3

性能优化配置

// AsyncConfig.java
@Configuration
@EnableAsync
public class AsyncConfig {
    @Bean("predictionExecutor")
    public Executor predictionExecutor() {
        ThreadPoolTaskExecutor executor = new ThreadPoolTaskExecutor();
        executor.setCorePoolSize(10);
        executor.setMaxPoolSize(50);
        executor.setQueueCapacity(1000);
        executor.setThreadNamePrefix("prediction-");
        executor.initialize();
        return executor;
    }
}

测试用例

// PredictionControllerTest.java
@SpringBootTest
@AutoConfigureMockMvc
public class PredictionControllerTest {
    @Autowired
    private MockMvc mockMvc;
    @Test
    public void testPrediction() throws Exception {
        String requestJson = """
            {
                "features": {
                    "feature1": 1.5,
                    "feature2": 2.3,
                    "feature3": 0.8
                }
            }
            """;
        mockMvc.perform(post("/api/v1/predict")
                .contentType(MediaType.APPLICATION_JSON)
                .content(requestJson))
                .andExpect(status().isOk())
                .andExpect(jsonPath("$.prediction").isNotEmpty());
    }
}

这个完整的Java模型部署案例包括了:

  1. 多种模型格式支持(PMML、ONNX、TensorFlow)
  2. RESTful API接口
  3. 批量预测支持
  4. 模型版本管理
  5. 性能优化
  6. 容器化部署
  7. 完整的测试覆盖

您可以根据实际需求选择合适的模型格式和部署方式。

上一篇Java评分卡案例

下一篇PMML案例

抱歉,评论功能暂时关闭!