java案例如何利用友谊赛数据做预测?

wen java案例 6

本文目录导读:

java案例如何利用友谊赛数据做预测?

  1. 完整Java预测案例
  2. 关键点说明

我来介绍一个利用友谊赛数据做足球预测的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>

关键点说明

  1. 数据准备:需要收集足够的友谊赛历史数据
  2. 特征工程:包括球队胜率、进球数、失球数、历史交锋等
  3. 模型选择:使用多层感知机(MLP)进行分类预测
  4. 评估验证:建议使用交叉验证评估模型性能
  5. 持续更新:模型需要定期用新数据重新训练

这个案例展示了完整的预测流程,实际应用中可能需要根据具体需求调整特征和模型参数。

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