Java AI预测案例

wen java案例 2

本文目录导读:

Java AI预测案例

  1. 案例:基于线性回归的房价预测
  2. 使用建议

我来为您提供一个Java AI预测的完整案例,使用机器学习模型进行简单的数值预测。

案例:基于线性回归的房价预测

添加Maven依赖

<dependencies>
    <!-- Apache Commons Math 用于数学计算 -->
    <dependency>
        <groupId>org.apache.commons</groupId>
        <artifactId>commons-math3</artifactId>
        <version>3.6.1</version>
    </dependency>
    <!-- 可选:如果使用深度学习 -->
    <dependency>
        <groupId>org.deeplearning4j</groupId>
        <artifactId>deeplearning4j-core</artifactId>
        <version>1.0.0-M1.1</version>
    </dependency>
</dependencies>

简单线性回归预测模型

import org.apache.commons.math3.stat.regression.SimpleRegression;
public class HousePricePredictor {
    private SimpleRegression regression;
    public HousePricePredictor() {
        this.regression = new SimpleRegression();
    }
    // 训练模型(房屋面积 -> 价格)
    public void train(double[][] trainingData) {
        for (double[] data : trainingData) {
            regression.addData(data[0], data[1]);
        }
    }
    // 预测价格
    public double predictPrice(double area) {
        return regression.predict(area);
    }
    // 获取模型评估指标
    public String getModelStats() {
        return String.format(
            "斜率: %.2f, 截距: %.2f, R²: %.4f",
            regression.getSlope(),
            regression.getIntercept(),
            regression.getRSquare()
        );
    }
    public static void main(String[] args) {
        // 训练数据:[面积(平方米), 价格(万元)]
        double[][] houseData = {
            {50, 150}, {80, 240}, {100, 300},
            {120, 360}, {150, 450}, {200, 600}
        };
        HousePricePredictor predictor = new HousePricePredictor();
        predictor.train(houseData);
        // 预测
        double[] areas = {90, 130, 180};
        for (double area : areas) {
            double price = predictor.predictPrice(area);
            System.out.printf("面积 %.0f㎡ -> 预测价格: %.2f万元%n", area, price);
        }
        // 显示模型信息
        System.out.println("\n模型统计: " + predictor.getModelStats());
    }
}

使用DeepLearning4j进行神经网络预测

import org.deeplearning4j.nn.conf.MultiLayerConfiguration;
import org.deeplearning4j.nn.conf.NeuralNetConfiguration;
import org.deeplearning4j.nn.conf.layers.DenseLayer;
import org.deeplearning4j.nn.conf.layers.OutputLayer;
import org.deeplearning4j.nn.multilayer.MultiLayerNetwork;
import org.deeplearning4j.nn.weights.WeightInit;
import org.nd4j.linalg.activations.Activation;
import org.nd4j.linalg.api.ndarray.INDArray;
import org.nd4j.linalg.factory.Nd4j;
import org.nd4j.linalg.learning.config.Adam;
import org.nd4j.linalg.lossfunctions.LossFunctions;
public class NeuralNetworkPredictor {
    private MultiLayerNetwork model;
    public NeuralNetworkPredictor() {
        // 构建神经网络配置
        MultiLayerConfiguration config = new NeuralNetConfiguration.Builder()
            .seed(12345)
            .weightInit(WeightInit.XAVIER)
            .updater(new Adam(0.01))
            .list()
            .layer(0, new DenseLayer.Builder()
                .nIn(1)        // 输入特征数
                .nOut(64)      // 隐藏层神经元
                .activation(Activation.RELU)
                .build())
            .layer(1, new DenseLayer.Builder()
                .nIn(64)
                .nOut(32)
                .activation(Activation.RELU)
                .build())
            .layer(2, new OutputLayer.Builder(
                LossFunctions.LossFunction.MSE)
                .nIn(32)
                .nOut(1)       // 输出1个值
                .activation(Activation.IDENTITY)
                .build())
            .build();
        model = new MultiLayerNetwork(config);
        model.init();
    }
    // 训练模型
    public void train(float[] inputs, float[] outputs, int epochs) {
        INDArray features = Nd4j.create(inputs, new int[]{inputs.length, 1});
        INDArray labels = Nd4j.create(outputs, new int[]{outputs.length, 1});
        for (int i = 0; i < epochs; i++) {
            model.fit(features, labels);
            if (i % 1000 == 0) {
                double loss = model.score();
                System.out.printf("Epoch %d, Loss: %.4f%n", i, loss);
            }
        }
    }
    // 预测
    public float predict(float input) {
        INDArray inputArray = Nd4j.create(new float[]{input}, new int[]{1, 1});
        INDArray output = model.output(inputArray, false);
        return output.getFloat(0);
    }
    public static void main(String[] args) {
        NeuralNetworkPredictor predictor = new NeuralNetworkPredictor();
        // 准备训练数据(非线性关系)
        float[] areas = {50, 80, 100, 120, 150, 200};
        float[] prices = {150, 280, 400, 520, 700, 1000};
        // 训练模型
        System.out.println("开始训练神经网络...");
        predictor.train(areas, prices, 5000);
        // 预测
        float[] testAreas = {90, 130, 180};
        System.out.println("\n预测结果:");
        for (float area : testAreas) {
            float predictedPrice = predictor.predict(area);
            System.out.printf("面积 %.0f㎡ -> 预测价格: %.2f万元%n", area, predictedPrice);
        }
    }
}

更完整的AI预测框架

import java.util.*;
import java.util.stream.*;
public class AdvancedPredictor {
    // 数据归一化
    public static class DataNormalizer {
        private double min, max;
        public void fit(double[] data) {
            min = Arrays.stream(data).min().orElse(0);
            max = Arrays.stream(data).max().orElse(1);
        }
        public double normalize(double value) {
            return (value - min) / (max - min);
        }
        public double denormalize(double value) {
            return value * (max - min) + min;
        }
    }
    // KNN预测器
    public static class KNNPredictor {
        private List<double[]> trainingData;
        private int k;
        public KNNPredictor(int k) {
            this.k = k;
            this.trainingData = new ArrayList<>();
        }
        public void addTrainingData(double[] data) {
            trainingData.add(data);
        }
        public double predict(double[] features) {
            // 计算所有距离
            List<double[]> distances = trainingData.stream()
                .map(data -> new double[]{
                    euclideanDistance(features, Arrays.copyOf(data, data.length - 1)),
                    data[data.length - 1]  // 标签值
                })
                .sorted(Comparator.comparingDouble(a -> a[0]))
                .collect(Collectors.toList());
            // 取前k个最近邻的平均值
            return distances.stream()
                .limit(k)
                .mapToDouble(d -> d[1])
                .average()
                .orElse(0);
        }
        private double euclideanDistance(double[] a, double[] b) {
            double sum = 0;
            for (int i = 0; i < a.length; i++) {
                sum += Math.pow(a[i] - b[i], 2);
            }
            return Math.sqrt(sum);
        }
    }
    public static void main(String[] args) {
        // 示例:使用KNN预测
        System.out.println("=== KNN预测示例 ===");
        KNNPredictor knn = new KNNPredictor(3);
        // 准备训练数据(特征:面积,卧室数;标签:价格)
        knn.addTrainingData(new double[]{80, 2, 200});
        knn.addTrainingData(new double[]{100, 3, 300});
        knn.addTrainingData(new double[]{120, 3, 360});
        knn.addTrainingData(new double[]{150, 4, 500});
        knn.addTrainingData(new double[]{200, 4, 650});
        // 预测
        double[] testFeatures = {110, 3};
        double predictedPrice = knn.predict(testFeatures);
        System.out.printf("面积110㎡, 3卧室 -> 预测价格: %.2f万元%n", predictedPrice);
        // 示例:数据归一化
        System.out.println("\n=== 数据归一化示例 ===");
        DataNormalizer normalizer = new DataNormalizer();
        double[] prices = {150, 280, 400, 520, 700, 1000};
        normalizer.fit(prices);
        System.out.println("原始数据: " + Arrays.toString(prices));
        double[] normalized = Arrays.stream(prices)
            .map(normalizer::normalize)
            .toArray();
        System.out.println("归一化后: " + Arrays.toString(normalized));
    }
}

模型性能评估

public class ModelEvaluator {
    // 计算均方误差
    public static double meanSquaredError(double[] actual, double[] predicted) {
        double sum = 0;
        for (int i = 0; i < actual.length; i++) {
            sum += Math.pow(actual[i] - predicted[i], 2);
        }
        return sum / actual.length;
    }
    // 计算平均绝对误差
    public static double meanAbsoluteError(double[] actual, double[] predicted) {
        double sum = 0;
        for (int i = 0; i < actual.length; i++) {
            sum += Math.abs(actual[i] - predicted[i]);
        }
        return sum / actual.length;
    }
    // 计算R²决定系数
    public static double rSquared(double[] actual, double[] predicted) {
        double mean = Arrays.stream(actual).average().orElse(0);
        double ssRes = 0;  // 残差平方和
        double ssTot = 0;  // 总平方和
        for (int i = 0; i < actual.length; i++) {
            ssRes += Math.pow(actual[i] - predicted[i], 2);
            ssTot += Math.pow(actual[i] - mean, 2);
        }
        return 1 - (ssRes / ssTot);
    }
    // 交叉验证
    public static double crossValidation(List<double[]> data, int folds) {
        Collections.shuffle(data);
        double totalScore = 0;
        int foldSize = data.size() / folds;
        for (int i = 0; i < folds; i++) {
            // 分割训练集和测试集
            List<double[]> testSet = data.subList(i * foldSize, (i + 1) * foldSize);
            List<double[]> trainSet = new ArrayList<>(data);
            trainSet.removeAll(testSet);
            // 训练模型
            KNNPredictor predictor = new KNNPredictor(3);
            for (double[] trainData : trainSet) {
                predictor.addTrainingData(trainData);
            }
            // 评估
            double[] actual = new double[testSet.size()];
            double[] predicted = new double[testSet.size()];
            for (int j = 0; j < testSet.size(); j++) {
                double[] testData = testSet.get(j);
                double[] features = Arrays.copyOf(testData, testData.length - 1);
                actual[j] = testData[testData.length - 1];
                predicted[j] = predictor.predict(features);
            }
            totalScore += rSquared(actual, predicted);
        }
        return totalScore / folds;
    }
}

使用建议

  1. 数据预处理:确保数据质量,处理缺失值和异常值
  2. 特征工程:选择相关特征,进行特征缩放
  3. 模型选择:根据问题复杂度选择合适的算法
  4. 评估验证:使用交叉验证评估模型性能
  5. 调参优化:调整超参数以获得更好的预测效果

这个案例包含了从基础统计到神经网络的不同预测方法,您可以根据具体需求选择合适的实现。

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