java案例如何利用历史大数据建模预测?

wen java案例 2

本文目录导读:

java案例如何利用历史大数据建模预测?

  1. 整体架构设计
  2. 完整案例:电商销售预测
  3. 高级优化建议
  4. 性能优化建议
  5. 部署到生产环境

我来给你一个完整的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();
        }
    }
}

这个完整的案例涵盖了数据预处理、特征工程、模型训练、评估和部署的全流程,你可以根据具体业务需求调整特征选择、模型参数和预测目标。

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