Java案例:利用历史大数据建模预测
下面我以一个电商销量预测为例,完整演示Java如何从历史大数据出发,完成数据处理、特征工程、模型训练、预测评估的全流程。

整体流程
历史数据(CSV/DB/HDFS)
→ 数据清洗
→ 特征工程
→ 划分训练/测试集
→ 训练模型
→ 模型评估
→ 在线预测
常用Java技术栈:
| 环节 | 技术选型 |
|---|---|
| 大数据存储 | HDFS、Hive、HBase |
| 批处理 | MapReduce、Spark(Java API)、Flink |
| 特征/建模 | Spark MLlib、Smile、Weka、DL4J |
| 在线服务 | Spring Boot + PMML/模型文件 |
| 流式预测 | Flink + Kafka |
案例背景
目标:根据某商品过去每天的价格、促销、天气、星期、历史销量,预测未来某天的销量。
历史数据样例(sales.csv):
date,price,promo,weather,weekday,sales
2023-01-01,29.9,0,2,7,120
2023-01-02,29.9,1,1,1,260
...
代码实现(使用 Smile 库,纯Java)
Smile 是一个高性能 Java 机器学习库,支持回归、分类、聚类等,Maven 一行引入即可。
引入依赖
<dependency>
<groupId>com.github.haifengl</groupId>
<artifactId>smile-core</artifactId>
<version>3.0.2</version>
</dependency>
数据加载与特征工程
import smile.data.DataFrame;
import smile.data.type.StructType;
import smile.io.Read;
import smile.data.formula.Formula;
import smile.data.vector.DoubleVector;
import java.nio.file.Paths;
import java.time.LocalDate;
public class SalesDataLoader {
public static DataFrame load(String path) throws Exception {
DataFrame df = Read.csv(Paths.get(path),
CSVOptions.DEFAULT.withHeader(true));
// 特征:价格、是否促销、天气编码、星期几、月份
DoubleVector dayOfMonth = DoubleVector.of("dayOfMonth",
df.stream().mapToDouble(r -> {
LocalDate d = LocalDate.parse(r.getString("date"));
return d.getDayOfMonth();
}).toArray());
df = df.merge(dayOfMonth);
return df;
}
}
训练回归模型(GBDT)
import smile.data.formula.Formula;
import smile.regression.GradientTreeBoost;
import smile.validation.CrossValidation;
import smile.validation.RegressionMetrics;
public class SalesTrainer {
public static GradientTreeBoost train(DataFrame df) {
// 特征列:price, promo, weather, weekday, dayOfMonth
// 目标列:sales
Formula formula = Formula.lhs("sales");
// 5折交叉验证
CrossValidation cv = CrossValidation.regression(
5, formula, df,
(f, d) -> GradientTreeBoost.fit(f, d, 200, 10));
System.out.println("RMSE = " + cv.rmse());
System.out.println("MAE = " + cv.mae());
// 用全量数据训练最终模型
GradientTreeBoost model = GradientTreeBoost.fit(formula, df);
System.out.println("Feature importance: " + model.importance());
return model;
}
}
保存与加载模型
import java.io.*;
public class ModelIO {
public static void save(GradientTreeBoost model, String file) throws IOException {
try (ObjectOutputStream oos =
new ObjectOutputStream(new FileOutputStream(file))) {
oos.writeObject(model);
}
}
public static GradientTreeBoost load(String file) throws Exception {
try (ObjectInputStream ois =
new ObjectInputStream(new FileInputStream(file))) {
return (GradientTreeBoost) ois.readObject();
}
}
}
在线预测(Spring Boot 接口示例)
@RestController
@RequestMapping("/predict")
public class SalesPredictController {
private GradientTreeBoost model;
@PostConstruct
public void init() throws Exception {
model = ModelIO.load("sales_model.bin");
}
@PostMapping
public double predict(@RequestBody SalesInput input) {
// 构造与训练时一致的 DataFrame 行
Tuple row = Tuple.of(
input.getPrice(),
input.getPromo(),
input.getWeather(),
input.getWeekday(),
input.getDayOfMonth()
);
// 通过 DataFrame 封装再预测
return model.predict(DataFrame.of(row, schema));
}
}
返回示例:
{ "predictedSales": 213.7 }
升级到真正的大数据场景
当数据量达到 TB 级 时,改用 Spark MLlib(Java) 分布式训练:
SparkSession spark = SparkSession.builder()
.appName("SalesForecast").master("yarn").getOrCreate();
Dataset<Row> df = spark.read().option("header", true)
.csv("hdfs://cluster/sales/*.csv");
// 特征向量
VectorAssembler assembler = new VectorAssembler()
.setInputCols(new String[]{"price","promo","weather","weekday"})
.setOutputCol("features");
// GBDT 回归
GBTRegressor gbt = new GBTRegressor()
.setLabelCol("sales")
.setFeaturesCol("features")
.setMaxIter(100);
Pipeline pipeline = new Pipeline().setStages(
new PipelineStage[]{assembler, gbt});
PipelineModel model = pipeline.fit(df);
// 保存到 HDFS
model.write().overwrite().save("hdfs://cluster/models/sales-gbt");
// 预测
model.transform(df).select("sales","prediction").show();
主要区别:
| 维度 | 单机Smile | Spark MLlib |
|---|---|---|
| 数据量 | < 10GB | TB~PB |
| 训练方式 | 内存 | 分布式 |
| 特征规模 | 万级 | 百万级 |
| 接入数据源 | CSV | HDFS/Hive/Kafka |
流式实时预测(Flink 示例)
若需要实时预测(比如用户点击即预测价格):
DataStream<SalesInput> stream = env.addSource(
new FlinkKafkaConsumer<>("sales-input", schema, props));
stream.map(input -> {
// 加载模型(广播变量或本地缓存)
return model.predict(input.toVector());
}).addSink(new FlinkKafkaProducer<>("sales-output", ...));
关键经验总结
- 特征 > 算法:日期、节假日、滞后特征(过去7天均值)、滑窗统计往往比模型本身更重要。
- 时间序列切分:不能随机划分,要按时间滑窗训练,否则会数据泄漏。
- 模型持久化:训练与推理分离,模型文件可放 HDFS 或 S3,Spring Boot 加载。
- 监控与再训练:定期用新数据重训,监控 RMSE 漂移。
- 技术选型:
- 中小规模 → Smile / Weka
- 大数据 → Spark MLlib
- 实时 → Flink + PMML/ONNX
- 深度学习 → DJL / DL4J
完整演示项目结构建议
sales-forecast/
├── data/sales.csv
├── src/main/java/com/demo/
│ ├── SalesDataLoader.java
│ ├── SalesTrainer.java
│ ├── ModelIO.java
│ ├── SalesPredictController.java
│ └── SalesApplication.java
└── pom.xml
如果你有具体的场景(比如金融风控、设备故障预测、用户流失),我可以针对性给出特征设计和模型选择方案,需要我继续展开哪一块?