本文目录导读:

我来介绍一个利用友谊赛数据做足球预测的Java案例,我会设计一个相对完整的示例,包含数据准备、特征提取、模型训练和预测。
完整Java预测案例
数据模型类
import java.time.LocalDate;
import java.util.ArrayList;
import java.util.List;
// 比赛数据模型
public class Match {
private LocalDate date;
private String homeTeam;
private String awayTeam;
private int homeGoals;
private int awayGoals;
private String competition; // "FRIENDLY" 或其他
// 构造函数、getter、setter
public Match(LocalDate date, String homeTeam, String awayTeam,
int homeGoals, int awayGoals, String competition) {
this.date = date;
this.homeTeam = homeTeam;
this.awayTeam = awayTeam;
this.homeGoals = homeGoals;
this.awayGoals = awayGoals;
this.competition = competition;
}
// 判断是否友谊赛
public boolean isFriendly() {
return "FRIENDLY".equalsIgnoreCase(competition);
}
// getter和setter方法
public LocalDate getDate() { return date; }
public String getHomeTeam() { return homeTeam; }
public String getAwayTeam() { return awayTeam; }
public int getHomeGoals() { return homeGoals; }
public int getAwayGoals() { return awayGoals; }
public String getCompetition() { return competition; }
}
特征工程处理器
import java.util.HashMap;
import java.util.Map;
public class FeatureExtractor {
private Map<String, TeamStats> teamStatsMap;
public FeatureExtractor() {
this.teamStatsMap = new HashMap<>();
}
// 提取特征向量
public double[] extractFeatures(Match match, List<Match> historicalMatches) {
double[] features = new double[8];
// 1. 主队近期友谊赛胜率
features[0] = calculateRecentWinRate(match.getHomeTeam(), historicalMatches);
// 2. 客队近期友谊赛胜率
features[1] = calculateRecentWinRate(match.getAwayTeam(), historicalMatches);
// 3. 主队场均进球
features[2] = calculateAverageGoals(match.getHomeTeam(), historicalMatches, true);
// 4. 客队场均进球
features[3] = calculateAverageGoals(match.getAwayTeam(), historicalMatches, true);
// 5. 主队场均失球
features[4] = calculateAverageGoals(match.getHomeTeam(), historicalMatches, false);
// 6. 客队场均失球
features[5] = calculateAverageGoals(match.getAwayTeam(), historicalMatches, false);
// 7. 两队历史交锋胜负记录
features[6] = calculateHeadToHead(match, historicalMatches);
// 8. 主队是否拥有主场优势(友谊赛可能无主场优势,但保留)
features[7] = 1.0; // 默认有主场优势
return features;
}
// 计算近期胜率
private double calculateRecentWinRate(String team, List<Match> matches) {
List<Match> teamMatches = getTeamFriendlyMatches(team, matches);
if (teamMatches.isEmpty()) return 0.5; // 无数据时返回50%
long wins = teamMatches.stream()
.filter(m -> (m.getHomeTeam().equals(team) && m.getHomeGoals() > m.getAwayGoals()) ||
(m.getAwayTeam().equals(team) && m.getAwayGoals() > m.getHomeGoals()))
.count();
return (double) wins / teamMatches.size();
}
// 计算平均进球或失球
private double calculateAverageGoals(String team, List<Match> matches, boolean goalsScored) {
List<Match> teamMatches = getTeamFriendlyMatches(team, matches);
if (teamMatches.isEmpty()) return 1.0;
int totalGoals = 0;
for (Match m : teamMatches) {
if (m.getHomeTeam().equals(team)) {
totalGoals += goalsScored ? m.getHomeGoals() : m.getAwayGoals();
} else {
totalGoals += goalsScored ? m.getAwayGoals() : m.getHomeGoals();
}
}
return (double) totalGoals / teamMatches.size();
}
// 计算历史交锋
private double calculateHeadToHead(Match match, List<Match> matches) {
List<Match> h2hMatches = matches.stream()
.filter(m -> (m.getHomeTeam().equals(match.getHomeTeam()) &&
m.getAwayTeam().equals(match.getAwayTeam())) ||
(m.getHomeTeam().equals(match.getAwayTeam()) &&
m.getAwayTeam().equals(match.getHomeTeam())))
.limit(5)
.toList();
if (h2hMatches.isEmpty()) return 0.5;
long homeWins = h2hMatches.stream()
.filter(m -> m.getHomeGoals() > m.getAwayGoals())
.count();
return (double) homeWins / h2hMatches.size();
}
// 获取球队的友谊赛列表
private List<Match> getTeamFriendlyMatches(String team, List<Match> matches) {
return matches.stream()
.filter(m -> m.isFriendly() &&
(m.getHomeTeam().equals(team) || m.getAwayTeam().equals(team)))
.limit(20) // 最近20场友谊赛
.toList();
}
// 内部类存储球队统计数据
private class TeamStats {
int matchesPlayed;
int wins;
int draws;
int losses;
int goalsScored;
int goalsConceded;
double winRate;
double avgGoalsScored;
double avgGoalsConceded;
}
}
机器学习预测模型
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 FriendlyMatchPredictor {
private MultiLayerNetwork model;
private FeatureExtractor featureExtractor;
public FriendlyMatchPredictor() {
this.featureExtractor = new FeatureExtractor();
initializeModel();
}
// 初始化神经网络模型
private void initializeModel() {
int inputSize = 8; // 特征数量
int hiddenSize1 = 16;
int hiddenSize2 = 8;
int outputSize = 3; // 胜/平/负
MultiLayerConfiguration config = new NeuralNetConfiguration.Builder()
.weightInit(WeightInit.XAVIER)
.updater(new Adam(0.001))
.list()
.layer(0, new DenseLayer.Builder()
.nIn(inputSize)
.nOut(hiddenSize1)
.activation(Activation.RELU)
.build())
.layer(1, new DenseLayer.Builder()
.nIn(hiddenSize1)
.nOut(hiddenSize2)
.activation(Activation.RELU)
.build())
.layer(2, new OutputLayer.Builder(LossFunctions.LossFunction.NEGATIVELOGLIKELIHOOD)
.nIn(hiddenSize2)
.nOut(outputSize)
.activation(Activation.SOFTMAX)
.build())
.build();
model = new MultiLayerNetwork(config);
model.init();
}
// 训练模型
public void trainModel(List<Match> matches, double learningRate) {
List<Match> friendlyMatches = matches.stream()
.filter(Match::isFriendly)
.toList();
// 准备训练数据
List<double[]> inputData = new ArrayList<>();
List<Integer> labels = new ArrayList<>();
for (int i = 10; i < friendlyMatches.size(); i++) {
Match match = friendlyMatches.get(i);
List<Match> historicalMatches = friendlyMatches.subList(0, i);
// 提取特征
double[] features = featureExtractor.extractFeatures(match, historicalMatches);
inputData.add(features);
// 确定标签(1-主胜,2-平局,0-客胜)
int label;
if (match.getHomeGoals() > match.getAwayGoals()) {
label = 1; // 主胜
} else if (match.getHomeGoals() == match.getAwayGoals()) {
label = 2; // 平局
} else {
label = 0; // 客胜
}
labels.add(label);
}
// 转换为ND4J格式
double[][] featuresArray = inputData.toArray(new double[0][]);
INDArray inputMatrix = Nd4j.create(featuresArray);
INDArray labelsMatrix = Nd4j.zeros(inputData.size(), 3);
for (int i = 0; i < labels.size(); i++) {
labelsMatrix.putScalar(new int[]{i, labels.get(i)}, 1.0);
}
// 训练模型
for (int epoch = 0; epoch < 100; epoch++) {
model.fit(inputMatrix, labelsMatrix);
}
}
// 预测比赛结果
public PredictionResult predictMatch(Match match, List<Match> historicalMatches) {
double[] features = featureExtractor.extractFeatures(match, historicalMatches);
INDArray input = Nd4j.create(new double[][]{features});
INDArray output = model.output(input);
double homeWinProb = output.getDouble(0, 0);
double drawProb = output.getDouble(0, 1);
double awayWinProb = output.getDouble(0, 2);
// 确定预测结果
String prediction;
double confidence;
if (homeWinProb >= drawProb && homeWinProb >= awayWinProb) {
prediction = "主胜";
confidence = homeWinProb;
} else if (drawProb >= homeWinProb && drawProb >= awayWinProb) {
prediction = "平局";
confidence = drawProb;
} else {
prediction = "客胜";
confidence = awayWinProb;
}
return new PredictionResult(prediction, confidence,
homeWinProb, drawProb, awayWinProb);
}
// 预测结果类
public static class PredictionResult {
private String prediction;
private double confidence;
private double homeWinProb;
private double drawProb;
private double awayWinProb;
public PredictionResult(String prediction, double confidence,
double homeWinProb, double drawProb, double awayWinProb) {
this.prediction = prediction;
this.confidence = confidence;
this.homeWinProb = homeWinProb;
this.drawProb = drawProb;
this.awayWinProb = awayWinProb;
}
// getter方法
public String getPrediction() { return prediction; }
public double getConfidence() { return confidence; }
public double getHomeWinProb() { return homeWinProb; }
public double getDrawProb() { return drawProb; }
public double getAwayWinProb() { return awayWinProb; }
@Override
public String toString() {
return String.format("预测: %s | 置信度: %.1f%% | 主胜: %.1f%% 平局: %.1f%% 客胜: %.1f%%",
prediction, confidence * 100, homeWinProb * 100,
drawProb * 100, awayWinProb * 100);
}
}
}
主程序和测试示例
import java.time.LocalDate;
import java.util.ArrayList;
import java.util.List;
import java.util.Random;
public class Main {
public static void main(String[] args) {
// 模拟友谊赛数据
List<Match> friendlyMatches = generateFriendlyMatches();
// 创建预测器
FriendlyMatchPredictor predictor = new FriendlyMatchPredictor();
// 训练模型
System.out.println("开始训练模型...");
predictor.trainModel(friendlyMatches, 0.01);
System.out.println("模型训练完成!\n");
// 预测新比赛
Match newMatch = new Match(
LocalDate.now(),
"中国",
"日本",
0, 0,
"FRIENDLY"
);
// 使用历史数据预测
FriendlyMatchPredictor.PredictionResult result =
predictor.predictMatch(newMatch, friendlyMatches);
System.out.println("预测比赛: 中国 VS 日本");
System.out.println(result);
}
// 生成模拟友谊赛数据
private static List<Match> generateFriendlyMatches() {
List<Match> matches = new ArrayList<>();
Random random = new Random(42); // 固定种子便于复现
String[] teams = {"中国", "日本", "韩国", "澳大利亚", "伊朗", "沙特", "巴西", "德国"};
for (int i = 0; i < 500; i++) {
String homeTeam = teams[random.nextInt(teams.length)];
String awayTeam;
do {
awayTeam = teams[random.nextInt(teams.length)];
} while (awayTeam.equals(homeTeam));
int homeGoals = random.nextInt(4);
int awayGoals = random.nextInt(4);
// 30%概率设置为主场优势
if (random.nextDouble() < 0.3) {
homeGoals = Math.max(homeGoals, awayGoals + random.nextInt(2));
}
LocalDate date = LocalDate.now().minusDays(i * 20);
matches.add(new Match(date, homeTeam, awayTeam,
homeGoals, awayGoals, "FRIENDLY"));
}
return matches;
}
}
依赖配置(Maven)
<dependencies>
<!-- DL4J 深度学习库 -->
<dependency>
<groupId>org.deeplearning4j</groupId>
<artifactId>deeplearning4j-core</artifactId>
<version>1.0.0-M2.1</version>
</dependency>
<!-- ND4J 数值计算库 -->
<dependency>
<groupId>org.nd4j</groupId>
<artifactId>nd4j-native-platform</artifactId>
<version>1.0.0-M2.1</version>
</dependency>
<!-- 数据处理 -->
<dependency>
<groupId>org.apache.commons</groupId>
<artifactId>commons-lang3</artifactId>
<version>3.12.0</version>
</dependency>
</dependencies>
关键点说明
- 数据准备:需要收集足够的友谊赛历史数据
- 特征工程:包括球队胜率、进球数、失球数、历史交锋等
- 模型选择:使用多层感知机(MLP)进行分类预测
- 评估验证:建议使用交叉验证评估模型性能
- 持续更新:模型需要定期用新数据重新训练
这个案例展示了完整的预测流程,实际应用中可能需要根据具体需求调整特征和模型参数。