本文目录导读:

我来为您介绍一个完整的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模型部署案例包括了:
- 多种模型格式支持(PMML、ONNX、TensorFlow)
- RESTful API接口
- 批量预测支持
- 模型版本管理
- 性能优化
- 容器化部署
- 完整的测试覆盖
您可以根据实际需求选择合适的模型格式和部署方式。