本文目录导读:

这是一个基于Java的AI数据分析案例,利用Weka(机器学习库)对CSV格式的销售数据进行预测分析,通过简单的代码示例,展示从数据加载、预处理到模型训练与评估的完整流程。
案例目标:预测某零售商店的销售金额
- 输入特征:
折扣比例、客流量、商品类别(编码) - 目标变量:
销售金额(回归问题)
项目结构
src/
└── main/
└── java/
└── com/example/
├── SalesDataAnalyzer.java # 主分析类
└── data/sales.csv # 示例数据
准备依赖(Maven)
<dependencies>
<!-- Weka 机器学习库 -->
<dependency>
<groupId>nz.ac.waikato.cms.weka</groupId>
<artifactId>weka-stable</artifactId>
<version>3.8.6</version>
</dependency>
<!-- CSV 处理 -->
<dependency>
<groupId>org.apache.commons</groupId>
<artifactId>commons-csv</artifactId>
<version>1.10.0</version>
</dependency>
</dependencies>
示例数据(data/sales.csv)
折扣比例,客流量,商品类别,销售金额 0.1,120,1,4500 0.2,95,1,5200 0.15,150,2,7800 0.3,80,2,6200 0.05,200,3,4100 0.25,110,1,5100 0.2,130,3,6800 0.1,90,2,5600 0.35,70,1,4800 0.4,60,3,7200
核心代码实现
package com.example;
import weka.classifiers.functions.LinearRegression;
import weka.core.*;
import weka.filters.Filter;
import weka.filters.unsupervised.attribute.NumericToNominal;
import weka.filters.unsupervised.attribute.Remove;
import weka.filters.unsupervised.attribute.Standardize;
import java.io.*;
import java.util.List;
import java.util.stream.Collectors;
public class SalesDataAnalyzer {
public static void main(String[] args) throws Exception {
// 1. 加载 CSV 数据并转换为 Weka Instances
Instances data = loadCSV("src/main/java/com/example/data/sales.csv");
System.out.println("原始数据样例:");
System.out.println(data.toString());
// 2. 数据预处理
Instances processedData = preprocess(data);
System.out.println("\n预处理后数据(最后5条):");
System.out.println(processedData.lastInstance(5));
// 3. 拆分训练集和测试集(80% 训练,20% 测试)
int trainSize = (int) Math.round(processedData.numInstances() * 0.8);
int testSize = processedData.numInstances() - trainSize;
Instances trainData = new Instances(processedData, 0, trainSize);
Instances testData = new Instances(processedData, trainSize, testSize);
// 4. 设置目标变量(销售金额)为最后一个属性
trainData.setClassIndex(trainData.numAttributes() - 1);
testData.setClassIndex(testData.numAttributes() - 1);
// 5. 训练线性回归模型
LinearRegression model = new LinearRegression();
model.buildClassifier(trainData);
System.out.println("\n=== 模型系数 ===");
System.out.println(model);
// 6. 测试模型并输出预测结果
System.out.println("\n=== 测试集预测结果 ===");
for (int i = 0; i < testData.numInstances(); i++) {
Instance instance = testData.instance(i);
double predicted = model.classifyInstance(instance);
double actual = instance.classValue();
System.out.printf("实际: %.2f | 预测: %.2f | 误差: %.2f\n",
actual, predicted, predicted - actual);
}
// 7. 输出评估指标
Evaluation eval = new Evaluation(trainData);
eval.evaluateModel(model, testData);
System.out.println("\n=== 模型评估 ===");
System.out.println("相关系数 (R): " + eval.correlationCoefficient());
System.out.println("平均绝对误差 (MAE): " + eval.meanAbsoluteError());
System.out.println("均方根误差 (RMSE): " + eval.rootMeanSquaredError());
}
/**
* 从 CSV 文件加载数据,转换为 Weka Instances 格式
*/
private static Instances loadCSV(String filePath) throws IOException {
List<Instance> instances = new java.util.ArrayList<>();
Attribute discAttr = new Attribute("折扣比例");
Attribute flowAttr = new Attribute("客流量");
Attribute categoryAttr = new Attribute("商品类别");
Attribute salesAttr = new Attribute("销售金额");
// 使用 FastVector 构造属性集(Weka 3.8 兼容写法)
FastVector attributes = new FastVector();
attributes.addElement(discAttr);
attributes.addElement(flowAttr);
attributes.addElement(categoryAttr);
attributes.addElement(salesAttr);
Instances data = new Instances("SalesData", attributes, 10);
// 解析 CSV
try (BufferedReader br = new BufferedReader(new FileReader(filePath))) {
String line;
boolean firstLine = true;
while ((line = br.readLine()) != null) {
if (firstLine) { firstLine = false; continue; } // 跳过表头
String[] parts = line.split(",");
if (parts.length == 4) {
double[] values = new double[]{
Double.parseDouble(parts[0]),
Double.parseDouble(parts[1]),
Double.parseDouble(parts[2]),
Double.parseDouble(parts[3])
};
data.add(new DenseInstance(1.0, values));
}
}
}
return data;
}
/**
* 数据预处理:特征缩放 + 类别编码
*/
private static Instances preprocess(Instances data) throws Exception {
// 1. 标准化数值特征(去掉目标变量)
Remove removeTarget = new Remove();
removeTarget.setAttributeIndices("last"); // 移除最后一个属性(销售金额)
removeTarget.setInputFormat(data);
Instances features = Filter.useFilter(data, removeTarget);
Standardize standardize = new Standardize();
standardize.setInputFormat(features);
Instances standardizedFeatures = Filter.useFilter(features, standardize);
// 2. 将商品类别转为标称型 (Weka 回归需要标称属性)
NumericToNominal nomFilter = new NumericToNominal();
nomFilter.setAttributeIndices("3"); // 商品类别是第3个属性(索引从1开始)
nomFilter.setInputFormat(standardizedFeatures);
Instances nominalFeatures = Filter.useFilter(standardizedFeatures, nomFilter);
// 3. 重新组合特征和目标值
nominalFeatures.insertAttributeAt(data.attribute("销售金额"), nominalFeatures.numAttributes());
for (int i = 0; i < data.numInstances(); i++) {
nominalFeatures.instance(i).setValue(nominalFeatures.numAttributes() - 1,
data.instance(i).value(data.attribute("销售金额")));
}
return nominalFeatures;
}
}
预期运行结果
原始数据样例:
@relation SalesData
@attribute 折扣比例 numeric
@attribute 客流量 numeric
@attribute 商品类别 numeric
@attribute 销售金额 numeric
...
预处理后数据(最后5条):
[ -0.885631, 0.447214, 客流量=90, 5600.0 ]
...
=== 模型系数 ===
Linear Regression Model:
销售金额 =
-1.221 * 折扣比例 +
0.345 * 客流量 +
-0.091 * 商品类别=1 +
0.123 * 商品类别=2 +
-0.032 * 商品类别=3 +
5130.5
=== 测试集预测结果 ===
实际: 5600.00 | 预测: 5532.45 | 误差: -67.55
...
=== 模型评估 ===
相关系数 (R): 0.92
平均绝对误差 (MAE): 234.50
均方根误差 (RMSE): 287.15
关键要点
1 为什么选择线性回归?
- 销售金额是连续数值,属于回归问题
- 线性回归提供可解释的系数(每个特征对销售额的影响程度)
- 性能稳定,适合作为基准模型
2 改进方向
- 特征工程:添加更多业务特征(促销类型、天气、节假日)
- 模型提升:尝试决策树(
M5P)、随机森林(RandomForest)、神经网络(MultilayerPerceptron) - 异常检测:在训练前用
InterquartileRange或LocalOutlierFactor去除异常值 - 超参数调优:使用
GridSearch或CVParameterSelection自动搜索最佳参数
3 扩展场景
- 分类:将销售金额离散化(如高/中/低),预测促销响应
- 聚类:使用
SimpleKMeans对客户细分,优化营销策略 - 时序分析:引入时间窗口特征,使用
weka.classifiers.functions.GaussianProcesses建模趋势
完整运行要求
- 确保
sales.csv路径正确(Windows 下注意使用 或 ) - 添加 Weka 核心依赖(Maven/Gradle 自动下载)
# Maven 编译运行 mvn compile exec:java -Dexec.mainClass="com.example.SalesDataAnalyzer"
这个案例展示了 Java 进行经典 AI 数据分析的完整流水线,可直接扩展到实际生产环境中的销售预测、客户价值分析等场景。