本文目录导读:

我来给你一个完整的Java大数据预测建模实战方案,包含从数据准备到模型部署的全流程。
整体架构设计
graph LR
A[历史数据] --> B[数据清洗]
B --> C[特征工程]
C --> D[模型训练]
D --> E[模型评估]
E --> F[模型部署]
完整案例:电商销售预测
1 数据准备与预处理
import com.opencsv.CSVReader;
import java.io.FileReader;
import java.util.*;
public class DataPreprocessor {
// 数据清洗和标准化
public static List<SalesRecord> loadAndCleanData(String filePath) throws Exception {
List<SalesRecord> records = new ArrayList<>();
try (CSVReader reader = new CSVReader(new FileReader(filePath))) {
String[] nextLine;
// 跳过表头
reader.readNext();
while ((nextLine = reader.readNext()) != null) {
SalesRecord record = new SalesRecord();
record.setDate(nextLine[0]);
record.setProductId(nextLine[1]);
record.setQuantity(Integer.parseInt(nextLine[2]));
record.setPrice(Double.parseDouble(nextLine[3]));
record.setRevenue(Double.parseDouble(nextLine[4]));
// 数据过滤:去除异常值
if (isValidRecord(record)) {
records.add(record);
}
}
}
// 排序和去重
records.sort(Comparator.comparing(SalesRecord::getDate));
return records;
}
private static boolean isValidRecord(SalesRecord record) {
return record.getQuantity() > 0 && record.getPrice() > 0
&& record.getRevenue() >= 0;
}
}
2 特征工程
import java.time.LocalDate;
import java.time.format.DateTimeFormatter;
import java.util.*;
public class FeatureEngineering {
// 时间序列特征提取
public static Map<String, double[]> extractFeatures(List<SalesRecord> records) {
Map<String, double[]> features = new HashMap<>();
// 按产品分组
Map<String, List<SalesRecord>> productGroups = groupByProduct(records);
for (Map.Entry<String, List<SalesRecord>> entry : productGroups.entrySet()) {
String productId = entry.getKey();
List<SalesRecord> productRecords = entry.getValue();
// 提取时间特征
double[] timeFeatures = extractTimeFeatures(productRecords);
// 提取销量特征
double[] salesFeatures = extractSalesFeatures(productRecords);
// 提取价格特征
double[] priceFeatures = extractPriceFeatures(productRecords);
// 合并所有特征
double[] allFeatures = mergeFeatures(timeFeatures, salesFeatures, priceFeatures);
features.put(productId, allFeatures);
}
return features;
}
private static double[] extractTimeFeatures(List<SalesRecord> records) {
double[] features = new double[7];
DateTimeFormatter formatter = DateTimeFormatter.ofPattern("yyyy-MM-dd");
recordLoop: for (SalesRecord record : records) {
LocalDate date = LocalDate.parse(record.getDate(), formatter);
// 记录数不足以分析时跳过
if (records.size() < 4) continue recordLoop;
// 1. 月份
features[0] += date.getMonthValue();
// 2. 星期几
features[1] += date.getDayOfWeek().getValue();
// 3. 是否为周末
features[2] += (date.getDayOfWeek().getValue() > 5) ? 1 : 0;
// 4. 月份的天数
features[3] += date.lengthOfMonth();
// 5. 季度
features[4] += (date.getMonthValue() - 1) / 3 + 1;
// 6. 是否为月初前5天
features[5] += (date.getDayOfMonth() <= 5) ? 1 : 0;
// 7. 是否为月末后5天
features[6] += (date.getDayOfMonth() > 25) ? 1 : 0;
}
// 归一化
for (int i = 0; i < features.length; i++) {
features[i] /= records.size();
}
return features;
}
private static double[] extractSalesFeatures(List<SalesRecord> records) {
double[] features = new double[5];
// 计算历史统计数据
List<Integer> quantities = new ArrayList<>();
for (SalesRecord record : records) {
quantities.add(record.getQuantity());
}
if (quantities.size() > 0) {
// 1. 平均销量
features[0] = quantities.stream().mapToInt(Integer::intValue).average().orElse(0);
// 2. 销量标准差
double variance = quantities.stream()
.mapToDouble(q -> Math.pow(q - features[0], 2))
.average().orElse(0);
features[1] = Math.sqrt(variance);
// 3. 最大值
features[2] = quantities.stream().mapToInt(Integer::intValue).max().orElse(0);
// 4. 最小值
features[3] = quantities.stream().mapToInt(Integer::intValue).min().orElse(0);
// 5. 趋势(最近5天平均销量 - 整体平均销量)
int recentCount = Math.min(5, quantities.size());
double recentAvg = quantities.subList(quantities.size() - recentCount, quantities.size())
.stream().mapToInt(Integer::intValue).average().orElse(0);
features[4] = recentAvg - features[0];
}
return features;
}
private static double[] extractPriceFeatures(List<SalesRecord> records) {
double[] features = new double[3];
List<Double> prices = new ArrayList<>();
for (SalesRecord record : records) {
prices.add(record.getPrice());
}
if (prices.size() > 0) {
// 1. 平均价格
features[0] = prices.stream().mapToDouble(Double::doubleValue).average().orElse(0);
// 2. 价格变动趋势
if (prices.size() > 2) {
features[1] = (prices.get(prices.size()-1) - prices.get(prices.size()-2))
/ prices.get(prices.size()-2);
}
// 3. 价格指数(相对于历史平均)
features[2] = prices.get(prices.size()-1) / features[0] - 1;
}
return features;
}
private static double[] mergeFeatures(double[]... featureArrays) {
int totalLength = 0;
for (double[] arr : featureArrays) {
totalLength += arr.length;
}
double[] merged = new double[totalLength];
int index = 0;
for (double[] arr : featureArrays) {
System.arraycopy(arr, 0, merged, index, arr.length);
index += arr.length;
}
return merged;
}
}
3 机器学习模型(使用Weka)
<!-- Maven依赖 -->
<dependency>
<groupId>nz.ac.waikato.cms.weka</groupId>
<artifactId>weka-stable</artifactId>
<version>3.8.5</version>
</dependency>
import weka.classifiers.functions.LinearRegression;
import weka.classifiers.functions.SMOreg;
import weka.classifiers.trees.RandomForest;
import weka.core.*;
import java.util.*;
public class SalesPredictor {
// 构建数据集
public static Instances buildDataset(Map<String, double[]> features,
Map<String, List<Double>> targets) {
ArrayList<Attribute> attributes = new ArrayList<>();
// 添加特征属性
for (int i = 0; i < 15; i++) {
attributes.add(new Attribute("feature" + i));
}
// 添加目标属性(本期销量)
attributes.add(new Attribute("quantity"));
Instances dataset = new Instances("SalesDataset", attributes, 0);
dataset.setClassIndex(dataset.numAttributes() - 1);
// 添加数据实例
for (Map.Entry<String, double[]> entry : features.entrySet()) {
String productId = entry.getKey();
double[] featureVector = entry.getValue();
List<Double> targetValues = targets.get(productId);
// 每个目标值生成一个训练样本
if (targetValues != null) {
for (double target : targetValues) {
Instance instance = new DenseInstance(dataset.numAttributes());
for (int i = 0; i < featureVector.length; i++) {
instance.setValue(attributes.get(i), featureVector[i]);
}
instance.setValue(attributes.get(featureVector.length), target);
dataset.add(instance);
}
}
}
return dataset;
}
// 模型训练与预测
public static PredictionResult trainAndPredict(Instances trainingData,
double[] predictionFeatures) throws Exception {
PredictionResult result = new PredictionResult();
// 模型选择与配置
RandomForest rfModel = new RandomForest();
rfModel.setNumIterations(100);
rfModel.setMaxDepth(10);
// 交叉验证
weka.classifiers.Evaluation evaluation = new weka.classifiers.Evaluation(trainingData);
evaluation.crossValidateModel(rfModel, trainingData, 10, new Random(42));
result.setMeanAbsoluteError(evaluation.meanAbsoluteError());
result.setRootMeanSquaredError(evaluation.rootMeanSquaredError());
result.setR2Score(evaluation.correlationCoefficient() * evaluation.correlationCoefficient());
// 重新训练完整模型
rfModel.buildClassifier(trainingData);
// 构建预测实例
Instances unlabeled = new Instances(trainingData);
Instance predictionInstance = new DenseInstance(trainingData.numAttributes());
for (int i = 0; i < predictionFeatures.length; i++) {
predictionInstance.setValue(i, predictionFeatures[i]);
}
unlabeled.add(predictionInstance);
// 预测
double prediction = rfModel.classifyInstance(unlabeled.firstInstance());
result.setPredictedQuantity(prediction);
return result;
}
// 模型评估
public static void evaluateModel(Instances trainingData, Instances testData) throws Exception {
// 多种模型对比
String[] algorithms = {"RandomForest", "LinearRegression", "SMOreg"};
for (String algorithm : algorithms) {
weka.classifiers.Classifier classifier;
switch (algorithm) {
case "RandomForest":
classifier = new RandomForest();
break;
case "LinearRegression":
classifier = new LinearRegression();
break;
case "SMOreg":
classifier = new SMOreg();
break;
default:
throw new IllegalArgumentException("Unknown algorithm");
}
// 训练模型
classifier.buildClassifier(trainingData);
// 评估模型
weka.classifiers.Evaluation evaluation =
new weka.classifiers.Evaluation(trainingData);
evaluation.evaluateModel(classifier, testData);
// 输出评估指标
System.out.println(algorithm + " 模型评估:");
System.out.println("MAE: " + evaluation.meanAbsoluteError());
System.out.println("RMSE: " + evaluation.rootMeanSquaredError());
System.out.println("R2: " + evaluation.correlationCoefficient());
System.out.println("----------------------------------");
}
}
}
4 主程序集成
import java.util.*;
public class PredictionMain {
public static void main(String[] args) {
try {
// 1. 加载数据
System.out.println("开始加载历史数据...");
List<SalesRecord> historicalData =
DataPreprocessor.loadAndCleanData("sales_data.csv");
// 2. 特征工程
System.out.println("提取特征...");
Map<String, double[]> features =
FeatureEngineering.extractFeatures(historicalData);
// 3. 准备目标值(预测未来7天销量)
Map<String, List<Double>> targets = prepareTargets(historicalData);
// 4. 构建训练数据集
Instances trainingData =
SalesPredictor.buildDataset(features, targets);
// 5. 划分训练集和测试集
trainingData.randomize(new Random(42));
int trainSize = (int) Math.round(trainingData.numInstances() * 0.8);
int testSize = trainingData.numInstances() - trainSize;
Instances trainSet = new Instances(trainingData, 0, trainSize);
Instances testSet = new Instances(trainingData, trainSize, testSize);
// 6. 模型训练与评估
System.out.println("开始模型训练...");
SalesPredictor.evaluateModel(trainSet, testSet);
// 7. 对新产品进行预测
double[] newProductFeatures = generatePredictionFeatures();
PredictionResult result =
SalesPredictor.trainAndPredict(trainSet, newProductFeatures);
System.out.println("预测结果:");
System.out.println("预测销量: " + result.getPredictedQuantity());
System.out.println("MAE: " + result.getMeanAbsoluteError());
} catch (Exception e) {
e.printStackTrace();
}
}
private static Map<String, List<Double>> prepareTargets(List<SalesRecord> records) {
Map<String, List<Double>> targets = new HashMap<>();
// 实现目标值准备逻辑
// 可以预测未来7天、14天、30天的销量
return targets;
}
private static double[] generatePredictionFeatures() {
// 生成预测特征
return new double[15];
}
}
5 数据类定义
public class SalesRecord {
private String date;
private String productId;
private int quantity;
private double price;
private double revenue;
// Getters and Setters
public String getDate() { return date; }
public void setDate(String date) { this.date = date; }
public String getProductId() { return productId; }
public void setProductId(String productId) { this.productId = productId; }
public int getQuantity() { return quantity; }
public void setQuantity(int quantity) { this.quantity = quantity; }
public double getPrice() { return price; }
public void setPrice(double price) { this.price = price; }
public double getRevenue() { return revenue; }
public void setRevenue(double revenue) { this.revenue = revenue; }
}
public class PredictionResult {
private double predictedQuantity;
private double meanAbsoluteError;
private double rootMeanSquaredError;
private double r2Score;
// Getters and Setters
}
高级优化建议
1 模型优化技巧
public class ModelOptimizer {
// 网格搜索调参
public static void gridSearch(Instances data) throws Exception {
double[] minErrors = new double[]{Double.MAX_VALUE, Double.MAX_VALUE};
int[] bestParams = new int[2];
// 调整RandomForest参数
for (int iterations = 50; iterations <= 200; iterations += 50) {
for (int depth = 5; depth <= 20; depth += 5) {
RandomForest rf = new RandomForest();
rf.setNumIterations(iterations);
rf.setMaxDepth(depth);
weka.classifiers.Evaluation eval =
new weka.classifiers.Evaluation(data);
eval.crossValidateModel(rf, data, 5, new Random(42));
if (eval.rootMeanSquaredError() < minErrors[0]) {
minErrors[0] = eval.rootMeanSquaredError();
bestParams[0] = iterations;
bestParams[1] = depth;
}
}
}
System.out.println("最佳参数:iterations=" + bestParams[0] + ", depth=" + bestParams[1]);
}
// 特征重要性分析
public static void featureImportance(Instances data, RandomForest model) throws Exception {
RandomForest rf = new RandomForest();
rf.buildClassifier(data);
// 使用信息增益评估特征重要性
weka.attributeSelection.InfoGainAttributeEval eval =
new weka.attributeSelection.InfoGainAttributeEval();
eval.buildEvaluator(data);
for (int i = 0; i < data.numAttributes() - 1; i++) {
double gain = eval.evaluateAttribute(i);
System.out.println("Feature " + i + " importance: " + gain);
}
}
}
2 实时预测接口
@RestController
@RequestMapping("/api/prediction")
public class PredictionController {
private SalesPredictor predictor;
@PostMapping("/sales")
public ResponseEntity<Map<String, Object>> predictSales(@RequestBody PredictionRequest request) {
try {
// 构建预测特征
double[] features = buildFeaturesFromRequest(request);
// 加载历史数据
List<SalesRecord> historicalData =
DataPreprocessor.loadAndCleanData("sales_data.csv");
// 构建训练集
Instances trainingData = buildTrainingSet(historicalData);
// 训练模型并预测
PredictionResult result =
predictor.trainAndPredict(trainingData, features);
Map<String, Object> response = new HashMap<>();
response.put("predictedQuantity", result.getPredictedQuantity());
response.put("confidence", 1 - result.getMeanAbsoluteError() / 100);
response.put("modelMetrics", result);
return ResponseEntity.ok(response);
} catch (Exception e) {
return ResponseEntity.status(500).body(Collections.singletonMap("error", e.getMessage()));
}
}
private double[] buildFeaturesFromRequest(PredictionRequest request) {
return new double[]{
request.getProductId(),
request.getMonth(),
request.getMarketingSpend(),
request.getPromotionFlag(),
request.getHistoricalAvg()
};
}
}
性能优化建议
- 并行处理:使用Java的Parallel Stream或ExecutorService处理大规模数据
- 内存优化:使用数据结构缓存,避免重复计算
- 增量学习:对于持续更新的数据,实现增量模型更新
- 模型持久化:使用Weka的序列化功能保存训练好的模型,避免每次重新训练
部署到生产环境
public class ModelDeployment {
// 模型序列化保存
public static void saveModel(weka.classifiers.Classifier model, String path)
throws Exception {
weka.core.SerializationHelper.write(path, model);
}
// 加载模型
public static weka.classifiers.Classifier loadModel(String path)
throws Exception {
return (weka.classifiers.Classifier)
weka.core.SerializationHelper.read(path);
}
// 定时更新模型
@Scheduled(cron = "0 0 2 * * ?") // 每天凌晨2点更新
public void scheduledModelUpdate() {
try {
System.out.println("开始更新预测模型...");
// 加载最新数据
List<SalesRecord> latestData =
DataPreprocessor.loadAndCleanData("latest_sales_data.csv");
// 特征工程
Map<String, double[]> features =
FeatureEngineering.extractFeatures(latestData);
// 训练模型
Instances trainingData =
SalesPredictor.buildDataset(features, prepareTargets(latestData));
// 训练最佳模型
RandomForest rf = new RandomForest();
rf.buildClassifier(trainingData);
// 保存模型
saveModel(rf, "models/sales_predictor.model");
System.out.println("模型更新完成");
} catch (Exception e) {
e.printStackTrace();
}
}
}
这个完整的案例涵盖了数据预处理、特征工程、模型训练、评估和部署的全流程,你可以根据具体业务需求调整特征选择、模型参数和预测目标。